diff --git a/crates/tw-api/msg-codes.txt b/crates/tw-api/msg-codes.txt index ba95626..8196efe 100644 --- a/crates/tw-api/msg-codes.txt +++ b/crates/tw-api/msg-codes.txt @@ -282,7 +282,7 @@ gw.output_limit.withheld gw.plugin.answer_unreadable gw.plugin.api gw.plugin.bad_output -gw.plugin.cannot_read +gw.plugin.cannot_read_body gw.plugin.changed gw.plugin.cpu_limit gw.plugin.engine @@ -291,6 +291,8 @@ gw.plugin.file_changed gw.plugin.manifest gw.plugin.memory_limit gw.plugin.model_not_allowed +gw.plugin.not_applicable +gw.plugin.not_declared 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 cae8193..4259ec3 100644 --- a/crates/tw-api/src/lib.rs +++ b/crates/tw-api/src/lib.rs @@ -691,6 +691,12 @@ pub const MSG_CODES: &str = include_str!("../msg-codes.txt"); /// (`UpdatePluginConfirmed`,请求体同 [`PluginUpdate`])—— 它和装、换源码、批准一样 /// 不给网页调,桌面端在系统的确认框里点了头才发。同一版起 core 自带几个默认插件,第一次 /// 见到时装上、停用着,写配置的这一版来源是 [`ConfigOrigin::Defaults`]。 +/// +/// 33 起**插件说得出自己处理哪几种请求**:[`ManifestView`] 和 [`PluginView`] 多了 +/// `requests`([`RequestKind`]:对话、嵌入、旧版补全)。插件只处理声明了的那几种 —— +/// 不写是只有对话;嵌入和旧版补全要插件自己声明 —— 别的种类的请求不过它、不记录, +/// 它出错、文件变了也拦不着它们。嵌入和旧版补全的视图是一项输入一条消息,`ctx.format` +/// 多了 `openai_embeddings`、`openai_completions`、`gemini_embed`。 pub const CONTROL_API_VERSION: u32 = 33; #[derive(Debug, Clone, Serialize, Deserialize)] @@ -4870,6 +4876,23 @@ slug_enum! { } } +slug_enum! { + /// 一种请求。插件**只处理它声明了的那几种**(插件文件里 manifest 的 `requests`, + /// 不写就是只有 `conversation`):别的种类的请求原样过去,不记录,插件出了什么错也 + /// 和它们无关。图片、音频这些别的接口不属于任何一种,所有插件都不管。 + pub enum RequestKind { + /// 对话:Anthropic Messages、OpenAI Chat Completions、Responses、Gemini 的生成, + /// 连同它们的数 token 和压缩 + Conversation = "conversation", + /// 嵌入:OpenAI 的 `/v1/embeddings`,Gemini 的 `:embedContent`、 + /// `:batchEmbedContents`。插件只改得了每项输入的文字,回答钩子不在它上面跑 + Embeddings = "embeddings", + /// 旧版补全:OpenAI 的 `/v1/completions`。插件只改得了每段提示的文字和几个参数, + /// 回答钩子不在它上面跑 + Completions = "completions", + } +} + slug_enum! { /// 插件出错(运行出错、文件变了、加载不了)时这个请求怎么办。 pub enum OnError { @@ -5039,6 +5062,9 @@ pub struct ManifestView { pub name: String, pub description: Option, pub permissions: Vec, + /// 插件处理哪几种请求,按 [`RequestKind::ALL`] 的顺序。至少有一种;manifest 没写 + /// `requests` 时是 `["conversation"]` + pub requests: Vec, /// 插件建议的范围。装上时照它填 pub scope: PluginScope, pub reply_mode: ReplyMode, @@ -5075,6 +5101,9 @@ pub struct PluginView { pub on_error: OnError, /// 读不出 manifest 时是空的 pub permissions: Vec, + /// 插件处理哪几种请求,按 [`RequestKind::ALL`] 的顺序(见 [`ManifestView::requests`])。 + /// 读不出 manifest 时按出厂的算:`["conversation"]` —— 跑不了的插件拦的也就是这几种 + pub requests: Vec, /// 生效的范围(配置里的) pub scope: PluginScope, pub reply_mode: ReplyMode, @@ -5360,6 +5389,7 @@ mod tests { check(Guard::ALL, Guard::slug, Guard::from_slug); check(RuleAction::ALL, RuleAction::slug, RuleAction::from_slug); check(Permission::ALL, Permission::slug, Permission::from_slug); + check(RequestKind::ALL, RequestKind::slug, RequestKind::from_slug); check(OnError::ALL, OnError::slug, OnError::from_slug); check(ReplyMode::ALL, ReplyMode::slug, ReplyMode::from_slug); check(SettingKind::ALL, SettingKind::slug, SettingKind::from_slug); diff --git a/crates/tw-control/src/plugins.rs b/crates/tw-control/src/plugins.rs index 6bab3e0..b01c356 100644 --- a/crates/tw-control/src/plugins.rs +++ b/crates/tw-control/src/plugins.rs @@ -126,6 +126,7 @@ fn view(a: &Active, entry: &tw_config::Plugin) -> tw_api::PluginView { enabled: a.enabled, on_error: a.on_error, permissions: a.permissions.clone(), + requests: a.requests.clone(), scope: scope_view(&entry.scope), reply_mode: a.reply_mode, settings_schema: m.map(schema).unwrap_or_default(), @@ -171,6 +172,7 @@ fn manifest_view(m: &Manifest) -> tw_api::ManifestView { name: m.name.clone(), description: m.description.clone(), permissions: m.permissions.clone(), + requests: m.requests.clone(), scope: tw_api::PluginScope { clients: m.scope.clients.clone(), models: m.scope.models.clone(), @@ -1050,6 +1052,7 @@ async fn awaken(s: &ControlState, a: &Active) -> Result { on_error: a.on_error, scope: a.scope.clone(), permissions: m.permissions.clone(), + requests: m.requests.clone(), reply_mode: m.reply_mode, hooks: m.hooks, settings, diff --git a/crates/tw-control/src/plugins/defaults.rs b/crates/tw-control/src/plugins/defaults.rs index b68cb00..105a6ff 100644 --- a/crates/tw-control/src/plugins/defaults.rs +++ b/crates/tw-control/src/plugins/defaults.rs @@ -11,7 +11,8 @@ //! 哈希是发出去的那份字节的 —— 再记下来; //! - **给过、配置里还在、文件和批准的都还是给出去的那一份,而 core 带的已经是新版**: //! 换文件、底稿和配置里的哈希;开关、出错时怎么办、范围和还声明着的设置照旧,新声明的 -//! 设置取默认值;**新版要了旧版没要的权限就停用**;记下新版; +//! 设置取默认值;**新版要了旧版没要的权限、或者多处理了一种请求(`requests`),就停用**; +//! 记下新版; //! - **给过、配置里没有了**:用户删的。**不再加回去**; //! - **给过、文件被用户改过**(或者批准的已经是别的一份):不动。 //! @@ -84,7 +85,8 @@ pub struct Seeded { pub added: Vec, /// 换成了新版的 pub updated: Vec, - /// 换成新版时停用了的:开着,而新版要了旧版没要的权限。也在 `updated` 里 + /// 换成新版时停用了的:开着,而新版要了旧版没要的权限、或者多处理了一种请求。也在 + /// `updated` 里 pub disabled: Vec, /// 只记了一笔「给过了」的:用户自己的插件占着这个 id,或者新版已经装上了 pub marked: Vec, @@ -259,7 +261,8 @@ impl Seeder { continue; } }; - // 旧版要过哪些权限。读不出来就当新版多要了 —— 宁可停用 + // 旧版要过哪些权限、处理哪几种请求。读不出来就当新版多要了 —— 宁可停用。 + // 多处理一种请求和多要一个权限一样:插件看得到、改得了的东西变多了 let more = if p.enabled { let old = match bytes { Some(b) => compile(mgr, &b, false).await.ok(), @@ -267,6 +270,7 @@ impl Seeder { }; old.as_ref().is_none_or(|old| { new.permissions.iter().any(|x| !old.permissions.contains(x)) + || new.requests.iter().any(|k| !old.requests.contains(k)) }) } else { false diff --git a/crates/tw-control/tests/plugin_defaults.rs b/crates/tw-control/tests/plugin_defaults.rs index 815e272..681fc2a 100644 --- a/crates/tw-control/tests/plugin_defaults.rs +++ b/crates/tw-control/tests/plugin_defaults.rs @@ -501,6 +501,36 @@ async fn a_new_version_that_wants_more_permissions_comes_back_turned_off() { assert_eq!(v["permissions"], json!(["system", "messages"])); } +/// 新版多处理了一种请求(`requests` 多了嵌入),权限一样:和多要一个权限一样,换上但 +/// 停用 —— 插件看得到、改得了的东西变多了,要用户自己再打开 +#[tokio::test] +async fn a_new_version_that_handles_more_kinds_of_request_comes_back_turned_off() { + let b = bed(); + let scrub = |requests: Value| { + source( + json!({"name": "Scrub", "api": 1, "permissions": ["messages"], + "requests": requests}), + &["onRequest"], + ) + }; + let v1 = scrub(json!(["conversation"])); + seeder(&[("scrub", &v1)]).seed(&b.mgr).await; + b.customize("scrub", json!({})).await; + assert_eq!(b.plugin("scrub").await["requests"], json!(["conversation"])); + + let v2 = scrub(json!(["conversation", "embeddings"])); + let done = seeder(&[("scrub", &v2)]).seed(&b.mgr).await; + assert_eq!(ids(&done.updated), ["scrub"], "{done:?}"); + assert_eq!(ids(&done.disabled), ["scrub"], "{done:?}"); + let p = b.entry("scrub").unwrap(); + assert_eq!(p.sha256, sha(&v2)); + assert!(!p.enabled); + let v = b.plugin("scrub").await; + assert_eq!(v["status"], json!({"kind": "disabled"})); + assert_eq!(v["permissions"], json!(["messages"])); + assert_eq!(v["requests"], json!(["conversation", "embeddings"])); +} + /// 用户自己的插件正好用了一个默认插件的 id:只记一笔「给过了」,它的文件和配置都不动, /// 之后出了新版也不动 #[tokio::test] diff --git a/crates/tw-control/tests/plugins.rs b/crates/tw-control/tests/plugins.rs index 0d3d2f5..4e8efe1 100644 --- a/crates/tw-control/tests/plugins.rs +++ b/crates/tw-control/tests/plugins.rs @@ -1478,3 +1478,111 @@ export function onRequest(req, ctx) { assert_eq!(v["status"], json!({"kind": "ok"})); assert_eq!(v["settings_schema"][0]["default"], "tomorrow"); } + +/// 试跑记下的嵌入请求,从控制面一路到真的沙箱:声明了嵌入的插件跑在一项输入一条消息的 +/// 视图上,前后两份打着码,回答钩子不试;`inspect` 和列表说得出它处理哪几种请求。只处理 +/// 对话的插件说清它当时没跑 +#[tokio::test] +async fn a_trial_on_a_recorded_embeddings_request_runs_only_plugins_that_declare_embeddings() { + let b = bed_with("real", false); + let scrub = r#"export const manifest = { + name: "Scrub inputs", + api: 1, + permissions: ["messages"], + requests: ["conversation", "embeddings"], +}; +export function onRequest(req, ctx) { + console.log(ctx.format); + for (const m of req.messages) { + for (const p of m.parts) { + if (p.type === "text") p.text = p.text.replaceAll("PROJECT-X", "[removed]"); + } + } + return req; +} +"#; + let (_, v) = call( + &b.app, + "POST", + "/plugins/inspect", + Some(json!({"source": scrub})), + ) + .await; + assert_eq!( + v["manifest"]["requests"], + json!(["conversation", "embeddings"]), + "{v}" + ); + let id = b.install(scrub, json!({})).await; + assert_eq!( + b.plugin(&id).await["requests"], + json!(["conversation", "embeddings"]) + ); + let key = "sk-ant-api03-TRIALKEYAAAAAAAAAAAAAAAAAAAA"; + let request = json!({ + "model": "text-embedding-3-small", + "input": ["PROJECT-X roadmap", format!("key {key}"), [101, 102]] + }) + .to_string(); + let answer = + json!({"object": "list", "data": [], "model": "text-embedding-3-small"}).to_string(); + let mut embeddings = row(9, 1_000); + embeddings.path = "/v1/embeddings".into(); + embeddings.provider = "openai".into(); + embeddings.model = "text-embedding-3-small".into(); + embeddings.sent_model = "text-embedding-3-small".into(); + { + let g = b.store.lock().await; + g.db().insert(&embeddings).unwrap(); + g.record_body( + 1_000, + 9, + tw_store::Which::Request, + request.as_bytes(), + request.len(), + ); + g.record_body( + 1_000, + 9, + tw_store::Which::Response, + answer.as_bytes(), + answer.len(), + ); + } + let (st, v) = call( + &b.app, + "POST", + &format!("/plugins/{id}/trial"), + Some(json!({"request_id": 9})), + ) + .await; + assert_eq!(st, StatusCode::OK, "{v}"); + assert!(v["error"].is_null(), "{v}"); + assert!(v["reply"].is_null(), "{v}"); + assert_eq!(v["request"]["outcome"], "changed", "{v}"); + let after: Value = serde_json::from_str(v["request"]["after"].as_str().unwrap()).unwrap(); + assert_eq!(after["input"][0], "[removed] roadmap"); + assert_eq!(after["input"][2], json!([101, 102])); + assert!( + !v.to_string().contains("TRIALKEY"), + "a secret was shown: {v}" + ); + assert_eq!(v["logs"][0]["text"], "openai_embeddings", "{v}"); + + // 只处理对话的插件:当时它就不在这个请求的范围里 + let chat_only = r#"export const manifest = { name: "Chat only", api: 1, permissions: ["messages"] }; +export function onRequest(req) { return req; } +"#; + let other = b.install(chat_only, json!({})).await; + assert_eq!(b.plugin(&other).await["requests"], json!(["conversation"])); + let (st, v) = call( + &b.app, + "POST", + &format!("/plugins/{other}/trial"), + Some(json!({"request_id": 9})), + ) + .await; + assert_eq!(st, StatusCode::OK, "{v}"); + assert!(v["request"].is_null(), "{v}"); + assert_eq!(v["error"]["code"], "gw.plugin.not_declared", "{v}"); +} diff --git a/crates/tw-gateway/src/client_api.rs b/crates/tw-gateway/src/client_api.rs index 7c76fbd..37ce149 100644 --- a/crates/tw-gateway/src/client_api.rs +++ b/crates/tw-gateway/src/client_api.rs @@ -107,6 +107,26 @@ impl ClientApi { || p == "/backend-api/codex/responses/compact" } + /// 这个路径是不是嵌入:OpenAI 的 `/v1/embeddings`,Gemini 的 `:embedContent`、 + /// `:batchEmbedContents`。 + /// + /// 插件声明了 `embeddings` 才处理它们(见 [`crate::plugin::request::Shape`]) + pub fn embeds(path: &str) -> bool { + let p = path.trim_end_matches('/'); + p.strip_prefix("/v1").unwrap_or(p) == "/embeddings" + || (p.contains("/models/") + && (p.ends_with(":embedContent") || p.ends_with(":batchEmbedContents"))) + } + + /// 这个路径是不是 OpenAI 的旧版补全(`/v1/completions`)。Anthropic 的旧版补全 + /// (`/v1/complete`)不算:插件不管它。 + /// + /// 插件声明了 `completions` 才处理它(见 [`crate::plugin::request::Shape`]) + pub fn completes(path: &str) -> bool { + let p = path.trim_end_matches('/'); + p.strip_prefix("/v1").unwrap_or(p) == "/completions" + } + /// 转换库里对应的格式 pub fn dialect(&self) -> Dialect { match self { @@ -288,6 +308,35 @@ mod tests { } } + /// 嵌入、旧版补全各是哪几个路径:生成回答、数 token、别家的旧版补全都不算 + #[test] + fn embeddings_and_legacy_completions_are_told_apart_by_path() { + for (path, embeds, completes) in [ + ("/v1/embeddings", true, false), + ("/embeddings/", true, false), + ( + "/v1beta/models/gemini-embedding-001:embedContent", + true, + false, + ), + ( + "/v1beta/models/text-embedding-004:batchEmbedContents", + true, + false, + ), + ("/v1/completions", false, true), + ("/completions", false, true), + ("/v1/chat/completions", false, false), + ("/v1/complete", false, false), + ("/v1/messages", false, false), + ("/v1beta/models/gemini-2.5-pro:countTokens", false, false), + ("/v1/images/generations", false, false), + ] { + assert_eq!(ClientApi::embeds(path), embeds, "{path}"); + assert_eq!(ClientApi::completes(path), completes, "{path}"); + } + } + #[test] fn a_path_we_do_not_know_is_not_guessed() { for path in ["/v1/models", "/v1/files", "/healthz", "/v1/messagesx", "/"] { diff --git a/crates/tw-gateway/src/plugin/defaults/manifests.json b/crates/tw-gateway/src/plugin/defaults/manifests.json index 26c979b..23b674e 100644 --- a/crates/tw-gateway/src/plugin/defaults/manifests.json +++ b/crates/tw-gateway/src/plugin/defaults/manifests.json @@ -17,6 +17,9 @@ "reply_tool_calls" ], "reply_mode": "stream", + "requests": [ + "conversation" + ], "scope": { "clients": [], "models": [ @@ -43,6 +46,9 @@ "system" ], "reply_mode": "block", + "requests": [ + "conversation" + ], "scope": { "clients": [], "models": [], @@ -75,6 +81,9 @@ "reply_tool_calls" ], "reply_mode": "block", + "requests": [ + "conversation" + ], "scope": { "clients": [], "models": [], diff --git a/crates/tw-gateway/src/plugin/engine.rs b/crates/tw-gateway/src/plugin/engine.rs index 6d4dc61..1ce9565 100644 --- a/crates/tw-gateway/src/plugin/engine.rs +++ b/crates/tw-gateway/src/plugin/engine.rs @@ -26,6 +26,10 @@ pub struct Manifest { pub description: Option, /// 按 [`tw_api::Permission::ALL`] 的顺序,不重复 pub permissions: Vec, + /// 插件处理哪几种请求(manifest 的 `requests`,没写是只有对话)。按 + /// [`tw_api::RequestKind::ALL`] 的顺序,不重复、不空。**别的种类的请求不过它**(见 + /// [`crate::plugin::set::PluginSet::for_request`]) + pub requests: Vec, /// 插件建议的范围。装上时照它填进配置,之后以配置为准 pub scope: Scope, pub reply_mode: tw_api::ReplyMode, @@ -65,6 +69,9 @@ impl Hooks { } } +/// 不写 `requests` 的插件处理的那几种:只有对话。读不出 manifest 的插件也按它算 +pub const DEFAULT_REQUESTS: &[tw_api::RequestKind] = &[tw_api::RequestKind::Conversation]; + /// 编不成的原因。 #[derive(Debug, Clone, PartialEq, thiserror::Error)] pub enum LoadError { diff --git a/crates/tw-gateway/src/plugin/fake.rs b/crates/tw-gateway/src/plugin/fake.rs index 57a9b83..6cf3031 100644 --- a/crates/tw-gateway/src/plugin/fake.rs +++ b/crates/tw-gateway/src/plugin/fake.rs @@ -179,6 +179,42 @@ fn manifest_of(text: &str) -> Result { )); } + // 处理哪几种请求:没写是只有对话。每一种都要有权限碰得到它,每个权限都要用得上 + let requests = match &m["requests"] { + serde_json::Value::Null => crate::plugin::engine::DEFAULT_REQUESTS.to_vec(), + serde_json::Value::Array(a) if !a.is_empty() => { + let mut out = Vec::new(); + for r in a { + let kind = r + .as_str() + .and_then(tw_api::RequestKind::from_slug) + .ok_or_else(|| bad(format!("`{r}` in requests is not a kind of request")))?; + if out.contains(&kind) { + return Err(bad(format!("`{r}` is listed twice in requests"))); + } + out.push(kind); + } + out.sort_by_key(|k| tw_api::RequestKind::ALL.iter().position(|x| x == k)); + out + } + _ => return Err(bad("requests has to be a non-empty list")), + }; + let inputs = [tw_api::Permission::Messages, tw_api::Permission::Params]; + let conversation = requests.contains(&tw_api::RequestKind::Conversation); + if requests + .iter() + .any(|k| *k != tw_api::RequestKind::Conversation && !inputs.iter().any(|p| has(*p))) + { + return Err(bad( + "embeddings and completions requests need the permission messages or params", + )); + } + if !conversation && permissions.iter().any(|p| !inputs.contains(p)) { + return Err(bad( + "a permission other than messages and params needs \"conversation\" in requests", + )); + } + let list = |v: &serde_json::Value| -> Result, LoadError> { match v { serde_json::Value::Null => Ok(Vec::new()), @@ -238,6 +274,7 @@ fn manifest_of(text: &str) -> Result { api, description, permissions, + requests, scope, reply_mode, settings, @@ -289,6 +326,35 @@ mod tests { )); } + /// `requests` 照沙箱的规矩读:没写是只有对话;每一种都要有权限碰得到它 + #[test] + fn requests_default_to_conversations_and_follow_the_sandbox_rules() { + use tw_api::RequestKind::*; + let load = |m: serde_json::Value| FakeEngine.load(source(m, &["onRequest"]).as_bytes()); + let m = load(json!({"name": "n", "api": 1, "permissions": ["messages"]})).unwrap(); + assert_eq!(m.manifest().requests, [Conversation]); + let m = load(json!({"name": "n", "api": 1, "permissions": ["messages"], + "requests": ["embeddings", "conversation"]})) + .unwrap(); + assert_eq!(m.manifest().requests, [Conversation, Embeddings]); + for requests in [ + json!([]), + json!(["images"]), + json!(["completions", "completions"]), + ] { + let r = load(json!({"name": "n", "api": 1, "permissions": ["messages"], + "requests": requests})); + assert!(matches!(r, Err(LoadError::Manifest(_))), "{requests}"); + } + let only_system = load(json!({"name": "n", "api": 1, "permissions": ["system"], + "requests": ["conversation", "embeddings"]})); + assert!(matches!(only_system, Err(LoadError::Manifest(_)))); + let no_conversation = load(json!({"name": "n", "api": 1, + "permissions": ["system", "messages"], + "requests": ["embeddings"]})); + assert!(matches!(no_conversation, Err(LoadError::Manifest(_)))); + } + #[test] fn a_syntax_marker_is_a_syntax_error_on_its_line() { let src = format!( diff --git a/crates/tw-gateway/src/plugin/host.rs b/crates/tw-gateway/src/plugin/host.rs index 2386b15..0b12cc3 100644 --- a/crates/tw-gateway/src/plugin/host.rs +++ b/crates/tw-gateway/src/plugin/host.rs @@ -178,6 +178,7 @@ pub mod double { api: 1, description: None, permissions: Vec::new(), + requests: crate::plugin::engine::DEFAULT_REQUESTS.to_vec(), scope: Scope::default(), reply_mode: tw_api::ReplyMode::Block, settings: Vec::new(), @@ -197,6 +198,16 @@ pub mod double { self } + /// 处理哪几种请求(manifest 的 `requests`)。不调就是只有对话 + pub fn requests(mut self, kinds: &[tw_api::RequestKind]) -> Self { + self.manifest.requests = tw_api::RequestKind::ALL + .iter() + .copied() + .filter(|k| kinds.contains(k)) + .collect(); + self + } + pub fn mode(mut self, mode: tw_api::ReplyMode) -> Self { self.manifest.reply_mode = mode; self @@ -320,6 +331,7 @@ pub mod double { on_error: tw_api::OnError::Reject, scope: Scope::default(), permissions: m.permissions.clone(), + requests: m.requests.clone(), reply_mode: m.reply_mode, hooks: m.hooks, settings: Default::default(), diff --git a/crates/tw-gateway/src/plugin/load.rs b/crates/tw-gateway/src/plugin/load.rs index f81bf2c..2e759be 100644 --- a/crates/tw-gateway/src/plugin/load.rs +++ b/crates/tw-gateway/src/plugin/load.rs @@ -204,6 +204,11 @@ impl Plugins { on_error: p.on_error.into(), scope: scope_of(&p.scope), permissions: m.map(|m| m.permissions.clone()).unwrap_or_default(), + // 读不出 manifest 的按不写 `requests` 的算(见 `Active::requests`) + requests: m.map_or_else( + || crate::plugin::engine::DEFAULT_REQUESTS.to_vec(), + |m| m.requests.clone(), + ), reply_mode: m.map_or(tw_api::ReplyMode::Block, |m| m.reply_mode), hooks: m.map(|m| m.hooks).unwrap_or_default(), settings, @@ -402,6 +407,7 @@ fn placeholder(id: &str) -> Manifest { api: 1, description: None, permissions: Vec::new(), + requests: crate::plugin::engine::DEFAULT_REQUESTS.to_vec(), scope: Scope::default(), reply_mode: tw_api::ReplyMode::Block, settings: Vec::new(), diff --git a/crates/tw-gateway/src/plugin/manifests.rs b/crates/tw-gateway/src/plugin/manifests.rs index 38b607c..1150497 100644 --- a/crates/tw-gateway/src/plugin/manifests.rs +++ b/crates/tw-gateway/src/plugin/manifests.rs @@ -23,8 +23,8 @@ use crate::plugin::set::Scope; /// 文件名,在插件目录里。点开头:它不是插件 pub const FILE: &str = ".manifests.json"; -/// 这份格式自己的版本。**manifest 的读法或者这里的写法改了就加一** -const FORMAT: u32 = 1; +/// 这份格式自己的版本。**manifest 的读法或者这里的写法改了就加一**(2:多了 `requests`) +const FORMAT: u32 = 2; /// 缓存认的版本:core 的版本、沙箱的哈希、这份格式的版本,三样有一样不同就不认 pub fn version() -> String { @@ -58,6 +58,7 @@ pub(crate) struct Entry { #[serde(default, skip_serializing_if = "Option::is_none")] description: Option, permissions: Vec, + requests: Vec, scope: ScopeEntry, reply_mode: tw_api::ReplyMode, settings: Vec, @@ -97,6 +98,7 @@ impl From<&Manifest> for Entry { api: m.api, description: m.description.clone(), permissions: m.permissions.clone(), + requests: m.requests.clone(), scope: ScopeEntry { clients: m.scope.clients.clone(), models: m.scope.models.clone(), @@ -130,6 +132,7 @@ impl From<&Entry> for Manifest { api: e.api, description: e.description.clone(), permissions: e.permissions.clone(), + requests: e.requests.clone(), scope: Scope { clients: e.scope.clients.clone(), models: e.scope.models.clone(), @@ -286,6 +289,7 @@ mod tests { api: 1, description: Some("在系统提示里写上今天的日期".into()), permissions: vec![tw_api::Permission::System], + requests: vec![tw_api::RequestKind::Conversation], scope: Scope { clients: vec![], models: vec!["deepseek*".into()], diff --git a/crates/tw-gateway/src/plugin/mod.rs b/crates/tw-gateway/src/plugin/mod.rs index 2007ab2..429a911 100644 --- a/crates/tw-gateway/src/plugin/mod.rs +++ b/crates/tw-gateway/src/plugin/mod.rs @@ -22,8 +22,9 @@ //! - [`request`]:请求钩子。排在路由之后,**每发往一个上游跑一次**(契约附录二的 I7、 //! I8):按这一次的客户端、发出去的模型和上游挑插件,从客户端的原话起改;换上游从 //! 原话重来,同一家重发不重跑。改过的请求再过一遍内容审查,然后才转换格式、脱敏。 -//! **发往上游的每个请求体都过**:数 token、Responses 的压缩也改,插件看不懂的接口按 -//! 插件的 `on_error` 处置(见 [`request::Shape`])。 +//! **插件只处理它声明了的那几种请求**(manifest 的 `requests`):对话(连同数 token、 +//! Responses 的压缩)是不写也有的,嵌入和旧版补全要插件自己声明;别的接口所有插件都 +//! 不管。没声明的那种请求不过它、不记录,它出了错也拦不着(见 [`request::Shape`])。 //! - [`reply`]:回答钩子。排在格式转换之后、工具调用审查和输出长度之前(I7)—— //! 这两道防护看的就是插件改过的那一版。 //! - [`pool`]:插件调用都是阻塞的、吃 CPU 的,放在专用线程池上跑,不占 tokio 的线程。 @@ -50,7 +51,9 @@ pub mod set; pub mod trial; pub mod view; -pub use engine::{Engine, Hooks, LoadError, MAX_SOURCE, Manifest, SettingSpec, Unavailable}; +pub use engine::{ + DEFAULT_REQUESTS, Engine, Hooks, LoadError, MAX_SOURCE, Manifest, SettingSpec, Unavailable, +}; pub use host::{Invocation, PluginHost, ReplyHost, RequestOutcome, RunError, ToolCallOutcome}; pub use load::{Plugins, RUN_CHANNEL_CAP, RunRecord, RunSender}; pub use set::{Active, Broken, LogLine, LogRing, PluginRun, PluginSet, Scope, State, Stats}; @@ -148,6 +151,7 @@ mod tests { on_error: tw_api::OnError::Reject, scope: Scope::default(), permissions: vec![tw_api::Permission::System], + requests: vec![tw_api::RequestKind::Conversation], reply_mode: tw_api::ReplyMode::Block, hooks: Hooks { request: true, diff --git a/crates/tw-gateway/src/plugin/reply/mod.rs b/crates/tw-gateway/src/plugin/reply/mod.rs index a332ba7..862580d 100644 --- a/crates/tw-gateway/src/plugin/reply/mod.rs +++ b/crates/tw-gateway/src/plugin/reply/mod.rs @@ -238,7 +238,7 @@ impl Chain { ctx.client, ctx.model, ctx.requested_model, - ctx.dialect, + super::format_name(ctx.dialect), ctx.upstream, &a.settings, ); @@ -307,7 +307,7 @@ impl Chain { ctx.client, ctx.model, ctx.requested_model, - ctx.dialect, + super::format_name(ctx.dialect), ctx.upstream, settings, ); diff --git a/crates/tw-gateway/src/plugin/request.rs b/crates/tw-gateway/src/plugin/request.rs index 2463f7a..f1f9c0f 100644 --- a/crates/tw-gateway/src/plugin/request.rs +++ b/crates/tw-gateway/src/plugin/request.rs @@ -18,7 +18,8 @@ //! 1. 读出这一刻的请求(前一个插件改过的话就是改过的)的视图,按权限裁掉没给的部分; //! 2. 认得出的密钥换成占位符([`super::bridge`]); //! 3. 在插件线程池上调 `onRequest`; -//! 4. 核对交回来的东西([`super::view::check`]),占位符换回去,写回原文。 +//! 4. 核对交回来的东西([`super::view::Src::check`],按这种请求的规矩),占位符换回去, +//! 写回原文。 //! //! 插件 `reject` 了,或者出错而它的 `on_error` 是拒绝,**整个请求被拒**,不换下一家: //! 换一家,管它的还是这个插件。文件变了、装不上的插件跑不了,管得着这一次的同样按 @@ -26,15 +27,22 @@ //! //! # 哪些请求过插件 //! -//! **发往上游的每一个请求体都过**,不只生成回答的那些 —— 插件删掉的东西,不能从旁边的 -//! 接口漏出去(见 [`Shape`]): +//! 按客户端调的接口分成几种(见 [`Shape`]),**插件只处理它声明了的那几种**(manifest +//! 的 `requests`,不写是只有对话): //! -//! - 生成回答:上面说的那样; -//! - 数 token(Anthropic 的 `count_tokens`、Gemini 的 `:countTokens`、Responses 的 -//! `input_tokens`)、Responses 的压缩:请求体就是一段对话,插件照样看、照样改,上游数的、 -//! 压的是改过的那一份。网关自己估数、一个字节都不发的那几种不跑插件; -//! - 嵌入、旧版补全、认不出的接口:插件看不懂它们的请求体。管得着的插件按它的 -//! `on_error`:拒绝就拒掉整个请求,跳过就原样发、记一笔跳过。 +//! - 对话 —— 生成回答:上面说的那样;数 token(Anthropic 的 `count_tokens`、Gemini 的 +//! `:countTokens`、Responses 的 `input_tokens`)、Responses 的压缩:请求体就是一段对话, +//! 插件照样看、照样改,上游数的、压的是改过的那一份 —— 插件删掉的东西不能从旁边的接口 +//! 漏出去。网关自己估数、一个字节都不发的那几种不跑插件; +//! - 嵌入(OpenAI 的 `/v1/embeddings`,Gemini 的 `:embedContent`、`:batchEmbedContents`)、 +//! 旧版补全(OpenAI 的 `/v1/completions`):一项输入一条消息,只改得了文字(见 +//! [`super::view::inputs`])。回答钩子不在它们上面跑; +//! - 别的接口(图片、音频、认不出的):不属于任何一种,**所有插件都不管**。 +//! +//! 没声明这一种的插件**不在这一次的范围里**:请求原样过去,什么都不记,它出错、文件 +//! 变了、装不上也拦不着这种请求 —— 不管它的 `on_error` 是什么。声明了的那几种里,请求体 +//! 读不出来(不是 JSON 之类)时,管得着的插件按它的 `on_error`:拒绝就拒掉整个请求,跳过 +//! 就原样发、记一笔跳过(`gw.plugin.cannot_read_body`)。 use std::borrow::Cow; use std::sync::Arc; @@ -57,15 +65,19 @@ pub type Ran = (Arc, PluginRun, Vec); /// 插件怎么看一个请求体:按客户端调的接口分(见 [`crate::client_api::ClientApi`])。 #[derive(Debug, Clone, Copy, PartialEq, Eq)] pub enum Shape { - /// 生成回答 + /// 生成回答(一段对话) Generate, /// 请求体和生成回答同一种形状、却不生成回答的接口:数 token、Responses 的压缩(见 - /// [`crate::client_api::ClientApi::like_generation`])。插件照样看、照样改,**`params` - /// 里只写回模型名**:这些接口不收输出上限、温度这些参数(带上是一个 400),数出来的 - /// token 也和它们无关 + /// [`crate::client_api::ClientApi::like_generation`])。也算对话:插件照样看、照样改, + /// **`params` 里只写回模型名** —— 这些接口不收输出上限、温度这些参数(带上是一个 + /// 400),数出来的 token 也和它们无关 Alike, - /// 嵌入、旧版补全、认不出的接口:**插件看不懂这种请求体**(见 [`unreadable`]) - Opaque, + /// 嵌入:OpenAI 的 `/v1/embeddings`,Gemini 的 `:embedContent`、`:batchEmbedContents` + Embeddings, + /// 旧版补全:OpenAI 的 `/v1/completions` + Completions, + /// 别的接口(图片、音频、认不出的……):**不属于任何一种,插件一律不管** + Other, } impl Shape { @@ -73,13 +85,28 @@ impl Shape { pub fn of(path: &str) -> Shape { use crate::client_api::ClientApi; if ClientApi::of_path(path).is_none() { - Shape::Opaque + Shape::Other } else if ClientApi::generates(path) { Shape::Generate } else if ClientApi::like_generation(path) { Shape::Alike + } else if ClientApi::embeds(path) { + Shape::Embeddings + } else if ClientApi::completes(path) { + Shape::Completions } else { - Shape::Opaque + Shape::Other + } + } + + /// 插件怎么读这种请求体。`dialect` 是客户端的格式。**插件不管的接口是 None** + pub fn form(self, dialect: Dialect) -> Option { + match self { + Shape::Generate | Shape::Alike => Some(view::Form::Conversation(dialect)), + Shape::Embeddings if dialect == Dialect::Gemini => Some(view::Form::GeminiEmbed), + Shape::Embeddings => Some(view::Form::OpenaiEmbeddings), + Shape::Completions => Some(view::Form::OpenaiCompletions), + Shape::Other => None, } } } @@ -101,6 +128,8 @@ pub struct Hook<'a> { body: &'a Bytes, /// 原文解析出来的 JSON。第一次有插件要跑时才解析 parsed: Option>, + /// 原文读不读得成插件的视图:读不成时是原因。和 `parsed` 一起第一次要用时才看 + readable: Option>, /// 按原文编好号的那本账。第一次要用时才编 base: Option, } @@ -137,6 +166,9 @@ pub struct Changed { pub path: String, /// 插件改了 `params.model` 的话,发给这一家的新模型名 pub renamed: Option, + /// 嵌入、旧版补全:改前、改后请求防护各看哪一份(见 [`view::inputs::screenable`])。 + /// 它们没有中间表示,管线拿这两份比,只看插件加进来的。对话是 None:管线自己解码 + pub screen: Option<(tw_dialect::ir::Request, tw_dialect::ir::Request)>, } /// 插件换了发给这一家的模型名。 @@ -170,6 +202,7 @@ impl<'a> Hook<'a> { client, body, parsed: None, + readable: None, base: None, } } @@ -187,13 +220,19 @@ impl<'a> Hook<'a> { } /// 发往 `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); + let Some(form) = self.shape.form(self.dialect) else { + return Ok(out); + }; + let here = self + .set + .for_request(form.kind(), self.client, to.model, to.upstream); if here.is_empty() { // 回答钩子要这个请求的密钥映射:管这一次的里面有,就现在记账。不生成回答的 // 请求没有回答钩子可跑 @@ -207,19 +246,28 @@ impl<'a> Hook<'a> { } return Ok(out); } - if self.shape == Shape::Opaque { - // 空的请求体里没有插件能改的东西 - if !self.body.iter().all(u8::is_ascii_whitespace) { - out.runs = unreadable(self.set, self.client, self.path, to)?; - } - return Ok(out); - } let mut bridge = self.base(); let (body, dialect, client, shape) = (self.body, self.dialect, self.client, self.shape); let original = self .parsed .get_or_insert_with(|| serde_json::from_slice::(body).ok()) .as_ref(); + // 这种请求插件读得懂,这一个却读不成视图:管这一次的插件一个都跑不了,各按各的 + // `on_error`。回答钩子照样要这本账(生成回答的请求,上游也许认得它) + let path = self.path; + let readable = self + .readable + .get_or_insert_with(|| match original { + None => Err("the request body is not JSON".into()), + Some(v) => view::build(form, wrapped_count(dialect, path, v).unwrap_or(v), path) + .map(|_| ()), + }) + .clone(); + if let Err(why) = readable { + out.runs = unreadable(here, &why, to.attempt)?; + out.bridge = Some(bridge); + return Ok(out); + } let mut raw: Option> = original.map(Cow::Borrowed); let mut path = self.path.to_string(); // 发给这一家的模型名:前一个插件改了 `params.model`,后面的看到的就是新的 @@ -256,9 +304,8 @@ impl<'a> Hook<'a> { let Some(current) = raw.as_deref() else { return Err(Failure::Unreadable("the request body is not JSON".into())); }; - let conversation = wrapped_count(dialect, &path, current).unwrap_or(current); - let mut built = - view::build(dialect, conversation, &path).map_err(Failure::Unreadable)?; + let readable = wrapped_count(dialect, &path, current).unwrap_or(current); + let mut built = view::build(form, readable, &path).map_err(Failure::Unreadable)?; sending(&mut built.view, &model); let mut input = view::trim(&built.view, &a.permissions); bridge.hide_value(&mut input); @@ -266,7 +313,7 @@ impl<'a> Hook<'a> { client, &model, to.requested_model, - dialect, + form.name(), to.upstream, &a.settings, ); @@ -281,13 +328,10 @@ impl<'a> Hook<'a> { 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)?; + let mut edits = built + .src + .check(&input, &returned, &a.permissions) + .map_err(Failure::Edit)?; if shape == Shape::Alike { model_only(&mut edits); } @@ -357,6 +401,13 @@ impl<'a> Hook<'a> { } if changed && let Some(v) = raw { let value = v.into_owned(); + let screen = match (form, original) { + (view::Form::Conversation(_), _) | (_, None) => None, + (f, Some(before)) => Some(( + view::inputs::screenable(f, before, self.path), + view::inputs::screenable(f, &value, &path), + )), + }; match serde_json::to_vec(&value) { Ok(b) => { out.changed = Some(Changed { @@ -364,6 +415,7 @@ impl<'a> Hook<'a> { value, path, renamed: renamed_by.map(|by| Renamed { model, by }), + screen, }) } // 序列化不该失败;真失败了就当没改过,不发半个请求体 @@ -377,36 +429,27 @@ impl<'a> Hook<'a> { } } -/// 插件看不懂、却要发往上游的东西:[`Shape::Opaque`] 的请求体,不是 Responses 的 -/// WebSocket 连接上的帧。管这一次的插件一个都跑不了 —— 跑不了的插件(文件变了、装不上) -/// 照旧,能跑的按它的 `on_error`:拒绝就拒掉整个请求,跳过就原样发、记一笔跳过。`path` -/// 是客户端调的路径,报出来的就是它 -pub fn unreadable( - set: &PluginSet, - client: Option<&str>, - path: &str, - to: &Target<'_>, -) -> Result, Box> { +/// 这种请求插件读得懂、这一个的请求体却读不成视图(不是 JSON 之类):`here` 里管这一次 +/// 的插件一个都跑不了。跑不了的插件(文件变了、装不上)照旧,能跑的按它的 `on_error`: +/// 拒绝就拒掉整个请求,跳过就原样发、记一笔跳过。`why` 是读不成的原因 +fn unreadable(here: Vec>, why: &str, attempt: usize) -> Result, Box> { let mut runs = Vec::new(); - for a in set.for_request(client, to.model, to.upstream) { + for a in here { let (outcome, error, refusal) = match &a.state { - super::set::State::Broken(why) => { - let (outcome, refusal) = broken(a.on_error, &a.name, why); - (outcome, broken_reason(&a.name, why), refusal) + super::set::State::Broken(b) => { + let (outcome, refusal) = broken(a.on_error, &a.name, b); + (outcome, broken_reason(&a.name, b), refusal) } super::set::State::Ready(_) => { - let why = cannot_read(&a.name, path); + let why = cannot_read(&a.name, why); match a.on_error { OnError::Skip => (PluginOutcome::Skipped, why, None), OnError::Reject => (PluginOutcome::Error, why.clone(), Some(why)), } } }; - runs.push(( - a.clone(), - not_run(&a, outcome, error, to.attempt), - Vec::new(), - )); + let run = not_run(&a, outcome, error, attempt); + runs.push((a, run, Vec::new())); if let Some(why) = refusal { return Err(Box::new(Refused { why, runs })); } @@ -414,7 +457,7 @@ pub fn unreadable( Ok(runs) } -/// 没跑的一次:跑不了的插件,看不懂的请求 +/// 没跑的一次:跑不了的插件,读不出来的请求体 fn not_run(a: &Active, outcome: PluginOutcome, error: Msg, attempt: usize) -> PluginRun { PluginRun { plugin_id: a.id.clone(), @@ -427,11 +470,12 @@ fn not_run(a: &Active, outcome: PluginOutcome, error: Msg, attempt: usize) -> Pl } } -/// 插件看不懂这个接口的请求体(见 [`unreadable`]) -pub(super) fn cannot_read(plugin: &str, path: &str) -> Msg { +/// 插件声明了这种请求,这一个的请求体却读不成它的视图(见 [`unreadable`])。`detail` +/// 是读不成的原因 +pub(super) fn cannot_read(plugin: &str, detail: &str) -> Msg { msg!( - "gw.plugin.cannot_read", plugin = plugin, path = path => - "Plugin `{plugin}` cannot read requests to {path}." + "gw.plugin.cannot_read_body", plugin = plugin, detail = detail => + "Plugin `{plugin}` cannot read this request: {detail}" ) } @@ -527,12 +571,13 @@ pub fn asked_model(dialect: Dialect, path: &str, raw: Option<&Value>) -> String } /// 插件看到的 `ctx`。请求钩子和回答钩子是同一个样子:`model` 是发给上游的模型名, -/// `requested_model` 是客户端要的,`upstream` 是这一次发往的那一家 +/// `requested_model` 是客户端要的,`format` 是请求体的写法([`view::Form::name`]), +/// `upstream` 是这一次发往的那一家 pub fn ctx( client: Option<&str>, model: &str, requested_model: &str, - dialect: Dialect, + format: &str, upstream: &str, settings: &serde_json::Map, ) -> Value { @@ -540,7 +585,7 @@ pub fn ctx( "client": client, "model": model, "requested_model": requested_model, - "format": super::format_name(dialect), + "format": format, "upstream": upstream, "settings": settings, }) diff --git a/crates/tw-gateway/src/plugin/sandbox.rs b/crates/tw-gateway/src/plugin/sandbox.rs index 82cd7e2..c6bc668 100644 --- a/crates/tw-gateway/src/plugin/sandbox.rs +++ b/crates/tw-gateway/src/plugin/sandbox.rs @@ -171,8 +171,17 @@ fn permission(p: tw_plugin::Permission) -> tw_api::Permission { } } +fn request_kind(k: tw_plugin::RequestKind) -> tw_api::RequestKind { + match k { + tw_plugin::RequestKind::Conversation => tw_api::RequestKind::Conversation, + tw_plugin::RequestKind::Embeddings => tw_api::RequestKind::Embeddings, + tw_plugin::RequestKind::Completions => tw_api::RequestKind::Completions, + } +} + fn manifest(m: &tw_plugin::Manifest) -> Manifest { let granted: Vec = m.permissions.iter().copied().map(permission).collect(); + let handled: Vec = m.requests.iter().copied().map(request_kind).collect(); Manifest { name: m.name.clone(), api: m.api, @@ -183,6 +192,11 @@ fn manifest(m: &tw_plugin::Manifest) -> Manifest { .copied() .filter(|p| granted.contains(p)) .collect(), + requests: tw_api::RequestKind::ALL + .iter() + .copied() + .filter(|k| handled.contains(k)) + .collect(), scope: Scope { clients: m.scope.clients.clone(), models: m.scope.models.clone(), diff --git a/crates/tw-gateway/src/plugin/sandbox/tests/mod.rs b/crates/tw-gateway/src/plugin/sandbox/tests/mod.rs index aae02a5..2aea3ea 100644 --- a/crates/tw-gateway/src/plugin/sandbox/tests/mod.rs +++ b/crates/tw-gateway/src/plugin/sandbox/tests/mod.rs @@ -108,11 +108,31 @@ fn the_manifest_is_carried_over() { tool_call: true } ); + // 没写 `requests`:只处理对话 + assert_eq!(m.requests, [tw_api::RequestKind::Conversation]); // 哈希的就是交进来的那些字节(不变式 I9) let sha: [u8; 32] = Sha256::digest(BOTH.as_bytes()).into(); assert_eq!(h.sha256(), sha); } +/// 声明了的几种请求换过来,按 `RequestKind::ALL` 排 +#[test] +fn the_kinds_of_request_are_carried_over_in_order() { + let h = load( + r#"export const manifest = { name: "Inputs", api: 1, permissions: ["messages"], + requests: ["completions", "embeddings", "conversation"] }; + export function onRequest(req) {}"#, + ); + assert_eq!( + h.manifest().requests, + [ + tw_api::RequestKind::Conversation, + tw_api::RequestKind::Embeddings, + tw_api::RequestKind::Completions + ] + ); +} + #[test] fn a_request_hook_changes_the_view_and_its_log_comes_along() { let h = load(BOTH); diff --git a/crates/tw-gateway/src/plugin/set.rs b/crates/tw-gateway/src/plugin/set.rs index 026a621..7d31649 100644 --- a/crates/tw-gateway/src/plugin/set.rs +++ b/crates/tw-gateway/src/plugin/set.rs @@ -91,6 +91,10 @@ pub struct Active { pub scope: Scope, /// 读不出 manifest 时是空的 pub permissions: Vec, + /// 处理哪几种请求(manifest 的 `requests`)。**读不出 manifest 时按出厂的算**(只有 + /// 对话,[`crate::plugin::engine::DEFAULT_REQUESTS`]):说不出它声明过什么,就按不写 + /// `requests` 的插件对待 —— 它拦的是对话,嵌入、补全照常过去 + pub requests: Vec, pub reply_mode: ReplyMode, pub hooks: Hooks, /// 交给插件的设置:配置里写的盖在 manifest 的默认值上,键和类型都对过 @@ -117,6 +121,12 @@ impl Active { State::Broken(b) => Some(b), } } + + /// 处不处理这一种请求。**不处理的种类在它的范围之外**:那种请求不过它,它跑不了、 + /// 出了错也拦不着那种请求 + pub fn handles(&self, kind: tw_api::RequestKind) -> bool { + self.requests.contains(&kind) + } } /// 一份插件,按配置里的顺序 —— **也就是运行的顺序**。 @@ -143,30 +153,34 @@ impl PluginSet { self.plugins.is_empty() } - /// 发往一个上游之前要过一遍的插件,按顺序:启用的、管得着这一次的,**连同跑不了 - /// 的** —— 跑不了的由调用方照它的 `on_error` 拒绝请求或者跳过它(管得着就要处置, - /// 不管它有没有请求钩子:它一旦加载不了,回答那一段同样做不了)。能跑的只列有请求 - /// 钩子的。**管不着这一次的不算**:只管别的上游的插件坏了,拦不着发往这一家的请求。 + /// 发往一个上游之前要过一遍的插件,按顺序:启用的、处理 `kind` 这种请求的、管得着 + /// 这一次的,**连同跑不了的** —— 跑不了的由调用方照它的 `on_error` 拒绝请求或者跳过它 + /// (管得着就要处置,不管它有没有请求钩子:它一旦加载不了,回答那一段同样做不了)。 + /// 能跑的只列有请求钩子的。**管不着这一次的不算**:只管别的上游的插件坏了,拦不着发往 + /// 这一家的请求;只处理对话的插件坏了,拦不着嵌入和补全。 pub fn for_request( &self, + kind: tw_api::RequestKind, client: Option<&str>, model: &str, upstream: &str, ) -> Vec> { self.plugins .iter() - .filter(|p| p.enabled && p.scope.covers(client, model, upstream)) + .filter(|p| p.enabled && p.handles(kind) && 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.handles(tw_api::RequestKind::Conversation)) .filter(|p| p.scope.covers(client, model, upstream)) .cloned() .collect() @@ -335,6 +349,7 @@ mod tests { api: 1, description: None, permissions: Vec::new(), + requests: crate::plugin::engine::DEFAULT_REQUESTS.to_vec(), scope: Scope::default(), reply_mode: ReplyMode::Block, settings: Vec::new(), @@ -351,6 +366,7 @@ mod tests { on_error: OnError::Reject, scope, permissions: Vec::new(), + requests: crate::plugin::engine::DEFAULT_REQUESTS.to_vec(), reply_mode: ReplyMode::Block, hooks, settings: Default::default(), @@ -374,6 +390,8 @@ mod tests { tool_call: false, }; + const CONVERSATION: tw_api::RequestKind = tw_api::RequestKind::Conversation; + fn ids(v: &[Arc]) -> Vec<&str> { v.iter().map(|p| p.id.as_str()).collect() } @@ -414,7 +432,12 @@ mod tests { active("a-last", REQUEST, Scope::default(), None), ]); assert_eq!( - ids(&set.for_request(Some("claude-code"), "claude-opus-4-5", "anthropic")), + ids(&set.for_request( + CONVERSATION, + Some("claude-code"), + "claude-opus-4-5", + "anthropic" + )), ["b-first", "changed", "a-last"] ); } @@ -434,10 +457,13 @@ mod tests { active("everywhere", REQUEST, Scope::default(), None), ]); assert_eq!( - ids(&set.for_request(None, "m", "relay-a")), + ids(&set.for_request(CONVERSATION, None, "m", "relay-a")), ["only-a", "broken-a", "everywhere"] ); - assert_eq!(ids(&set.for_request(None, "m", "relay-b")), ["everywhere"]); + assert_eq!( + ids(&set.for_request(CONVERSATION, None, "m", "relay-b")), + ["everywhere"] + ); } /// 模型看的是发出去的那个:规则把 claude 改成 glm 发给中转,管 `glm-*` 的插件管这一次 @@ -449,13 +475,66 @@ mod tests { scope(&[], &["glm-*"], &[]), None, )]); - assert_eq!(ids(&set.for_request(None, "glm-4.6", "relay")), ["glm"]); + assert_eq!( + ids(&set.for_request(CONVERSATION, None, "glm-4.6", "relay")), + ["glm"] + ); assert!( - set.for_request(None, "claude-sonnet-4-5", "relay") + set.for_request(CONVERSATION, None, "claude-sonnet-4-5", "relay") .is_empty() ); } + /// 只列处理这一种请求的插件,**跑不了的也一样**:只处理对话的插件坏了,拦不着嵌入和 + /// 补全;声明了嵌入的坏了,拦的也只是嵌入(和对话,如果也声明了的话) + #[test] + fn the_request_list_has_only_plugins_that_handle_the_kind() { + use tw_api::RequestKind::*; + let with = |a: Arc, kinds: &[tw_api::RequestKind]| { + let mut a = Arc::try_unwrap(a).unwrap(); + a.requests = kinds.to_vec(); + Arc::new(a) + }; + let set = PluginSet::new(vec![ + active("chat", REQUEST, Scope::default(), None), + active( + "chat-broken", + REQUEST, + Scope::default(), + Some(Broken::Changed), + ), + with( + active("embeds", REQUEST, Scope::default(), None), + &[Conversation, Embeddings], + ), + with( + active( + "embeds-broken", + REQUEST, + Scope::default(), + Some(Broken::Changed), + ), + &[Embeddings], + ), + with( + active("completes", REQUEST, Scope::default(), None), + &[Completions], + ), + ]); + assert_eq!( + ids(&set.for_request(Embeddings, None, "m", "u")), + ["embeds", "embeds-broken"] + ); + assert_eq!( + ids(&set.for_request(Completions, None, "m", "u")), + ["completes"] + ); + assert_eq!( + ids(&set.for_request(Conversation, None, "m", "u")), + ["chat", "chat-broken", "embeds"] + ); + } + #[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 3ca36aa..2b35252 100644 --- a/crates/tw-gateway/src/plugin/trial.rs +++ b/crates/tw-gateway/src/plugin/trial.rs @@ -21,7 +21,7 @@ use tw_types::{Msg, msg}; use super::bridge::Bridge; use super::host::{PluginHost, RequestOutcome}; use super::pool::Pool; -use super::request::{Shape, rejected, request_unreadable}; +use super::request::{Shape, cannot_read, rejected, request_unreadable}; use super::set::LogLine; use super::view; @@ -67,6 +67,22 @@ fn pretty(v: &Value) -> String { serde_json::to_string_pretty(v).unwrap_or_default() } +/// 试跑的这个请求调的接口插件一律不管:当时没有插件跑在它上面 +fn not_applicable(path: &str) -> Msg { + msg!( + "gw.plugin.not_applicable", path = path => + "Plugins do not run on requests to {path}." + ) +} + +/// 插件没声明这种请求(manifest 的 `requests`):当时它不在这个请求的范围里 +fn not_declared(plugin: &str) -> Msg { + msg!( + "gw.plugin.not_declared", plugin = plugin => + "Plugin `{plugin}` does not handle this kind of request." + ) +} + /// 存下来的回答读不出来 fn answer_unreadable(detail: impl Into) -> Msg { msg!( @@ -143,27 +159,36 @@ async fn tried( let client = request.as_ref().and_then(|r| r.client); let name = host.manifest().name.clone(); - // ── 请求钩子:和这个请求当时一样看(见 [`super::request::Shape`]) - if let (Some(r), Some(d), true) = (&request, dialect, host.manifest().hooks.request) { - let shape = Shape::of(r.path); - match parsed.as_ref() { - _ if shape == Shape::Opaque => { - t.error = Some(super::request::cannot_read(&name, r.path)) + // ── 请求钩子:和这个请求当时一样看(见 [`super::request::Shape`])。插件不管的接口、 + // 插件没声明的那种请求,当时它就没跑:试也不试,说清为什么 + let shape = request.as_ref().map(|r| Shape::of(r.path)); + if let (Some(r), Some(shape), true) = (&request, shape, host.manifest().hooks.request) { + // 认不出的接口没有格式,插件也不管它 + let form = dialect.and_then(|d| Some((d, shape.form(d)?))); + match (form, parsed.as_ref()) { + (None, _) => t.error = Some(not_applicable(r.path)), + (Some((_, f)), _) if !host.manifest().requests.contains(&f.kind()) => { + t.error = Some(not_declared(&name)) } - None => t.error = Some(request_unreadable("the request body is not JSON")), - Some(raw) => { + (Some(_), None) => t.error = Some(cannot_read(&name, "the request body is not JSON")), + (Some((d, form)), Some(raw)) => { let mut masked = raw.clone(); bridge.hide_value(&mut masked); - let conversation = - super::request::wrapped_count(d, r.path, &masked).unwrap_or(&masked); - match view::build(d, conversation, r.path) { - Err(e) => t.error = Some(request_unreadable(e)), + let readable = super::request::wrapped_count(d, r.path, &masked).unwrap_or(&masked); + match view::build(form, readable, r.path) { + Err(e) => t.error = Some(cannot_read(&name, &e)), 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, &requested, d, &upstream, settings); + let ctx = super::request::ctx( + client, + &model, + &requested, + form.name(), + &upstream, + settings, + ); let (h, given) = (host.clone(), input.clone()); let before = pretty(&masked); let ran = pool.run(move || h.on_request(given, ctx)).await; @@ -187,12 +212,7 @@ async fn tried( Some(rejected(&name, reason)), ), Ok(RequestOutcome::Changed(out)) => { - match view::check( - &input, - &out, - &m.permissions, - built.src.hidden_tools(), - ) { + match built.src.check(&input, &out, &m.permissions) { Err(e) => { (Outcome::Error, before.clone(), Some(e.msg())) } @@ -238,11 +258,11 @@ async fn tried( } } - // ── 回答钩子 + // ── 回答钩子:只在生成回答的对话上跑(嵌入、补全、数 token 都不跑) let Some(reply) = reply else { return t; }; - if !host.manifest().hooks.on_reply() { + if !host.manifest().hooks.on_reply() || shape.is_some_and(|s| s != Shape::Generate) { return t; } // 回答按客户端的格式收成一整份 diff --git a/crates/tw-gateway/src/plugin/trial/tests/mod.rs b/crates/tw-gateway/src/plugin/trial/tests/mod.rs index 1d59c3c..def9b2f 100644 --- a/crates/tw-gateway/src/plugin/trial/tests/mod.rs +++ b/crates/tw-gateway/src/plugin/trial/tests/mod.rs @@ -235,9 +235,9 @@ fn stored<'a>(path: &'a str, body: &'a [u8]) -> StoredRequest<'a> { } /// 试跑数 token、嵌入这些请求,和它们当时一样看:Gemini 包着的数 token 改里面那一份、 -/// 只写回模型名;插件看不懂的请求体说清看不懂 +/// 只写回模型名;插件没声明的那种请求、插件不管的接口,说清当时它就没跑 #[tokio::test] -async fn a_trial_reads_counting_and_unreadable_requests_as_they_were_read() { +async fn a_trial_reads_counting_and_unhandled_requests_as_they_were_read() { let tune = || { Double::new("tune") .permit(&[Permission::System, Permission::Params]) @@ -271,6 +271,7 @@ async fn a_trial_reads_counting_and_unreadable_requests_as_they_were_read() { assert_eq!(inner["systemInstruction"]["parts"][0]["text"], "Be brief."); assert!(inner.get("generationConfig").is_none(), "{after}"); + // 只处理对话的插件当时不在嵌入请求的范围里 let embeddings = br#"{"model":"text-embedding-3-small","input":["hi"]}"#; let t = run( Arc::new(Pool::new(1, 4)), @@ -284,6 +285,100 @@ async fn a_trial_reads_counting_and_unreadable_requests_as_they_were_read() { assert!(t.request.is_none()); assert_eq!( t.error.map(|m| m.code).as_deref(), - Some("gw.plugin.cannot_read") + Some("gw.plugin.not_declared") ); + // 插件一律不管的接口 + let t = run( + Arc::new(Pool::new(1, 4)), + tune(), + &Default::default(), + rules(), + Some(stored("/v1/images/generations", br#"{"prompt":"a cat"}"#)), + None, + ) + .await; + assert!(t.request.is_none()); + let e = t.error.unwrap(); + assert_eq!( + (e.code.as_str(), e.arg("path")), + ("gw.plugin.not_applicable", "/v1/images/generations") + ); +} + +/// 去掉记号的插件,声明了嵌入和补全 +fn scrubbing() -> Arc { + Double::new("scrub") + .permit(&[Permission::Messages]) + .requests(&[ + tw_api::RequestKind::Embeddings, + tw_api::RequestKind::Completions, + ]) + .on_request(|mut view, ctx| { + assert_ne!(ctx["format"], "anthropic"); + for m in view["messages"].as_array_mut().unwrap() { + for p in m["parts"].as_array_mut().unwrap() { + if let Some(t) = p["text"].as_str() { + p["text"] = json!(t.replace("CLASSIFIED", "[removed]")); + } + } + } + Invocation::ok(RequestOutcome::Changed(view)) + }) + .into_host() +} + +/// 试跑存下来的嵌入、补全请求:和当时一样,一项输入一条消息;前后两份打过码;`ctx.format` +/// 说得出是哪一种 +#[tokio::test] +async fn a_trial_runs_on_stored_embeddings_and_completions_requests() { + let cases: [(&str, Value, &str, &str); 3] = [ + ( + "/v1/embeddings", + json!({ "model": "text-embedding-3-small", + "input": ["the CLASSIFIED plan", format!("key {KEY}"), [1, 2]] }), + "/input/0", + "the [removed] plan", + ), + ( + "/v1/completions", + json!({ "model": "gpt-3.5-turbo-instruct", "prompt": format!("Say hi to CLASSIFIED {KEY}"), + "max_tokens": 5 }), + "/prompt", + "Say hi to [removed] <>", + ), + ( + "/v1beta/models/gemini-embedding-001:batchEmbedContents", + json!({ "requests": [{ "model": "models/gemini-embedding-001", + "content": { "parts": [{ "text": format!("the CLASSIFIED plan {KEY}") }] } }] }), + "/requests/0/content/parts/0/text", + "the [removed] plan <>", + ), + ]; + for (path, body, at, want) in cases { + let bytes = body.to_string().into_bytes(); + let t = run( + Arc::new(Pool::new(1, 4)), + scrubbing(), + &Default::default(), + rules(), + Some(stored(path, &bytes)), + // 回答钩子不在嵌入、补全上跑:存着的回答试跑也不看 + Some(StoredReply { + body: br#"{"object":"list","data":[]}"#, + upstream: Dialect::Chat, + provider: "up", + }), + ) + .await; + assert_eq!(t.error, None, "{path}"); + assert!(t.reply.is_none(), "{path}"); + let req = t.request.unwrap(); + assert_eq!(req.outcome, Outcome::Changed, "{path}"); + let after: Value = serde_json::from_str(&req.after).unwrap(); + assert_eq!(after.pointer(at).unwrap(), want, "{path}"); + assert!( + !req.before.contains(KEY) && !req.after.contains(KEY), + "{path}" + ); + } } diff --git a/crates/tw-gateway/src/plugin/view/gemini.rs b/crates/tw-gateway/src/plugin/view/gemini.rs index aef7408..33b264e 100644 --- a/crates/tw-gateway/src/plugin/view/gemini.rs +++ b/crates/tw-gateway/src/plugin/view/gemini.rs @@ -61,12 +61,12 @@ fn name_in(v: &Value, camel: &str) -> String { field(v, camel).map_or_else(|| camel.to_string(), |(_, k)| k) } -fn fstr<'a>(v: &'a Value, camel: &str) -> Option<&'a str> { +pub(super) fn fstr<'a>(v: &'a Value, camel: &str) -> Option<&'a str> { field(v, camel).and_then(|(x, _)| x.as_str()) } /// `/v1beta/models/gemini-2.5-pro:generateContent` 里的模型 -fn path_model(path: &str) -> Option<&str> { +pub(super) fn path_model(path: &str) -> Option<&str> { let (_, rest) = path.split_once("/models/")?; let (model, _) = rest.rsplit_once(':')?; Some(model) diff --git a/crates/tw-gateway/src/plugin/view/inputs.rs b/crates/tw-gateway/src/plugin/view/inputs.rs new file mode 100644 index 0000000..502a834 --- /dev/null +++ b/crates/tw-gateway/src/plugin/view/inputs.rs @@ -0,0 +1,457 @@ +//! 嵌入和旧版补全的请求视图:**一项输入一条消息**。 +//! +//! - OpenAI 的嵌入(`/v1/embeddings`)的 `input`、旧版补全(`/v1/completions`)的 +//! `prompt`:一个字符串是一条消息;数组里每一项一条 —— 字符串是一段文字,一串 token +//! (数字数组)是一个只读的 `other` 部分(`label` 是 `tokens`)。整个就是一串 token +//! (数组里全是数字)的,是一条消息。 +//! - Gemini 的嵌入:`:embedContent` 的 `content`、`:batchEmbedContents` 里每个请求的 +//! `content`,一个 Content 一条消息,它的每个部分一个部分:文字是文字,别的只读。 +//! +//! 消息都是 `user`。没有 `system`、没有 `tools`。`params`:嵌入只有 `model`;补全是 +//! `model`、`max_tokens`、`temperature`、`top_p`、`stop`。`suffix`、`dimensions`、 +//! `taskType` 这些不给看、不动。 +//! +//! # 改写规则比对话严 +//! +//! **只有文字能改**:消息和部分不能加、不能删、不能挪 —— 上游按输入的先后一项一项地回 +//! (第几个向量、第几段补全),多一项少一项,客户端拿到的回答就对不上号了。核对在 +//! [`check`],写回([`apply`])再守一道。 +//! +//! 写回只碰改过的那几段文字(和改了的参数):别的字段原样留着。 + +use serde_json::{Map, Value, json}; + +use super::*; + +/// 视图里每条消息在原文里是哪一项,和它有几个部分。 +pub struct Src { + form: Form, + items: Vec<(Item, usize)>, + /// 补全的 `stop` 原来是一个字符串 + stop_string: bool, +} + +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +enum Item { + /// OpenAI:`input` / `prompt` 整个(一个字符串、一串 token、认不出的值) + Whole, + /// OpenAI:数组里的第几项 + At(usize), + /// Gemini:`:embedContent` 的 `content` + Content, + /// Gemini:`:batchEmbedContents` 里第几个请求的 `content` + Request(usize), +} + +impl Form { + /// OpenAI 那两种的输入写在哪个字段里 + fn field(self) -> &'static str { + match self { + Form::OpenaiCompletions => "prompt", + _ => "input", + } + } + + /// 视图的 `params` 里有的那几个 + fn params(self) -> &'static [&'static str] { + match self { + Form::OpenaiCompletions => &["model", "max_tokens", "temperature", "top_p", "stop"], + _ => &["model"], + } + } + + /// 报错时怎么称呼这种请求 + fn noun(self) -> &'static str { + match self { + Form::OpenaiCompletions => "a completions request", + _ => "an embeddings request", + } + } +} + +/// 一项输入(OpenAI 那两种)在视图里的那个部分 +fn openai_part(key: &str, v: &Value) -> Value { + match v { + Value::String(s) => part_text(key, s), + other => part_other(key, label(other)), + } +} + +/// 读不成文字的一项叫什么。一串 token 是 `tokens` +fn label(v: &Value) -> &str { + match v { + Value::Array(a) if a.iter().all(Value::is_number) => "tokens", + Value::Array(_) => "array", + Value::Object(o) => o.get("type").and_then(Value::as_str).unwrap_or("object"), + Value::Number(_) => "number", + Value::Bool(_) => "boolean", + Value::Null => "null", + Value::String(_) => "text", + } +} + +/// 一串 token:不空、全是数字的数组 +fn tokens(a: &[Value]) -> bool { + !a.is_empty() && a.iter().all(Value::is_number) +} + +pub fn build(form: Form, raw: &Value, path: &str) -> Result { + let mut view = Map::new(); + view.insert("format".into(), json!(form.name())); + let mut messages = Vec::new(); + let mut items = Vec::new(); + let mut push = |item: Item, parts: Vec| { + let i = messages.len(); + items.push((item, parts.len())); + messages.push(json!({ "key": msg_key(i), "role": "user", "parts": parts })); + }; + let mut params = Map::new(); + let mut stop_string = false; + match form { + Form::OpenaiEmbeddings | Form::OpenaiCompletions => { + let model = raw.get("model").and_then(Value::as_str).unwrap_or_default(); + view.insert("model".into(), json!(model)); + params.insert("model".into(), json!(model)); + let key = |i: usize| part_key(i, 0); + match raw.get(form.field()) { + None | Some(Value::Null) => {} + // 一串 token 是一项输入,不是每个数字一项 + Some(Value::Array(a)) if !tokens(a) => { + for (i, x) in a.iter().enumerate() { + push(Item::At(i), vec![openai_part(&key(i), x)]); + } + } + Some(whole) => push(Item::Whole, vec![openai_part(&key(0), whole)]), + } + if form == Form::OpenaiCompletions { + if let Some(n) = raw.get("max_tokens").and_then(Value::as_u64) { + params.insert("max_tokens".into(), json!(n)); + } + for k in ["temperature", "top_p"] { + if let Some(x) = raw.get(k).filter(|x| x.is_number()) { + params.insert(k.into(), x.clone()); + } + } + match raw.get("stop") { + Some(Value::String(s)) => { + stop_string = true; + params.insert("stop".into(), json!([s])); + } + Some(Value::Array(a)) => { + params.insert( + "stop".into(), + Value::Array(a.iter().filter(|s| s.is_string()).cloned().collect()), + ); + } + _ => {} + } + } + } + Form::GeminiEmbed => { + let model = gemini::path_model(path).ok_or_else(|| { + format!("the path {path} does not say which Gemini model to call") + })?; + view.insert("model".into(), json!(model)); + params.insert("model".into(), json!(model)); + if batch(path) { + for (i, r) in raw + .get("requests") + .and_then(Value::as_array) + .into_iter() + .flatten() + .enumerate() + { + push(Item::Request(i), content_parts(i, r.get("content"))); + } + } else if let Some(c) = raw.get("content") { + push(Item::Content, content_parts(0, Some(c))); + } + } + Form::Conversation(_) => return Err("a conversation is not a list of inputs".into()), + } + view.insert("messages".into(), Value::Array(messages)); + view.insert("params".into(), Value::Object(params)); + Ok(Built { + view: Value::Object(view), + src: super::Src::Inputs(Src { + form, + items, + stop_string, + }), + }) +} + +/// 请求防护看的那一份:每项输入是一条用户消息,里面是它读得出的文字。 +/// +/// 嵌入、补全没有中间表示(转换只为生成回答),插件改过之后再看一遍请求防护时 +/// (只看插件加进来的,见 [`crate::guard::screen_more`])就拿改前、改后各一份比。读不成 +/// 视图的是空的 +pub fn screenable(form: Form, raw: &Value, path: &str) -> tw_dialect::ir::Request { + use tw_dialect::ir::{Message, Part, Request, Role}; + let view = build(form, raw, path) + .map(|b| b.view) + .unwrap_or(Value::Null); + let messages = view["messages"] + .as_array() + .into_iter() + .flatten() + .map(|m| Message { + role: Role::User, + parts: m["parts"] + .as_array() + .into_iter() + .flatten() + .filter(|p| p["type"] == "text") + .filter_map(|p| p["text"].as_str()) + .map(|t| Part::Text(t.to_string())) + .collect(), + }) + .collect(); + Request { + messages, + ..Default::default() + } +} + +/// `:batchEmbedContents`(一个请求里好几项输入) +fn batch(path: &str) -> bool { + path.trim_end_matches('/').ends_with(":batchEmbedContents") +} + +/// Gemini 的一个 Content 的部分:有文字的是文字,别的只读 +fn content_parts(i: usize, content: Option<&Value>) -> Vec { + content + .and_then(|c| c.get("parts")) + .and_then(Value::as_array) + .into_iter() + .flatten() + .enumerate() + .map(|(j, p)| { + let key = part_key(i, j); + match gemini::fstr(p, "text") { + Some(t) => part_text(&key, t), + None => part_other( + &key, + p.as_object() + .and_then(|o| o.keys().next()) + .map_or("unknown", String::as_str), + ), + } + }) + .collect() +} + +/// 核对:先按对话的规矩(key、权限、只读的部分),再加上这一种的:**消息和部分一个都 +/// 不能多、不能少**,参数只能改这种请求有的那几个。 +pub fn check( + src: &Src, + input: &Value, + output: &Value, + perms: &[Permission], +) -> Result { + let edits = super::check(input, output, perms, &[])?; + if let Some(ms) = &edits.messages { + fixed(src, ms)?; + } + if let Some(p) = &edits.params { + params_allowed(src.form, p)?; + } + Ok(edits) +} + +/// 消息、部分都是原来那些,一一对应、先后不变 +fn fixed(src: &Src, ms: &[MsgEdit]) -> Result<(), EditError> { + let noun = src.form.noun(); + let added = || { + bad(format!( + "messages cannot be added to {noun}: each message is one input, and the answer \ + comes back input by input" + )) + }; + for (k, e) in ms.iter().enumerate() { + let (from, parts) = match e { + MsgEdit::Insert { .. } => return Err(added()), + MsgEdit::Keep { from, parts } => (*from, parts), + }; + if from != k { + return Err(removed(noun)); + } + let Some(&(_, count)) = src.items.get(from) else { + return Err(added()); + }; + let Some(parts) = parts else { continue }; + let kept = parts.len() == count + && parts + .iter() + .enumerate() + .all(|(j, p)| matches!(p, PartEdit::Keep { from, .. } if *from == j)); + if !kept { + return Err(bad(format!( + "parts cannot be added to or removed from the messages of {noun}; only the text \ + of a text part can change" + ))); + } + } + if ms.len() != src.items.len() { + return Err(removed(noun)); + } + Ok(()) +} + +fn removed(noun: &str) -> EditError { + bad(format!( + "messages cannot be removed from {noun}: each message is one input, and the answer \ + comes back input by input" + )) +} + +/// 参数只改这种请求有的那几个:嵌入只有 `model` +fn params_allowed(form: Form, p: &ParamsEdit) -> Result<(), EditError> { + let changed = [ + ("max_tokens", p.max_tokens.is_some()), + ("temperature", p.temperature.is_some()), + ("top_p", p.top_p.is_some()), + ("stop", p.stop.is_some()), + ]; + match changed + .iter() + .find(|(k, c)| *c && !form.params().contains(k)) + { + Some((k, _)) => Err(bad(format!( + "{} has no `params.{k}`; only {} can change", + form.noun(), + form.params() + .iter() + .map(|k| format!("`params.{k}`")) + .collect::>() + .join(", ") + ))), + None => Ok(()), + } +} + +/// 写回:改过的文字落回原来那一项,参数照这种请求的写法。返回新的路径(Gemini 换了 +/// 模型时)。 +pub fn apply( + raw: &mut Value, + src: &Src, + edits: &Edits, + path: &str, +) -> Result, EditError> { + let Some(obj) = raw.as_object_mut() else { + return Err(bad("the request body is not a JSON object")); + }; + if edits.system.is_some() || edits.tools.is_some() { + return Err(bad(format!( + "{} has no system prompt and no tools", + src.form.noun() + ))); + } + if let Some(ms) = &edits.messages { + fixed(src, ms)?; + for (k, e) in ms.iter().enumerate() { + let MsgEdit::Keep { + parts: Some(parts), .. + } = e + else { + continue; + }; + for (j, p) in parts.iter().enumerate() { + if let PartEdit::Keep { + change: Some(change), + .. + } = p + { + let Change::Text(t) = change else { + return Err(bad("only the text of an input can change")); + }; + set_text(obj, src, src.items[k].0, j, t)?; + } + } + } + } + let Some(p) = &edits.params else { + return Ok(None); + }; + params_allowed(src.form, p)?; + match src.form { + Form::GeminiEmbed => { + let Some(m) = &p.model else { + return Ok(None); + }; + // 路径上是模型,请求体里的 `model`(`models/…`)也写着它:两处对得上上游才收 + let named = json!(format!("models/{m}")); + if batch(path) { + for r in obj + .get_mut("requests") + .and_then(Value::as_array_mut) + .into_iter() + .flatten() + { + if let Some(slot) = r.get_mut("model").filter(|x| x.is_string()) { + *slot = named.clone(); + } + } + } else if let Some(slot) = obj.get_mut("model").filter(|x| x.is_string()) { + *slot = named; + } + Ok(Some(crate::forward::gemini_path_with_model(path, m))) + } + _ => { + if let Some(m) = &p.model { + obj.insert("model".into(), json!(m)); + } + anthropic::set_opt(obj, "max_tokens", p.max_tokens.map(|o| o.map(Value::from))); + anthropic::set_opt( + obj, + "temperature", + p.temperature.map(|o| o.map(Value::from)), + ); + anthropic::set_opt(obj, "top_p", p.top_p.map(|o| o.map(Value::from))); + anthropic::set_opt( + obj, + "stop", + p.stop.clone().map(|o| { + o.map(|s| match s.as_slice() { + [one] if src.stop_string => json!(one), + _ => json!(s), + }) + }), + ); + Ok(None) + } + } +} + +/// 第 `part` 个部分的文字换成 `t`。**只换得了原来就是文字的那一项** +fn set_text( + obj: &mut Map, + src: &Src, + item: Item, + part: usize, + t: &str, +) -> Result<(), EditError> { + let slot = match item { + Item::Whole => obj.get_mut(src.form.field()), + Item::At(i) => obj.get_mut(src.form.field()).and_then(|v| v.get_mut(i)), + Item::Content => obj + .get_mut("content") + .and_then(|c| c.get_mut("parts")) + .and_then(|p| p.get_mut(part)) + .and_then(|p| p.get_mut("text")), + Item::Request(i) => obj + .get_mut("requests") + .and_then(|r| r.get_mut(i)) + .and_then(|r| r.get_mut("content")) + .and_then(|c| c.get_mut("parts")) + .and_then(|p| p.get_mut(part)) + .and_then(|p| p.get_mut("text")), + }; + match slot { + Some(v @ Value::String(_)) => { + *v = json!(t); + Ok(()) + } + _ => Err(bad("only the text of an input can change")), + } +} diff --git a/crates/tw-gateway/src/plugin/view/mod.rs b/crates/tw-gateway/src/plugin/view/mod.rs index 9182270..6118db8 100644 --- a/crates/tw-gateway/src/plugin/view/mod.rs +++ b/crates/tw-gateway/src/plugin/view/mod.rs @@ -19,6 +19,12 @@ //! 认识、有没有重复、留下来的有没有挪位置、只读的东西改没改。核对只看视图本身, //! 和格式无关;写回时格式自己的限制(Anthropic 的消息里没有 system 角色)由各格式 //! 报。 +//! +//! # 不只对话([`Form`]) +//! +//! 嵌入和旧版补全也有视图([`inputs`]):一项输入一条消息,只有文字能改,消息和部分 +//! 一个都不能多、不能少。数据面一律经 [`Src::check`] 核对 —— 它按这种请求的规矩来; +//! [`check`] 本身是对话的规矩。 use std::collections::{HashMap, HashSet}; @@ -32,6 +38,7 @@ use super::bridge::Bridge; pub mod anthropic; pub mod chat; pub mod gemini; +pub mod inputs; pub mod responses; mod segments; @@ -111,6 +118,47 @@ impl Role { } } +/// 插件读得懂的一种请求体:一段对话(客户端的四种格式各一种),或者一张输入的单子 +/// (嵌入、旧版补全,见 [`inputs`])。 +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub enum Form { + /// 生成回答,和形状一样的数 token、压缩 + Conversation(Dialect), + /// OpenAI 的 `/v1/embeddings` + OpenaiEmbeddings, + /// OpenAI 的 `/v1/completions` + OpenaiCompletions, + /// Gemini 的 `:embedContent`、`:batchEmbedContents` + GeminiEmbed, +} + +impl From for Form { + fn from(d: Dialect) -> Form { + Form::Conversation(d) + } +} + +impl Form { + /// 插件那一侧的写法:视图的 `format`、`ctx.format` + pub fn name(self) -> &'static str { + match self { + Form::Conversation(d) => crate::plugin::format_name(d), + Form::OpenaiEmbeddings => "openai_embeddings", + Form::OpenaiCompletions => "openai_completions", + Form::GeminiEmbed => "gemini_embed", + } + } + + /// 这是哪一种请求(manifest 的 `requests` 里的那个词) + pub fn kind(self) -> tw_api::RequestKind { + match self { + Form::Conversation(_) => tw_api::RequestKind::Conversation, + Form::OpenaiEmbeddings | Form::GeminiEmbed => tw_api::RequestKind::Embeddings, + Form::OpenaiCompletions => tw_api::RequestKind::Completions, + } + } +} + /// 一份读好的请求:完整的视图(还没按权限裁),和写回时要用的位置。 pub struct Built { pub view: Value, @@ -123,6 +171,8 @@ pub enum Src { Chat(chat::Src), Responses(responses::Src), Gemini(gemini::Src), + /// 嵌入、旧版补全 + Inputs(inputs::Src), } impl Src { @@ -133,23 +183,41 @@ impl Src { Src::Chat(s) => &s.hidden_tools, Src::Responses(s) => &s.hidden_tools, Src::Gemini(s) => &s.hidden_tools, + Src::Inputs(_) => &[], + } + } + + /// 核对插件交回来的东西,**按这种请求的规矩**:对话是 [`check`];嵌入、旧版补全在 + /// 那之上再加一层([`inputs::check`]:只有文字能改,消息和部分不增不减)。数据面 + /// 一律走这里 + pub fn check( + &self, + input: &Value, + output: &Value, + perms: &[Permission], + ) -> Result { + match self { + Src::Inputs(s) => inputs::check(s, input, output, perms), + _ => check(input, output, perms, self.hidden_tools()), } } } -/// 把客户端发来的请求读成视图。`path` 是客户端请求的路径(Gemini 的模型写在里面)。 +/// 把客户端发来的请求读成视图。`form` 是这种请求体怎么读(一段对话的话就是客户端的 +/// 格式,[`Dialect`] 直接转得过来),`path` 是客户端请求的路径(Gemini 的模型写在里面)。 /// -/// 不是 JSON 对象、或者不是这四种格式的,读不出来。 -pub fn build(dialect: Dialect, raw: &Value, path: &str) -> Result { +/// 不是 JSON 对象、或者不是这几种请求体的,读不出来。 +pub fn build(form: impl Into
, raw: &Value, path: &str) -> Result { if !raw.is_object() { return Err("the request body is not a JSON object".into()); } - match dialect { - Dialect::Anthropic => Ok(anthropic::build(raw)), - Dialect::Chat => Ok(chat::build(raw)), - Dialect::Responses => Ok(responses::build(raw)), - Dialect::Gemini => gemini::build(raw, path), - Dialect::Bedrock => Err("Bedrock is not a client format".into()), + match form.into() { + Form::Conversation(Dialect::Anthropic) => Ok(anthropic::build(raw)), + Form::Conversation(Dialect::Chat) => Ok(chat::build(raw)), + Form::Conversation(Dialect::Responses) => Ok(responses::build(raw)), + Form::Conversation(Dialect::Gemini) => gemini::build(raw, path), + Form::Conversation(Dialect::Bedrock) => Err("Bedrock is not a client format".into()), + f => inputs::build(f, raw, path), } } @@ -165,6 +233,7 @@ pub fn apply( Src::Chat(s) => chat::apply(raw, s, edits).map(|_| None), Src::Responses(s) => responses::apply(raw, s, edits).map(|_| None), Src::Gemini(s) => gemini::apply(raw, s, edits, path), + Src::Inputs(s) => inputs::apply(raw, s, edits, path), } } @@ -372,6 +441,7 @@ pub fn check( } "system" => { need(perms, Permission::System, "system")?; + given(input, "system")?; let Some(s) = v.as_str() else { return Err(bad("`system` must be a string")); }; @@ -381,14 +451,17 @@ pub fn check( } "messages" => { need(perms, Permission::Messages, "messages")?; + given(input, "messages")?; edits.messages = check_messages(input, v)?; } "tools" => { need(perms, Permission::Tools, "tools")?; + given(input, "tools")?; edits.tools = check_tools(input, v, hidden_tools)?; } "params" => { need(perms, Permission::Params, "params")?; + given(input, "params")?; let p = check_params(input, v)?; if !p.is_empty() { edits.params = Some(p); @@ -411,6 +484,15 @@ fn need(perms: &[Permission], p: Permission, section: &str) -> Result<(), EditEr } } +/// 这一节插件拿到过没有。**交回来的不能凭空多出一节**:嵌入、补全的视图里本来就没有 +/// 系统提示和工具,权限再全也加不进去 +fn given(input: &Value, section: &str) -> Result<(), EditError> { + match input.get(section) { + Some(_) => Ok(()), + None => Err(bad(format!("this request has no `{section}`"))), + } +} + fn only_fields(o: &Map, allowed: &[&str], what: &str) -> Result<(), EditError> { match o.keys().find(|k| !allowed.contains(&k.as_str())) { Some(k) => Err(bad(format!("{what} has an unknown field `{k}`"))), diff --git a/crates/tw-gateway/src/plugin/view/tests/mod.rs b/crates/tw-gateway/src/plugin/view/tests/mod.rs index 0c1e8cd..4d5dc20 100644 --- a/crates/tw-gateway/src/plugin/view/tests/mod.rs +++ b/crates/tw-gateway/src/plugin/view/tests/mod.rs @@ -32,7 +32,7 @@ fn edit( Ok((next, p)) } -fn view_of(d: Dialect, raw: &Value, path: &str) -> Value { +fn view_of(d: impl Into, raw: &Value, path: &str) -> Value { build(d, raw, path).expect("builds").view } @@ -1221,3 +1221,487 @@ fn placeholders_in_edits_are_revealed_before_write_back() { format!("my key is {KEY} (rotated)") ); } + +// ───────────────────────────────────────────────────────── 嵌入、旧版补全 + +/// 和 [`edit`] 一样,但按这种请求自己的规矩核对([`Src::check`],数据面走的就是它)。 +/// 插件有 `messages` 和 `params` 两个权限 +fn edit_inputs( + form: Form, + raw: &Value, + path: &str, + f: impl FnOnce(&mut Value), +) -> Result<(Value, Option), EditError> { + let built = build(form, raw, path).expect("builds"); + let perms = [Permission::Messages, Permission::Params]; + let input = trim(&built.view, &perms); + let mut out = input.clone(); + f(&mut out); + let edits = built.src.check(&input, &out, &perms)?; + let mut next = raw.clone(); + let p = apply(&mut next, &built.src, &edits, path)?; + Ok((next, p)) +} + +fn embeddings() -> Value { + json!({ + "model": "text-embedding-3-small", + "input": ["the SECRET plan", "a second line", [9906, 1917]], + "encoding_format": "float", + "dimensions": 256, + "user": "u-1" + }) +} + +const EMBEDDINGS: &str = "/v1/embeddings"; + +fn completions() -> Value { + json!({ + "model": "gpt-3.5-turbo-instruct", + "prompt": ["Say hi to SECRET", [9906, 1917], "def f():"], + "suffix": "\n# end", + "max_tokens": 16, + "temperature": 0.5, + "stop": "\n\n", + "logprobs": 2, + "echo": false + }) +} + +const COMPLETIONS: &str = "/v1/completions"; + +fn gemini_embed() -> Value { + json!({ + "model": "models/gemini-embedding-001", + "content": { "parts": [{ "text": "the SECRET plan" }, { "text": "more" }] }, + "taskType": "RETRIEVAL_DOCUMENT", + "title": "Plans", + "outputDimensionality": 768 + }) +} + +const GEMINI_EMBED: &str = "/v1beta/models/gemini-embedding-001:embedContent"; + +fn gemini_batch() -> Value { + json!({ + "requests": [ + { "model": "models/gemini-embedding-001", "taskType": "RETRIEVAL_QUERY", + "content": { "parts": [{ "text": "first SECRET" }] } }, + { "model": "models/gemini-embedding-001", + "content": { "parts": [ + { "text": "second" }, + { "inlineData": { "mimeType": "image/png", "data": "iVBORw0KGgo=" } } + ] } } + ] + }) +} + +const GEMINI_BATCH: &str = "/v1beta/models/gemini-embedding-001:batchEmbedContents"; + +/// 四种输入的单子各一份:(写法, 请求体, 路径, 第一段文字在原文里的位置) +fn input_samples() -> Vec<(Form, Value, &'static str, &'static str)> { + vec![ + (Form::OpenaiEmbeddings, embeddings(), EMBEDDINGS, "/input/0"), + ( + Form::OpenaiCompletions, + completions(), + COMPLETIONS, + "/prompt/0", + ), + ( + Form::GeminiEmbed, + gemini_embed(), + GEMINI_EMBED, + "/content/parts/0/text", + ), + ( + Form::GeminiEmbed, + gemini_batch(), + GEMINI_BATCH, + "/requests/0/content/parts/0/text", + ), + ] +} + +fn parts_of(v: &Value) -> Vec> { + v["messages"] + .as_array() + .unwrap() + .iter() + .map(|m| { + assert_eq!(m["role"], "user", "{m}"); + m["parts"] + .as_array() + .unwrap() + .iter() + .map(|p| { + let t = p["type"].as_str().unwrap().to_string(); + let what = p["text"].as_str().or(p["label"].as_str()).unwrap(); + (t, what.to_string()) + }) + .collect() + }) + .collect() +} + +fn pairs(list: &[&[(&str, &str)]]) -> Vec> { + list.iter() + .map(|m| { + m.iter() + .map(|(a, b)| (a.to_string(), b.to_string())) + .collect() + }) + .collect() +} + +/// 一项输入一条 `user` 消息:文字是文字,一串 token 是只读的 `other`;没有系统提示和 +/// 工具;参数只有这种请求有的那几个 +#[test] +fn inputs_read_as_one_user_message_each() { + let v = view_of(Form::OpenaiEmbeddings, &embeddings(), EMBEDDINGS); + assert_eq!(v["format"], "openai_embeddings"); + assert_eq!(v["model"], "text-embedding-3-small"); + assert_eq!( + parts_of(&v), + pairs(&[ + &[("text", "the SECRET plan")], + &[("text", "a second line")], + &[("other", "tokens")] + ]) + ); + assert_eq!(v["params"], json!({ "model": "text-embedding-3-small" })); + assert!(v.get("system").is_none() && v.get("tools").is_none(), "{v}"); + + let v = view_of(Form::OpenaiCompletions, &completions(), COMPLETIONS); + assert_eq!(v["format"], "openai_completions"); + assert_eq!( + parts_of(&v), + pairs(&[ + &[("text", "Say hi to SECRET")], + &[("other", "tokens")], + &[("text", "def f():")] + ]) + ); + // `suffix`、`logprobs` 这些不给看 + assert_eq!( + v["params"], + json!({ "model": "gpt-3.5-turbo-instruct", "max_tokens": 16, "temperature": 0.5, "stop": ["\n\n"] }) + ); + + let v = view_of(Form::GeminiEmbed, &gemini_embed(), GEMINI_EMBED); + assert_eq!(v["format"], "gemini_embed"); + assert_eq!(v["model"], "gemini-embedding-001"); + assert_eq!( + parts_of(&v), + pairs(&[&[("text", "the SECRET plan"), ("text", "more")]]) + ); + assert_eq!(v["params"], json!({ "model": "gemini-embedding-001" })); + + let v = view_of(Form::GeminiEmbed, &gemini_batch(), GEMINI_BATCH); + assert_eq!( + parts_of(&v), + pairs(&[ + &[("text", "first SECRET")], + &[("text", "second"), ("other", "inlineData")] + ]) + ); +} + +/// 一个字符串是一条消息,一串 token(全是数字的数组)也是一条,几串 token 是几条 +#[test] +fn a_single_input_and_token_inputs() { + let one = json!({ "model": "m", "input": "hello" }); + let v = view_of(Form::OpenaiEmbeddings, &one, EMBEDDINGS); + assert_eq!(parts_of(&v), pairs(&[&[("text", "hello")]])); + let (out, _) = edit_inputs(Form::OpenaiEmbeddings, &one, EMBEDDINGS, |v| { + msgs(v)[0]["parts"][0]["text"] = json!("hi") + }) + .unwrap(); + // 还是一个字符串 + assert_eq!(out, json!({ "model": "m", "input": "hi" })); + + let tokens = json!({ "model": "m", "prompt": [1, 2, 3] }); + let v = view_of(Form::OpenaiCompletions, &tokens, COMPLETIONS); + assert_eq!(parts_of(&v), pairs(&[&[("other", "tokens")]])); + let many = json!({ "model": "m", "prompt": [[1, 2], [3]] }); + let v = view_of(Form::OpenaiCompletions, &many, COMPLETIONS); + assert_eq!( + parts_of(&v), + pairs(&[&[("other", "tokens")], &[("other", "tokens")]]) + ); + // 一串 token 只读 + let r = edit_inputs(Form::OpenaiCompletions, &tokens, COMPLETIONS, |v| { + msgs(v)[0]["parts"][0]["label"] = json!("text") + }); + assert!(matches!(r, Err(EditError::PermissionViolation(_))), "{r:?}"); + // 没有输入:一条消息都没有 + let none = json!({ "model": "m" }); + assert!( + view_of(Form::OpenaiCompletions, &none, COMPLETIONS)["messages"] + .as_array() + .unwrap() + .is_empty() + ); +} + +/// 改一段文字:写回之后,除了那一段,**一个字节都不差** +#[test] +fn editing_one_input_changes_only_that_text() { + for (form, raw, path, at) in input_samples() { + let (out, new_path) = edit_inputs(form, &raw, path, |v| { + let t = v["messages"][0]["parts"][0]["text"] + .as_str() + .unwrap() + .replace("SECRET", "[removed]"); + v["messages"][0]["parts"][0]["text"] = json!(t); + }) + .unwrap(); + assert_eq!(new_path, None, "{form:?}"); + let mut want = raw.clone(); + let was = want.pointer(at).unwrap().as_str().unwrap().to_string(); + *want.pointer_mut(at).unwrap() = json!(was.replace("SECRET", "[removed]")); + assert_ne!(want, raw); + assert_eq!(out.to_string(), want.to_string(), "{form:?}"); + } +} + +/// 原样交回是没改;只改了参数也只动参数 +#[test] +fn returning_an_inputs_view_untouched_changes_nothing() { + for (form, raw, path, _) in input_samples() { + let built = build(form, &raw, path).unwrap(); + let perms = [Permission::Messages, Permission::Params]; + let input = trim(&built.view, &perms); + let edits = built.src.check(&input, &input, &perms).unwrap(); + assert!(edits.is_empty(), "{form:?}: {edits:?}"); + } +} + +/// 消息、部分不能加、不能删、不能挪;只读的不能改;没有的那几节交回来也不收 +#[test] +fn inputs_cannot_be_added_removed_or_reordered() { + use EditError::*; + let kind = |r: Result<(Value, Option), EditError>| match r { + Err(PermissionViolation(_)) => "permission", + Err(BadOutput(_)) => "bad", + Ok(_) => "ok", + }; + for (form, raw, path, _) in input_samples() { + let run = |f: &dyn Fn(&mut Value)| kind(edit_inputs(form, &raw, path, |v| f(v))); + let case = format!("{form:?} {path}"); + // 加一条、删一条、挪一条 + assert_eq!( + run(&|v| msgs(v) + .push(json!({ "role": "user", "parts": [{ "type": "text", "text": "more" }] }))), + "bad", + "{case}" + ); + assert_eq!( + run(&|v| msgs(v).insert( + 0, + json!({ "role": "user", "parts": [{ "type": "text", "text": "first" }] }) + )), + "bad", + "{case}" + ); + assert_eq!( + run(&|v| { + msgs(v).remove(0); + }), + "bad", + "{case}" + ); + if v_len(&raw, form, path) > 1 { + assert_eq!( + run(&|v| { + let last = msgs(v).len() - 1; + msgs(v).remove(last); + }), + "bad", + "{case}" + ); + assert_eq!(run(&|v| msgs(v).swap(0, 1)), "bad", "{case}"); + } + // 部分也一样 + assert_eq!( + run(&|v| v["messages"][0]["parts"] + .as_array_mut() + .unwrap() + .push(json!({ "type": "text", "text": "more" }))), + "bad", + "{case}" + ); + assert_eq!( + run(&|v| v["messages"][0]["parts"].as_array_mut().unwrap().clear()), + "bad", + "{case}" + ); + // 角色只读 + assert_eq!( + run(&|v| v["messages"][0]["role"] = json!("assistant")), + "permission", + "{case}" + ); + // 没有系统提示、没有工具:权限再全也加不进去 + let all = all(); + let built = build(form, &raw, path).unwrap(); + let input = trim(&built.view, &all); + for (k, x) in [("system", json!("be nice")), ("tools", json!([]))] { + let mut out = input.clone(); + out[k] = x; + let r = built.src.check(&input, &out, &all); + assert!(matches!(r, Err(BadOutput(_))), "{case} {k}: {r:?}"); + } + } +} + +fn v_len(raw: &Value, form: Form, path: &str) -> usize { + view_of(form, raw, path)["messages"] + .as_array() + .unwrap() + .len() +} + +/// 只读的那几项(一串 token、图片)不能改;嵌入的参数只有模型名 +#[test] +fn read_only_items_and_params_an_input_list_does_not_have() { + let r = edit_inputs(Form::OpenaiEmbeddings, &embeddings(), EMBEDDINGS, |v| { + v["messages"][2]["parts"][0] = json!({ "key": "m2.p0", "type": "text", "text": "x" }) + }); + assert!(matches!(r, Err(EditError::PermissionViolation(_))), "{r:?}"); + let r = edit_inputs(Form::GeminiEmbed, &gemini_batch(), GEMINI_BATCH, |v| { + v["messages"][1]["parts"][1]["label"] = json!("text") + }); + assert!(matches!(r, Err(EditError::PermissionViolation(_))), "{r:?}"); + for (form, raw, path) in [ + (Form::OpenaiEmbeddings, embeddings(), EMBEDDINGS), + (Form::GeminiEmbed, gemini_embed(), GEMINI_EMBED), + ] { + for (k, x) in [ + ("max_tokens", json!(10)), + ("temperature", json!(0.1)), + ("top_p", json!(0.5)), + ("stop", json!(["x"])), + ] { + let r = edit_inputs(form, &raw, path, |v| v["params"][k] = x.clone()); + assert!( + matches!(r, Err(EditError::BadOutput(_))), + "{form:?} {k}: {r:?}" + ); + } + } +} + +/// 补全的参数写回原来的写法:`stop` 原来是一个字符串还写成字符串;去掉的就去掉 +#[test] +fn completions_params_are_written_back_in_their_own_fields() { + let (out, _) = edit_inputs(Form::OpenaiCompletions, &completions(), COMPLETIONS, |v| { + v["params"]["stop"] = json!(["END"]); + v["params"]["max_tokens"] = json!(64); + v["params"].as_object_mut().unwrap().remove("temperature"); + v["params"]["top_p"] = json!(0.9); + v["params"]["model"] = json!("davinci-002"); + }) + .unwrap(); + let mut want = completions(); + want["stop"] = json!("END"); + want["max_tokens"] = json!(64); + want.as_object_mut().unwrap().remove("temperature"); + want["top_p"] = json!(0.9); + want["model"] = json!("davinci-002"); + assert_eq!(out, want); + // 嵌入的模型名 + let (out, path) = edit_inputs(Form::OpenaiEmbeddings, &embeddings(), EMBEDDINGS, |v| { + v["params"]["model"] = json!("text-embedding-3-large") + }) + .unwrap(); + assert_eq!(path, None); + assert_eq!(out["model"], "text-embedding-3-large"); +} + +/// Gemini 换模型:路径换,请求体里写着的 `models/…` 跟着换(批量的每一个请求都换) +#[test] +fn a_gemini_embedding_model_change_moves_the_path_and_the_named_models() { + let (out, path) = edit_inputs(Form::GeminiEmbed, &gemini_batch(), GEMINI_BATCH, |v| { + v["params"]["model"] = json!("text-embedding-004") + }) + .unwrap(); + assert_eq!( + path.as_deref(), + Some("/v1beta/models/text-embedding-004:batchEmbedContents") + ); + for r in out["requests"].as_array().unwrap() { + assert_eq!(r["model"], "models/text-embedding-004"); + } + let (out, path) = edit_inputs(Form::GeminiEmbed, &gemini_embed(), GEMINI_EMBED, |v| { + v["params"]["model"] = json!("text-embedding-004") + }) + .unwrap(); + assert_eq!( + path.as_deref(), + Some("/v1beta/models/text-embedding-004:embedContent") + ); + assert_eq!(out["model"], "models/text-embedding-004"); + // 没写 `model` 的不加 + let mut bare = gemini_embed(); + bare.as_object_mut().unwrap().remove("model"); + let (out, _) = edit_inputs(Form::GeminiEmbed, &bare, GEMINI_EMBED, |v| { + v["params"]["model"] = json!("text-embedding-004") + }) + .unwrap(); + assert!(out.get("model").is_none(), "{out}"); +} + +/// 写回再守一道:哪怕拿对话的规矩核对过(加了、删了消息),写回时也不收 +#[test] +fn write_back_refuses_edits_that_change_the_number_of_inputs() { + let raw = embeddings(); + let built = build(Form::OpenaiEmbeddings, &raw, EMBEDDINGS).unwrap(); + let perms = [Permission::Messages]; + let input = trim(&built.view, &perms); + let mut out = input.clone(); + msgs(&mut out).remove(1); + // 对话的规矩允许删消息 + let edits = check(&input, &out, &perms, &[]).unwrap(); + let mut next = raw.clone(); + let r = apply(&mut next, &built.src, &edits, EMBEDDINGS); + assert!(matches!(r, Err(EditError::BadOutput(_))), "{r:?}"); +} + +/// 写回之后再读一遍还读得出来,什么样的返回值都不会让核对和写回 panic +#[test] +fn random_garbage_on_input_lists_never_panics() { + let mut rng = Rng(7); + for _ in 0..1000 { + for (form, raw, path, _) in input_samples() { + let built = build(form, &raw, path).unwrap(); + let perms = all(); + let input = trim(&built.view, &perms); + let mut out = input.clone(); + if let Some(ms) = out["messages"].as_array_mut() + && !ms.is_empty() + { + let i = rng.below(ms.len()); + match rng.below(5) { + 0 => { + ms.remove(i); + } + 1 => ms[i]["parts"][0]["text"] = json!(rng.text()), + 2 => ms[i]["parts"] = json!([]), + 3 => ms[i]["key"] = json!("m9"), + _ => ms.swap(0, i), + } + } + if rng.chance(30) { + out["params"]["model"] = json!(rng.text()); + } + if let Ok(edits) = built.src.check(&input, &out, &perms) { + let mut next = raw.clone(); + if let Ok(p) = apply(&mut next, &built.src, &edits, path) { + let path = p.unwrap_or(path.to_string()); + build(form, &next, &path).unwrap(); + } + } + } + } +} diff --git a/crates/tw-gateway/src/server/pipeline/plug.rs b/crates/tw-gateway/src/server/pipeline/plug.rs index a515ce4..4e3bdab 100644 --- a/crates/tw-gateway/src/server/pipeline/plug.rs +++ b/crates/tw-gateway/src/server/pipeline/plug.rs @@ -16,8 +16,9 @@ //! 换一家,管它的还是这些插件。 //! //! **发往上游的每一跳都过这一步**,不只生成回答的:数 token、Responses 的压缩一样过插件 -//! (插件删掉的东西不能从这些接口漏出去),插件看不懂的接口按插件的 `on_error` 处置(见 -//! [`crate::plugin::request::Shape`])。网关自己估数、不发出去的那一跳到不了这里。 +//! (插件删掉的东西不能从这些接口漏出去),嵌入和旧版补全过声明了它们的插件,别的接口 +//! 插件不管(见 [`crate::plugin::request::Shape`])。网关自己估数、不发出去的那一跳到不了 +//! 这里。 use bytes::Bytes; @@ -94,15 +95,21 @@ pub(super) async fn attempt( let decoded = req.api.filter(|_| reading.generates).map(|api| { tw_dialect::convert::decode(api.dialect(), &c.value, &c.path, req.query.as_deref()) }); - // 请求防护:只看插件加进来的。解不开的不看 —— 和开头那一遍一样,同格式直通照样发 - if let (Some(Ok(before)), Some(Ok(after))) = (&reading.decoded, &decoded) + // 请求防护:只看插件加进来的。解不开的不看 —— 和开头那一遍一样,同格式直通照样发。 + // 嵌入、旧版补全没有中间表示:比的是改前改后每项输入的文字(见 `Changed::screen`) + let screened = match (&reading.decoded, &decoded, &c.screen) { + (Some(Ok(before)), Some(Ok(after)), _) => Some((&before.request, &after.request)), + (_, _, Some((before, after))) => Some((before, after)), + _ => None, + }; + if let Some((before, after)) = screened && let Some(why) = crate::guard::screen_more( &state.bus, started.id, &provider.name, &crate::guard::Screen::of(rt), - &before.request, - &after.request, + before, + after, ) { return Err(why); diff --git a/crates/tw-gateway/src/server/upgrade.rs b/crates/tw-gateway/src/server/upgrade.rs index 9b7e153..cbe48aa 100644 --- a/crates/tw-gateway/src/server/upgrade.rs +++ b/crates/tw-gateway/src/server/upgrade.rs @@ -135,60 +135,17 @@ pub(super) async fn ws_upgrade( .await .map_err(|e| GatewayError::config(crate::state::credential_failed(e, &name)))?; let (id, ending) = open(&choice, &name, provider.billing.into()); - // 插件:升级那一刻的那一份表,一条连接用到底。**插件只看得懂 Responses 的 WebSocket** - // (每个 `response.create` 是一次请求);别的路径上的帧插件看不懂,管得着的插件按它的 - // `on_error` —— 拒绝就不接这条连接,跳过就记一笔、这条连接不过插件 - let hint = crate::hint::client_hint(&headers); - let plugins = if rt.plugins.is_empty() { - None - } else if crate::client_api::ClientApi::of_path(uri.path()) + // 插件:升级那一刻的那一份表,一条连接用到底。**插件只管 Responses 的 WebSocket**(每个 + // `response.create` 是一次对话请求);别的路径上的连接(比如 Realtime 的 `/v1/realtime`) + // 不属于插件处理的任何一种请求,所有插件都不管:原样接上,什么都不记 + let responses = crate::client_api::ClientApi::of_path(uri.path()) == Some(crate::client_api::ClientApi::OpenaiResponses) - && crate::client_api::ClientApi::generates(uri.path()) - { - Some(crate::ws::Plugins { - pool: state.plugin_pool.clone(), - set: rt.plugins.clone(), - client: hint.clone(), - }) - } else { - let to = crate::plugin::request::Target { - upstream: &name, - // 升级请求没有正文:说不出是哪个模型 - model: "", - requested_model: "", - attempt: 0, - }; - let hop_started = std::time::Instant::now(); - match crate::plugin::request::unreadable(&rt.plugins, hint.as_deref(), uri.path(), &to) { - Ok(runs) => { - crate::plugin::request::record(&state, id, &runs); - None - } - Err(refused) => { - crate::plugin::request::record(&state, id, &refused.runs); - let err = GatewayError::denied(refused.why); - // 和 HTTP 那条路一样:没发出去的这一跳在尝试链上,原因就是拒绝它的那句话 - state.bus.emit(tw_api::Event::RequestRouted { - id, - route: choice.route, - rule: choice.rule, - group: choice.group, - rewritten_by: Vec::new(), - denied_by: None, - affinity: None, - attempts: vec![crate::server::hop_failed( - &name, - None, - err.detail.clone(), - hop_started, - )], - billing: tw_api::Billing::PerToken, - }); - ending.failed(err.source.into(), err.detail.clone()); - return Err(err); - } - } - }; + && crate::client_api::ClientApi::generates(uri.path()); + let plugins = (responses && !rt.plugins.is_empty()).then(|| crate::ws::Plugins { + pool: state.plugin_pool.clone(), + set: rt.plugins.clone(), + client: crate::hint::client_hint(&headers), + }); let upstream = crate::ws::Upstream { url: crate::ws::upstream_url(&provider.base_url, uri.path(), query.as_deref()), headers: upstream_headers, diff --git a/crates/tw-gateway/src/ws.rs b/crates/tw-gateway/src/ws.rs index 4332945..0721ab8 100644 --- a/crates/tw-gateway/src/ws.rs +++ b/crates/tw-gateway/src/ws.rs @@ -26,9 +26,9 @@ //! (`response.created` 到 `response.completed`)起一组回答钩子的实例,排在占位符 //! 还原之后、工具墙之前。插件出错而策略是拒绝时,切掉的是那一次回答,连接照常。 //! -//! **插件只看得懂 Responses 的 WebSocket。**别的路径上(比如 Realtime 的 `/v1/realtime`) -//! 的帧插件看不懂,升级时就按管得着的插件的 `on_error` 处置:拒绝就不接这条连接,跳过就 -//! 记一笔、这条连接不过插件(见 `server::upgrade`、[`crate::plugin::request::unreadable`])。 +//! **插件只管 Responses 的 WebSocket**(每个 `response.create` 是一次对话请求)。别的路径 +//! 上的连接(比如 Realtime 的 `/v1/realtime`)不属于插件处理的任何一种请求:所有插件都 +//! 不管,原样接上,什么都不记(见 `server::upgrade`)。 //! //! # 两条明说的边界 //! diff --git a/crates/tw-gateway/tests/plugins_js.rs b/crates/tw-gateway/tests/plugins_js.rs index e6a2804..951356c 100644 --- a/crates/tw-gateway/tests/plugins_js.rs +++ b/crates/tw-gateway/tests/plugins_js.rs @@ -383,3 +383,123 @@ export function onRequest(req) { assert_eq!(sent[0][k].to_string(), request[k].to_string(), "{k}"); } } + +/// 声明了嵌入和补全的插件,在真的沙箱里:`requests` 读得出来,一项输入一条消息,上游收到 +/// 的是改过的那一份,一串 token 原样。只处理对话的那一个(一跑就抛错、出错时拒绝)**不跑 +/// 在这些请求上**,也拦不着它们 +#[tokio::test] +async fn a_javascript_plugin_that_declares_embeddings_rewrites_their_inputs() { + const SCRUB: &str = r#" +export const manifest = { + name: "Scrub", api: 1, permissions: ["messages"], + requests: ["conversation", "embeddings", "completions"], +}; +export function onRequest(req, ctx) { + console.log(ctx.format + " " + req.messages.length + " " + req.format); + for (const m of req.messages) { + for (const p of m.parts) { + if (p.type === "text") p.text = p.text.replaceAll("PROJECT-X", "[removed]"); + } + } + return req; +} +"#; + const STRICT: &str = r#" +export const manifest = { name: "Strict", api: 1, permissions: ["messages"] }; +export function onRequest(req) { + throw new Error("conversations only"); +} +"#; + let seen = Arc::new(Mutex::new(Vec::new())); + let up = upstream(seen.clone()).await; + let tmp = tempfile::tempdir().unwrap(); + std::fs::create_dir_all(tmp.path().join("plugins")).unwrap(); + std::fs::write(tmp.path().join("plugins/scrub.js"), SCRUB).unwrap(); + std::fs::write(tmp.path().join("plugins/strict.js"), STRICT).unwrap(); + let plugin = |id: &str, src: &str| Plugin { + id: id.into(), + file: format!("plugins/{id}.js"), + sha256: sha256_hex(src.as_bytes()), + enabled: true, + on_error: PluginOnError::Reject, + scope: Default::default(), + settings: Default::default(), + }; + let cfg = Config { + version: 1, + listen: Listen::default(), + clients: vec![Client { + name: "claude-code".into(), + key: "tw-testkey".into(), + ..Default::default() + }], + providers: vec![Provider { + name: "openai".into(), + base_url: format!("http://{up}"), + key: Some("sk-upstream".into()), + protocol: Some(Protocol::OpenaiChat), + ..Default::default() + }], + plugins: vec![plugin("scrub", SCRUB), plugin("strict", STRICT)], + ..Default::default() + }; + let state = tw_gateway::AppState::new(cfg).unwrap(); + state.set_config_dir(tmp.path().to_path_buf()); + let rt = state.runtime(); + let scrub = rt.plugins.get("scrub").cloned().unwrap(); + assert!(scrub.ready().is_some(), "{:?}", scrub.broken()); + assert_eq!( + scrub.requests, + [ + tw_api::RequestKind::Conversation, + tw_api::RequestKind::Embeddings, + tw_api::RequestKind::Completions + ] + ); + let strict = rt.plugins.get("strict").cloned().unwrap(); + assert_eq!(strict.requests, [tw_api::RequestKind::Conversation]); + let addr = tw_gateway::serve(state.clone(), ([127, 0, 0, 1], 0).into()) + .await + .unwrap(); + tokio::time::sleep(Duration::from_millis(30)).await; + let post = |path: &'static str, body: Value| async move { + reqwest::Client::new() + .post(format!("http://{addr}{path}")) + .header("authorization", "Bearer tw-testkey") + .header("content-type", "application/json") + .body(body.to_string()) + .send() + .await + .unwrap() + .status() + }; + let status = post( + "/v1/embeddings", + json!({ "model": "text-embedding-3-small", "input": ["PROJECT-X plan", [101, 102]] }), + ) + .await; + assert_eq!(status, 200); + let status = post( + "/v1/completions", + json!({ "model": "gpt-3.5-turbo-instruct", "prompt": "Summarize PROJECT-X", "max_tokens": 8 }), + ) + .await; + assert_eq!(status, 200); + let sent = seen.lock().unwrap().clone(); + assert_eq!(sent.len(), 2); + assert_eq!(sent[0]["input"], json!(["[removed] plan", [101, 102]])); + assert_eq!(sent[1]["prompt"], "Summarize [removed]"); + assert_eq!(sent[1]["max_tokens"], 8); + let logs: Vec = scrub.logs.lines().into_iter().map(|l| l.text).collect(); + assert_eq!( + logs, + [ + "openai_embeddings 2 openai_embeddings", + "openai_completions 1 openai_completions" + ] + ); + let st = scrub.stats.view(); + assert_eq!((st.calls, st.changed, st.errors), (2, 2, 0)); + // 只处理对话的那一个一次都没跑 + assert_eq!(strict.stats.view(), tw_api::PluginStats::default()); +} diff --git a/crates/tw-gateway/tests/plugins_request.rs b/crates/tw-gateway/tests/plugins_request.rs index 9b3c914..c322802 100644 --- a/crates/tw-gateway/tests/plugins_request.rs +++ b/crates/tw-gateway/tests/plugins_request.rs @@ -1644,83 +1644,523 @@ async fn counting_and_compacting_take_only_the_model_from_params() { } } -/// 插件看不懂的请求体(嵌入、认不出的接口):管得着的插件按它的 `on_error` —— 拒绝就不发, -/// 跳过就原样发、记一笔跳过。范围外的插件、空的请求体不算 +// ───────────────────────────────────────────────────────── 嵌入、旧版补全 + +const EMBED_GEMINI: &str = "/v1beta/models/gemini-embedding-001:embedContent"; +const EMBED_GEMINI_BATCH: &str = "/v1beta/models/gemini-embedding-001:batchEmbedContents"; + +fn embeddings_body() -> Value { + json!({ "model": "text-embedding-3-small", "dimensions": 256, + "input": [format!("Plan {MARK}"), "unrelated", [9906, 1917]] }) +} + +fn completions_body() -> Value { + json!({ "model": "gpt-3.5-turbo-instruct", "max_tokens": 16, "suffix": " end", + "prompt": [format!("Plan {MARK}"), [9906, 1917], format!("{MARK} notes")] }) +} + +fn gemini_embed_body() -> Value { + json!({ "model": "models/gemini-embedding-001", "taskType": "RETRIEVAL_DOCUMENT", + "content": { "parts": [{ "text": format!("Plan {MARK}") }] } }) +} + +fn gemini_batch_body() -> Value { + json!({ "requests": [ + { "model": "models/gemini-embedding-001", + "content": { "parts": [{ "text": format!("Plan {MARK}") }] } }, + { "model": "models/gemini-embedding-001", + "content": { "parts": [{ "text": "second" }, { "text": format!("{MARK} notes") }] } } + ] }) +} + +/// 四种非对话的请求体,各配一家同格式的上游:(路径, 协议, 请求体) +fn input_cases() -> Vec<(&'static str, Protocol, Value)> { + vec![ + ("/v1/embeddings", Protocol::OpenaiChat, embeddings_body()), + ("/v1/completions", Protocol::OpenaiChat, completions_body()), + (EMBED_GEMINI, Protocol::Gemini, gemini_embed_body()), + (EMBED_GEMINI_BATCH, Protocol::Gemini, gemini_batch_body()), + ] +} + +/// 把 [`MARK`] 从每项输入的文字里删掉的插件,**声明了嵌入和补全**。`ctx.format` 记下来 +fn scrub_inputs(formats: Arc>>) -> Double { + Double::new("scrub inputs") + .permit(&[Permission::Messages]) + .requests(&[ + tw_api::RequestKind::Conversation, + tw_api::RequestKind::Embeddings, + tw_api::RequestKind::Completions, + ]) + .on_request(move |mut view, ctx| { + formats + .lock() + .unwrap() + .push(ctx["format"].as_str().unwrap().to_string()); + assert_eq!(view["format"], ctx["format"]); + for m in view["messages"].as_array_mut().unwrap() { + assert_eq!(m["role"], "user"); + for p in m["parts"].as_array_mut().unwrap() { + if p["type"] == "text" { + let t = p["text"].as_str().unwrap().replace(MARK, "[removed]"); + p["text"] = json!(t); + } + } + } + Invocation::ok(RequestOutcome::Changed(view)) + }) +} + +/// 声明了嵌入、补全的插件:上游收到的是删过记号的那一份,**除了那几段文字一个字节都 +/// 不差**(客户端发来的就是排好序的紧凑 JSON,改过的请求体也是这么写的)。一串 token 原样; +/// 改过的那一份照样存下来,运行照样记在请求上 #[tokio::test] -async fn requests_plugins_cannot_read_follow_on_error() { +async fn a_plugin_that_declares_embeddings_and_completions_scrubs_their_inputs() { + for (path, protocol, body) in input_cases() { + let up = Upstream::default(); + let base = start_upstream(up.clone()).await; + let formats = Arc::new(Mutex::new(Vec::new())); + let mut gw = gateway( + vec![provider("same", base, protocol)], + SecurityMode::Off, + vec![entry("scrub", scrub_inputs(formats.clone()))], + ) + .await; + let (status, answer) = post(&gw, path, &body).await; + assert_eq!(status, 200, "{path}: {answer}"); + assert_eq!(up.hits(), 1, "{path}"); + let (got_path, _) = up.seen.lock().unwrap()[0].clone(); + assert_eq!(got_path, path); + let raw = String::from_utf8(up.raw.lock().unwrap()[0].to_vec()).unwrap(); + assert_eq!( + raw, + body.to_string().replace(MARK, "[removed]"), + "{path}: only the inputs' text may differ" + ); + assert_eq!( + gw.runs(), + [("scrub".to_string(), "changed".to_string(), 0)], + "{path}" + ); + let after = gw.after_plugins().await.expect("the body after plugins"); + assert!(!after.to_string().contains(MARK), "{path}: {after}"); + let want = match path { + "/v1/embeddings" => "openai_embeddings", + "/v1/completions" => "openai_completions", + _ => "gemini_embed", + }; + assert_eq!(formats.lock().unwrap().as_slice(), [want], "{path}"); + } +} + +/// **没声明的那种请求不在插件的范围里**:原样发、什么都不记、不算一次调用 —— 出错时拒绝 +/// 也一样。插件一律不管的接口(认不出的、空正文的)也是这样 +#[tokio::test] +async fn kinds_a_plugin_did_not_declare_pass_through_unrecorded() { + for on_error in [OnError::Reject, OnError::Skip] { + // 只处理对话的插件,跑一次就失败:在范围里的话,拒绝档下请求就被拒了 + let calls = Arc::new(AtomicUsize::new(0)); + let c = calls.clone(); + let failing = Double::new("failing") + .permit(&[Permission::Messages]) + .on_request(move |_, _| { + c.fetch_add(1, Ordering::SeqCst); + Invocation::err(RunError::Threw { + message: "nope".into(), + stack: None, + }) + }); + for (path, protocol, body) in input_cases() { + let up = Upstream::default(); + let base = start_upstream(up.clone()).await; + let gw = gateway( + vec![provider("same", base, protocol)], + SecurityMode::Off, + vec![entry_with("failing", failing.clone(), |a| { + a.on_error = on_error + })], + ) + .await; + let (status, answer) = post(&gw, path, &body).await; + assert_eq!(status, 200, "{on_error:?} {path}: {answer}"); + // 原样:客户端发来的那些字节 + assert_eq!( + up.raw.lock().unwrap()[0], + Bytes::from(body.to_string()), + "{on_error:?} {path}" + ); + tokio::time::sleep(Duration::from_millis(50)).await; + assert!(gw.runs().is_empty(), "{on_error:?} {path}: {:?}", gw.runs()); + assert_eq!(gw.stats("failing"), tw_api::PluginStats::default()); + } + assert_eq!(calls.load(Ordering::SeqCst), 0, "{on_error:?}"); + } + + // 插件一律不管的接口:认不出的路径,和没有正文的那种(取消一次 Responses 的回答) let up = Upstream::default(); let base = start_upstream(up.clone()).await; - let embeddings = - json!({ "model": "text-embedding-3-small", "input": [format!("Plan {MARK}")] }); - let providers = || vec![provider("chat", base, Protocol::OpenaiChat)]; - - let mut gw = gateway( - providers(), + let gw = gateway( + vec![provider("chat", base, Protocol::OpenaiChat)], SecurityMode::Off, - vec![entry("scrub", scrub())], + vec![entry("scrub", scrub_inputs(Default::default()))], ) .await; - let (status, body) = post(&gw, "/v1/embeddings", &embeddings).await; - assert_eq!(status, 403, "{body}"); - assert_eq!( - body["error"]["message"], - "[ThinkWatch] Plugin `Plugin scrub` cannot read requests to /v1/embeddings." - ); - assert_eq!(up.hits(), 0); - assert_eq!(gw.runs(), [("scrub".to_string(), "error".to_string(), 0)]); - assert!(gw.after_plugins().await.is_none()); - // 认不出的接口也一样 - let (status, body) = post( + let (status, _) = post( &gw, "/v1/rerank", &json!({ "model": "rerank-1", "query": MARK }), ) .await; - assert_eq!(status, 403, "{body}"); - assert_eq!(up.hits(), 0); + assert_eq!(status, 200); + assert!(String::from_utf8_lossy(&up.raw.lock().unwrap()[0]).contains(MARK)); + let up = Upstream::default(); + let base = start_upstream(up.clone()).await; + let gw2 = gateway( + vec![provider("responses", base, Protocol::OpenaiResponses)], + SecurityMode::Off, + vec![entry("scrub", scrub_inputs(Default::default()))], + ) + .await; + let r = reqwest::Client::new() + .post(format!("http://{}/v1/responses/resp_1/cancel", gw2.addr)) + .header("authorization", "Bearer tw-testkey") + .send() + .await + .unwrap(); + assert_eq!(r.status(), 200); + assert_eq!(up.hits(), 1); + tokio::time::sleep(Duration::from_millis(50)).await; + assert!(gw.runs().is_empty() && gw2.runs().is_empty()); +} + +/// 跑不了的插件(文件变了)只拦它声明过的那几种:只处理对话的拦不着嵌入,声明了嵌入的 +/// 照它的 `on_error` 拒掉嵌入请求 +#[tokio::test] +async fn a_broken_plugin_only_rejects_the_kinds_it_declared() { + let up = Upstream::default(); + let base = start_upstream(up.clone()).await; + let changed = |kinds: &[tw_api::RequestKind]| { + let mut a = double::active( + "old", + Double::new("Old") + .permit(&[Permission::Messages]) + .requests(kinds), + ); + a.state = PluginState::Broken(Broken::Changed); + Arc::new(a) + }; + let providers = || vec![provider("chat", base, Protocol::OpenaiChat)]; - // 跳过:原样发,记一笔跳过(不算一次调用) + // 只处理对话:嵌入照常,对话被拒 let gw = gateway( providers(), SecurityMode::Off, - vec![entry_with("scrub", scrub(), |a| a.on_error = OnError::Skip)], + vec![changed(&[tw_api::RequestKind::Conversation])], ) .await; - let (status, _) = post(&gw, "/v1/embeddings", &embeddings).await; + let (status, answer) = post(&gw, "/v1/embeddings", &embeddings_body()).await; + assert_eq!(status, 200, "{answer}"); + let (status, answer) = post(&gw, "/v1/completions", &completions_body()).await; + assert_eq!(status, 200, "{answer}"); + tokio::time::sleep(Duration::from_millis(50)).await; + assert!(gw.runs().is_empty(), "{:?}", gw.runs()); + let chat = json!({ "model": "gpt-5", "messages": [{ "role": "user", "content": "hi" }] }); + let (status, answer) = post(&gw, "/v1/chat/completions", &chat).await; + assert_eq!(status, 403, "{answer}"); + assert_eq!(gw.runs(), [("old".to_string(), "error".to_string(), 0)]); + assert_eq!(up.hits(), 2); + + // 声明了嵌入:嵌入被拒,补全照常 + let gw = gateway( + providers(), + SecurityMode::Off, + vec![changed(&[tw_api::RequestKind::Embeddings])], + ) + .await; + let (status, answer) = post(&gw, "/v1/embeddings", &embeddings_body()).await; + assert_eq!(status, 403, "{answer}"); + assert!( + answer["error"]["message"] + .as_str() + .unwrap() + .contains("changed on disk"), + "{answer}" + ); + let (status, _) = post(&gw, "/v1/completions", &completions_body()).await; assert_eq!(status, 200); - assert_eq!(up.hits(), 1); - assert_eq!(gw.runs(), [("scrub".to_string(), "skipped".to_string(), 0)]); - assert_eq!(gw.stats("scrub").calls, 0); + assert_eq!(up.hits(), 3); +} - // 范围外的插件不算 +/// 嵌入、补全上的插件和对话上的一样:按发出去的模型、上游挑;只看到占位符;出错、 +/// `reject` 按 `on_error` 拒掉整个请求;不能多一项、少一项输入;请求体读不出来时按 +/// `on_error` —— 拒绝就不发(`gw.plugin.cannot_read_body`),跳过就原样发、记一笔跳过 +#[tokio::test] +async fn embeddings_follow_the_same_scope_placeholders_and_on_error() { + const EMBED: &str = "/v1/embeddings"; + let up = Upstream::default(); + let base = start_upstream(up.clone()).await; + let providers = || vec![provider("a", base, Protocol::OpenaiChat)]; + let kinds = [tw_api::RequestKind::Embeddings]; + + // 范围外:别的模型的插件不跑 let gw = gateway( providers(), SecurityMode::Off, - vec![entry_with("scrub", scrub(), |a| { - a.scope.models = vec!["claude-*".into()] + vec![entry_with("scrub", scrub_inputs(Default::default()), |a| { + a.scope.models = vec!["text-embedding-3-large".into()] })], ) .await; - let (status, _) = post(&gw, "/v1/embeddings", &embeddings).await; + let (status, _) = post(&gw, EMBED, &embeddings_body()).await; assert_eq!(status, 200); + assert!(String::from_utf8_lossy(&up.raw.lock().unwrap()[0]).contains(MARK)); + tokio::time::sleep(Duration::from_millis(50)).await; assert!(gw.runs().is_empty()); - // 空的请求体里没有插件能改的东西(取消一次 Responses 的回答) + // 占位符:插件看不到真的密钥,改过的地方换回去再发 + let saw = Arc::new(Mutex::new(String::new())); + let s = saw.clone(); + let checked = Double::new("checked") + .permit(&[Permission::Messages]) + .requests(&kinds) + .on_request(move |mut view, ctx| { + *s.lock().unwrap() = view.to_string(); + assert_eq!( + (ctx["upstream"].as_str(), ctx["format"].as_str()), + (Some("a"), Some("openai_embeddings")) + ); + let t = view["messages"][0]["parts"][0]["text"] + .as_str() + .unwrap() + .to_string(); + view["messages"][0]["parts"][0]["text"] = json!(format!("{t} (checked)")); + Invocation::ok(RequestOutcome::Changed(view)) + }); + up.raw.lock().unwrap().clear(); + let gw = gateway( + providers(), + SecurityMode::Off, + vec![entry("checked", checked)], + ) + .await; + let body = + json!({ "model": "text-embedding-3-small", "input": format!("my key is {USER_KEY}") }); + let (status, _) = post(&gw, EMBED, &body).await; + assert_eq!(status, 200); + let saw = saw.lock().unwrap().clone(); + assert!(!saw.contains(USER_KEY), "the plugin saw the key: {saw}"); + assert!(saw.contains("<>"), "{saw}"); + let sent: Value = serde_json::from_slice(&up.raw.lock().unwrap()[0]).unwrap(); + assert_eq!(sent["input"], format!("my key is {USER_KEY} (checked)")); + + // 出错、拒绝、多加一项输入:整个请求不发 + let failing = Double::new("failing") + .permit(&[Permission::Messages]) + .requests(&kinds) + .on_request(|_, _| { + Invocation::err(RunError::Threw { + message: "nope".into(), + stack: None, + }) + }); + let refusing = Double::new("refusing") + .permit(&[Permission::Messages]) + .requests(&kinds) + .on_request(|_, _| Invocation::ok(RequestOutcome::Rejected("not embedded".into()))); + let adding = Double::new("adding") + .permit(&[Permission::Messages]) + .requests(&kinds) + .on_request(|mut view, _| { + view["messages"] + .as_array_mut() + .unwrap() + .push(json!({ "role": "user", "parts": [{ "type": "text", "text": "one more" }] })); + Invocation::ok(RequestOutcome::Changed(view)) + }); + for (e, code) in [ + (entry("failing", failing), "gw.plugin.request_failed"), + (entry("refusing", refusing), "gw.plugin.rejected"), + (entry("adding", adding), "gw.plugin.request_failed"), + ] { + up.raw.lock().unwrap().clear(); + let gw = gateway(providers(), SecurityMode::Off, vec![e]).await; + let rx = gw.state.bus.subscribe(); + let (status, body) = post(&gw, EMBED, &embeddings_body()).await; + assert_eq!(status, 403, "{code}: {body}"); + assert!(up.raw.lock().unwrap().is_empty(), "{code}"); + let attempts = routed(rx).await; + assert_eq!( + attempts[0].error.as_ref().map(|e| e.code.as_str()), + Some(code) + ); + } + + // 请求体读不出来:拒绝就不发,跳过就原样发、记一笔跳过(不算一次调用) + let not_json = |gw: SocketAddr| { + reqwest::Client::new() + .post(format!("http://{gw}/v1/embeddings")) + .header("x-api-key", "tw-testkey") + .header("content-type", "application/json") + .body("{ not json") + }; + let gw = gateway( + providers(), + SecurityMode::Off, + vec![entry("scrub", scrub_inputs(Default::default()))], + ) + .await; + up.raw.lock().unwrap().clear(); + let r = not_json(gw.addr).send().await.unwrap(); + assert_eq!(r.status(), 403); + let answer: Value = r.json().await.unwrap(); + assert_eq!( + answer["error"]["message"], + "[ThinkWatch] Plugin `Plugin scrub` cannot read this request: the request body is not JSON" + ); + assert!(up.raw.lock().unwrap().is_empty()); + assert_eq!(gw.runs(), [("scrub".to_string(), "error".to_string(), 0)]); + let gw = gateway( + providers(), + SecurityMode::Off, + vec![entry_with("scrub", scrub_inputs(Default::default()), |a| { + a.on_error = OnError::Skip + })], + ) + .await; + let r = not_json(gw.addr).send().await.unwrap(); + assert_eq!(r.status(), 200); + assert_eq!(up.raw.lock().unwrap()[0], Bytes::from("{ not json")); + assert_eq!(gw.runs(), [("scrub".to_string(), "skipped".to_string(), 0)]); + assert_eq!(gw.stats("scrub").calls, 0); +} + +/// 一串 token 只读:插件原样交回,上游收到的是客户端的原话;改它就是越权 +#[tokio::test] +async fn token_id_inputs_are_read_only() { let up = Upstream::default(); let base = start_upstream(up.clone()).await; + let body = json!({ "model": "gpt-3.5-turbo-instruct", "prompt": [[1, 2, 3], [4, 5]] }); + let saw = Arc::new(Mutex::new(Value::Null)); + let s = saw.clone(); + let look = Double::new("look") + .permit(&[Permission::Messages]) + .requests(&[tw_api::RequestKind::Completions]) + .on_request(move |view, _| { + *s.lock().unwrap() = view.clone(); + Invocation::ok(RequestOutcome::Changed(view)) + }); let gw = gateway( - vec![provider("responses", base, Protocol::OpenaiResponses)], + vec![provider("a", base, Protocol::OpenaiChat)], SecurityMode::Off, - vec![entry("scrub", scrub())], + vec![entry("look", look)], ) .await; - let r = reqwest::Client::new() - .post(format!("http://{}/v1/responses/resp_1/cancel", gw.addr)) - .header("authorization", "Bearer tw-testkey") - .send() - .await - .unwrap(); - assert_eq!(r.status(), 200); + let (status, _) = post(&gw, "/v1/completions", &body).await; + assert_eq!(status, 200); + assert_eq!(up.raw.lock().unwrap()[0], Bytes::from(body.to_string())); + let saw = saw.lock().unwrap().clone(); + let parts: Vec<&Value> = saw["messages"] + .as_array() + .unwrap() + .iter() + .map(|m| &m["parts"][0]) + .collect(); + assert_eq!(parts.len(), 2); + for p in parts { + assert_eq!( + (p["type"].as_str(), p["label"].as_str()), + (Some("other"), Some("tokens")) + ); + } + tokio::time::sleep(Duration::from_millis(50)).await; + assert_eq!( + gw.runs(), + [("look".to_string(), "unchanged".to_string(), 0)] + ); + + let forge = Double::new("forge") + .permit(&[Permission::Messages]) + .requests(&[tw_api::RequestKind::Completions]) + .on_request(|mut view, _| { + view["messages"][0]["parts"][0] = + json!({ "key": "m0.p0", "type": "text", "text": "hi" }); + Invocation::ok(RequestOutcome::Changed(view)) + }); + let gw = gateway( + vec![provider("a", base, Protocol::OpenaiChat)], + SecurityMode::Off, + vec![entry("forge", forge)], + ) + .await; + let (status, answer) = post(&gw, "/v1/completions", &body).await; + assert_eq!(status, 403, "{answer}"); + assert!( + answer["error"]["message"] + .as_str() + .unwrap() + .contains("no permission to change"), + "{answer}" + ); assert_eq!(up.hits(), 1); - assert!(gw.runs().is_empty(), "{:?}", gw.runs()); +} + +/// 插件改过的嵌入请求再看一遍请求防护,**只看插件加进来的**:插件写进来的命中拦下整个 +/// 请求;客户端原话里就有的(嵌入开头不看请求防护)不因为插件改了别的一项被拦 +#[tokio::test] +async fn screening_sees_the_inputs_a_plugin_wrote() { + let up = Upstream::default(); + let base = start_upstream(up.clone()).await; + let security = Security { + content: tw_config::ContentPolicy { + mode: 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 write = |text: &'static str| { + Double::new("write") + .permit(&[Permission::Messages]) + .requests(&[tw_api::RequestKind::Embeddings]) + .on_request(move |mut view, _| { + view["messages"][1]["parts"][0]["text"] = json!(text); + Invocation::ok(RequestOutcome::Changed(view)) + }) + }; + let gw = gateway_with( + vec![provider("a", base, Protocol::OpenaiChat)], + SecurityMode::Off, + vec![entry("write", write("the forbidden-plan"))], + security.clone(), + ) + .await; + let rx = gw.state.bus.subscribe(); + let (status, body) = post(&gw, "/v1/embeddings", &embeddings_body()).await; + assert_eq!(status, 403, "{body}"); + assert!(up.raw.lock().unwrap().is_empty()); + let attempts = routed(rx).await; + assert_eq!( + attempts[0].error.as_ref().map(|e| e.code.as_str()), + Some("gw.content.refused") + ); + + let gw = gateway_with( + vec![provider("a", base, Protocol::OpenaiChat)], + SecurityMode::Off, + vec![entry("write", write("harmless"))], + security, + ) + .await; + let mut body = embeddings_body(); + body["input"][0] = json!("the forbidden-plan, as the client wrote it"); + let (status, answer) = post(&gw, "/v1/embeddings", &body).await; + assert_eq!(status, 200, "{answer}"); + let sent: Value = serde_json::from_slice(&up.raw.lock().unwrap()[0]).unwrap(); + assert_eq!(sent["input"][1], "harmless"); } diff --git a/crates/tw-gateway/tests/plugins_ws.rs b/crates/tw-gateway/tests/plugins_ws.rs index 1533e60..c439912 100644 --- a/crates/tw-gateway/tests/plugins_ws.rs +++ b/crates/tw-gateway/tests/plugins_ws.rs @@ -431,11 +431,10 @@ async fn realtime_upstream() -> (SocketAddr, Arc>>) { (addr, seen) } -/// 不是 Responses 的 WebSocket(比如 Realtime 的 `/v1/realtime`):插件看不懂它的帧。管得着 -/// 的插件按它的 `on_error` 在升级时就处置 —— 拒绝就不接这条连接,跳过就接上、帧原样过去、 -/// 记一笔跳过 +/// 不是 Responses 的 WebSocket(比如 Realtime 的 `/v1/realtime`):不属于插件处理的任何一种 +/// 请求,**所有插件都不管** —— 出错时拒绝的、跳过的都一样:接上,帧原样过去,什么都不记 #[tokio::test] -async fn a_websocket_plugins_cannot_read_follows_on_error_at_the_upgrade() { +async fn a_websocket_plugins_do_not_handle_passes_through_unrecorded() { use tokio_tungstenite::tungstenite::client::IntoClientRequest; let look = || { Double::new("look") @@ -455,55 +454,30 @@ async fn a_websocket_plugins_cannot_read_follows_on_error_at_the_upgrade() { "content": [{ "type": "input_text", "text": "secret plan" }] } }) .to_string(); - // 拒绝:升级不成,一个字节都没到上游 - let (up, seen) = realtime_upstream().await; - let (gw, runs) = gateway_with(up, vec![entry("look", look())], Default::default()).await; - match tokio_tungstenite::connect_async(request(gw)).await { - Err(tokio_tungstenite::tungstenite::Error::Http(r)) => { - assert_eq!(r.status(), 403); - let body = - String::from_utf8_lossy(r.body().as_deref().unwrap_or_default()).into_owned(); - assert!( - body.contains("Plugin `Plugin look` cannot read requests to /v1/realtime."), - "{body}" - ); - } - other => panic!("the upgrade went through: {:?}", other.map(|_| ())), + for on_error in [tw_api::OnError::Reject, tw_api::OnError::Skip] { + let (up, seen) = realtime_upstream().await; + let (gw, runs) = gateway_with( + up, + vec![entry_with("look", look(), |a| a.on_error = on_error)], + Default::default(), + ) + .await; + let (mut c, _) = tokio_tungstenite::connect_async(request(gw)) + .await + .unwrap_or_else(|e| panic!("{on_error:?}: the upgrade was refused: {e}")); + c.send(WsMsg::Text(item.clone().into())).await.unwrap(); + let back = tokio::time::timeout(Duration::from_secs(3), c.next()) + .await + .expect("no echo") + .unwrap() + .unwrap(); + assert_eq!(back.into_text().unwrap().as_str(), item, "{on_error:?}"); + assert_eq!( + seen.lock().unwrap().as_slice(), + std::slice::from_ref(&item), + "{on_error:?}" + ); + tokio::time::sleep(Duration::from_millis(100)).await; + assert!(runs.lock().unwrap().is_empty(), "{on_error:?}"); } - tokio::time::sleep(Duration::from_millis(100)).await; - assert!(seen.lock().unwrap().is_empty()); - let outcomes: Vec = runs - .lock() - .unwrap() - .iter() - .map(|r| r.run.outcome.slug().to_string()) - .collect(); - assert_eq!(outcomes, ["error"]); - - // 跳过:接上,帧原样过去,这条连接上记一笔跳过 - let (up, seen) = realtime_upstream().await; - let (gw, runs) = gateway_with( - up, - vec![entry_with("look", look(), |a| { - a.on_error = tw_api::OnError::Skip - })], - Default::default(), - ) - .await; - let (mut c, _) = tokio_tungstenite::connect_async(request(gw)).await.unwrap(); - c.send(WsMsg::Text(item.clone().into())).await.unwrap(); - let back = tokio::time::timeout(Duration::from_secs(3), c.next()) - .await - .expect("no echo") - .unwrap() - .unwrap(); - assert_eq!(back.into_text().unwrap().as_str(), item); - assert_eq!(seen.lock().unwrap().as_slice(), [item]); - let outcomes: Vec = runs - .lock() - .unwrap() - .iter() - .map(|r| r.run.outcome.slug().to_string()) - .collect(); - assert_eq!(outcomes, ["skipped"]); } diff --git a/crates/tw-plugin/src/lib.rs b/crates/tw-plugin/src/lib.rs index 7ca2828..a3fb97a 100644 --- a/crates/tw-plugin/src/lib.rs +++ b/crates/tw-plugin/src/lib.rs @@ -199,6 +199,8 @@ pub struct Manifest { pub api: u32, pub description: Option, pub permissions: BTreeSet, + /// 插件处理哪几种请求(清单里的 `requests`)。没写是只有对话 + pub requests: BTreeSet, pub scope: Scope, pub reply_mode: ReplyMode, /// 按作者写的先后 @@ -255,6 +257,56 @@ impl Permission { } } +/// 一种请求(清单里 `requests` 的一项)。**插件只处理它声明了的那几种**:别的种类的 +/// 请求不过它,出了什么错也和它无关 +#[derive( + Debug, Clone, Copy, PartialEq, Eq, PartialOrd, Ord, Hash, serde::Serialize, serde::Deserialize, +)] +#[serde(rename_all = "snake_case")] +pub enum RequestKind { + /// 对话:Anthropic Messages、OpenAI Chat、Responses、Gemini 的生成,连同它们的数 token + /// 和压缩。不写 `requests` 时就是只有它 + Conversation, + /// 嵌入:OpenAI 的 `/v1/embeddings`,Gemini 的 `:embedContent`、`:batchEmbedContents` + Embeddings, + /// 旧版补全:OpenAI 的 `/v1/completions` + Completions, +} + +impl RequestKind { + pub const ALL: [RequestKind; 3] = [ + RequestKind::Conversation, + RequestKind::Embeddings, + RequestKind::Completions, + ]; + + /// 清单和控制面里都是这个词 + pub fn as_str(self) -> &'static str { + match self { + RequestKind::Conversation => "conversation", + RequestKind::Embeddings => "embeddings", + RequestKind::Completions => "completions", + } + } + + pub fn from_manifest(s: &str) -> Option { + RequestKind::ALL.into_iter().find(|k| k.as_str() == s) + } + + /// 管得着这种请求的权限:它的视图里有的那几节,加上只在对话上跑的回答钩子。 + /// + /// 嵌入和旧版补全的视图只有 `messages`(每项输入一条)和 `params`,没有系统提示、 + /// 没有工具,回答钩子也不在它们上面跑 + pub fn reached_by(self) -> &'static [Permission] { + match self { + RequestKind::Conversation => &Permission::ALL, + RequestKind::Embeddings | RequestKind::Completions => { + &[Permission::Messages, Permission::Params] + } + } + } +} + /// 清单里的 `match`。每一项是带 `*` 的通配;空的表示不限 #[derive(Debug, Clone, Default, PartialEq, Eq)] pub struct Scope { diff --git a/crates/tw-plugin/src/manifest.rs b/crates/tw-plugin/src/manifest.rs index 73c812a..dee458b 100644 --- a/crates/tw-plugin/src/manifest.rs +++ b/crates/tw-plugin/src/manifest.rs @@ -7,7 +7,9 @@ use std::collections::BTreeSet; use serde_json::{Map, Value}; -use crate::{Hooks, LoadError, Manifest, Permission, ReplyMode, Scope, SettingKind, SettingSpec}; +use crate::{ + Hooks, LoadError, Manifest, Permission, ReplyMode, RequestKind, Scope, SettingKind, SettingSpec, +}; const MAX_NAME: usize = 64; const MAX_DESCRIPTION: usize = 500; @@ -97,7 +99,14 @@ pub(crate) fn parse(info: &[u8]) -> Result { for key in m.keys() { if !matches!( key.as_str(), - "name" | "api" | "description" | "permissions" | "match" | "reply" | "settings" + "name" + | "api" + | "description" + | "permissions" + | "requests" + | "match" + | "reply" + | "settings" ) { return Err(err(format!("the manifest has an unknown field `{key}`"))); } @@ -138,6 +147,7 @@ pub(crate) fn parse(info: &[u8]) -> Result { }; let permissions = permissions(m.get("permissions"))?; + let requests = requests(m.get("requests"))?; let scope = scope(m.get("match"))?; let reply_mode = match m.get("reply") { None | Some(Value::Null) => ReplyMode::Block, @@ -148,12 +158,14 @@ pub(crate) fn parse(info: &[u8]) -> Result { let settings = settings(m.get("settings"), info.settings_order.as_deref())?; check_hooks(&hooks, &permissions, reply_mode)?; + check_requests(&requests, &permissions)?; Ok(Manifest { name, api, description, permissions, + requests, scope, reply_mode, settings, @@ -192,6 +204,70 @@ fn permissions(v: Option<&Value>) -> Result, LoadError> { Ok(out) } +/// 插件处理哪几种请求。**不写就是只有对话**,写了就得是一张非空的单子 +fn requests(v: Option<&Value>) -> Result, LoadError> { + let list = match v { + None | Some(Value::Null) => return Ok(BTreeSet::from([RequestKind::Conversation])), + Some(Value::Array(a)) => a, + Some(_) => return Err(err("`requests` must be a list")), + }; + if list.is_empty() { + return Err(err( + "`requests` must list at least one kind of request; leave it out to handle conversations only", + )); + } + let mut out = BTreeSet::new(); + for r in list { + let Some(s) = r.as_str() else { + return Err(err("every entry of `requests` must be a string")); + }; + let Some(kind) = RequestKind::from_manifest(s) else { + return Err(err(format!( + "`requests` lists \"{s}\", which is not a kind of request; the kinds are {}", + RequestKind::ALL + .iter() + .map(|k| format!("\"{}\"", k.as_str())) + .collect::>() + .join(", ") + ))); + }; + if !out.insert(kind) { + return Err(err(format!("\"{s}\" is listed twice in `requests`"))); + } + } + Ok(out) +} + +/// 声明的请求和申请的权限对得上:**每一种请求都有权限碰得到它的视图**,**每个权限都在 +/// 某一种声明了的请求上用得着** —— 和钩子、权限一一对应是同一个道理,什么都不白要。 +/// +/// 嵌入和旧版补全的视图只有 `messages` 和 `params`,回答钩子也只在对话上跑:只要了 +/// `system` 的插件处理不了嵌入,不处理对话的插件用不着 `system`、`tools` 和回答钩子 +fn check_requests( + requests: &BTreeSet, + perms: &BTreeSet, +) -> Result<(), LoadError> { + for kind in requests { + if !kind.reached_by().iter().any(|p| perms.contains(p)) { + return Err(err(format!( + "`requests` lists \"{}\", but none of the permissions applies to those requests; \ + they show only \"messages\" and \"params\"", + kind.as_str() + ))); + } + } + for p in perms { + if !requests.iter().any(|k| k.reached_by().contains(p)) { + return Err(err(format!( + "permission \"{}\" only applies to conversations, and `requests` does not list \ + \"conversation\"", + p.manifest_name() + ))); + } + } + Ok(()) +} + fn scope(v: Option<&Value>) -> Result { let m = match v { None | Some(Value::Null) => return Ok(Scope::default()), diff --git a/crates/tw-plugin/tests/loading.rs b/crates/tw-plugin/tests/loading.rs index b5b4e0d..2f6327e 100644 --- a/crates/tw-plugin/tests/loading.rs +++ b/crates/tw-plugin/tests/loading.rs @@ -10,7 +10,7 @@ use std::time::Instant; use common::*; use serde_json::json; use sha2::{Digest, Sha256}; -use tw_plugin::{LoadError, Permission, RequestOutcome}; +use tw_plugin::{LoadError, Permission, RequestKind, RequestOutcome}; fn load_err(src: &[u8]) -> LoadError { let t = Instant::now(); @@ -226,6 +226,92 @@ fn manifests_that_break_the_rules_are_load_errors() { } } +// ── 处理哪几种请求(`requests`)───────────────────────────────────── + +#[test] +fn requests_default_to_conversations_and_can_list_more_kinds() { + let p = load_source(&plugin( + r#"{ name: "x", api: 1, permissions: ["messages"] }"#, + ON_REQUEST, + )); + assert_eq!( + p.manifest().requests, + BTreeSet::from([RequestKind::Conversation]) + ); + let p = load_source(&plugin( + r#"{ name: "x", api: 1, permissions: ["messages", "params"], + requests: ["embeddings", "conversation", "completions"] }"#, + ON_REQUEST, + )); + assert_eq!(p.manifest().requests, BTreeSet::from(RequestKind::ALL)); + // 只处理嵌入的:权限只有嵌入的视图里有的那几节 + let p = load_source(&plugin( + r#"{ name: "x", api: 1, permissions: ["messages"], requests: ["embeddings"] }"#, + ON_REQUEST, + )); + assert_eq!( + p.manifest().requests, + BTreeSet::from([RequestKind::Embeddings]) + ); +} + +#[test] +fn requests_that_break_the_rules_are_load_errors() { + let both = format!("{ON_REQUEST}\n{ON_TEXT}"); + let cases: &[(&str, &str, &str)] = &[ + ( + "an empty list", + r#"{ name: "x", api: 1, permissions: ["messages"], requests: [] }"#, + ON_REQUEST, + ), + ( + "a kind that does not exist", + r#"{ name: "x", api: 1, permissions: ["messages"], requests: ["images"] }"#, + ON_REQUEST, + ), + ( + "a string instead of a list", + r#"{ name: "x", api: 1, permissions: ["messages"], requests: "embeddings" }"#, + ON_REQUEST, + ), + ( + "an entry that is not a string", + r#"{ name: "x", api: 1, permissions: ["messages"], requests: [1] }"#, + ON_REQUEST, + ), + ( + "a kind listed twice", + r#"{ name: "x", api: 1, permissions: ["messages"], requests: ["embeddings", "embeddings"] }"#, + ON_REQUEST, + ), + // 嵌入的视图里没有系统提示:只要了 `system` 的插件碰不到嵌入请求里的任何东西 + ( + "a kind none of the permissions reaches", + r#"{ name: "x", api: 1, permissions: ["system"], requests: ["conversation", "embeddings"] }"#, + ON_REQUEST, + ), + // 不处理对话,`system` 就白要了 + ( + "a permission that only applies to conversations, without conversations", + r#"{ name: "x", api: 1, permissions: ["system", "messages"], requests: ["completions"] }"#, + ON_REQUEST, + ), + // 回答钩子只在对话上跑 + ( + "reply hooks without conversations", + r#"{ name: "x", api: 1, permissions: ["messages", "reply.text"], requests: ["embeddings"] }"#, + &both, + ), + ]; + for (what, manifest, hooks) in cases { + match rt().load(plugin(manifest, hooks).as_bytes()) { + Err(LoadError::Manifest(why)) => assert!(why.contains("requests"), "{what}: {why}"), + Err(e) => panic!("{what}: {e:?}"), + Ok(_) => panic!("{what}: loaded"), + } + } +} + #[test] fn an_unsupported_api_version_is_named() { let e = load_err( diff --git a/docs/config.md b/docs/config.md index 2627c88..202f26b 100644 --- a/docs/config.md +++ b/docs/config.md @@ -1020,6 +1020,14 @@ 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. +A plugin handles the kinds of request its code declares: conversations +(Anthropic Messages, OpenAI Chat Completions and Responses, and Gemini, +including their token counts and compaction), embeddings (`/v1/embeddings`, +Gemini `:embedContent` and `:batchEmbedContents`) and legacy completions +(`/v1/completions`). A plugin that declares none handles conversations only. +Requests of a kind a plugin does not handle pass without it, whatever its +`on_error`. Other endpoints, such as images and audio, pass without any plugin. + diff --git a/docs/config.zh-CN.md b/docs/config.zh-CN.md index 516f061..0d0c7d3 100644 --- a/docs/config.zh-CN.md +++ b/docs/config.zh-CN.md @@ -835,6 +835,8 @@ default_route: default 插件在路由之后改写请求,请求每发往一个上游改写一次。故障转移到另一个上游时,从客户端发来的原样重新开始;插件看得到这一次发往哪个上游、用哪个模型名。路由、模型准入和会话归组看的都是客户端发来的原样。 +插件处理它在代码里声明的那几种请求:对话(Anthropic Messages、OpenAI Chat Completions 和 Responses、Gemini,连同它们的数 token 和压缩)、嵌入(`/v1/embeddings`、Gemini 的 `:embedContent` 和 `:batchEmbedContents`)和旧版补全(`/v1/completions`)。没有声明的插件只处理对话。插件不处理的那种请求不经过它,不论 `on_error` 怎么设。其他接口(图片、音频等)不经过任何插件。 +