diff --git a/crates/tw-api/msg-codes.txt b/crates/tw-api/msg-codes.txt index d67fd24..b78a5fd 100644 --- a/crates/tw-api/msg-codes.txt +++ b/crates/tw-api/msg-codes.txt @@ -288,6 +288,7 @@ gw.plugin.failed gw.plugin.file_changed gw.plugin.manifest gw.plugin.memory_limit +gw.plugin.model_not_allowed gw.plugin.not_located gw.plugin.nothing_to_try gw.plugin.output_limit diff --git a/crates/tw-api/src/lib.rs b/crates/tw-api/src/lib.rs index 5686737..b208e9e 100644 --- a/crates/tw-api/src/lib.rs +++ b/crates/tw-api/src/lib.rs @@ -678,9 +678,10 @@ pub const MSG_CODES: &str = include_str!("../msg-codes.txt"); /// /// **33 起有脚本插件**:`/plugins` 一组端点(列表、试编、装、改、换源码、看改动、批准、 /// 排顺序、删、试跑、日志),事件多了 [`Event::PluginFailed`](插件在请求上出错,或者 -/// 文件变了、加载不了而停用),[`RequestDetail`] 多了 `plugins`(每一次运行)和 -/// `request_after_plugins`(插件改过的请求体),[`HistoryRow`] 多了 `plugin_changed`。 -/// 装、换源码、批准三个端点不给网页调:要在系统的确认框里点头。照 32 写的界面看不到插件。 +/// 文件变了、加载不了而停用),[`RequestDetail`] 多了 `plugins`(每一次运行,带着跑在 +/// 尝试链的第几跳)和 `request_after_plugins`(插件改过的请求体),[`HistoryRow`] 多了 +/// `plugin_changed`。装、换源码、批准三个端点不给网页调:要在系统的确认框里点头。照 32 +/// 写的界面看不到插件。 pub const CONTROL_API_VERSION: u32 = 33; #[derive(Debug, Clone, Serialize, Deserialize)] @@ -3696,10 +3697,12 @@ pub struct RequestDetail { pub row: HistoryRow, /// 客户端发来的原样 pub request_body: Option, - /// 插件改过之后、发往上游的那一份。**只有插件改了请求才有** + /// 插件改过之后、发往上游的那一份:最后发出去的那一跳收到的(回答的那一家收到的就是 + /// 它)。**只有插件改了那一跳的请求才有** pub request_after_plugins: Option, pub response_body: Option, - /// 插件在这个请求上的每一次运行,按先后(请求钩子在前,回答钩子在后) + /// 插件在这个请求上的每一次运行,按先后:每一跳的请求钩子,回答那一跳的回答钩子。 + /// 按 [`PluginRunView::attempt`] 对着尝试链分组 pub plugins: Vec, /// 这个请求还在跑。**记录在结局到了才落库**,这时的 `row` 是到目前为止 /// 知道的那些:开始时的身份和上游,响应头到了就有状态码,路由走完就有 @@ -4990,9 +4993,9 @@ impl SettingValue { pub struct PluginScope { /// 客户端应用:`claude-code`、`codex`……(请求记录上的 `client_hint`) pub clients: Vec, - /// 客户端要的模型 + /// 发给上游的模型:路由规则改了名的,按改名之后的 pub models: Vec, - /// 服务回答的上游。**只管回答那一段**:改请求时还没选上游 + /// 发往的上游。**请求和回答都按它**:请求钩子排在路由之后,每发往一个上游跑一次 pub upstreams: Vec, } @@ -5220,6 +5223,9 @@ pub struct PluginRunView { /// 当时的名字。**插件写的字** pub plugin_name: String, pub hook: PluginHook, + /// 跑在尝试链上的第几跳(从 0 起,对着 [`RoutingView::attempts`])。请求钩子每发往一个 + /// 上游跑一次,故障转移换了上游就多一组;回答钩子跑在回答的那一跳上 + pub attempt: u32, pub outcome: PluginOutcome, /// 出错、拒绝的原因 pub error: Option, diff --git a/crates/tw-config/src/plugins.rs b/crates/tw-config/src/plugins.rs index fc99388..80de75d 100644 --- a/crates/tw-config/src/plugins.rs +++ b/crates/tw-config/src/plugins.rs @@ -83,10 +83,10 @@ pub struct PluginScope { /// 客户端应用:`claude-code`、`codex`…… #[serde(default, skip_serializing_if = "Vec::is_empty")] pub clients: Vec, - /// 客户端要的模型 + /// 发给上游的模型:路由规则改了名的,按改名之后的 #[serde(default, skip_serializing_if = "Vec::is_empty")] pub models: Vec, - /// 服务回答的上游。只管回答那一段:改请求时还没选上游 + /// 发往的上游。请求和回答都按它:请求钩子每发往一个上游跑一次 #[serde(default, skip_serializing_if = "Vec::is_empty")] pub upstreams: Vec, } diff --git a/crates/tw-config/tests/manual/schema.rs b/crates/tw-config/tests/manual/schema.rs index 0819690..b29acca 100644 --- a/crates/tw-config/tests/manual/schema.rs +++ b/crates/tw-config/tests/manual/schema.rs @@ -1605,8 +1605,8 @@ pub fn sections() -> Vec
{ Kind::Strs, Def::Is("[]"), t( - "Models the client asks for, as model ids or globs (`claude-*`). `[]`: every model.", - "客户端请求的模型,写模型 ID 或通配(`claude-*`)。`[]`:所有模型。", + "Models sent to the upstream, as model ids or globs (`claude-*`). When a routing rule renames the model, the new name is the one that matches. `[]`: every model.", + "发给上游的模型,写模型 ID 或通配(`claude-*`)。路由规则改了模型名的,按改名之后的匹配。`[]`:所有模型。", ), ), row( @@ -1614,8 +1614,8 @@ pub fn sections() -> Vec
{ Kind::Strs, Def::Is("[]"), t( - "Upstreams whose answers the plugin handles, by name or glob. It applies to answers only: a request is changed before an upstream is chosen. `[]`: every upstream.", - "插件处理哪些上游的回答,写名字或通配。只作用于回答:请求在选定上游之前就已改写。`[]`:所有上游。", + "Upstreams the plugin handles, by name or glob, for requests and answers alike. `[]`: every upstream.", + "插件处理哪些上游,写名字或通配,请求和回答都按它。`[]`:所有上游。", ), ), ], diff --git a/crates/tw-control/src/lib.rs b/crates/tw-control/src/lib.rs index cac80af..6663c66 100644 --- a/crates/tw-control/src/lib.rs +++ b/crates/tw-control/src/lib.rs @@ -1056,6 +1056,13 @@ async fn request_detail( .map_err(records)? .into_iter() .map(|r| tw_api::PluginRunView { + // 第几跳记在 `detail` 里(数据面每一次运行都写) + attempt: r + .detail + .as_deref() + .and_then(|d| serde_json::from_str::(d).ok()) + .and_then(|d| d.get("attempt").and_then(serde_json::Value::as_u64)) + .unwrap_or(0) as u32, plugin_id: r.plugin_id, plugin_name: r.plugin_name, hook: r.hook, diff --git a/crates/tw-control/src/plugins.rs b/crates/tw-control/src/plugins.rs index 9b7cadd..a124814 100644 --- a/crates/tw-control/src/plugins.rs +++ b/crates/tw-control/src/plugins.rs @@ -889,7 +889,8 @@ fn refused(why: Msg) -> tw_api::PluginTrialResult { /// 试跑本身在数据面那一侧(视图、写回、占位符都在 [`tw_gateway::plugin::trial`])。 /// /// 存下来的回答是上游的原话:回答它的那一家说什么格式,看服务它的那一跳转换过没有, -/// 和会话记录读回答是同一个办法 +/// 和会话记录读回答是同一个办法。插件的 `ctx` 按这一行的路由给:回答它的那一家,和发给 +/// 那一家的模型名 async fn run_trial( s: &ControlState, active: &Active, @@ -916,6 +917,8 @@ async fn run_trial( query: None, body, client: row.client_hint.as_deref(), + upstream: &row.provider, + sent_model: &row.sent_model, }), reply .as_deref() diff --git a/crates/tw-control/tests/plugins.rs b/crates/tw-control/tests/plugins.rs index aa61ece..ea445a5 100644 --- a/crates/tw-control/tests/plugins.rs +++ b/crates/tw-control/tests/plugins.rs @@ -805,10 +805,17 @@ async fn a_request_shows_its_plugin_runs_and_the_body_after_them() { let g = b.store.lock().await; g.db().insert(&row(1, 1_000)).unwrap(); g.db().insert(&row(2, 2_000)).unwrap(); - g.record_plugin_run(&run_row(1, 1_000, tw_api::PluginOutcome::Changed)); + // 故障转移过一次:第 0 跳、第 1 跳各跑一次请求钩子,回答钩子跑在回答的第 1 跳上 + let mut first = run_row(1, 1_000, tw_api::PluginOutcome::Changed); + first.detail = Some(r#"{"attempt":0,"changed":["system"]}"#.into()); + g.record_plugin_run(&first); + let mut second = run_row(1, 1_100, tw_api::PluginOutcome::Changed); + second.detail = Some(r#"{"attempt":1,"changed":["system"]}"#.into()); + g.record_plugin_run(&second); let mut reply = run_row(1, 1_500, tw_api::PluginOutcome::Error); reply.hook = tw_api::PluginHook::Reply; reply.error = Some(tw_types::msg!("gw.plugin.failed" => "The plugin failed.")); + reply.detail = Some(r#"{"attempt":1,"text_calls":1}"#.into()); g.record_plugin_run(&reply); g.record_plugin_run(&run_row(2, 2_000, tw_api::PluginOutcome::Unchanged)); g.record_body( @@ -824,12 +831,17 @@ async fn a_request_shows_its_plugin_runs_and_the_body_after_them() { let (st, d) = call(&b.app, "GET", "/request/1", None).await; assert_eq!(st, StatusCode::OK, "{d}"); let runs = d["plugins"].as_array().unwrap(); - assert_eq!(runs.len(), 2); + assert_eq!(runs.len(), 3); assert_eq!(runs[0]["hook"], "request"); assert_eq!(runs[0]["outcome"], "changed"); assert_eq!(runs[0]["cpu_us"], 120); - assert_eq!(runs[1]["hook"], "reply"); - assert_eq!(runs[1]["error"]["code"], "gw.plugin.failed"); + let attempts: Vec = runs + .iter() + .map(|r| r["attempt"].as_u64().unwrap()) + .collect(); + assert_eq!(attempts, [0, 1, 1]); + assert_eq!(runs[2]["hook"], "reply"); + assert_eq!(runs[2]["error"]["code"], "gw.plugin.failed"); assert_eq!(d["row"]["plugin_changed"], true); let after = d["request_after_plugins"]["text"].as_str().unwrap(); assert!(after.contains("today is Friday"), "{after}"); diff --git a/crates/tw-gateway/src/bodies.rs b/crates/tw-gateway/src/bodies.rs index 37b7838..1c83f32 100644 --- a/crates/tw-gateway/src/bodies.rs +++ b/crates/tw-gateway/src/bodies.rs @@ -56,9 +56,10 @@ pub enum BodyKind { Request, Response, /// 插件改过之后的请求体(`Request` 存的是客户端发来的那一份)。**只有插件真的改了 - /// 才存**,挨着 `Request` 放。交来的是要发出去的那一份(插件交回的占位符已经换回 - /// 原值,见 [`crate::plugin::request`]),带着这个请求的 [`Redaction`]:落盘前和别的 - /// 正文一样换掉、打码([`BodyRecord::for_disk`]) + /// 才存**,挨着 `Request` 放。请求钩子每一跳跑一次,存的是最后发出去的那一跳收到的 + /// 那一份 —— 回答的那一家收到的就是它(客户端那种格式、转换之前,插件交回的占位符 + /// 已经换回原值,见 [`crate::plugin::request`]),带着那一跳的 [`Redaction`]:落盘前和 + /// 别的正文一样换掉、打码([`BodyRecord::for_disk`]) AfterPlugins, } diff --git a/crates/tw-gateway/src/guard.rs b/crates/tw-gateway/src/guard.rs index f949c70..75c9f34 100644 --- a/crates/tw-gateway/src/guard.rs +++ b/crates/tw-gateway/src/guard.rs @@ -130,6 +130,21 @@ pub fn replace( (bytes::Bytes::from(r.text), r.ledger) } +/// `after` 里 `before` 没有的那些值:插件写进请求里的(见 [`crate::plugin::request`])。 +/// +/// 按规则和打过码的样子比:同一个值在两份里打出来的码一样。客户端原话里就有的值,开头 +/// 那一遍已经报过了,插件改过的那一份里再出现不再报一次。 +pub fn more_found(before: &[Finding], after: Vec) -> Vec { + after + .into_iter() + .filter(|f| { + !before + .iter() + .any(|b| b.rule == f.rule && b.masked == f.masked) + }) + .collect() +} + /// 找到的东西写成事件里的样子。 pub fn items(found: &[Finding]) -> Vec { found @@ -194,6 +209,64 @@ pub fn screen( refusal } +/// 插件改过的请求再看一遍:**只看插件加进来的。** +/// +/// 客户端的原话在开头已经看过([`screen`]),该报的报了、该拒的拒了;插件改过的那一份 +/// 要是整个再报一遍,同一处藏匿字符、同一条命中会在安全日志里出现两次。所以两份都扫, +/// 原话里就有的那几处减掉,剩下的照 [`screen`] 的规矩报、下结论 —— 拦截档下原话里命中 +/// 「拦」的请求走不到这一步,这一遍拒不拒只看插件加进来的。 +pub fn screen_more( + bus: &tw_observe::EventBus, + id: u64, + provider: &str, + s: &Screen, + before: &tw_dialect::ir::Request, + after: &tw_dialect::ir::Request, +) -> Option { + let hidden = if s.hidden_mode.detects() { + let was = tw_guard::hidden::scan_request(before, &s.hidden); + let mut now = tw_guard::hidden::scan_request(after, &s.hidden); + // 一种藏法在一个地方合成一条:插件往同一处又藏了几个,那一条就变了,整条再报 + now.retain(|n| { + !was.iter().any(|w| { + w.kind == n.kind + && w.in_tool_result == n.in_tool_result + && w.example == n.example + && w.revealed == n.revealed + && n.count <= w.count + }) + }); + now + } else { + Vec::new() + }; + let mut refusal = hidden_found(bus, id, provider, s.hidden_mode, &hidden); + if s.content_mode.detects() && !s.content.is_empty() { + let mut was = s.content.scan_request(before); + let mut hits = s.content.scan_request(after); + // 一处一处地减:原话里有一处,插件那一版里同样的一处就不是新的。位置不比 —— + // 插件在前面加了字,后面的位置都挪了 + hits.retain(|h| { + match was.iter().position(|w| { + w.rule == h.rule + && w.custom == h.custom + && w.action == h.action + && w.snippet == h.snippet + && w.in_tool_result == h.in_tool_result + }) { + Some(i) => { + was.swap_remove(i); + false + } + None => true, + } + }); + let refused = content_matched(bus, id, provider, s.content_mode, &hits); + refusal = refusal.or(refused); + } + refusal +} + /// 没法按消息结构读的正文(解不开的 WebSocket 帧):**只查藏匿字符** —— 它在任何 /// 地方都没有正当用途;内容规则按整段原文查的话,系统提示里的话也会被当成调用方的。 pub fn screen_text( diff --git a/crates/tw-gateway/src/plugin/mod.rs b/crates/tw-gateway/src/plugin/mod.rs index 72247d3..c69edec 100644 --- a/crates/tw-gateway/src/plugin/mod.rs +++ b/crates/tw-gateway/src/plugin/mod.rs @@ -19,7 +19,9 @@ //! 什么都没改时一个字节都不动。 //! - [`bridge`]:插件永远看不到真的密钥(不变式 I5)。进插件之前按出站脱敏的规则把 //! 认得出的密钥换成占位符,出来之后换回去;**不看脱敏开在哪一档**。 -//! - [`request`]:请求钩子。一个客户端请求只跑一次(I8),排在内容审查和路由之前(I7)。 +//! - [`request`]:请求钩子。排在路由之后,**每发往一个上游跑一次**(契约附录二的 I7、 +//! I8):按这一次的客户端、发出去的模型和上游挑插件,从客户端的原话起改;换上游从 +//! 原话重来,同一家重发不重跑。改过的请求再过一遍内容审查,然后才转换格式、脱敏。 //! - [`reply`]:回答钩子。排在格式转换之后、工具调用审查和输出长度之前(I7)—— //! 这两道防护看的就是插件改过的那一版。 //! - [`pool`]:插件调用都是阻塞的、吃 CPU 的,放在专用线程池上跑,不占 tokio 的线程。 diff --git a/crates/tw-gateway/src/plugin/reply/mod.rs b/crates/tw-gateway/src/plugin/reply/mod.rs index e10406f..f6abd96 100644 --- a/crates/tw-gateway/src/plugin/reply/mod.rs +++ b/crates/tw-gateway/src/plugin/reply/mod.rs @@ -61,11 +61,15 @@ pub struct Call { pub struct ReplyCtx<'a> { pub dialect: Dialect, pub client: Option<&'a str>, - /// 客户端要的模型 + /// 发给回答它的那一家的模型名:路由规则、请求钩子改过的是改过之后的 pub model: &'a str, + /// 客户端要的模型 + pub requested_model: &'a str, /// 回答它的那一家 pub upstream: &'a str, pub request_id: u64, + /// 回答它的那一跳是尝试链上的第几跳。记在每一次运行的 `detail` 里 + pub attempt: usize, } /// 一个插件在这次回答里的状态。 @@ -116,6 +120,7 @@ pub struct Chain { bridge: Bridge, dialect: Dialect, request_id: u64, + attempt: usize, recorded: bool, } @@ -165,7 +170,7 @@ impl Chain { ctx: &ReplyCtx<'_>, ) -> Result, GatewayError> { let mut stages = Vec::new(); - // 跑不了的插件在请求钩子那一步已经按 `on_error` 处理过了:这里只有能跑的 + // 跑不了的插件在回答它的那一次发出去之前已经按 `on_error` 处理过了:这里只有能跑的 for a in set.for_reply(ctx.client, ctx.model, ctx.upstream) { let Some(host) = a.ready().cloned() else { continue; @@ -174,8 +179,9 @@ impl Chain { let c = super::request::ctx( ctx.client, ctx.model, + ctx.requested_model, ctx.dialect, - Some(ctx.upstream), + ctx.upstream, &a.settings, ); let made = state @@ -195,7 +201,7 @@ impl Chain { outcome: PluginOutcome::Error, error: Some(why.clone()), cpu_us: 0, - detail: None, + detail: Some(json!({ "attempt": ctx.attempt })), }; state.plugin_ran(ctx.request_id, &a, run, Vec::new()); if a.on_error == OnError::Reject { @@ -207,6 +213,7 @@ impl Chain { bridge, dialect: ctx.dialect, request_id: ctx.request_id, + attempt: ctx.attempt, recorded: false, }; started.finish(); @@ -225,6 +232,7 @@ impl Chain { bridge, dialect: ctx.dialect, request_id: ctx.request_id, + attempt: ctx.attempt, recorded: false, })) } @@ -244,8 +252,9 @@ impl Chain { let c = super::request::ctx( ctx.client, ctx.model, + ctx.requested_model, ctx.dialect, - Some(ctx.upstream), + ctx.upstream, settings, ); let instance = pool @@ -261,6 +270,7 @@ impl Chain { bridge: Bridge::new(Arc::new(tw_guard::redact::rules::RuleSet::none())), dialect: ctx.dialect, request_id: ctx.request_id, + attempt: ctx.attempt, recorded: false, })) } @@ -638,6 +648,7 @@ impl Chain { error: s.error.clone(), cpu_us: s.cpu.as_micros().min(u64::MAX as u128) as u64, detail: Some(json!({ + "attempt": self.attempt, "text_calls": c.text_calls, "text_changed": c.text_changed, "tool_calls": c.tool_calls, diff --git a/crates/tw-gateway/src/plugin/reply/tests/mod.rs b/crates/tw-gateway/src/plugin/reply/tests/mod.rs index 94ba374..b08eae9 100644 --- a/crates/tw-gateway/src/plugin/reply/tests/mod.rs +++ b/crates/tw-gateway/src/plugin/reply/tests/mod.rs @@ -51,8 +51,10 @@ async fn chain_with( dialect, client: None, model: "m", + requested_model: "m", upstream: "u", request_id: 1, + attempt: 0, }, ) .await diff --git a/crates/tw-gateway/src/plugin/request.rs b/crates/tw-gateway/src/plugin/request.rs index 014d09d..1637a4a 100644 --- a/crates/tw-gateway/src/plugin/request.rs +++ b/crates/tw-gateway/src/plugin/request.rs @@ -1,11 +1,17 @@ -//! 请求钩子:客户端的请求发出去之前,按顺序交给范围内的插件改。 +//! 请求钩子:请求发往一个上游之前,按顺序交给管这一次的插件改。 //! -//! # 位置和次数 +//! # 位置和次数(契约附录二) //! -//! 排在本地应答之后、读路由事实之前(不变式 I7):插件改过的请求体重新解码,路由、 -//! 内容审查、会话指纹、出站脱敏看到的都是改过的那一份;`params.model` 改了,路由 -//! 就按新的模型走。**一个客户端请求只跑一次**(I8):故障转移、OAuth 重试、去封存 -//! 重发用的都是这一份结果。 +//! 排在路由之后:路由、模型准入、会话指纹看的都是客户端的原话,插件左右不了请求去 +//! 哪一家。**每发往一个上游跑一次**:管线每试一家([`Hook::attempt`]),按这一次的 +//! 客户端、发给这一家的模型名和这一家挑出管它的插件,从客户端的原话起改 —— 换到下一 +//! 家时重新从原话起,给上一家的改动到不了下一家。同一家重发(OAuth 换 token、去封存) +//! 用这一跳已经定好的请求体,不重跑。 +//! +//! 插件看到的 `model`(视图里的和 `params.model`)、`ctx.model` 都是**发给这一家的 +//! 模型名**(路由规则改写之后的),`ctx.requested_model` 是客户端要的,`ctx.upstream` +//! 是这一家。插件改了 `params.model`,只是换掉发给这一家的名字:不重新路由,也不再对 +//! 一遍上游的模型清单。 //! //! # 每个插件一步 //! @@ -14,12 +20,11 @@ //! 3. 在插件线程池上调 `onRequest`; //! 4. 核对交回来的东西([`super::view::check`]),占位符换回去,写回原文。 //! -//! 插件 `reject` 了,这个请求就被拒;出错了(沙箱报错、交回来的东西不合规矩)按它的 -//! `on_error`:拒绝这个请求,或者跳过它接着往下走。文件变了、装不上的插件跑不了, -//! 范围内的请求同样按 `on_error` 处理 —— **只看客户端和模型**:它要是只有回答钩子、 -//! 只管某几家上游,这时候还不知道会去哪一家,宁可多拦(这是插件坏着的时候,用户 -//! 会收到通知)。 +//! 插件 `reject` 了,或者出错而它的 `on_error` 是拒绝,**整个请求被拒**,不换下一家: +//! 换一家,管它的还是这个插件。文件变了、装不上的插件跑不了,管得着这一次的同样按 +//! `on_error` 处理;只管别的上游、别的模型的,这一次不算它。 +use std::borrow::Cow; use std::sync::Arc; use bytes::Bytes; @@ -34,40 +39,304 @@ use super::pool::Pool; use super::set::{Active, Broken, LogLine, PluginRun, PluginSet}; use super::view; -/// 请求钩子跑完之后交回管线的东西。 +/// 一个插件在这一次上的运行,连同它写的日志。 +pub type Ran = (Arc, PluginRun, Vec); + +/// 一个请求上的请求钩子。管线每发往一个上游调一次 [`Hook::attempt`],**每次都从客户端 +/// 的原话起**;原文只解析一次、密钥只编一次号,几次尝试共用。 +pub struct Hook<'a> { + set: &'a PluginSet, + rules: Arc, + /// 客户端的格式 + dialect: Dialect, + /// 客户端调的路径(Gemini 的模型在里面) + path: &'a str, + /// 客户端是哪个应用(请求那一行上记的那个,认不出是 `None`) + client: Option<&'a str>, + /// 客户端发来的原文 + body: &'a Bytes, + /// 原文解析出来的 JSON。第一次有插件要跑时才解析 + parsed: Option>, + /// 按原文编好号的那本账。第一次要用时才编 + base: Option, +} + +/// 这一次发往哪儿。 +pub struct Target<'a> { + pub upstream: &'a str, + /// 发给它的模型名:路由规则改写过的是改写之后的 + pub model: &'a str, + /// 客户端要的模型 + pub requested_model: &'a str, + /// 尝试链上的第几跳(从 0 起)。记在每一次运行的 `detail` 里,界面按它分组 + pub attempt: usize, +} + +/// 一次尝试上请求钩子跑完之后交回管线的东西。 #[derive(Default)] pub struct Plugged { - /// 每个跑过(或者该跑没跑)的插件一条,按顺序,连同它写的日志。**请求的号这时 - /// 还没发**,开始事件之后再交出去(见 [`record`]) - pub runs: Vec<(Arc, PluginRun, Vec)>, - /// 插件改过之后的请求体和路径(Gemini 换了模型时路径也变) - pub body: Option, - pub path: Option, - /// 这个请求的密钥映射。回答钩子接着用它:同一个值在两头是同一个占位符 + /// 每个跑过(或者该跑没跑)的插件一条,按顺序 + pub runs: Vec, + /// 插件改过的话,改过之后的请求 + pub changed: Option, + /// 这一次的密钥映射:跑过插件、或者管这一次的插件里有回答钩子时才有。回答钩子 + /// 接着用它:同一个值在两头是同一个占位符 pub bridge: Option, - /// 客户端要的模型(插件改之前的) - pub model: String, } -impl Plugged { - pub fn changed(&self) -> bool { - self.body.is_some() - } +/// 插件改过之后的请求:客户端那种格式,占位符已经换回原值。 +pub struct Changed { + pub body: Bytes, + /// 改过之后的 JSON。管线要重新解码它(内容审查、格式转换) + pub value: Value, + /// 调的路径。Gemini 换了模型时是新的 + pub path: String, + /// 插件改了 `params.model` 的话,发给这一家的新模型名 + pub renamed: Option, } -/// 请求被插件拒了:回给客户端的错误,和到这一步为止的记录。 +/// 插件换了发给这一家的模型名。 +pub struct Renamed { + pub model: String, + /// 最后改它的那个插件的名字(报错时说是谁)。**插件写的字** + pub by: String, +} + +/// 请求被插件拒了:回给客户端的那句话,和到这一步为止的记录。 pub struct Refused { pub why: Msg, - pub plugged: Plugged, + pub runs: Vec, } -/// 这个请求是谁发的、要什么:插件的 `ctx` 和范围都看它。 -pub struct Asked<'a> { - pub dialect: Dialect, - /// 客户端调的路径(Gemini 的模型在里面) - pub path: &'a str, - /// 客户端是哪个应用(请求那一行上记的那个,认不出是 `None`) - pub client: Option<&'a str>, +impl<'a> Hook<'a> { + pub fn new( + set: &'a PluginSet, + rules: Arc, + dialect: Dialect, + path: &'a str, + client: Option<&'a str>, + body: &'a Bytes, + ) -> Self { + Self { + set, + rules, + dialect, + path, + client, + body, + parsed: None, + base: None, + } + } + + /// 按客户端原文编好号的那本账(第一次调时才编) + fn base(&mut self) -> Bridge { + let (rules, body) = (&self.rules, self.body); + self.base + .get_or_insert_with(|| { + let mut b = Bridge::new(rules.clone()); + b.learn(body); + b + }) + .clone() + } + + /// 发往 `to` 之前跑一遍管这一次的插件。**从客户端的原话起**,上一次尝试改过什么 + /// 都不算。管这一次的一个都没有时什么都不做,连请求体都不解析。 + pub async fn attempt(&mut self, pool: &Pool, to: &Target<'_>) -> Result> { + let mut out = Plugged::default(); + if self.set.is_empty() { + return Ok(out); + } + 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() + { + out.bridge = Some(self.base()); + } + return Ok(out); + } + let mut bridge = self.base(); + let (body, dialect, client) = (self.body, self.dialect, self.client); + let original = self + .parsed + .get_or_insert_with(|| serde_json::from_slice::(body).ok()) + .as_ref(); + let mut raw: Option> = original.map(Cow::Borrowed); + let mut path = self.path.to_string(); + // 发给这一家的模型名:前一个插件改了 `params.model`,后面的看到的就是新的 + let mut model = to.model.to_string(); + let mut renamed_by: Option = None; + let mut changed = false; + for a in here { + 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 })), + }; + out.runs.push((a.clone(), run, Vec::new())); + if let Some(why) = refusal { + return Err(Box::new(Refused { + why, + runs: out.runs, + })); + } + continue; + } + super::set::State::Ready(h) => h.clone(), + }; + let mut run = PluginRun { + plugin_id: a.id.clone(), + plugin_name: a.name.clone(), + hook: PluginHook::Request, + outcome: PluginOutcome::Unchanged, + error: None, + cpu_us: 0, + detail: Some(json!({ "attempt": to.attempt })), + }; + let mut logs = Vec::new(); + let result: Result, Failure> = async { + let Some(current) = raw.as_deref() else { + return Err(Failure::Unreadable("the request body is not JSON".into())); + }; + let mut built = + view::build(dialect, current, &path).map_err(Failure::Unreadable)?; + sending(&mut built.view, &model); + let mut input = view::trim(&built.view, &a.permissions); + bridge.hide_value(&mut input); + let ctx = ctx( + client, + &model, + to.requested_model, + dialect, + to.upstream, + &a.settings, + ); + let (h, given) = (host.clone(), input.clone()); + let inv = pool + .run(move || h.on_request(given, ctx)) + .await + .map_err(|e| Failure::Run(RunError::Trap(e.to_string())))?; + run.cpu_us = inv.cpu.as_micros().min(u64::MAX as u128) as u64; + logs = inv.logs; + match inv.result.map_err(Failure::Run)? { + RequestOutcome::Unchanged => Ok(None), + RequestOutcome::Rejected(reason) => Err(Failure::Rejected(reason)), + RequestOutcome::Changed(returned) => { + let mut edits = view::check( + &input, + &returned, + &a.permissions, + built.src.hidden_tools(), + ) + .map_err(Failure::Edit)?; + if edits.is_empty() { + return Ok(None); + } + let sections = sections(&edits); + 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)?; + Ok(Some(Rewritten { + value: next, + path: new_path, + model: new_model, + sections, + })) + } + } + } + .await; + match result { + Ok(None) => {} + Ok(Some(r)) => { + raw = Some(Cow::Owned(r.value)); + if let Some(p) = r.path { + path = p; + } + if let Some(m) = r.model { + model = m; + renamed_by = Some(a.name.clone()); + } + changed = true; + run.outcome = PluginOutcome::Changed; + run.detail = Some(json!({ "attempt": to.attempt, "changed": r.sections })); + } + Err(Failure::Rejected(reason)) => { + run.outcome = PluginOutcome::Rejected; + run.error = Some(reason_msg(&reason)); + out.runs.push((a.clone(), run, logs)); + return Err(Box::new(Refused { + why: rejected(&a.name, reason), + runs: out.runs, + })); + } + Err(f) => { + let why = f.msg(); + run.outcome = PluginOutcome::Error; + run.error = Some(why.clone()); + out.runs.push((a.clone(), run, logs)); + if a.on_error == OnError::Reject { + return Err(Box::new(Refused { + why: msg!( + "gw.plugin.request_failed", + plugin = a.name.clone(), detail = why.text => + "Plugin `{plugin}` failed, so the request was not sent: {detail}" + ), + runs: out.runs, + })); + } + continue; + } + } + out.runs.push((a.clone(), run, logs)); + } + if changed && let Some(v) = raw { + let value = v.into_owned(); + match serde_json::to_vec(&value) { + Ok(b) => { + out.changed = Some(Changed { + body: Bytes::from(b), + value, + path, + renamed: renamed_by.map(|by| Renamed { model, by }), + }) + } + // 序列化不该失败;真失败了就当没改过,不发半个请求体 + Err(e) => { + tracing::error!("the request changed by plugins could not be serialized: {e}") + } + } + } + out.bridge = Some(bridge); + Ok(out) + } +} + +/// 视图里的模型名换成发给这一家的那个(`model` 和 `params.model`):插件看到的就是要发出去 +/// 的,和 `ctx.model` 一致。原文里写的是客户端要的那个,路由规则的改写在格式转换那一步才 +/// 落到请求体上 +pub(super) fn sending(view: &mut Value, model: &str) { + let Some(o) = view.as_object_mut() else { + return; + }; + o.insert("model".into(), json!(model)); + if let Some(p) = o.get_mut("params").and_then(Value::as_object_mut) { + p.insert("model".into(), json!(model)); + } } /// 客户端要的模型:请求体里的 `model`,Gemini 写在路径里。 @@ -85,185 +354,26 @@ pub fn asked_model(dialect: Dialect, path: &str, raw: Option<&Value>) -> String .to_string() } -/// 插件看到的 `ctx` +/// 插件看到的 `ctx`。请求钩子和回答钩子是同一个样子:`model` 是发给上游的模型名, +/// `requested_model` 是客户端要的,`upstream` 是这一次发往的那一家 pub fn ctx( client: Option<&str>, model: &str, + requested_model: &str, dialect: Dialect, - upstream: Option<&str>, + upstream: &str, settings: &serde_json::Map, ) -> Value { json!({ "client": client, "model": model, + "requested_model": requested_model, "format": super::format_name(dialect), "upstream": upstream, "settings": settings, }) } -/// 跑请求钩子。范围内一个插件都没有时什么都不做,连请求体都不解析。 -pub async fn run( - pool: &Pool, - set: &PluginSet, - rules: &Arc, - asked: &Asked<'_>, - body: &Bytes, -) -> Result> { - let mut out = Plugged::default(); - if set.is_empty() { - return Ok(out); - } - let parsed = serde_json::from_slice::(body).ok(); - let model = asked_model(asked.dialect, asked.path, parsed.as_ref()); - out.model = model.clone(); - let here = set.for_request(asked.client, &model); - // 回答钩子要用这个请求的密钥映射:范围里有回答钩子的话,现在就记账 - let later = set.all().iter().any(|a| { - a.enabled - && a.ready().is_some() - && a.hooks.on_reply() - && a.scope.covers_request(asked.client, &model) - }); - if here.is_empty() && !later { - return Ok(out); - } - let mut bridge = Bridge::new(rules.clone()); - bridge.learn(body); - if here.is_empty() { - out.bridge = Some(bridge); - return Ok(out); - } - let mut raw = parsed; - let mut path = asked.path.to_string(); - let mut changed = false; - for a in here { - 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: None, - }; - out.runs.push((a.clone(), run, Vec::new())); - if let Some(why) = refusal { - out.bridge = Some(bridge); - return Err(Box::new(Refused { why, plugged: out })); - } - continue; - } - super::set::State::Ready(h) => h.clone(), - }; - let mut run = PluginRun { - plugin_id: a.id.clone(), - plugin_name: a.name.clone(), - hook: PluginHook::Request, - outcome: PluginOutcome::Unchanged, - error: None, - cpu_us: 0, - detail: None, - }; - let mut logs = Vec::new(); - let result: Result, Failure> = async { - let Some(current) = raw.as_ref() else { - return Err(Failure::Unreadable("the request body is not JSON".into())); - }; - let built = view::build(asked.dialect, current, &path).map_err(Failure::Unreadable)?; - let mut input = view::trim(&built.view, &a.permissions); - bridge.hide_value(&mut input); - let ctx = ctx(asked.client, &model, asked.dialect, None, &a.settings); - let (h, given) = (host.clone(), input.clone()); - let inv = pool - .run(move || h.on_request(given, ctx)) - .await - .map_err(|e| Failure::Run(RunError::Trap(e.to_string())))?; - run.cpu_us = inv.cpu.as_micros().min(u64::MAX as u128) as u64; - logs = inv.logs; - match inv.result.map_err(Failure::Run)? { - RequestOutcome::Unchanged => Ok(None), - RequestOutcome::Rejected(reason) => Err(Failure::Rejected(reason)), - RequestOutcome::Changed(returned) => { - let mut edits = - view::check(&input, &returned, &a.permissions, built.src.hidden_tools()) - .map_err(Failure::Edit)?; - if edits.is_empty() { - return Ok(None); - } - let sections = sections(&edits); - edits.reveal(&bridge); - let mut next = current.clone(); - let new_path = - view::apply(&mut next, &built.src, &edits, &path).map_err(Failure::Edit)?; - Ok(Some((next, new_path, sections))) - } - } - } - .await; - match result { - Ok(None) => {} - Ok(Some((next, new_path, sections))) => { - raw = Some(next); - if let Some(p) = new_path { - path = p; - } - changed = true; - run.outcome = PluginOutcome::Changed; - run.detail = Some(json!({ "changed": sections })); - } - Err(Failure::Rejected(reason)) => { - run.outcome = PluginOutcome::Rejected; - run.error = Some(reason_msg(&reason)); - out.runs.push((a.clone(), run, logs)); - out.bridge = Some(bridge); - return Err(Box::new(Refused { - why: rejected(&a.name, reason), - plugged: out, - })); - } - Err(f) => { - let why = f.msg(); - run.outcome = PluginOutcome::Error; - run.error = Some(why.clone()); - out.runs.push((a.clone(), run, logs)); - if a.on_error == OnError::Reject { - out.bridge = Some(bridge); - return Err(Box::new(Refused { - why: msg!( - "gw.plugin.request_failed", - plugin = a.name.clone(), detail = why.text => - "Plugin `{plugin}` failed, so the request was not sent: {detail}" - ), - plugged: out, - })); - } - continue; - } - } - out.runs.push((a.clone(), run, logs)); - } - if changed && let Some(v) = &raw { - match serde_json::to_vec(v) { - Ok(b) => { - out.body = Some(Bytes::from(b)); - if path != asked.path { - out.path = Some(path); - } - } - // 序列化不该失败;真失败了就当没改过,不发半个请求体 - Err(e) => { - tracing::error!("the request changed by plugins could not be serialized: {e}") - } - } - } - out.bridge = Some(bridge); - Ok(out) -} - /// 跑不了的插件:记成什么,要不要拒掉这个请求 fn broken(on_error: OnError, name: &str, why: &Broken) -> (PluginOutcome, Option) { if on_error == OnError::Skip { @@ -330,8 +440,17 @@ fn sections(e: &view::Edits) -> Vec<&'static str> { s } -/// 一个插件改过的请求:新的原文、新的路径(换了的话)、改了哪几部分 -type Rewritten = (Value, Option, Vec<&'static str>); +/// 一个插件改过的请求 +struct Rewritten { + /// 新的原文 + value: Value, + /// 新的路径(换了的话) + path: Option, + /// 新的模型名(改了 `params.model` 的话) + model: Option, + /// 改了哪几部分 + sections: Vec<&'static str>, +} /// 一个插件没跑成。 enum Failure { @@ -354,11 +473,9 @@ impl Failure { } } -/// 请求那一行有了号之后,把请求钩子的记录交出去:每一次运行(计数、日志、失败的通知, -/// 见 [`crate::AppState::plugin_ran`])。改过的请求体由开始事件那一步另外交给存请求体的 -/// 那一层 -pub fn record(state: &crate::AppState, id: u64, plugged: &Plugged) { - for (a, run, logs) in &plugged.runs { +/// 把这一次的运行记到请求上(计数、日志、失败的通知,见 [`crate::AppState::plugin_ran`]) +pub fn record(state: &crate::AppState, id: u64, runs: &[Ran]) { + for (a, run, logs) in runs { state.plugin_ran(id, a, run.clone(), logs.clone()); } } diff --git a/crates/tw-gateway/src/plugin/set.rs b/crates/tw-gateway/src/plugin/set.rs index 063c5a8..026a621 100644 --- a/crates/tw-gateway/src/plugin/set.rs +++ b/crates/tw-gateway/src/plugin/set.rs @@ -18,26 +18,26 @@ use crate::plugin::host::PluginHost; /// 一个插件管哪些请求。**每张单子里都是 `*` 通配**(不分大小写,和路由规则同一种), /// 空着是「都管」。 +/// +/// **按每一次发往上游来看**(契约附录二):请求钩子排在路由之后,每试一家上游跑一次, +/// 那时这一次的客户端、发出去的模型和上游都定了 —— 请求钩子和回答钩子看的是同三样。 #[derive(Debug, Clone, Default, PartialEq, Eq)] pub struct Scope { /// 客户端应用:`claude-code`、`codex`……(请求记录上的 `client_hint`) pub clients: Vec, - /// 客户端要的模型 + /// **发给上游的模型**:路由规则改写过的是改写之后的那个,不是客户端写的 pub models: Vec, - /// 服务这个回答的上游。**只管回答那一段** —— 请求钩子跑的时候还没选上游 + /// 这一次发往的上游 pub upstreams: Vec, } impl Scope { - /// 请求钩子管不管这个请求。认不出是哪个应用(`client` 是 None)时,只有不挑 - /// 应用的插件管它 - pub fn covers_request(&self, client: Option<&str>, model: &str) -> bool { - listed(&self.clients, client) && listed(&self.models, Some(model)) - } - - /// 回答钩子管不管这个回答 - pub fn covers_reply(&self, client: Option<&str>, model: &str, upstream: &str) -> bool { - self.covers_request(client, model) && listed(&self.upstreams, Some(upstream)) + /// 管不管发往 `upstream`、模型名是 `model` 的这一次。认不出是哪个应用(`client` + /// 是 None)时,只有不挑应用的插件管它 + pub fn covers(&self, client: Option<&str>, model: &str, upstream: &str) -> bool { + listed(&self.clients, client) + && listed(&self.models, Some(model)) + && listed(&self.upstreams, Some(upstream)) } } @@ -143,25 +143,31 @@ impl PluginSet { self.plugins.is_empty() } - /// 这个请求上要过一遍的插件,按顺序:启用的、范围管得着的,**连同跑不了的** —— - /// 跑不了的由调用方照它的 `on_error` 拒绝请求或者跳过它(管得着就要处置,不管它 - /// 有没有请求钩子:它一旦加载不了,回答那一段同样做不了)。能跑的只列有请求钩子的。 - pub fn for_request(&self, client: Option<&str>, model: &str) -> Vec> { + /// 发往一个上游之前要过一遍的插件,按顺序:启用的、管得着这一次的,**连同跑不了 + /// 的** —— 跑不了的由调用方照它的 `on_error` 拒绝请求或者跳过它(管得着就要处置, + /// 不管它有没有请求钩子:它一旦加载不了,回答那一段同样做不了)。能跑的只列有请求 + /// 钩子的。**管不着这一次的不算**:只管别的上游的插件坏了,拦不着发往这一家的请求。 + pub fn for_request( + &self, + client: Option<&str>, + model: &str, + upstream: &str, + ) -> Vec> { self.plugins .iter() - .filter(|p| p.enabled && p.scope.covers_request(client, model)) + .filter(|p| p.enabled && p.scope.covers(client, model, upstream)) .filter(|p| p.ready().is_none() || p.hooks.request) .cloned() .collect() } - /// 这个回答上要过一遍的插件,按顺序:启用的、能跑的、有回答钩子的、范围管得着的 - /// (连同上游)。**跑不了的不在这里**:它们在请求那一段已经处置过了。 + /// 这个回答上要过一遍的插件,按顺序:启用的、能跑的、有回答钩子的、管得着回答它的 + /// 那一次的。**跑不了的不在这里**:它们在那一次发出去之前已经处置过了。 pub fn for_reply(&self, client: Option<&str>, model: &str, upstream: &str) -> Vec> { self.plugins .iter() .filter(|p| p.enabled && p.ready().is_some() && p.hooks.on_reply()) - .filter(|p| p.scope.covers_reply(client, model, upstream)) + .filter(|p| p.scope.covers(client, model, upstream)) .cloned() .collect() } @@ -375,23 +381,22 @@ mod tests { #[test] fn an_empty_list_covers_everything_and_globs_ignore_case() { let all = Scope::default(); - assert!(all.covers_request(None, "anything")); - assert!(all.covers_reply(Some("codex"), "gpt-5", "openai")); + assert!(all.covers(None, "anything", "anywhere")); + assert!(all.covers(Some("codex"), "gpt-5", "openai")); let s = scope(&["claude-*"], &["Claude-Sonnet-*"], &["anthropic"]); - assert!(s.covers_request(Some("claude-code"), "claude-sonnet-4-5")); - assert!(!s.covers_request(Some("codex"), "claude-sonnet-4-5")); - assert!(!s.covers_request(Some("claude-code"), "gpt-5")); - assert!(s.covers_reply(Some("claude-code"), "claude-sonnet-4-5", "anthropic")); - assert!(!s.covers_reply(Some("claude-code"), "claude-sonnet-4-5", "relay")); + assert!(s.covers(Some("claude-code"), "claude-sonnet-4-5", "anthropic")); + assert!(!s.covers(Some("codex"), "claude-sonnet-4-5", "anthropic")); + assert!(!s.covers(Some("claude-code"), "gpt-5", "anthropic")); + assert!(!s.covers(Some("claude-code"), "claude-sonnet-4-5", "relay")); } /// 认不出是哪个应用的请求,挑应用的插件不管它 —— 管了就等于对每个不认识的 /// 客户端都改请求 #[test] fn an_unknown_client_is_covered_only_by_plugins_that_do_not_pick_clients() { - assert!(scope(&[], &[], &[]).covers_request(None, "m")); - assert!(!scope(&["*"], &[], &[]).covers_request(None, "m")); + assert!(scope(&[], &[], &[]).covers(None, "m", "u")); + assert!(!scope(&["*"], &[], &[]).covers(None, "m", "u")); } /// 请求钩子那一段:按配置的顺序;跑不了的也在(调用方照 `on_error` 处置), @@ -409,11 +414,48 @@ mod tests { active("a-last", REQUEST, Scope::default(), None), ]); assert_eq!( - ids(&set.for_request(Some("claude-code"), "claude-opus-4-5")), + ids(&set.for_request(Some("claude-code"), "claude-opus-4-5", "anthropic")), ["b-first", "changed", "a-last"] ); } + /// 发往哪一家定了才挑插件:只管某一家的,发往别家时不跑;**坏了的也一样** —— 它 + /// 只拦发往它那一家的请求,不再因为「还不知道去哪儿」把别家的也拦下 + #[test] + fn the_request_list_follows_the_upstream_of_the_attempt() { + let set = PluginSet::new(vec![ + active("only-a", REQUEST, scope(&[], &[], &["relay-a"]), None), + active( + "broken-a", + REQUEST, + scope(&[], &[], &["relay-a"]), + Some(Broken::Changed), + ), + active("everywhere", REQUEST, Scope::default(), None), + ]); + assert_eq!( + ids(&set.for_request(None, "m", "relay-a")), + ["only-a", "broken-a", "everywhere"] + ); + assert_eq!(ids(&set.for_request(None, "m", "relay-b")), ["everywhere"]); + } + + /// 模型看的是发出去的那个:规则把 claude 改成 glm 发给中转,管 `glm-*` 的插件管这一次 + #[test] + fn models_match_the_model_sent_upstream() { + let set = PluginSet::new(vec![active( + "glm", + REQUEST, + scope(&[], &["glm-*"], &[]), + None, + )]); + assert_eq!(ids(&set.for_request(None, "glm-4.6", "relay")), ["glm"]); + assert!( + set.for_request(None, "claude-sonnet-4-5", "relay") + .is_empty() + ); + } + #[test] fn the_reply_list_has_only_ready_plugins_with_reply_hooks_for_that_upstream() { let set = PluginSet::new(vec![ diff --git a/crates/tw-gateway/src/plugin/trial.rs b/crates/tw-gateway/src/plugin/trial.rs index 0fab027..f9ec8da 100644 --- a/crates/tw-gateway/src/plugin/trial.rs +++ b/crates/tw-gateway/src/plugin/trial.rs @@ -6,6 +6,10 @@ //! //! 给人看的前后两份都是**换过占位符的**:插件本来就只看得到占位符,界面上显示的也 //! 不该是真值。试跑不进统计、不进日志圈、不留请求记录,日志交给调用方。 +//! +//! `ctx` 按那一行记下的路由给:`upstream` 是回答它的那一家,`model` 是发给那一家的 +//! 模型名,`requested_model` 是客户端要的 —— 和那个请求当时跑插件时看到的一样 +//! (契约附录二)。 use std::sync::Arc; @@ -21,12 +25,16 @@ use super::request::{rejected, request_unreadable}; use super::set::LogLine; use super::view; -/// 存下来的请求:客户端调的路径、查询串、请求体,和请求那一行上记的客户端。 +/// 存下来的请求:客户端调的路径、查询串、请求体,和请求那一行上记的客户端、路由。 pub struct StoredRequest<'a> { pub path: &'a str, pub query: Option<&'a str>, pub body: &'a [u8], pub client: Option<&'a str>, + /// 回答它的那一家(请求那一行的 `provider`)。没发出去的是空的 + pub upstream: &'a str, + /// 发给那一家的模型名(请求那一行的 `sent_model`)。空的话按客户端要的那个 + pub sent_model: &'a str, } /// 存下来的回答:上游的原话(流或者整包),它是什么格式、哪一家回的。 @@ -115,10 +123,23 @@ async fn tried( if let Some(r) = &request { bridge.learn(r.body); } - let model = match (&request, dialect) { + let requested = match (&request, dialect) { (Some(r), Some(d)) => super::request::asked_model(d, r.path, parsed.as_ref()), _ => String::new(), }; + // 发给上游的模型名和上游:那一行记下的路由 + let model = request + .as_ref() + .map(|r| r.sent_model) + .filter(|m| !m.is_empty()) + .map_or_else(|| requested.clone(), str::to_string); + let upstream = request + .as_ref() + .map(|r| r.upstream) + .filter(|u| !u.is_empty()) + .or(reply.as_ref().map(|r| r.provider)) + .unwrap_or_default() + .to_string(); let client = request.as_ref().and_then(|r| r.client); let name = host.manifest().name.clone(); @@ -131,10 +152,12 @@ async fn tried( bridge.hide_value(&mut masked); match view::build(d, &masked, r.path) { Err(e) => t.error = Some(request_unreadable(e)), - Ok(built) => { + Ok(mut built) => { let m = host.manifest(); + super::request::sending(&mut built.view, &model); let input = view::trim(&built.view, &m.permissions); - let ctx = super::request::ctx(client, &model, d, None, settings); + let ctx = + super::request::ctx(client, &model, &requested, d, &upstream, settings); let (h, given) = (host.clone(), input.clone()); let before = pretty(&masked); let ran = pool.run(move || h.on_request(given, ctx)).await; @@ -238,8 +261,10 @@ async fn tried( dialect: client_dialect, client, model: &model, + requested_model: &requested, upstream: reply.provider, request_id: 0, + attempt: 0, }; let mut chain = match super::reply::Chain::trial(pool, host.clone(), settings, &ctx).await { Ok(Some(c)) => c, diff --git a/crates/tw-gateway/src/plugin/trial/tests/mod.rs b/crates/tw-gateway/src/plugin/trial/tests/mod.rs index b651bde..75b3663 100644 --- a/crates/tw-gateway/src/plugin/trial/tests/mod.rs +++ b/crates/tw-gateway/src/plugin/trial/tests/mod.rs @@ -76,6 +76,8 @@ async fn a_trial_shows_both_sides_masked_and_leaves_no_trace() { query: None, body: &body, client: Some("claude-code"), + upstream: "anthropic", + sent_model: "claude-sonnet-4-5", }), Some(StoredReply { body: &answer, @@ -130,6 +132,8 @@ async fn an_answer_from_another_format_is_read_in_the_clients_format() { query: None, body: &body, client: None, + upstream: "anthropic", + sent_model: "claude-sonnet-4-5", }), Some(StoredReply { body: chat.as_bytes(), @@ -164,6 +168,8 @@ async fn a_rejection_is_reported_without_an_after() { query: None, body: &body, client: None, + upstream: "anthropic", + sent_model: "claude-sonnet-4-5", }), None, ) @@ -175,3 +181,43 @@ async fn a_rejection_is_reported_without_an_after() { assert_eq!(e.code, "gw.plugin.rejected"); assert!(e.text.contains("not today")); } + +/// `ctx` 按那一行记下的路由给:回答它的那一家、发给它的模型名、客户端要的模型 —— 视图里 +/// 的模型名和 `ctx.model` 是同一个 +#[tokio::test] +async fn a_trial_gives_the_plugin_the_routing_the_request_had() { + let pool = Arc::new(Pool::new(1, 4)); + let body = request(); + let saw = Arc::new(std::sync::Mutex::new(Value::Null)); + let s = saw.clone(); + let look = Double::new("look") + .permit(&[Permission::Params]) + .on_request(move |view, ctx| { + *s.lock().unwrap() = json!({ "view": view, "ctx": ctx }); + Invocation::ok(RequestOutcome::Unchanged) + }) + .into_host(); + let t = run( + pool, + look, + &Default::default(), + rules(), + Some(StoredRequest { + path: "/v1/messages", + query: None, + body: &body, + client: Some("claude-code"), + upstream: "relay", + sent_model: "glm-4.6", + }), + None, + ) + .await; + assert_eq!(t.error, None); + let saw = saw.lock().unwrap().clone(); + assert_eq!(saw["ctx"]["upstream"], "relay"); + assert_eq!(saw["ctx"]["model"], "glm-4.6"); + assert_eq!(saw["ctx"]["requested_model"], "claude-sonnet-4-5"); + assert_eq!(saw["view"]["model"], "glm-4.6"); + assert_eq!(saw["view"]["params"]["model"], "glm-4.6"); +} diff --git a/crates/tw-gateway/src/server.rs b/crates/tw-gateway/src/server.rs index f117f03..7f53eda 100644 --- a/crates/tw-gateway/src/server.rs +++ b/crates/tw-gateway/src/server.rs @@ -198,7 +198,6 @@ async fn passthrough( dialect, started, from, - before: None, }; let result = pipeline::pipeline(state, rt, req, live, &mut ending).await; if let Some(end) = ending.take() { diff --git a/crates/tw-gateway/src/server/pipeline.rs b/crates/tw-gateway/src/server/pipeline.rs index 6beb6cf..8ad791a 100644 --- a/crates/tw-gateway/src/server/pipeline.rs +++ b/crates/tw-gateway/src/server/pipeline.rs @@ -5,8 +5,8 @@ //! 本地应答在准入之前(离线也要能答),准入在路由之前(列表即承诺), //! 并发闸门在路由之后(被规则挡下的不用先排队)。 //! -//! 发出开始事件之后的两段各自一个子模块:[`hop`] 依次试候选上游, -//! [`relay`] 把选中那一家的响应交给客户端。 +//! 发出开始事件之后的两段各自一个子模块:[`hop`] 依次试候选上游(每一跳先过插件的 +//! 请求钩子,见 [`plug`]),[`relay`] 把选中那一家的响应交给客户端。 use std::sync::Arc; @@ -22,6 +22,7 @@ use tw_types::msg; mod hop; mod opening; +mod plug; mod relay; /// 256 MiB。大到能装下几张 4K 图的 base64(膨胀 33%),小到失控的 @@ -41,40 +42,13 @@ pub(super) struct Inbound { pub(super) dialect: tw_dialect::ir::Dialect, pub(super) started: std::time::Instant, pub(super) from: Sender, - /// 插件改过这个请求时,客户端发来的原样(见 [`Before`])。没改过是 None - pub(super) before: Option, -} - -/// 插件改过的请求,客户端发来时的样子。 -/// -/// **请求记录存的是它**(客户端发了什么),插件改过的那一份另外交给插件的记录; -/// 开始事件和结局里的模型名也是客户端要的那一个 —— 和路由规则改写模型时一样, -/// 实际发出去的模型记在尝试链的每一跳上。 -pub(super) struct Before { - pub(super) body: Bytes, - pub(super) path: String, - pub(super) model: String, -} - -impl Inbound { - /// 客户端要的模型 - pub(super) fn asked_model<'a>(&'a self, reading: &'a crate::client_api::Reading) -> &'a str { - self.before - .as_ref() - .map_or(reading.facts.model.as_str(), |b| b.model.as_str()) - } - - /// 客户端调的路径 - fn asked_path(&self) -> &str { - self.before - .as_ref() - .map_or(self.uri.path(), |b| b.path.as_str()) - } } /// 发出开始事件之后,后面几步都要用的。 struct Started { id: u64, + /// 开始的时刻:请求那一行的 `at_ms`。之后才交去存的正文(插件改过的请求)挂在它上面 + at_ms: u64, /// 熔断过滤之后的候选,按顺序试 alive: Vec, /// 第一阶段的结论。路由事件在它上面补上第二阶段和尝试链 @@ -84,6 +58,8 @@ struct Started { /// 出站脱敏的账本:拦截档下按客户端原文编好了号,每一跳接着它换(见 /// [`crate::guard::look`])。别的档位是空的 ledger: tw_guard::redact::replace::Ledger, + /// 出站脱敏在客户端原文里找到的。插件改过的那一跳只再报插件写进来的(见 [`plug`]) + found: Vec, } pub(super) async fn pipeline( @@ -111,38 +87,8 @@ pub(super) async fn pipeline( ))); } - // 管线第 1.5 步:插件的请求钩子。插件表跟着运行时走:**整个请求是同一份**, - // 回答钩子用的也是它 - let plugin_set = rt.plugins.clone(); - let (req, mut plugged) = match request_plugins(&state, &rt, &plugin_set, req).await { - Ok(done) => done, - // 插件拒了这个请求:**照样开始**,流量里要有这一行,插件的记录挂在它上面。 - // 路由还没跑,没有路由事件 - Err(refused) => { - let (req, refused) = *refused; - let (reading, fp) = read(&req, intent); - let choice = Choice { - route: rt.engine.route_of(&req.client_name).to_string(), - ..Default::default() - }; - let to = ("", tw_api::Billing::PerToken); - let (_, ledger) = look(&rt, &req, &refused.plugged); - let redaction = redaction(&rt, ledger); - open( - &state, - &req, - &reading, - &choice, - to, - fp.as_deref(), - &refused.plugged, - ending, - redaction, - ); - return Err(GatewayError::denied(refused.why)); - } - }; - + // 管线第 2 步:读出路由事实,路由。**看的是客户端的原话**:插件的请求钩子排在路由 + // 之后(每发往一个上游跑一次,见 `plug`),左右不了请求去哪一家 let (reading, fp) = read(&req, intent); let conv = conversation(&rt, &req, &reading, fp.as_deref()); let (choice, decision) = match route(&state, &rt, &req, &reading, conv.as_ref())? { @@ -152,16 +98,15 @@ pub(super) async fn pipeline( Routed::Refused(choice, why) => { let to = ("", tw_api::Billing::PerToken); // 一个字节都没发出去,也没什么可报的;存下来的请求照样按这一档换、打码 - let (_, ledger) = look(&rt, &req, &plugged); + let (_, ledger) = look(&rt, &req); let redaction = redaction(&rt, ledger); - let id = open( + let (id, _) = open( &state, &req, &reading, &choice, to, fp.as_deref(), - &plugged, ending, redaction, ); @@ -188,13 +133,24 @@ pub(super) async fn pipeline( choice, &decision, fp.as_deref(), - &plugged, ending, ); screen(&state, &rt, &reading, &started)?; - let answer = hop::try_upstreams(&state, &rt, &req, &reading, &decision, &started).await?; + // 插件的请求钩子在每一跳里跑(见 `plug`):从客户端的原话起改,几跳共用原文的解析 + // 和密钥的编号。插件表跟着运行时走:**整个请求是同一份**,回答钩子用的也是它 + let hint = crate::hint::client_hint(&req.headers); + let mut hook = crate::plugin::request::Hook::new( + &rt.plugins, + rt.redact.clone(), + req.dialect, + req.uri.path(), + hint.as_deref(), + &req.body, + ); + let answer = + hop::try_upstreams(&state, &rt, &req, &reading, &decision, &started, &mut hook).await?; // 网关估的数不是哪一家回答的:不记这段对话留在哪一家 - let served = match answer { + let mut served = match answer { hop::Answer::Served(served) => *served, hop::Answer::Estimated(body) => { let ending = ending @@ -203,9 +159,16 @@ pub(super) async fn pipeline( return Ok(estimated(&state, &req, started.id, body, ending)); } }; + // 这一跳的账比开头那本多了号(插件往请求里写了新的值,拦截档下接着编了号):回答 + // 落盘时按这一本换,回显的占位符和存下来的请求对得上 + if served.ledger.len() != started.ledger.len() + && let Some(e) = ending.as_mut() + { + e.redact_with(redaction(&rt, served.ledger.clone())); + } // 回答钩子:上游回了成功的回答才有。**在交出结局之前起实例**:起不来而策略是拒绝时, // 这个请求按返回的错误收场,客户端还一个字节都没收到 - let reply_plugins = match plugged.bridge.take() { + let reply_plugins = match served.bridge.take() { Some(bridge) if reading.generates && served.upstream.status().is_success() => { // 拦截档下回答里的占位符是这一跳编的(接着请求的账),换回占位符时用同一本 let bridge = if served.ledger.is_empty() { @@ -213,15 +176,16 @@ pub(super) async fn pipeline( } else { bridge.with_ledger(served.ledger.clone()) }; - let hint = crate::hint::client_hint(&req.headers); let ctx = crate::plugin::reply::ReplyCtx { dialect: req.dialect, client: hint.as_deref(), - model: req.asked_model(&reading), + model: &served.model, + requested_model: &reading.facts.model, upstream: &served.provider.name, request_id: started.id, + attempt: served.attempt, }; - crate::plugin::reply::Chain::start(&state, &plugin_set, bridge, &ctx).await? + crate::plugin::reply::Chain::start(&state, &rt.plugins, bridge, &ctx).await? } _ => None, }; @@ -640,7 +604,6 @@ fn start( choice: Choice, decision: &tw_engine::Decision, fp: Option<&str>, - plugged: &crate::plugin::request::Plugged, ending: &mut Option, ) -> Started { // 熔断过滤。**只有一个候选时完全旁路**,全都熔断时 fail-open —— @@ -669,15 +632,14 @@ fn start( // 是同一条记录**,差别只在换没换 —— 真正的替换在每一跳发出去之前做, // 那一跳的请求体可能是转换过格式的。拦截档下账本在这里就编好号:每一跳、 // 存下来的那份请求都按它换,同一个值处处是同一个占位符 - let (found, ledger) = look(rt, req, plugged); - let id = open( + let (found, ledger) = look(rt, req); + let (id, at_ms) = open( state, req, reading, &choice, (first, billing.into()), fp, - plugged, ending, redaction(rt, ledger.clone()), ); @@ -693,32 +655,27 @@ fn start( } Started { id, + at_ms, alive, choice, conversation: crate::affinity::identity(&req.headers, fp), ledger, + found, } } -/// 出站脱敏看一遍要发出去的请求(见 [`crate::guard::look`]):插件改过的话就是改过的 -/// 那一份。 +/// 出站脱敏看一遍客户端发来的原文(见 [`crate::guard::look`])。 /// -/// **插件看过这个请求的话,接着插件看到的那本账编号**:插件拿到的占位符是按客户端 -/// 原文编的(见 [`crate::plugin::bridge`]),同一个值在插件那儿、在每一跳、在存下来的 -/// 请求和回答里都是同一个号。 +/// 插件的密钥映射按同一份原文、同一个找法编号(见 [`crate::plugin::bridge`]):插件看到的 +/// 占位符和这本账里的是同一个号。插件往某一跳写进新的值,那一跳接着编(见 [`plug`])。 fn look( rt: &Runtime, req: &Inbound, - plugged: &crate::plugin::request::Plugged, ) -> ( Vec, tw_guard::redact::replace::Ledger, ) { - let mode = rt.config.security.redact.mode; - match plugged.bridge.as_ref() { - Some(b) => crate::guard::look_from(mode, &rt.redact, &req.body, b.ledger().clone()), - None => crate::guard::look(mode, &rt.redact, &req.body), - } + crate::guard::look(rt.config.security.redact.mode, &rt.redact, &req.body) } /// 这个请求的正文落盘之前怎么换、怎么打码:此刻生效的规则,和这个请求的账本。 @@ -730,8 +687,8 @@ fn redaction(rt: &Runtime, ledger: tw_guard::redact::replace::Ledger) -> crate:: } /// 发 `RequestStarted`、把这个请求欠着的结局放进 `ending`、把请求体交去留档, -/// 交回这个请求的号。`to` 是要发往的那一家和它怎么收钱;一家都不会去的(被规则 -/// 拒绝了)是空的名字。`redaction` 是请求体、响应体落盘之前怎么换、打码。 +/// 交回这个请求的号和开始的时刻。`to` 是要发往的那一家和它怎么收钱;一家都不会去的 +/// (被规则拒绝了)是空的名字。`redaction` 是请求体、响应体落盘之前怎么换、打码。 /// /// **会话在这里定**(见 [`crate::session::Sessions`]):开始事件带着它,落库的 /// 那一行记的也是它。 @@ -743,10 +700,9 @@ fn open( choice: &Choice, to: (&str, tw_api::Billing), fp: Option<&str>, - plugged: &crate::plugin::request::Plugged, ending: &mut Option, redaction: crate::bodies::Redaction, -) -> u64 { +) -> (u64, u64) { let facts = &reading.facts; let id = state.bus.next_id(); let at_ms = now_ms(); @@ -765,9 +721,9 @@ fn open( rewritten_by: choice.rewritten_by.clone(), provider: to.0.to_string(), billing: to.1, - model: req.asked_model(reading).to_string(), + model: facts.model.clone(), method: "POST".to_string(), - path: req.asked_path().to_string(), + path: req.uri.path().to_string(), // 路由已经估过的那个数,不再算一遍。**它也就是发给上游的那一份的估算**:之后每 // 一跳只会改模型名、输出上限和推理开关(规则)、换一种写法(格式转换)、把几个值 // 换成占位符(脱敏)—— 前两样不动这个数,脱敏差出的几个 token 在估算本身的误差 @@ -780,7 +736,7 @@ fn open( let mut end = crate::ending::Ending::new( state.bus.clone(), id, - req.asked_model(reading).to_string(), + facts.model.clone(), req.started, at_ms as i64, sink.clone(), @@ -794,97 +750,20 @@ fn open( // `bodies::offer`)。**交出去的是原文**:换掉、打码在落盘那一头做,不占 // 转发这条路(见 `crate::bodies`)。 // - // 存的是**客户端发来的那一份**;插件改过的话,改过之后的另存一份,换掉、打码的 - // 规矩一样 - let body = req.before.as_ref().map_or(&req.body, |b| &b.body); + // 存的是**客户端发来的那一份**。插件改过的话,回答它的那一跳收到的那一份在试完 + // 上游之后另存(见 `hop`),换掉、打码的规矩一样 crate::bodies::offer( &sink, crate::bodies::BodyRecord::new( id, at_ms as i64, crate::bodies::BodyKind::Request, - body.clone(), - body.len(), - redaction.clone(), + req.body.clone(), + req.body.len(), + redaction, ), ); - if let Some(after) = &plugged.body { - crate::bodies::offer( - &sink, - crate::bodies::BodyRecord::new( - id, - at_ms as i64, - crate::bodies::BodyKind::AfterPlugins, - after.clone(), - after.len(), - redaction, - ), - ); - } - crate::plugin::request::record(state, id, plugged); - id -} - -/// 管线第 1.5 步:插件的请求钩子(见 [`crate::plugin::request`])。 -/// -/// **只给生成回答的请求跑**:计 token、嵌入这些接口没有「一次回答」可言。插件改过 -/// 请求的话,交回的 `Inbound` 带着改过的请求体(Gemini 换了模型时还有新的路径), -/// 客户端发来的原样留在 `before` 里。 -async fn request_plugins( - state: &AppState, - rt: &Runtime, - set: &crate::plugin::PluginSet, - mut req: Inbound, -) -> Result< - (Inbound, crate::plugin::request::Plugged), - Box<(Inbound, crate::plugin::request::Refused)>, -> { - let Some(api) = req - .api - .filter(|_| crate::client_api::ClientApi::generates(req.uri.path())) - else { - return Ok((req, Default::default())); - }; - if set.is_empty() { - return Ok((req, Default::default())); - } - let hint = crate::hint::client_hint(&req.headers); - let path = req.uri.path().to_string(); - let asked = crate::plugin::request::Asked { - dialect: api.dialect(), - path: &path, - client: hint.as_deref(), - }; - match crate::plugin::request::run(&state.plugin_pool, set, &rt.redact, &asked, &req.body).await - { - Ok(plugged) => { - if let Some(body) = plugged.body.clone() { - let new_path = plugged.path.clone(); - req.before = Some(Before { - body: std::mem::replace(&mut req.body, body), - path: path.clone(), - model: plugged.model.clone(), - }); - if let Some(p) = new_path { - req.uri = with_path(&req.uri, &p); - } - } - Ok((req, plugged)) - } - Err(refused) => Err(Box::new((req, *refused))), - } -} - -/// 换掉路径,查询串照旧 -fn with_path(uri: &axum::http::Uri, path: &str) -> axum::http::Uri { - let pq = match uri.query() { - Some(q) => format!("{path}?{q}"), - None => path.to_string(), - }; - axum::http::Uri::builder() - .path_and_query(pq) - .build() - .unwrap_or_else(|_| uri.clone()) + (id, at_ms) } /// 请求防护:调用方发来的正文里(连同工具结果)有没有藏起来的字符、有没有命中 @@ -892,6 +771,8 @@ fn with_path(uri: &axum::http::Uri, path: &str) -> axum::http::Uri { /// /// **在开始事件之后**:记录要挂在这个请求上,拒掉的请求也要在流量里留一行 —— /// 被拒是一次来源为 `denied` 的失败。**在尝试上游之前**:拒掉的一个字节都不发。 +/// 看的是客户端的原话;插件在某一跳改过的请求,在那一跳再看一遍插件加进来的 +/// (见 [`plug`])。 /// /// 按解码出来的消息看,所以只有生成回答的请求才看:计 token、嵌入这些接口没有 /// 「调用方的消息」可言;解不开的体也不看 —— 同格式直通照样发,上游可能认得它。 diff --git a/crates/tw-gateway/src/server/pipeline/hop.rs b/crates/tw-gateway/src/server/pipeline/hop.rs index 02f804c..2a3cf04 100644 --- a/crates/tw-gateway/src/server/pipeline/hop.rs +++ b/crates/tw-gateway/src/server/pipeline/hop.rs @@ -5,6 +5,10 @@ //! //! 「尝试链」要留下来:用户能看见故障转移在替他工作,**这是信任的来源**。 //! 一个静默切换过的请求和一个一次就成的请求,在用户眼里应该是不同的。 +//! +//! 每一跳先过插件的请求钩子([`super::plug`]):发往哪一家、发什么模型名这时都定了, +//! 管这一跳的插件从客户端的原话起改,这一跳的转换、脱敏、发送用改过的那一份。换到下一 +//! 家时从原话重来;同一家重发(OAuth 换 token、去封存)用这一跳定好的请求体,不重跑。 use bytes::Bytes; @@ -20,6 +24,13 @@ use tw_types::msg; pub(super) struct Served<'a> { pub(super) upstream: reqwest::Response, pub(super) provider: &'a tw_config::Provider, + /// 发给它的模型名:路由规则、插件改过的是改过之后的。回答钩子的 `ctx.model` 和范围看它 + pub(super) model: String, + /// 它是尝试链上的第几跳。回答钩子的运行记录按它分组 + pub(super) attempt: usize, + /// 这一跳的密钥映射:跑过插件、或者管这一跳的插件里有回答钩子时才有(见 + /// [`super::plug`]) + pub(super) bridge: Option, /// 成功那一次的脱敏账本。**必须是成功那一次的** —— 每一跳都接着原文那本账换, /// 而那一跳发出去的体(可能转换过格式)里还有原文没有的值时,号是那一跳新发的 pub(super) ledger: tw_guard::redact::replace::Ledger, @@ -36,6 +47,35 @@ pub(super) enum Answer<'a> { Estimated(Bytes), } +/// 这一跳的客户端那种格式的请求:插件在这一跳改过的话是改过的那一份,没改过就是 +/// 客户端的原话。转换、参数改写都从它起。 +struct Asked<'r> { + body: &'r Bytes, + path: &'r str, + decoded: Option<&'r Result>, +} + +impl<'r> Asked<'r> { + fn of( + req: &'r Inbound, + reading: &'r crate::client_api::Reading, + plugged: &'r super::plug::Plugged, + ) -> Self { + match &plugged.rewritten { + Some(r) => Asked { + body: &r.body, + path: &r.path, + decoded: r.decoded.as_ref(), + }, + None => Asked { + body: &req.body, + path: req.uri.path(), + decoded: reading.decoded.as_ref(), + }, + } + } +} + /// 这一跳要发出去的东西。 struct Outbound { body: Bytes, @@ -62,6 +102,7 @@ pub(super) async fn try_upstreams<'a>( reading: &crate::client_api::Reading, decision: &tw_engine::Decision, started: &Started, + hook: &mut crate::plugin::request::Hook<'_>, ) -> Result, GatewayError> { let id = started.id; // 数 token(见 `crate::count`):选中的那一家数不了就由网关估,**不换模型** @@ -80,8 +121,11 @@ pub(super) async fn try_upstreams<'a>( let mut rewritten_by = started.choice.rewritten_by.clone(); // 第二阶段拒绝了它的那条规则 let mut denied_by: Option = None; - // 不再试下一家的原因:第二阶段拒绝了,或者规则求不了值。**路由事件照样要发** + // 不再试下一家的原因:第二阶段拒绝了、规则求不了值、插件拒绝了。**路由事件照样要发** let mut halt: Option = None; + // 最后发出去的那一跳,插件改过的话改过之后的请求和那一跳的账:存下来的「插件改过的 + // 请求」就是它 —— 回答的那一家收到的那一份 + let mut after_plugins: Option<(Bytes, tw_guard::redact::replace::Ledger)> = None; for (i, name) in started.alive.iter().enumerate() { // 后面没有别的候选了 @@ -125,10 +169,11 @@ pub(super) async fn try_upstreams<'a>( // **在循环里面,因为故障转移换了 provider 之后必须重算**。 // 否则「走中转的一律脱敏」这条规则,在从官方转移到中转时会漏掉 // —— 而那正是最需要它的时刻。 - let effective_set = match rt - .engine - .phase_two(&reading.facts, &provider.name, &decision.set) - { + let mut effective_set = match rt.engine.phase_two( + &reading.facts, + &provider.name, + &decision.set, + ) { Ok(tw_engine::Outcome2::Proceed { set, rewritten_by: more, @@ -166,19 +211,18 @@ pub(super) async fn try_upstreams<'a>( } }; - // 这一跳要发的模型名:规则或者插件改写过、和客户端要的不一样的才记(见 - // `AttemptView::model`) - let model = Some( - effective_set - .model - .clone() - .unwrap_or_else(|| reading.facts.model.clone()), - ) - .filter(|m| m != req.asked_model(reading)); + // 这一跳要发的模型名:规则改写过的是改写之后的 + let sent = effective_set + .model + .clone() + .unwrap_or_else(|| reading.facts.model.clone()); + // 改写过、和客户端要的不一样的才记(见 `AttemptView::model`) + let asked_other = |m: &String| *m != reading.facts.model; // 数 token 不换模型:另一个模型的 tokenizer 数出来的不是这个数 if counting { + let model = Some(sent.clone()).filter(asked_other); match &count_model { - None => count_model = Some(model.clone()), + None => count_model = Some(model), Some(first) if *first != model => { attempts.pop(); continue; @@ -187,7 +231,41 @@ pub(super) async fn try_upstreams<'a>( } } - let out = match prepare(state, req, reading, provider, &effective_set, id) { + // 插件的请求钩子:管这一跳的从客户端的原话起改。**拒绝的是整个请求**,不换下一家 + let plugged = match super::plug::attempt( + state, + rt, + req, + reading, + started, + hook, + provider, + &sent, + chain.len(), + ) + .await + { + Ok(p) => p, + Err(why) => { + // 这一跳没有发出去。**它在尝试链上**,原因就是拒绝它的那句话 + chain.push(hop_failed( + &provider.name, + Some(sent.clone()).filter(asked_other), + why.clone(), + hop_started, + )); + halt = Some(GatewayError::denied(why)); + break; + } + }; + // 插件换了发给这一家的模型名:和规则改写的一样,只是盖过它 + if let Some(m) = &plugged.model { + effective_set.model = Some(m.clone()); + } + let model = Some(plugged.model.clone().unwrap_or(sent)).filter(asked_other); + + let asked = Asked::of(req, reading, &plugged); + let out = match prepare(state, req, reading, &asked, provider, &effective_set, id) { Ok(out) => out, Err(err) => { chain.push(hop_failed( @@ -205,13 +283,14 @@ pub(super) async fn try_upstreams<'a>( let unsealed = unseal_upfront(state, req, started, provider, &out); // 出站脱敏的拦截档:换掉**这一跳真正发出去的那一份**(可能转换过 // 格式)。规则是全局的,每一跳换掉的是同一批东西;**接着原文那本账换**, - // 同一个值在每一跳、在存下来的那份请求里都是同一个占位符 - let (body, ledger) = crate::guard::replace( - rt.config.security.redact.mode, - &rt.redact, - unsealed, - &started.ledger, - ); + // 同一个值在每一跳、在存下来的那份请求里都是同一个占位符。插件改过的一跳接着 + // 插件那本账(插件写进来的新值在那里编好了号) + let seed = plugged + .rewritten + .as_ref() + .map_or(&started.ledger, |r| &r.ledger); + let (body, ledger) = + crate::guard::replace(rt.config.security.redact.mode, &rt.redact, unsealed, seed); // 用这个 provider 自己的 Client —— 它带着该走的代理。**在取密钥 // 之前拿到**:OAuth 换 token 也要走这条代理。 @@ -273,6 +352,17 @@ pub(super) async fn try_upstreams<'a>( attempt = attempts.len(), "forwarding" ); + // 这一跳要发出去了:插件改过的话,它收到的就是改过的那一份 + after_plugins = plugged + .rewritten + .as_ref() + .map(|r| (r.body.clone(), ledger.clone())); + // 这一跳接下了的话,回答钩子要的 + let sent_model = effective_set + .model + .clone() + .unwrap_or_else(|| reading.facts.model.clone()); + let (attempt, bridge) = (chain.len(), plugged.bridge); let sent = send( state, @@ -400,6 +490,9 @@ pub(super) async fn try_upstreams<'a>( served = Some(Served { upstream: r, provider, + model: sent_model, + attempt, + bridge, ledger, session: out.session, }); @@ -496,6 +589,9 @@ pub(super) async fn try_upstreams<'a>( served = Some(Served { upstream: r, provider, + model: sent_model, + attempt, + bridge, ledger, session: out.session, }); @@ -543,6 +639,25 @@ pub(super) async fn try_upstreams<'a>( (Some(s), None) => s.provider.billing, (None, None) => Default::default(), }; + // 插件改过的请求:最后发出去的那一跳收到的那一份(回答的那一家收到的就是它)。 + // 挂在请求那一行上,落盘前按那一跳的账换、打码 + if let Some((body, ledger)) = after_plugins { + let len = body.len(); + crate::bodies::offer( + &state.body_sink(), + crate::bodies::BodyRecord::new( + id, + started.at_ms as i64, + crate::bodies::BodyKind::AfterPlugins, + body, + len, + crate::bodies::Redaction { + rules: rt.redact.clone(), + ledger, + }, + ), + ); + } let choice = &started.choice; state.bus.emit(tw_api::Event::RequestRouted { id, @@ -711,12 +826,14 @@ fn unsendable_tool( Some(GatewayError::new(crate::error::Source::Request, msg)) } -/// 把客户端的请求改成这一跳要发的样子:同格式时只做参数改写, -/// 跨格式时转换。转换不了就换下一家:同格式的上游可能还在后面。 +/// 把这一跳的请求(客户端那种格式,插件改过的话是改过的,见 [`Asked`])改成要发的 +/// 样子:同格式时只做参数改写,跨格式时转换。转换不了就换下一家:同格式的上游可能 +/// 还在后面。 fn prepare( state: &AppState, req: &Inbound, reading: &crate::client_api::Reading, + asked: &Asked<'_>, provider: &tw_config::Provider, effective_set: &tw_engine::SetAction, id: u64, @@ -730,7 +847,7 @@ fn prepare( // 不止生成请求:数 token 这样的请求一样带着它的头,也可能带着会话日志 let harness = reading.harness.is_some(); let to_deepseek = tw_dialect::official::is_deepseek_host(&provider.base_url); - let mut path = req.uri.path().to_string(); + let mut path = asked.path.to_string(); let mut query = req.query.clone(); let mut session: Option = None; let body = match target { @@ -738,7 +855,7 @@ fn prepare( // 参数改写。**只在这里动 body,而且只动被点名的那几个字段** —— // 出站直通说过任何 body 改写都可能是缓存杀手,所以这是 // 一个用户显式要求的例外,不是默认行为。 - let out = forward::apply_set(&req.body, effective_set, client_dialect); + let out = forward::apply_set(asked.body, effective_set, client_dialect); if let (Some(tw_dialect::ir::Dialect::Gemini), Some(m)) = (client_dialect, &effective_set.model) { @@ -792,7 +909,7 @@ fn prepare( } if chatgpt { // 客户端要整包,后端只给流:由网关收齐。收齐要知道客户端的格式,所以要一个会话 - if let Some(Ok(d)) = &reading.decoded + if let Some(Ok(d)) = asked.decoded && !d.request.stream { session = Some( @@ -811,7 +928,7 @@ fn prepare( out } Some(dialect) => { - let d = match &reading.decoded { + let d = match asked.decoded { Some(Ok(d)) => d, other => { let why = match other { @@ -887,7 +1004,7 @@ fn prepare( } else if harness && to_deepseek { // 转换成另一种格式发给 DeepSeek 官方:直连时它收得到的扩展照样带上 Bytes::from( - tw_dialect::harness::carry(&req.body, &p.body).unwrap_or(p.body.clone()), + tw_dialect::harness::carry(asked.body, &p.body).unwrap_or(p.body.clone()), ) } else { Bytes::from(p.body.clone()) diff --git a/crates/tw-gateway/src/server/pipeline/plug.rs b/crates/tw-gateway/src/server/pipeline/plug.rs new file mode 100644 index 0000000..f063885 --- /dev/null +++ b/crates/tw-gateway/src/server/pipeline/plug.rs @@ -0,0 +1,158 @@ +//! 管线第 5 步里每一跳的头一步:插件的请求钩子(见 [`crate::plugin::request`])。 +//! +//! 排在路由之后、这一跳的格式转换之前(契约附录二)。按这一跳的上游和发给它的模型名 +//! 挑出管它的插件,**从客户端的原话起改**:上一跳改过什么都不带过来。插件改过的请求在 +//! 这一跳发出去之前还要过几道: +//! +//! - **请求防护再看一遍**,只看插件加进来的(客户端的原话在开头看过、报过了,见 +//! [`crate::guard::screen_more`])。拦下就拒绝整个请求,不换下一家 —— 和开头那一遍 +//! 一样; +//! - 插件换了发出去的模型名:**密钥的模型范围照样管**。规则改写的模型名要过这一关,插件 +//! 改的也要;上游的模型清单不再对(契约附录二); +//! - 出站脱敏接着插件那本账编号:插件写进来的新值拿到新的号,报一条记录; +//! - 重新解码:格式转换用改过的这一份。 +//! +//! 插件拒绝了、出错而策略是拒绝、或者上面哪一道没过,**整个请求被拒**,不换下一家: +//! 换一家,管它的还是这些插件。 + +use bytes::Bytes; + +use super::{Inbound, Started}; +use crate::state::{AppState, Runtime}; +use tw_types::{Msg, msg}; + +/// 插件在这一跳上改过的请求,和发出去之前要用的。 +pub(super) struct Rewritten { + /// 客户端那种格式,占位符已经换回原值 + pub(super) body: Bytes, + /// 调的路径(Gemini 换了模型时是新的) + pub(super) path: String, + /// 改过的请求解码出来的中间表示:格式转换用它 + pub(super) decoded: Option>, + /// 这一跳出站脱敏接着编号的账:拦截档下是插件那本账接着编的(插件写进来的新值有了 + /// 新的号),别的档位是空的 + pub(super) ledger: tw_guard::redact::replace::Ledger, +} + +/// 请求钩子在这一跳上的结果。 +#[derive(Default)] +pub(super) struct Plugged { + /// 插件改过的话,改过的请求 + pub(super) rewritten: Option, + /// 插件改了 `params.model` 的话,发给这一家的新模型名 + pub(super) model: Option, + /// 这一跳的密钥映射。回答钩子接着用(见 [`crate::plugin::request::Plugged::bridge`]) + pub(super) bridge: Option, +} + +/// 发往 `provider` 之前跑一遍管这一跳的插件。`model` 是发给它的模型名(路由规则改写 +/// 之后的),`attempt` 是这一跳在尝试链上的位置。每一次运行当场记到请求上。 +/// +/// `Err` 是拒绝整个请求时告诉客户端的那句话。 +#[allow(clippy::too_many_arguments)] +pub(super) async fn attempt( + state: &AppState, + rt: &Runtime, + req: &Inbound, + reading: &crate::client_api::Reading, + started: &Started, + hook: &mut crate::plugin::request::Hook<'_>, + provider: &tw_config::Provider, + 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, + requested_model: &reading.facts.model, + attempt, + }; + let p = match hook.attempt(&state.plugin_pool, &to).await { + Ok(p) => p, + Err(refused) => { + crate::plugin::request::record(state, started.id, &refused.runs); + return Err(refused.why); + } + }; + crate::plugin::request::record(state, started.id, &p.runs); + let mut out = Plugged { + bridge: p.bridge, + ..Default::default() + }; + let Some(c) = p.changed else { + return Ok(out); + }; + 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()); + // 请求防护:只看插件加进来的。解不开的不看 —— 和开头那一遍一样,同格式直通照样发 + if let (Some(Ok(before)), Ok(after)) = (&reading.decoded, &decoded) + && let Some(why) = crate::guard::screen_more( + &state.bus, + started.id, + &provider.name, + &crate::guard::Screen::of(rt), + &before.request, + &after.request, + ) + { + return Err(why); + } + // 出站脱敏:接着插件看到的那本账编号(同一个值还是同一个号),插件写进来的新值报一条 + let mode = rt.config.security.redact.mode; + let seed = out + .bridge + .as_ref() + .map_or_else(|| started.ledger.clone(), |b| b.ledger().clone()); + let (found, ledger) = crate::guard::look_from(mode, &rt.redact, &c.body, seed); + let more = crate::guard::more_found(&started.found, found); + if !more.is_empty() { + state.bus.emit(tw_api::Event::SecretsFound { + id: started.id, + provider: provider.name.clone(), + replaced: mode.acts(), + items: crate::guard::items(&more), + at_ms: crate::server::now_ms(), + }); + } + out.model = c.renamed.map(|r| r.model); + out.rewritten = Some(Rewritten { + body: c.body, + path: c.path, + decoded: Some(decoded), + ledger, + }); + Ok(out) +} + +/// 插件换上的模型名,这把密钥用不用得了。**和路由规则改写的模型名过同一关**(见 +/// `super::admit` 里的说法):密钥的模型范围管的是发出去的模型,谁改的都一样。 +fn allowed(rt: &Runtime, req: &Inbound, r: &crate::plugin::request::Renamed) -> Result<(), Msg> { + let allow = rt + .config + .clients + .iter() + .find(|c| c.name == req.client_name) + .and_then(|c| c.allow.as_deref()); + match allow { + Some(patterns) + if !patterns + .iter() + .any(|p| tw_engine::rule::glob_match(p, &r.model)) => + { + Err(msg!( + "gw.plugin.model_not_allowed", + plugin = r.by.clone(), model = r.model.clone(), key = req.client_name.clone() => + "Plugin `{plugin}` changed the model to {model}, which gateway key `{key}` may not \ + use, so the request was not sent." + )) + } + _ => Ok(()), + } +} diff --git a/crates/tw-gateway/src/server/pipeline/relay.rs b/crates/tw-gateway/src/server/pipeline/relay.rs index ae4eaa7..2941c4a 100644 --- a/crates/tw-gateway/src/server/pipeline/relay.rs +++ b/crates/tw-gateway/src/server/pipeline/relay.rs @@ -38,6 +38,7 @@ pub(super) fn respond( provider, ledger, session, + .. } = served; let status = StatusCode::from_u16(upstream.status().as_u16()).unwrap_or(StatusCode::BAD_GATEWAY); diff --git a/crates/tw-gateway/src/ws.rs b/crates/tw-gateway/src/ws.rs index 4ca9d10..cd92326 100644 --- a/crates/tw-gateway/src/ws.rs +++ b/crates/tw-gateway/src/ws.rs @@ -19,10 +19,12 @@ //! 一次回答数,超了只切掉那一次回答(替它发 `response.failed`),连接照常。 //! //! 脚本插件也在这条路上跑(见 [`crate::plugin`]):客户端发来的每个 -//! `response.create` 是一次请求,排在请求防护和脱敏之前过请求钩子;上游每一次回答 +//! `response.create` 是一次请求。**这条路只有一跳**(升级时就连定了那一家,不换), +//! 所以每个 `response.create` 过一遍请求钩子:上游是这条连接连的那一家,模型名是这一帧 +//! 写的(WebSocket 上没有规则改写)。位置和 HTTP 那条路的一跳一样 —— 请求防护先看 +//! 客户端的原话,插件改过的再看一遍插件加进来的,然后才脱敏、发出。上游每一次回答 //! (`response.created` 到 `response.completed`)起一组回答钩子的实例,排在占位符 -//! 还原之后、工具墙之前 —— 和 HTTP 那条路的位置一样。插件出错而策略是拒绝时,切掉 -//! 的是那一次回答,连接照常。 +//! 还原之后、工具墙之前。插件出错而策略是拒绝时,切掉的是那一次回答,连接照常。 //! //! # 两条明说的边界 //! @@ -152,8 +154,10 @@ struct Pipes { id: u64, /// 范围里可能有插件时才有 plugins: Option, - /// 最近一次 `response.create` 要的模型和它的密钥映射:回答钩子用 - asked_model: String, + /// 最近一次 `response.create`:客户端要的模型、发出去的模型(插件可能换了它)和 + /// 它的密钥映射。回答钩子用 + requested_model: String, + sent_model: String, bridge: Option, /// 这一次回答的回答钩子 reply: Option, @@ -232,7 +236,8 @@ pub async fn proxy( provider: upstream.provider.name, id, plugins, - asked_model: String::new(), + requested_model: String::new(), + sent_model: String::new(), bridge: None, reply: None, }; @@ -376,7 +381,15 @@ async fn pump( let Some(Ok(m)) = msg else { break End::Closed }; let out = match m { Message::Text(t) => { - // 插件的请求钩子:排在请求防护和脱敏之前,它们看的是插件改过的那一版 + // 请求防护在插件和脱敏之前:看的是客户端的原话 + if let Some(why) = screen_frame(&state, p, t.as_str()) { + let _ = c_tx.send(Message::Text( + format!("[ThinkWatch] {}", why.text).into(), + )).await; + break End::Cut(why); + } + // 插件的请求钩子:改过的那一版再过一遍请求防护(只看插件加进来的), + // 脱敏换的是改过的那一版 let t = match plugin_request(&state, p, t.as_str()).await { Ok(t) => t, Err(why) => { @@ -386,13 +399,6 @@ async fn pump( break End::Cut(why); } }; - // 请求防护在脱敏之前:看的是客户端的原话 - if let Some(why) = screen_frame(&state, p, t.as_str()) { - let _ = c_tx.send(Message::Text( - format!("[ThinkWatch] {}", why.text).into(), - )).await; - break End::Cut(why); - } let mode = p.rules.redact_mode; let found = crate::guard::find(mode, &p.rules.redact, t.as_bytes()); if found.is_empty() { @@ -624,36 +630,89 @@ async fn fail_response(p: &mut Pipes, c_tx: &mut ClientSink, why: Msg) -> Flow { /// 一次 `response.create` 过插件的请求钩子。返回要发给上游的那一帧(插件改过的话是 /// 改过的),被拒了返回告诉客户端的那句话。别的帧原样。 +/// +/// 这条路只有一跳:上游是这条连接连的那一家,发给它的模型名就是这一帧写的,运行记在 +/// 第 0 跳上。插件改过的那一版**再过一遍请求防护**,只看插件加进来的(客户端的原话已经 +/// 在 [`screen_frame`] 看过了)。 async fn plugin_request(state: &AppState, p: &mut Pipes, text: &str) -> Result { let Some(pc) = p.plugins.as_ref() else { return Ok(text.to_string()); }; - let create = serde_json::from_str::(text) + let Some(frame) = serde_json::from_str::(text) .ok() - .is_some_and(|v| v.get("type").and_then(|t| t.as_str()) == Some("response.create")); - if !create { + .filter(|v| v.get("type").and_then(|t| t.as_str()) == Some("response.create")) + else { return Ok(text.to_string()); - } - let asked = crate::plugin::request::Asked { - dialect: tw_dialect::ir::Dialect::Responses, - path: "/responses", - client: pc.client.as_deref(), }; + let requested = frame + .get("model") + .and_then(|m| m.as_str()) + .unwrap_or_default() + .to_string(); let body = bytes::Bytes::copy_from_slice(text.as_bytes()); - match crate::plugin::request::run(&pc.pool, &pc.set, &p.rules.redact, &asked, &body).await { - Ok(plugged) => { - crate::plugin::request::record(state, p.id, &plugged); - p.asked_model = plugged.model.clone(); - p.bridge = plugged.bridge.clone(); - Ok(match &plugged.body { - Some(b) => String::from_utf8_lossy(b).into_owned(), - None => text.to_string(), - }) - } + let mut hook = crate::plugin::request::Hook::new( + &pc.set, + p.rules.redact.clone(), + tw_dialect::ir::Dialect::Responses, + "/responses", + pc.client.as_deref(), + &body, + ); + let to = crate::plugin::request::Target { + upstream: &p.provider, + model: &requested, + requested_model: &requested, + attempt: 0, + }; + let plugged = match hook.attempt(&pc.pool, &to).await { + Ok(plugged) => plugged, Err(refused) => { - crate::plugin::request::record(state, p.id, &refused.plugged); - Err(refused.why) + crate::plugin::request::record(state, p.id, &refused.runs); + return Err(refused.why); + } + }; + crate::plugin::request::record(state, p.id, &plugged.runs); + let sent = plugged + .changed + .as_ref() + .and_then(|c| c.renamed.as_ref()) + .map_or_else(|| requested.clone(), |r| r.model.clone()); + p.bridge = plugged.bridge; + let out = match plugged.changed { + None => text.to_string(), + Some(c) => { + if let Some(why) = screen_changed(state, p, &frame, &c.value) { + return Err(why); + } + String::from_utf8_lossy(&c.body).into_owned() + } + }; + p.requested_model = requested; + p.sent_model = sent; + Ok(out) +} + +/// 插件改过的那一帧再看一遍请求防护:**只看插件加进来的**(见 +/// [`crate::guard::screen_more`])。两份都要按 Responses 解得开;解不开的只查藏匿字符, +/// 和 [`screen_frame`] 一样 +fn screen_changed( + state: &AppState, + p: &Pipes, + before: &serde_json::Value, + after: &serde_json::Value, +) -> Option { + let s = &p.rules.screen; + if !s.hidden_mode.detects() && !s.content_mode.detects() { + return None; + } + let decode = |v: &serde_json::Value| { + tw_dialect::convert::decode(tw_dialect::ir::Dialect::Responses, v, "/responses", None).ok() + }; + match (decode(before), decode(after)) { + (Some(b), Some(a)) => { + crate::guard::screen_more(&state.bus, p.id, &p.provider, s, &b.request, &a.request) } + _ => crate::guard::screen_text(&state.bus, p.id, &p.provider, s, &after.to_string()), } } @@ -670,9 +729,11 @@ async fn start_reply(state: &AppState, p: &mut Pipes) -> Result<(), Msg> { let ctx = crate::plugin::reply::ReplyCtx { dialect: tw_dialect::ir::Dialect::Responses, client: pc.client.as_deref(), - model: &p.asked_model, + model: &p.sent_model, + requested_model: &p.requested_model, upstream: &p.provider, request_id: p.id, + attempt: 0, }; match crate::plugin::reply::Chain::start(state, &pc.set, bridge, &ctx).await { Ok(Some(chain)) => { diff --git a/crates/tw-gateway/tests/plugins_reply.rs b/crates/tw-gateway/tests/plugins_reply.rs index 88d223a..3b5e671 100644 --- a/crates/tw-gateway/tests/plugins_reply.rs +++ b/crates/tw-gateway/tests/plugins_reply.rs @@ -58,6 +58,15 @@ fn provider(base: SocketAddr, protocol: Protocol) -> Provider { } async fn gateway(p: Provider, security: Security, entries: Vec>) -> Gw { + gateway_routed(p, security, Vec::new(), entries).await +} + +async fn gateway_routed( + p: Provider, + security: Security, + routes: Vec, + entries: Vec>, +) -> Gw { let cfg = Config { version: 1, listen: Listen::default(), @@ -68,6 +77,7 @@ async fn gateway(p: Provider, security: Security, entries: Vec>) -> }], providers: vec![p], security, + routes, ..Default::default() }; let state = tw_gateway::AppState::new(cfg).unwrap(); @@ -235,6 +245,57 @@ async fn a_converted_stream_is_rewritten_in_the_clients_format() { assert_eq!(ctx["format"], "anthropic"); assert_eq!(ctx["upstream"], "up"); assert_eq!(ctx["model"], "claude-sonnet-4-5"); + assert_eq!(ctx["requested_model"], "claude-sonnet-4-5"); +} + +/// 规则把模型改了名发给上游:回答钩子的 `ctx.model` 是发出去的那个,`requested_model` 是 +/// 客户端要的;范围里的模型也按发出去的那个对 +#[tokio::test] +async fn reply_hooks_see_and_are_scoped_by_the_model_sent_upstream() { + let up = upstream("text/event-stream", anthropic_sse(&["hello"], None)).await; + let saw = Arc::new(Mutex::new(Value::Null)); + let s = saw.clone(); + let spy = Double::new("spy") + .permit(&[Permission::ReplyText]) + .on_reply(true, false, false, move |ctx| { + *s.lock().unwrap() = ctx; + Ok(Box::new(Closures { + text: Box::new(|t| Invocation::ok(Some(t.to_uppercase()))), + end: Box::new(|| Invocation::ok(None)), + tool: Box::new(|_| Invocation::ok(ToolCallOutcome::Unchanged)), + })) + }); + let rename = vec![tw_engine::RouteSet::default_with(vec![tw_engine::Rule { + name: "rename".into(), + when: Default::default(), + to: Some("up".into()), + set: Some(tw_engine::SetAction { + model: Some("glm-4.6".into()), + ..Default::default() + }), + deny: None, + }])]; + let gw = gateway_routed( + provider(up, Protocol::Anthropic), + Security::default(), + rename, + vec![ + entry_with("spy", spy, |a| a.scope.models = vec!["glm-*".into()]), + entry_with("asked", upper(), |a| { + a.scope.models = vec!["claude-*".into()] + }), + ], + ) + .await; + let (status, body) = post(&gw, "/v1/messages", &ask(true)).await; + assert_eq!(status, 200, "{body}"); + // 只有管发出去的那个模型的插件跑了(`upper` 跑了也是大写,所以看计数) + assert_eq!(anthropic_text(&body), "HELLO"); + assert_eq!(gw.stats("asked").calls, 0); + let ctx = saw.lock().unwrap().clone(); + assert_eq!(ctx["model"], "glm-4.6"); + assert_eq!(ctx["requested_model"], "claude-sonnet-4-5"); + assert_eq!(ctx["upstream"], "up"); } #[tokio::test] diff --git a/crates/tw-gateway/tests/plugins_request.rs b/crates/tw-gateway/tests/plugins_request.rs index bfeeb47..6800b1a 100644 --- a/crates/tw-gateway/tests/plugins_request.rs +++ b/crates/tw-gateway/tests/plugins_request.rs @@ -3,6 +3,9 @@ //! 插件用替身(`tw_gateway::plugin::host::double`):钩子是 Rust 闭包。要证明的是 //! 接线 —— 跑在哪一步、跑几次、改过的请求谁看得见、拒绝和出错怎么回给客户端、 //! 插件看到的是不是占位符 —— 这些和插件用什么语言写无关。 +//! +//! 请求钩子排在路由之后、**每发往一个上游跑一次**(契约附录二):换上游从客户端的 +//! 原话重来,范围按这一次的上游和发出去的模型名算,同一家重发不重跑。 use std::net::SocketAddr; use std::sync::atomic::{AtomicUsize, Ordering}; @@ -23,19 +26,39 @@ use tw_gateway::plugin::{ const USER_KEY: &str = "sk-ant-api03-USERSOWNKEYAAAAAAAAAAAAAA"; -/// 假上游:记下收到的每一个请求(路径和正文),按格式回一个最简单的回答。 -/// `fail_first` 次请求回 500,用来看故障转移。 +/// 假上游:记下收到的每一个请求(路径和正文,正文连同原样的字节),按格式回一个最简单 +/// 的回答。`fail_first` 次请求回 500,用来看故障转移。 #[derive(Clone, Default)] struct Upstream { seen: Arc>>, + raw: Arc>>, fail_first: Arc, } +impl Upstream { + /// 先回 `n` 次 500 + fn failing(n: usize) -> Self { + let u = Upstream::default(); + u.fail_first.store(n, Ordering::SeqCst); + u + } + + /// 收到的第 `i` 个请求的系统提示,整个写成一串(Anthropic 的 system 可能是几块) + fn system(&self, i: usize) -> String { + self.seen.lock().unwrap()[i].1["system"].to_string() + } + + fn hits(&self) -> usize { + self.seen.lock().unwrap().len() + } +} + async fn start_upstream(u: Upstream) -> SocketAddr { async fn answer(State(u): State, uri: Uri, body: Bytes) -> axum::response::Response { let path = uri.path().to_string(); let v: Value = serde_json::from_slice(&body).unwrap_or(Value::Null); u.seen.lock().unwrap().push((path.clone(), v)); + u.raw.lock().unwrap().push(body.clone()); if u.fail_first.load(Ordering::SeqCst) > 0 { u.fail_first.fetch_sub(1, Ordering::SeqCst); return axum::response::Response::builder() @@ -78,6 +101,7 @@ struct Gw { addr: SocketAddr, state: tw_gateway::AppState, bodies: tokio::sync::mpsc::Receiver, + runs: Arc>>, } impl Gw { @@ -86,6 +110,35 @@ impl Gw { self.state.runtime().plugins.get(id).unwrap().stats.view() } + /// 记下的每一次运行:`(插件, 结局, 跑在第几跳)`,按先后 + fn runs(&self) -> Vec<(String, String, u64)> { + self.runs + .lock() + .unwrap() + .iter() + .map(|r| { + ( + r.run.plugin_id.clone(), + r.run.outcome.slug().to_string(), + r.run + .detail + .as_ref() + .and_then(|d| d["attempt"].as_u64()) + .expect("every run says which attempt it ran on"), + ) + }) + .collect() + } + + /// 插件改过之后存下来的那一份请求体(没有就是 None) + async fn after_plugins(&mut self) -> Option { + self.bodies() + .await + .into_iter() + .find(|b| b.kind == tw_gateway::bodies::BodyKind::AfterPlugins) + .map(|b| serde_json::from_slice(&b.body).unwrap()) + } + /// 交去存的正文,一直收到 `wait` 里再没有新的为止 async fn bodies(&mut self) -> Vec { let mut out = Vec::new(); @@ -122,7 +175,20 @@ async fn gateway_with( mode, ..Default::default() }; - let cfg = Config { + gateway_of( + Config { + providers, + security, + ..config() + }, + entries, + ) + .await +} + +/// 一个密钥(`claude-code`),别的都空着 +fn config() -> Config { + Config { version: 1, listen: Listen::default(), clients: vec![Client { @@ -130,14 +196,24 @@ async fn gateway_with( key: "tw-testkey".into(), ..Default::default() }], - providers, - security, ..Default::default() - }; + } +} + +async fn gateway_of(cfg: Config, entries: Vec>) -> Gw { let state = tw_gateway::AppState::new(cfg).unwrap(); state.swap_plugins(PluginSet::new(entries)); let (tx, bodies) = tw_gateway::bodies::channel(); state.set_body_sink(tx); + let runs: Arc>> = Arc::default(); + let (rtx, mut rrx) = tokio::sync::mpsc::channel(tw_gateway::plugin::RUN_CHANNEL_CAP); + state.set_plugin_sink(rtx); + let r = runs.clone(); + tokio::spawn(async move { + while let Some(rec) = rrx.recv().await { + r.lock().unwrap().push(rec); + } + }); let addr = tw_gateway::serve(state.clone(), ([127, 0, 0, 1], 0).into()) .await .unwrap(); @@ -146,6 +222,7 @@ async fn gateway_with( addr, state, bodies, + runs, } } @@ -263,8 +340,11 @@ async fn an_unchanged_result_sends_the_original_bytes() { .await .unwrap(); assert_eq!(r.status(), 200); - let seen = up.seen.lock().unwrap().clone(); - assert_eq!(seen[0].1, serde_json::from_str::(raw).unwrap()); + // 一个字节都不差:插件跑在这一跳上,没改就发原样 + assert_eq!( + up.raw.lock().unwrap()[0], + Bytes::from_static(raw.as_bytes()) + ); let st = gw.stats("same"); assert_eq!((st.calls, st.changed), (1, 0)); let mut gw = gw; @@ -276,27 +356,38 @@ async fn an_unchanged_result_sends_the_original_bytes() { ); } +/// 插件换的模型名只换掉发给这一家的名字:**不重新路由**(换成 `gpt-5`,按规则本该去 b, +/// 还是发给 a),记录上客户端要的那个照旧,尝试链上记的是发出去的那个 #[tokio::test] -async fn a_new_model_is_what_routing_uses_and_the_record_keeps_the_asked_one() { +async fn a_new_model_renames_what_this_upstream_gets_and_the_record_keeps_the_asked_one() { let up = Upstream::default(); let base = start_upstream(up.clone()).await; + let other = Upstream::default(); + let other_base = start_upstream(other.clone()).await; let swap = Double::new("swap model") .permit(&[Permission::Params]) .on_request(|mut view, ctx| { assert_eq!(ctx["model"], "claude-sonnet-4-5"); - view["params"]["model"] = json!("claude-opus-4-5"); + view["params"]["model"] = json!("gpt-5"); Invocation::ok(RequestOutcome::Changed(view)) }); - let gw = gateway( - vec![provider("a", base, Protocol::Anthropic)], - SecurityMode::Off, + let gw = gateway_of( + Config { + providers: vec![ + provider("a", base, Protocol::Anthropic), + provider("b", other_base, Protocol::Anthropic), + ], + routes: by_model(None), + ..config() + }, vec![entry("swap", swap)], ) .await; let rx = gw.state.bus.subscribe(); let (status, _) = post(&gw, "/v1/messages", &anthropic_body()).await; assert_eq!(status, 200); - assert_eq!(up.seen.lock().unwrap()[0].1["model"], "claude-opus-4-5"); + assert_eq!((up.hits(), other.hits()), (1, 0), "the plugin re-routed it"); + assert_eq!(up.seen.lock().unwrap()[0].1["model"], "gpt-5"); let mut rx = rx; let mut started_model = None; let mut attempt_model = None; @@ -310,7 +401,95 @@ async fn a_new_model_is_what_routing_uses_and_the_record_keeps_the_asked_one() { } } assert_eq!(started_model.as_deref(), Some("claude-sonnet-4-5")); - assert_eq!(attempt_model.as_deref(), Some("claude-opus-4-5")); + assert_eq!(attempt_model.as_deref(), Some("gpt-5")); +} + +/// `claude-*` 发给 a(`set_model` 给了的话改名发),别的发给 b +fn by_model(set_model: Option<&str>) -> Vec { + vec![tw_engine::RouteSet::default_with(vec![ + tw_engine::Rule { + name: "claude".into(), + when: tw_engine::rule::When { + model: Some("claude-*".into()), + ..Default::default() + }, + to: Some("a".into()), + set: set_model.map(|m| tw_engine::SetAction { + model: Some(m.into()), + ..Default::default() + }), + deny: None, + }, + tw_engine::Rule { + name: "rest".into(), + when: Default::default(), + to: Some("b".into()), + set: None, + deny: None, + }, + ])] +} + +/// 插件看到的模型名是发给这一家的(规则改写之后的):`ctx.model`、视图里的 `model` 和 +/// `params.model` 都是它;`ctx.requested_model` 是客户端要的,`ctx.upstream` 是这一家。 +/// 插件再换名字,盖过规则的改写 +#[tokio::test] +async fn the_hook_sees_the_upstream_the_sent_model_and_the_asked_one() { + let up = Upstream::default(); + let base = start_upstream(up.clone()).await; + let saw = Arc::new(Mutex::new(Value::Null)); + let s = saw.clone(); + let look = Double::new("look") + .permit(&[Permission::System, Permission::Params]) + .on_request(move |mut view, ctx| { + *s.lock().unwrap() = json!({ "view": view.clone(), "ctx": ctx }); + view["system"] = json!("short"); + Invocation::ok(RequestOutcome::Changed(view)) + }); + let gw = gateway_of( + Config { + providers: vec![provider("a", base, Protocol::Anthropic)], + routes: by_model(Some("glm-4.6")), + ..config() + }, + vec![entry("look", look)], + ) + .await; + let (status, _) = post(&gw, "/v1/messages", &anthropic_body()).await; + assert_eq!(status, 200); + let saw = saw.lock().unwrap().clone(); + assert_eq!(saw["ctx"]["upstream"], "a"); + assert_eq!(saw["ctx"]["model"], "glm-4.6"); + assert_eq!(saw["ctx"]["requested_model"], "claude-sonnet-4-5"); + assert_eq!(saw["ctx"]["client"], "claude-code"); + assert_eq!(saw["view"]["model"], "glm-4.6"); + assert_eq!(saw["view"]["params"]["model"], "glm-4.6"); + // 没改模型名:规则的改写照常落到请求体上 + let sent = up.seen.lock().unwrap()[0].1.clone(); + assert_eq!(sent["model"], "glm-4.6"); + assert_eq!(sent["system"][0]["text"], "short"); + + // 插件换了名字:发出去的是插件的,不是规则的 + let up = Upstream::default(); + let base = start_upstream(up.clone()).await; + let rename = Double::new("rename") + .permit(&[Permission::Params]) + .on_request(|mut view, _| { + view["params"]["model"] = json!("glm-5"); + Invocation::ok(RequestOutcome::Changed(view)) + }); + let gw = gateway_of( + Config { + providers: vec![provider("a", base, Protocol::Anthropic)], + routes: by_model(Some("glm-4.6")), + ..config() + }, + vec![entry("rename", rename)], + ) + .await; + let (status, _) = post(&gw, "/v1/messages", &anthropic_body()).await; + assert_eq!(status, 200); + assert_eq!(up.seen.lock().unwrap()[0].1["model"], "glm-5"); } #[tokio::test] @@ -539,17 +718,12 @@ async fn out_of_scope_plugins_do_not_run_and_count_tokens_is_left_alone() { assert_eq!(calls.load(Ordering::SeqCst), 1); } -#[tokio::test] -async fn the_plugin_runs_once_even_when_the_request_fails_over() { - let up = Upstream::default(); - up.fail_first.store(1, Ordering::SeqCst); - let base = start_upstream(up.clone()).await; - let calls = Arc::new(AtomicUsize::new(0)); - let c = calls.clone(); - let counting = Double::new("count") +/// 把这一次的上游写进系统提示的插件,数自己跑了几次 +fn tag(calls: Arc) -> Double { + Double::new("tag") .permit(&[Permission::System]) - .on_request(move |mut view, _| { - c.fetch_add(1, Ordering::SeqCst); + .on_request(move |mut view, ctx| { + calls.fetch_add(1, Ordering::SeqCst); // 跑在插件线程上,不在 tokio 的线程上 assert!( std::thread::current() @@ -557,26 +731,316 @@ async fn the_plugin_runs_once_even_when_the_request_fails_over() { .unwrap_or_default() .starts_with("tw-plugin-") ); - view["system"] = json!("changed"); + let s = view["system"].as_str().unwrap().to_string(); + view["system"] = json!(format!( + "{s}\n\n[for {}]", + ctx["upstream"].as_str().unwrap() + )); Invocation::ok(RequestOutcome::Changed(view)) - }); - let gw = gateway( - vec![ - provider("a", base, Protocol::Anthropic), - provider("b", base, Protocol::Anthropic), - ], + }) +} + +/// a 先回 500,b 接下:两家各一个假上游 +async fn a_fails_then_b() -> (Upstream, Upstream, Vec) { + let a = Upstream::failing(1); + let b = Upstream::default(); + let providers = vec![ + provider("a", start_upstream(a.clone()).await, Protocol::Anthropic), + provider("b", start_upstream(b.clone()).await, Protocol::Anthropic), + ]; + (a, b, providers) +} + +/// 故障转移从客户端的原话重来:插件每一跳跑一次,给 a 的改动到不了 b。存下来的「插件 +/// 改过的请求」是回答的那一家(b)收到的那一份 +#[tokio::test] +async fn failing_over_starts_again_from_the_clients_original_request() { + let (a, b, providers) = a_fails_then_b().await; + let calls = Arc::new(AtomicUsize::new(0)); + let mut gw = gateway( + providers, SecurityMode::Off, - vec![entry("count", counting)], + vec![entry("tag", tag(calls.clone()))], ) .await; let (status, _) = post(&gw, "/v1/messages", &anthropic_body()).await; assert_eq!(status, 200); - let seen = up.seen.lock().unwrap().clone(); - assert_eq!(seen.len(), 2, "one failure, one success"); - // 两跳发出去的是同一份改过的请求,插件只跑了一次 - assert_eq!(seen[0].1, seen[1].1); - assert_eq!(seen[1].1["system"][0]["text"], "changed"); - assert_eq!(calls.load(Ordering::SeqCst), 1); + assert_eq!((a.hits(), b.hits()), (1, 1), "one failure, one success"); + assert!(a.system(0).contains("[for a]"), "{}", a.system(0)); + let to_b = b.system(0); + assert!(to_b.contains("[for b]"), "{to_b}"); + assert!(!to_b.contains("[for a]"), "a's edit reached b: {to_b}"); + // 缓存断点还在原来那一块上 + assert_eq!( + b.seen.lock().unwrap()[0].1["system"][0]["cache_control"], + json!({ "type": "ephemeral" }) + ); + assert_eq!(calls.load(Ordering::SeqCst), 2); + // 每一跳一条,带着第几跳 + assert_eq!( + gw.runs(), + [ + ("tag".to_string(), "changed".to_string(), 0), + ("tag".to_string(), "changed".to_string(), 1) + ] + ); + let after = gw.after_plugins().await.expect("the body after plugins"); + assert!(after["system"].to_string().contains("[for b]"), "{after}"); + assert!(!after["system"].to_string().contains("[for a]"), "{after}"); +} + +/// 只管 a 的插件只在发往 a 时跑;发往 b 的是客户端的原话。反过来也一样 +#[tokio::test] +async fn a_plugin_scoped_to_one_upstream_runs_only_for_it() { + for only in ["a", "b"] { + let (a, b, providers) = a_fails_then_b().await; + let calls = Arc::new(AtomicUsize::new(0)); + let mut gw = gateway( + providers, + SecurityMode::Off, + vec![entry_with("tag", tag(calls.clone()), |e| { + e.scope.upstreams = vec![only.into()] + })], + ) + .await; + let (status, _) = post(&gw, "/v1/messages", &anthropic_body()).await; + assert_eq!(status, 200, "{only}"); + assert_eq!(calls.load(Ordering::SeqCst), 1, "{only}"); + let (tagged, untouched) = if only == "a" { (&a, &b) } else { (&b, &a) }; + assert!( + tagged.system(0).contains(&format!("[for {only}]")), + "{only}: {}", + tagged.system(0) + ); + assert_eq!( + untouched.seen.lock().unwrap()[0].1, + anthropic_body(), + "{only}: the plugin ran for an upstream outside its scope" + ); + let attempt = if only == "a" { 0 } else { 1 }; + assert_eq!( + gw.runs(), + [("tag".to_string(), "changed".to_string(), attempt)] + ); + // 回答的那一家收到的没被改过,就没有「插件改过的请求」 + let after = gw.after_plugins().await; + if only == "a" { + assert_eq!(after, None, "b got the original"); + } else { + assert!(after.is_some()); + } + } +} + +/// 坏了的插件只拦管得着的那一跳:它只管 a,请求只去 b 时照常;要发往它管的那一家时 +/// 拒绝整个请求(不换下一家),发往别家的那一跳照常发过 +#[tokio::test] +async fn a_broken_plugin_rejects_only_attempts_in_its_scope() { + let broken = |upstream: &str| { + let mut a = double::active("old", Double::new("Old")); + a.state = PluginState::Broken(Broken::Changed); + a.scope.upstreams = vec![upstream.into()]; + Arc::new(a) + }; + // 只去 b:管 a 的坏插件不拦它 + let b = Upstream::default(); + let gw = gateway( + vec![provider( + "b", + start_upstream(b.clone()).await, + Protocol::Anthropic, + )], + SecurityMode::Off, + vec![broken("a")], + ) + .await; + let (status, body) = post(&gw, "/v1/messages", &anthropic_body()).await; + assert_eq!(status, 200, "{body}"); + assert_eq!(b.hits(), 1); + assert!(gw.runs().is_empty(), "{:?}", gw.runs()); + + // a 失败、换到 b:管 b 的坏插件拦下整个请求,b 一个字节都没收到 + let (a, b, providers) = a_fails_then_b().await; + let gw = gateway(providers, SecurityMode::Off, vec![broken("b")]).await; + let rx = gw.state.bus.subscribe(); + let (status, body) = post(&gw, "/v1/messages", &anthropic_body()).await; + assert_eq!(status, 403, "{body}"); + assert!( + body["error"]["message"] + .as_str() + .unwrap() + .contains("changed on disk"), + "{body}" + ); + assert_eq!((a.hits(), b.hits()), (1, 0)); + assert_eq!(gw.runs(), [("old".to_string(), "error".to_string(), 1)]); + // 尝试链上:a 回了 500,b 被插件拦下 + let attempts = routed(rx).await; + assert_eq!(attempts.len(), 2, "{attempts:?}"); + assert_eq!( + (attempts[0].provider.as_str(), attempts[0].status), + ("a", Some(500)) + ); + assert_eq!(attempts[1].provider, "b"); + assert_eq!( + attempts[1].error.as_ref().map(|e| e.code.as_str()), + Some("gw.plugin.changed") + ); +} + +/// 插件出错(策略是拒绝)、插件 `reject`:拒绝的是整个请求,不换到下一家 +#[tokio::test] +async fn a_plugin_refusal_rejects_the_whole_request_without_failing_over() { + let only_a = |rejects: bool| { + Double::new("picky") + .permit(&[Permission::System]) + .on_request(move |_, ctx| { + if ctx["upstream"] != "a" { + return Invocation::ok(RequestOutcome::Unchanged); + } + if rejects { + Invocation::ok(RequestOutcome::Rejected("not for a".into())) + } else { + Invocation::err(RunError::Threw { + message: "only on a".into(), + stack: None, + }) + } + }) + }; + for (rejects, code) in [ + (true, "gw.plugin.rejected"), + (false, "gw.plugin.request_failed"), + ] { + let a = Upstream::default(); + let b = Upstream::default(); + let gw = gateway( + vec![ + provider("a", start_upstream(a.clone()).await, Protocol::Anthropic), + provider("b", start_upstream(b.clone()).await, Protocol::Anthropic), + ], + SecurityMode::Off, + vec![entry("picky", only_a(rejects))], + ) + .await; + let rx = gw.state.bus.subscribe(); + let (status, body) = post(&gw, "/v1/messages", &anthropic_body()).await; + assert_eq!(status, 403, "{rejects}: {body}"); + assert_eq!((a.hits(), b.hits()), (0, 0), "{rejects}"); + let attempts = routed(rx).await; + assert_eq!(attempts.len(), 1, "{rejects}: {attempts:?}"); + assert_eq!(attempts[0].provider, "a"); + assert_eq!( + attempts[0].error.as_ref().map(|e| e.code.as_str()), + Some(code), + "{rejects}" + ); + } +} + +/// 这个请求的尝试链(路由事件里的) +async fn routed( + mut rx: tokio::sync::broadcast::Receiver, +) -> Vec { + while let Ok(Ok(ev)) = tokio::time::timeout(Duration::from_millis(500), rx.recv()).await { + if let tw_api::Event::RequestRouted { attempts, .. } = ev { + return attempts; + } + } + panic!("no routing event"); +} + +/// OAuth 上游回 401:换一个 token、向同一家再发一次。**用这一跳定好的请求体,插件不重跑** +#[tokio::test] +async fn a_same_upstream_oauth_retry_does_not_run_the_hook_again() { + // token 端点:每次换发一个新的(at-1、at-2……) + let issued = Arc::new(AtomicUsize::new(0)); + let i = issued.clone(); + let tokens = Router::new().route( + "/token", + axum::routing::post(move || { + let n = i.fetch_add(1, Ordering::SeqCst) + 1; + async move { + axum::Json(json!({ + "access_token": format!("at-{n}"), "token_type": "Bearer", "expires_in": 3600 + })) + } + }), + ); + let l = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap(); + let token_addr = l.local_addr().unwrap(); + tokio::spawn(async move { axum::serve(l, tokens).await.unwrap() }); + // 上游:配置里那个 token(at-0)在别处被吊销了,回 401;换来的照常回答 + let bodies: Arc>> = Arc::default(); + let seen = bodies.clone(); + let app = Router::new().fallback(axum::routing::post( + move |headers: axum::http::HeaderMap, body: Bytes| { + let seen = seen.clone(); + async move { + seen.lock().unwrap().push(body); + // Anthropic 的上游,token 放在 x-api-key 里 + let first = ["x-api-key", "authorization"].iter().any(|h| { + headers + .get(*h) + .and_then(|v| v.to_str().ok()) + .is_some_and(|v| v.ends_with("at-0")) + }); + let (status, reply) = if first { + (401, json!({ "type": "error", "error": { "type": "authentication_error", "message": "revoked" } })) + } else { + (200, json!({ "id": "msg_1", "type": "message", "role": "assistant", "model": "m", + "content": [{ "type": "text", "text": "ok" }], "stop_reason": "end_turn", + "usage": { "input_tokens": 1, "output_tokens": 1 } })) + }; + axum::response::Response::builder() + .status(status) + .header("content-type", "application/json") + .body(axum::body::Body::from(reply.to_string())) + .unwrap() + } + }, + )); + let l = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap(); + let up_addr = l.local_addr().unwrap(); + tokio::spawn(async move { axum::serve(l, app).await.unwrap() }); + let oauth = Provider { + name: "a".into(), + base_url: format!("http://{up_addr}"), + oauth: Some(tw_config::OAuth { + access: Some("at-0".into()), + expires_at: None, + refresh: "rt-test".into(), + endpoint: format!("http://{token_addr}/token"), + client_id: Some("tw-test".into()), + client_secret: None, + refresh_before: Some("5m".into()), + }), + protocol: Some(Protocol::Anthropic), + ..Default::default() + }; + let calls = Arc::new(AtomicUsize::new(0)); + let gw = gateway( + vec![oauth], + SecurityMode::Off, + vec![entry("tag", tag(calls.clone()))], + ) + .await; + let (status, body) = post(&gw, "/v1/messages", &anthropic_body()).await; + assert_eq!(status, 200, "{body}"); + let bodies = bodies.lock().unwrap().clone(); + assert_eq!( + bodies.len(), + 2, + "refused, then sent again with a fresh token" + ); + assert_eq!(bodies[0], bodies[1], "the retry was not the same request"); + assert!(String::from_utf8_lossy(&bodies[1]).contains("[for a]")); + assert_eq!( + calls.load(Ordering::SeqCst), + 1, + "the hook ran again for the retry" + ); + assert_eq!(issued.load(Ordering::SeqCst), 1); } /// 插件看到的是占位符,和脱敏开在哪一档无关;它改过的地方占位符换回真值,然后才轮到 @@ -712,9 +1176,77 @@ async fn screening_sees_the_body_after_plugins() { security, ) .await; + let rx = gw.state.bus.subscribe(); let (status, body) = post(&gw, "/v1/messages", &anthropic_body()).await; assert_eq!(status, 403, "{body}"); assert!(up.seen.lock().unwrap().is_empty()); + // 拦在这一跳上:尝试链上看得出本来要发给谁、为什么没发 + let attempts = routed(rx).await; + assert_eq!(attempts.len(), 1, "{attempts:?}"); + assert_eq!( + attempts[0].error.as_ref().map(|e| e.code.as_str()), + Some("gw.content.refused") + ); +} + +/// 插件改过的请求再看一遍时,只报插件加进来的:客户端原话里就有的那一处命中(观察档) +/// 已经在开头报过,不因为插件改了系统提示再报一次。插件写进来的新密钥报一条,记在这一跳 +/// 的上游上 +#[tokio::test] +async fn after_a_plugin_only_what_it_added_is_reported_again() { + const WRITTEN: &str = "ghp_BBBBBBBBBBBBBBBBBBBBBBBBBBBBBBBBBBBB"; + let up = Upstream::default(); + let base = start_upstream(up.clone()).await; + let security = Security { + content: tw_config::ContentPolicy { + mode: SecurityMode::Observe, + custom: vec![tw_config::CustomContentRule { + name: "no plan".into(), + pattern: "forbidden-plan".into(), + matching: Default::default(), + action: tw_config::ContentAction::Block, + disabled: false, + }], + ..Default::default() + }, + ..Default::default() + }; + let note = Double::new("note") + .permit(&[Permission::System]) + .on_request(|mut view, _| { + let s = view["system"].as_str().unwrap().to_string(); + view["system"] = json!(format!("{s}\n\nCI token: {WRITTEN}")); + Invocation::ok(RequestOutcome::Changed(view)) + }); + let gw = gateway_with( + vec![provider("a", base, Protocol::Anthropic)], + SecurityMode::Observe, + vec![entry("note", note)], + security, + ) + .await; + let mut rx = gw.state.bus.subscribe(); + let mut body = anthropic_body(); + body["messages"][0]["content"] = json!(format!("the forbidden-plan, key {USER_KEY}")); + let (status, answer) = post(&gw, "/v1/messages", &body).await; + assert_eq!(status, 200, "{answer}"); + let (mut matched, mut secrets) = (0, Vec::new()); + while let Ok(Ok(ev)) = tokio::time::timeout(Duration::from_millis(500), rx.recv()).await { + match ev { + tw_api::Event::ContentMatched { .. } => matched += 1, + tw_api::Event::SecretsFound { + provider, items, .. + } => secrets.push((provider, items.len(), items[0].masked.clone())), + _ => {} + } + } + assert_eq!(matched, 1, "the client's own match was reported again"); + assert_eq!(secrets.len(), 2, "{secrets:?}"); + // 开头那一条是客户端的那把;插件写进来的那一把另报一条,只有它 + assert!(secrets[0].2.starts_with("sk-an"), "{secrets:?}"); + assert_eq!(secrets[1].0, "a"); + assert_eq!(secrets[1].1, 1, "{secrets:?}"); + assert!(secrets[1].2.starts_with("ghp_"), "{secrets:?}"); } #[tokio::test] diff --git a/crates/tw-gateway/tests/plugins_security.rs b/crates/tw-gateway/tests/plugins_security.rs index 7ccf8c8..1e204e7 100644 --- a/crates/tw-gateway/tests/plugins_security.rs +++ b/crates/tw-gateway/tests/plugins_security.rs @@ -12,10 +12,8 @@ //! - 没有插件改动的请求一个字节都不变;WebSocket(Codex 的 Responses WebSocket)那一路 //! 同样看占位符、同样过工具调用审查、拒绝了不发给上游。 //! -//! 标了 `#[ignore]` 的有两类,断言写的都是该有的样子: -//! - `pending`:等「先路由、再跑请求钩子」(契约附录二)落地后打开; -//! - 还没解决的问题:插件写下的占位符会被换回真值(契约 I5 的写法),计 token 的请求不经过 -//! 请求钩子,插件改的 `params.model` 不再对照密钥的模型范围。 +//! 标了 `#[ignore]` 的两条是**还没解决的问题**,断言写的是该有的样子:插件写下的占位符会被 +//! 换回真值(契约 I5 的写法),计 token 的请求不经过请求钩子。 mod plugin_harness; @@ -379,10 +377,9 @@ export function onReplyText(text) { return text.repeat(50); }"#; } #[tokio::test] -#[ignore = "addendum 2: a model a plugin writes into params.model is sent without checking it \ - against the key's model list; see the track 4 report"] async fn a_plugin_cannot_switch_to_a_model_the_key_may_not_use() { - // 密钥只许用 claude-sonnet-*;插件把模型换成 opus + // 密钥只许用 claude-sonnet-*;插件把模型换成 opus。上游的模型清单不再对(契约附录二), + // 密钥的模型范围照样管:路由规则改的名字要过这一关,插件改的也要 let to_opus = r#" export const manifest = { name: "换模型", api: 1, permissions: ["params"] }; export function onRequest(req) { req.params.model = "claude-opus-4-1"; return req; }"#; @@ -402,8 +399,6 @@ export function onRequest(req) { req.params.model = "claude-opus-4-1"; return re // ── I8:每次发往上游跑一次,换上游就从原始请求重来 ───────────────── -const PENDING: &str = "pending: route-first request hooks (contract addendum 2)"; - /// 每次运行写下这一次发往的上游和一个不会重复的记号 const NONCE: &str = r#" export const manifest = { name: "记号", api: 1, permissions: ["system"] }; @@ -424,9 +419,7 @@ async fn failing_over(plugins: Vec) -> (Upstream, Upstream, Gateway) { } #[tokio::test] -#[ignore = "pending: route-first request hooks (contract addendum 2)"] async fn failing_over_starts_again_from_the_clients_original_request() { - let _ = PENDING; let (dead, up, gw) = failing_over(vec![Plug::new("nonce", NONCE)]).await; let r = gw.ask(plain("你好", false)).await; assert_eq!(r.status, 200, "{}", r.body); @@ -487,7 +480,6 @@ async fn sending_again_without_sealed_reasoning_reuses_the_request_hook_result() } #[tokio::test] -#[ignore = "pending: route-first request hooks (contract addendum 2)"] async fn a_request_hook_runs_only_for_the_upstreams_in_its_scope() { let (dead, up, gw) = failing_over(vec![Plug::new("nonce", NONCE).upstreams(&["second"])]).await; let r = gw.ask(plain("你好", false)).await; @@ -509,7 +501,6 @@ async fn a_request_hook_runs_only_for_the_upstreams_in_its_scope() { } #[tokio::test] -#[ignore = "pending: route-first request hooks (contract addendum 2)"] async fn a_broken_plugin_refuses_only_the_attempts_in_its_scope() { // 文件变了的插件,范围只有 second:发往 relay 的请求照常,不被它拒 let up = Upstream::start(vec![Answer::Text("好的".into())]).await; @@ -526,7 +517,6 @@ async fn a_broken_plugin_refuses_only_the_attempts_in_its_scope() { } #[tokio::test] -#[ignore = "pending: route-first request hooks (contract addendum 2)"] async fn a_plugin_failure_refuses_the_whole_request_without_failing_over() { // 插件只在发往 relay 时出错:拒绝的是整个请求,不会换到 second 去 let throws = r#" @@ -579,7 +569,6 @@ fn split_by_model( } #[tokio::test] -#[ignore = "pending: route-first request hooks (contract addendum 2)"] async fn a_model_a_plugin_writes_renames_what_is_sent_without_rerouting() { // 路由按客户端的原话选了 relay;插件把模型改成 gpt-5,请求照样发给 relay,只是名字换了 let rename = r#" @@ -603,7 +592,6 @@ export function onRequest(req) { req.params.model = "gpt-5"; return req; }"#; } #[tokio::test] -#[ignore = "pending: route-first request hooks (contract addendum 2)"] async fn the_request_hook_sees_the_upstream_and_both_model_names() { // 规则把 claude-sonnet-4-5 改名成 relay-sonnet 发给 relay:ctx.model 是改名之后的, // ctx.requested_model 是客户端要的,ctx.upstream 是这一跳的上游 diff --git a/crates/tw-gateway/tests/plugins_ws.rs b/crates/tw-gateway/tests/plugins_ws.rs index fa06d3d..d1ec351 100644 --- a/crates/tw-gateway/tests/plugins_ws.rs +++ b/crates/tw-gateway/tests/plugins_ws.rs @@ -1,5 +1,6 @@ //! WebSocket 那条路上的插件:一次 `response.create` 一次请求钩子,上游的每一次回答 -//! 一组回答钩子。和 HTTP 那条路同样的位置、同样的规矩。 +//! 一组回答钩子。和 HTTP 那条路的一跳同样的位置、同样的规矩:这条路只有一跳(升级时 +//! 连定的那一家),`ctx.upstream` 就是它,运行记在第 0 跳上。 use std::net::SocketAddr; use std::sync::{Arc, Mutex}; @@ -16,7 +17,8 @@ use tw_config::{Client, Config, Listen, Provider}; use tw_gateway::plugin::host::double; use tw_gateway::plugin::host::double::{Closures, Double}; use tw_gateway::plugin::{ - Active, Invocation, PluginSet, RequestOutcome, RunError, ToolCallOutcome, + Active, Broken, Invocation, PluginSet, RequestOutcome, RunError, RunRecord, + State as PluginState, ToolCallOutcome, }; /// 假上游:记下收到的每一帧,每个 `response.create` 回一次完整的回答 @@ -94,6 +96,17 @@ async fn answer(mut sock: WebSocket, seen: Arc>>) { } async fn gateway(up: SocketAddr, entries: Vec>) -> SocketAddr { + gateway_with(up, entries, tw_config::Security::default()) + .await + .0 +} + +/// 网关,连同记下的每一次插件运行 +async fn gateway_with( + up: SocketAddr, + entries: Vec>, + security: tw_config::Security, +) -> (SocketAddr, Arc>>) { let cfg = Config { version: 1, listen: Listen::default(), @@ -109,15 +122,25 @@ async fn gateway(up: SocketAddr, entries: Vec>) -> SocketAddr { protocol: Some(tw_config::Protocol::OpenaiResponses), ..Default::default() }], + security, ..Default::default() }; let state = tw_gateway::AppState::new(cfg).unwrap(); 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, ([127, 0, 0, 1], 0).into()) .await .unwrap(); tokio::time::sleep(Duration::from_millis(40)).await; - addr + (addr, runs) } type Socket = @@ -165,8 +188,13 @@ async fn one_answer(c: &mut Socket) -> Vec { } fn entry(id: &str, d: Double) -> Arc { + entry_with(id, d, |_| {}) +} + +fn entry_with(id: &str, d: Double, f: impl FnOnce(&mut Active)) -> Arc { let mut a = double::active(id, d); a.name = format!("Plugin {id}"); + f(&mut a); Arc::new(a) } @@ -179,6 +207,10 @@ async fn each_response_create_goes_through_the_request_hook_and_each_answer_thro .on_request(|mut view, ctx| { assert_eq!(ctx["format"], "openai_responses"); assert_eq!(ctx["client"], "codex"); + // 这条路只有一跳:上游是这条连接连的那一家,模型名就是这一帧写的 + assert_eq!(ctx["upstream"], "up"); + assert_eq!(ctx["model"], "gpt-5.1-codex"); + assert_eq!(ctx["requested_model"], "gpt-5.1-codex"); view["system"] = json!("You are Codex. Today is Friday."); Invocation::ok(RequestOutcome::Changed(view)) }) @@ -256,3 +288,118 @@ async fn a_failing_reply_plugin_fails_that_answer_and_the_connection_stays() { assert!(!frames.iter().any(|f| f.contains("hel")), "{frames:?}"); } } + +/// 范围按这条连接连的那一家算:只管别家的插件不跑,只管别家的坏插件也不拦;管这一家的 +/// 照常跑,运行记在第 0 跳上,回答钩子的 `ctx` 和请求钩子的一样 +#[tokio::test] +async fn scope_follows_the_upstream_of_the_connection() { + let (up, seen) = upstream().await; + let reply_ctx = Arc::new(Mutex::new(Value::Null)); + let rc = reply_ctx.clone(); + let here = Double::new("here") + .permit(&[Permission::System, Permission::ReplyText]) + .on_request(|mut view, _| { + view["system"] = json!("for up"); + Invocation::ok(RequestOutcome::Changed(view)) + }) + .on_reply(true, false, false, move |ctx| { + *rc.lock().unwrap() = ctx; + Ok(Box::new(Closures { + text: Box::new(|_| Invocation::ok(None)), + end: Box::new(|| Invocation::ok(None)), + tool: Box::new(|_| Invocation::ok(ToolCallOutcome::Unchanged)), + })) + }); + let elsewhere = Double::new("elsewhere") + .permit(&[Permission::System]) + .on_request(|_, _| panic!("ran for an upstream outside its scope")); + let broken_elsewhere = { + let mut a = double::active("old", Double::new("Old")); + a.state = PluginState::Broken(Broken::Changed); + a.scope.upstreams = vec!["relay-*".into()]; + Arc::new(a) + }; + let (gw, runs) = gateway_with( + up, + vec![ + entry("here", here), + entry_with("elsewhere", elsewhere, |a| { + a.scope.upstreams = vec!["relay-*".into()] + }), + broken_elsewhere, + ], + Default::default(), + ) + .await; + let mut c = connect(gw).await; + c.send(create("hi")).await.unwrap(); + let frames = one_answer(&mut c).await; + assert!( + frames.last().unwrap().contains("response.completed"), + "{frames:?}" + ); + assert_eq!(seen.lock().unwrap()[0]["instructions"], "for up"); + let ctx = reply_ctx.lock().unwrap().clone(); + assert_eq!(ctx["upstream"], "up"); + assert_eq!(ctx["model"], "gpt-5.1-codex"); + assert_eq!(ctx["requested_model"], "gpt-5.1-codex"); + tokio::time::sleep(Duration::from_millis(50)).await; + let runs: Vec<(String, String, u64)> = runs + .lock() + .unwrap() + .iter() + .map(|r| { + ( + r.run.plugin_id.clone(), + r.run.hook.slug().to_string(), + r.run.detail.as_ref().unwrap()["attempt"].as_u64().unwrap(), + ) + }) + .collect(); + assert_eq!( + runs, + [ + ("here".to_string(), "request".to_string(), 0), + ("here".to_string(), "reply".to_string(), 0) + ] + ); +} + +/// 插件往 `response.create` 里加的内容照样过请求防护:拦下就切断,上游什么都没收到 +#[tokio::test] +async fn content_a_plugin_adds_to_a_response_create_is_screened() { + let (up, seen) = upstream().await; + let adds = Double::new("adds") + .permit(&[Permission::Messages]) + .on_request(|mut view, _| { + view["messages"][0]["parts"][0]["text"] = json!("the forbidden-plan"); + Invocation::ok(RequestOutcome::Changed(view)) + }); + let security = tw_config::Security { + content: tw_config::ContentPolicy { + mode: tw_config::SecurityMode::Enforce, + custom: vec![tw_config::CustomContentRule { + name: "no plan".into(), + pattern: "forbidden-plan".into(), + matching: Default::default(), + action: tw_config::ContentAction::Block, + disabled: false, + }], + ..Default::default() + }, + ..Default::default() + }; + let (gw, _) = gateway_with(up, vec![entry("adds", adds)], security).await; + let mut c = connect(gw).await; + c.send(create("hi")).await.unwrap(); + let frames = one_answer(&mut c).await; + assert!( + frames.iter().any(|f| f.contains("no plan")), + "the client was not told why: {frames:?}" + ); + assert!( + seen.lock().unwrap().is_empty(), + "{:?}", + seen.lock().unwrap() + ); +} diff --git a/docs/config.md b/docs/config.md index 97774c9..2627c88 100644 --- a/docs/config.md +++ b/docs/config.md @@ -1014,6 +1014,12 @@ Neither keeps the rest of the configuration from taking effect. Plugins run in the order of this list. +A plugin changes a request after routing, each time the request is sent to an +upstream. A request that fails over to another upstream starts again from what +the client sent, and the plugin sees which upstream and which model name the +request goes to. Routing, model checks and session grouping use what the +client sent. + @@ -1034,8 +1040,8 @@ Plugins run in the order of this list. | Field | Type | Default | Description | |---|---|---|---| | `clients` | list of strings | `[]` | Client apps (`claude-code`, `codex`, …), as names or globs. `[]`: every client, including requests whose app is not recognised. | -| `models` | list of strings | `[]` | Models the client asks for, as model ids or globs (`claude-*`). `[]`: every model. | -| `upstreams` | list of strings | `[]` | Upstreams whose answers the plugin handles, by name or glob. It applies to answers only: a request is changed before an upstream is chosen. `[]`: every upstream. | +| `models` | list of strings | `[]` | Models sent to the upstream, as model ids or globs (`claude-*`). When a routing rule renames the model, the new name is the one that matches. `[]`: every model. | +| `upstreams` | list of strings | `[]` | Upstreams the plugin handles, by name or glob, for requests and answers alike. `[]`: every upstream. | ```yaml diff --git a/docs/config.zh-CN.md b/docs/config.zh-CN.md index 9551bf5..516f061 100644 --- a/docs/config.zh-CN.md +++ b/docs/config.zh-CN.md @@ -833,6 +833,8 @@ default_route: default 插件按本列表的顺序运行。 +插件在路由之后改写请求,请求每发往一个上游改写一次。故障转移到另一个上游时,从客户端发来的原样重新开始;插件看得到这一次发往哪个上游、用哪个模型名。路由、模型准入和会话归组看的都是客户端发来的原样。 + @@ -853,8 +855,8 @@ default_route: default | 字段 | 类型 | 默认值 | 说明 | |---|---|---|---| | `clients` | 字符串列表 | `[]` | 客户端应用(`claude-code`、`codex` 等),写名字或通配。`[]`:所有客户端,包括认不出应用的请求。 | -| `models` | 字符串列表 | `[]` | 客户端请求的模型,写模型 ID 或通配(`claude-*`)。`[]`:所有模型。 | -| `upstreams` | 字符串列表 | `[]` | 插件处理哪些上游的回答,写名字或通配。只作用于回答:请求在选定上游之前就已改写。`[]`:所有上游。 | +| `models` | 字符串列表 | `[]` | 发给上游的模型,写模型 ID 或通配(`claude-*`)。路由规则改了模型名的,按改名之后的匹配。`[]`:所有模型。 | +| `upstreams` | 字符串列表 | `[]` | 插件处理哪些上游,写名字或通配,请求和回答都按它。`[]`:所有上游。 | ```yaml