diff --git a/crates/tw-dialect/src/params.rs b/crates/tw-dialect/src/params.rs index 2cedb307..78d24b45 100644 --- a/crates/tw-dialect/src/params.rs +++ b/crates/tw-dialect/src/params.rs @@ -86,15 +86,97 @@ pub fn set_max_output_tokens(dialect: Dialect, body: &mut Value, n: u64) { } /// 把最大输出 token 数限制在 `cap` 以内:客户端写的比它大、或者没写,就写成 `cap`; -/// 写的不比它大就不动。返回改没改。 -pub fn cap_max_output_tokens(dialect: Dialect, body: &mut Value, cap: u64) -> bool { - if !body.is_object() || max_output_tokens(dialect, body).is_some_and(|n| n <= cap) { +/// 写的不比它大就不动。返回改没改。`official` 是这个请求要发往的是不是厂商官方的端点 +/// ([`crate::official::is_official_host`],和转换时的 [`crate::ir::Target::official`] +/// 一个意思)。 +/// +/// 和 [`set_max_output_tokens`] 不一样的几处,都是因为「限制」要对上游真的管用: +/// +/// - **Chat 两个名字都写了的,各自压到 `cap` 以内。**只看解码器先认的那一个 +/// (`max_completion_tokens`)的话,`{"max_completion_tokens": 100, "max_tokens": 9999}` +/// 原样过去,只认 `max_tokens` 的上游等于没限。比 `cap` 小的那个不动,不替客户端放宽。 +/// 值是 `null` 的算没写。 +/// - **Chat 两个都没写时,官方端点写 `max_completion_tokens`,别家写 `max_tokens`**: +/// OpenAI 官方的推理模型不认 `max_tokens`,带着它整个请求 400;兼容实现大多只认 +/// `max_tokens`。和编码器选名字的办法一样。 +/// - **Anthropic 开了思考的**(`thinking.type` 是 `enabled`,带 `budget_tokens`):思考 +/// 用的 token 算在 `max_tokens` 里,所以 `budget_tokens` 必须小于 `max_tokens`,否则整个 +/// 请求 400。把 `max_tokens` 压到不比预算大时,`cap` 大于 1024(预算的下限)就把预算一并 +/// 压到 `cap - 1`;不大于 1024 的话预算没有合法的值可取,去掉 `thinking` —— 不思考地 +/// 回答,总比整个请求被拒强。Converse(`additionalModelRequestFields.thinking`)背后的 +/// Claude 同理。预算本来就比限制小的不动;客户端自己写的 `max_tokens` 就不比预算大的, +/// 是它自己的请求不合法,也不动。 +pub fn cap_max_output_tokens(dialect: Dialect, body: &mut Value, cap: u64, official: bool) -> bool { + let Some(obj) = body.as_object_mut() else { + return false; + }; + if dialect == Dialect::Chat { + return cap_chat(obj, cap, official); + } + if max_output_tokens(dialect, body).is_some_and(|n| n <= cap) { return false; } set_max_output_tokens(dialect, body, cap); + let thinking = match dialect { + Dialect::Anthropic => body.as_object_mut(), + Dialect::Bedrock => body + .get_mut("additionalModelRequestFields") + .and_then(Value::as_object_mut), + _ => None, + }; + if let Some(holder) = thinking { + fit_thinking(holder, cap); + } true } +/// Chat 的两个名字(见 [`cap_max_output_tokens`]) +fn cap_chat(obj: &mut Map, cap: u64, official: bool) -> bool { + const NAMES: [&str; 2] = ["max_completion_tokens", "max_tokens"]; + let written: Vec<&str> = NAMES + .into_iter() + .filter(|k| obj.get(*k).is_some_and(|v| !v.is_null())) + .collect(); + if written.is_empty() { + let name = if official { NAMES[0] } else { NAMES[1] }; + obj.insert(name.into(), Value::from(cap)); + return true; + } + let mut changed = false; + for k in written { + // 写的不是一个非负整数的,和比上限大的一样换成上限:上游读不懂它,就等于没限 + if !obj.get(k).and_then(Value::as_u64).is_some_and(|n| n <= cap) { + obj.insert(k.into(), Value::from(cap)); + changed = true; + } + } + changed +} + +/// `max_tokens` 刚压到 `cap` 之后,让思考的预算还小于它(见 [`cap_max_output_tokens`])。 +/// `holder` 是装着 `thinking` 的那个对象 +fn fit_thinking(holder: &mut Map, cap: u64) { + /// Anthropic 允许的最小思考预算 + const BUDGET_MIN: u64 = 1024; + let Some(thinking) = holder.get_mut("thinking").and_then(Value::as_object_mut) else { + return; + }; + if thinking.get("type").and_then(Value::as_str) != Some("enabled") { + return; + } + let Some(budget) = thinking.get("budget_tokens").and_then(Value::as_u64) else { + return; + }; + if budget < cap { + return; + } + if cap > BUDGET_MIN { + thinking.insert("budget_tokens".into(), Value::from(cap - 1)); + } else { + holder.remove("thinking"); + } +} + /// 改要的模型。Gemini 和 Bedrock 的模型写在路径里,请求体里没有它,什么都不做(路径 /// 由调用方改)。 pub fn set_model(dialect: Dialect, body: &mut Value, model: &str) { @@ -275,22 +357,309 @@ mod tests { #[test] fn a_cap_lowers_or_fills_and_leaves_a_smaller_value_alone() { - for d in ALL { - let mut v = json!({}); - assert!(cap_max_output_tokens(d, &mut v, 100), "{d:?}"); - assert_eq!(max_output_tokens(d, &v), Some(100)); - set_max_output_tokens(d, &mut v, 50); - assert!(!cap_max_output_tokens(d, &mut v, 100), "{d:?}"); - assert_eq!(max_output_tokens(d, &v), Some(50)); - set_max_output_tokens(d, &mut v, 500); - assert!(cap_max_output_tokens(d, &mut v, 100), "{d:?}"); - assert_eq!(max_output_tokens(d, &v), Some(100)); + for official in [false, true] { + for d in ALL { + let mut v = json!({}); + assert!(cap_max_output_tokens(d, &mut v, 100, official), "{d:?}"); + assert_eq!(max_output_tokens(d, &v), Some(100)); + set_max_output_tokens(d, &mut v, 50); + assert!(!cap_max_output_tokens(d, &mut v, 100, official), "{d:?}"); + assert_eq!(max_output_tokens(d, &v), Some(50)); + set_max_output_tokens(d, &mut v, 500); + assert!(cap_max_output_tokens(d, &mut v, 100, official), "{d:?}"); + assert_eq!(max_output_tokens(d, &v), Some(100)); + } } let mut not_an_object = json!([1]); - assert!(!cap_max_output_tokens(Dialect::Chat, &mut not_an_object, 1)); + assert!(!cap_max_output_tokens( + Dialect::Chat, + &mut not_an_object, + 1, + false + )); assert_eq!(not_an_object, json!([1])); } + /// 四种客户端格式(外加 Converse)各压一次:`(格式, 请求, 官方端点吗)` → 压完的样子 + fn capped(d: Dialect, mut v: Value, cap: u64, official: bool) -> (bool, Value) { + let changed = cap_max_output_tokens(d, &mut v, cap, official); + (changed, v) + } + + #[test] + fn a_chat_request_that_names_both_fields_is_capped_in_both() { + // 只看先认的那一个的话,`max_tokens` 原样过去,只认它的上游等于没限 + let both = json!({"max_completion_tokens": 50, "max_tokens": 5000}); + for official in [false, true] { + assert_eq!( + capped(Dialect::Chat, both.clone(), 100, official), + ( + true, + json!({"max_completion_tokens": 50, "max_tokens": 100}) + ) + ); + // 反过来也一样;比上限小的那个不替客户端放宽 + assert_eq!( + capped( + Dialect::Chat, + json!({"max_completion_tokens": 5000, "max_tokens": 50}), + 100, + official + ), + ( + true, + json!({"max_completion_tokens": 100, "max_tokens": 50}) + ) + ); + // 两个都不大:不动 + let small = json!({"max_completion_tokens": 50, "max_tokens": 60}); + assert_eq!( + capped(Dialect::Chat, small.clone(), 100, official), + (false, small) + ); + // 读不懂的值等于没限,换成上限;`null` 算没写 + assert_eq!( + capped( + Dialect::Chat, + json!({"max_completion_tokens": "lots", "max_tokens": 50}), + 100, + official + ), + ( + true, + json!({"max_completion_tokens": 100, "max_tokens": 50}) + ) + ); + assert_eq!( + capped( + Dialect::Chat, + json!({"max_completion_tokens": null, "max_tokens": 50}), + 100, + official + ), + ( + false, + json!({"max_completion_tokens": null, "max_tokens": 50}) + ) + ); + } + // 别的格式只有一个名字,照常压 + for (d, v, want) in [ + ( + Dialect::Anthropic, + json!({"max_tokens": 5000}), + json!({"max_tokens": 100}), + ), + ( + Dialect::Responses, + json!({"max_output_tokens": 5000}), + json!({"max_output_tokens": 100}), + ), + ( + Dialect::Gemini, + json!({"generationConfig": {"maxOutputTokens": 5000}}), + json!({"generationConfig": {"maxOutputTokens": 100}}), + ), + ( + Dialect::Bedrock, + json!({"inferenceConfig": {"maxTokens": 5000}}), + json!({"inferenceConfig": {"maxTokens": 100}}), + ), + ] { + assert_eq!(capped(d, v, 100, false), (true, want), "{d:?}"); + } + } + + #[test] + fn a_chat_request_with_no_limit_gets_the_field_its_upstream_reads() { + // OpenAI 官方的推理模型带着 `max_tokens` 整个请求 400;兼容实现大多只认 `max_tokens` + assert_eq!( + capped(Dialect::Chat, json!({"model": "o3"}), 100, true), + (true, json!({"model": "o3", "max_completion_tokens": 100})) + ); + assert_eq!( + capped(Dialect::Chat, json!({"model": "o3"}), 100, false), + (true, json!({"model": "o3", "max_tokens": 100})) + ); + // 写成 `null` 的等于没写 + assert_eq!( + capped(Dialect::Chat, json!({"max_tokens": null}), 100, true), + ( + true, + json!({"max_tokens": null, "max_completion_tokens": 100}) + ) + ); + assert_eq!( + capped(Dialect::Chat, json!({"max_tokens": null}), 100, false), + (true, json!({"max_tokens": 100})) + ); + // 别的格式只有一个名字,和发往哪儿无关 + for official in [false, true] { + for (d, want) in [ + (Dialect::Anthropic, json!({"max_tokens": 100})), + (Dialect::Responses, json!({"max_output_tokens": 100})), + ( + Dialect::Gemini, + json!({"generationConfig": {"maxOutputTokens": 100}}), + ), + ( + Dialect::Bedrock, + json!({"inferenceConfig": {"maxTokens": 100}}), + ), + ] { + assert_eq!(capped(d, json!({}), 100, official), (true, want), "{d:?}"); + } + } + } + + #[test] + fn thinking_still_fits_under_a_lowered_anthropic_limit() { + let thinking = |budget: u64| json!({"type": "enabled", "budget_tokens": budget}); + // 预算不比新的上限小:压到上限减一 + assert_eq!( + capped( + Dialect::Anthropic, + json!({"max_tokens": 32000, "thinking": thinking(16000)}), + 8000, + true + ), + ( + true, + json!({"max_tokens": 8000, "thinking": thinking(7999)}) + ) + ); + // 没写 max_tokens 的,填上上限也一样 + assert_eq!( + capped( + Dialect::Anthropic, + json!({"thinking": thinking(8000)}), + 8000, + false + ), + ( + true, + json!({"max_tokens": 8000, "thinking": thinking(7999)}) + ) + ); + // 上限只比最小预算多一:预算正好是最小值 + assert_eq!( + capped( + Dialect::Anthropic, + json!({"max_tokens": 4096, "thinking": thinking(2048)}), + 1025, + false + ), + ( + true, + json!({"max_tokens": 1025, "thinking": thinking(1024)}) + ) + ); + // 上限不比最小预算大:开不了思考,去掉它,请求照样能答 + for cap in [1024, 600] { + assert_eq!( + capped( + Dialect::Anthropic, + json!({"max_tokens": 4096, "thinking": thinking(2048), "x": 1}), + cap, + false + ), + (true, json!({"max_tokens": cap, "x": 1})), + "{cap}" + ); + } + // 预算本来就小于上限的、没开思考的、自适应的:不动思考 + for t in [ + thinking(2000), + json!({"type": "disabled"}), + json!({"type": "adaptive"}), + ] { + assert_eq!( + capped( + Dialect::Anthropic, + json!({"max_tokens": 32000, "thinking": t.clone()}), + 8000, + false + ), + (true, json!({"max_tokens": 8000, "thinking": t})) + ); + } + // 客户端自己写的 max_tokens 没超上限:不是我们改出来的,不碰 + let own = json!({"max_tokens": 2000, "thinking": thinking(4000)}); + assert_eq!( + capped(Dialect::Anthropic, own.clone(), 8000, false), + (false, own) + ); + // Converse 背后的 Claude 同理 + assert_eq!( + capped( + Dialect::Bedrock, + json!({ + "inferenceConfig": {"maxTokens": 32000}, + "additionalModelRequestFields": {"thinking": thinking(16000)} + }), + 8000, + false + ), + ( + true, + json!({ + "inferenceConfig": {"maxTokens": 8000}, + "additionalModelRequestFields": {"thinking": thinking(7999)} + }) + ) + ); + } + + #[test] + fn other_formats_keep_their_reasoning_settings_when_capped() { + // 思考的预算只有 Anthropic(和 Converse 上的 Claude)要小于输出上限;别的格式的 + // 推理开关不归这个函数管 + for (d, v, want) in [ + ( + Dialect::Chat, + json!({"max_tokens": 32000, "reasoning_effort": "high"}), + json!({"max_tokens": 8000, "reasoning_effort": "high"}), + ), + ( + Dialect::Responses, + json!({"max_output_tokens": 32000, "reasoning": {"effort": "high"}}), + json!({"max_output_tokens": 8000, "reasoning": {"effort": "high"}}), + ), + ( + Dialect::Gemini, + json!({"generationConfig": { + "maxOutputTokens": 32000, + "thinkingConfig": {"thinkingBudget": 16000} + }}), + json!({"generationConfig": { + "maxOutputTokens": 8000, + "thinkingConfig": {"thinkingBudget": 16000} + }}), + ), + ] { + assert_eq!(capped(d, v, 8000, true), (true, want), "{d:?}"); + } + } + + #[test] + fn the_routing_rule_setter_is_not_the_cap() { + // 桌面版路由规则的 `set` 照旧:两个名字写了哪个改哪个(往大改也改),都没写写 + // `max_tokens`,不碰思考 + let mut v = json!({"max_completion_tokens": 50, "max_tokens": 5000}); + set_max_output_tokens(Dialect::Chat, &mut v, 9000); + assert_eq!( + v, + json!({"max_completion_tokens": 9000, "max_tokens": 9000}) + ); + let mut v = + json!({"max_tokens": 32000, "thinking": {"type": "enabled", "budget_tokens": 16000}}); + set_max_output_tokens(Dialect::Anthropic, &mut v, 8000); + assert_eq!( + v, + json!({"max_tokens": 8000, "thinking": {"type": "enabled", "budget_tokens": 16000}}) + ); + } + #[test] fn the_model_lives_in_the_body_except_where_it_lives_in_the_path() { for d in ALL { diff --git a/crates/tw-gateway/src/guard.rs b/crates/tw-gateway/src/guard.rs index b313c958..d4eb6c2c 100644 --- a/crates/tw-gateway/src/guard.rs +++ b/crates/tw-gateway/src/guard.rs @@ -22,11 +22,13 @@ //! 拦截档下 [`look`] 按客户端原文里出现的先后给找到的值编好号,每一跳都接着这本账换 //! (见 [`tw_guard::redact::flow`])。 +use std::collections::HashSet; + use tw_config::SecurityMode as Mode; use tw_guard::content::Screening; use tw_guard::redact::flow; use tw_guard::redact::replace::Ledger; -use tw_guard::redact::rules::{Finding, Hit, RuleSet}; +use tw_guard::redact::rules::{Finding, Hit, Rule, RuleSet}; /// 按规则找一遍,**不算我们自己的占位符,也不进 base64 载荷**(见 [`flow::hits`])。 pub fn hits(text: &str, rules: &RuleSet) -> Vec { @@ -67,22 +69,37 @@ pub fn replace( /// `after` 里 `before` 没有的那些值:插件写进请求里的(见 [`crate::plugin::request`])。 /// /// 按规则和打过码的样子比:同一个值在两份里打出来的码一样。客户端原话里就有的值,开头 -/// 那一遍已经报过了,插件改过的那一份里再出现不再报一次。 +/// 那一遍已经报过了,插件改过的那一份里再出现不再报一次。**查表比**,不两两比:两份里 +/// 各有几万个值时,后者是平方级的。 pub fn more_found(before: &[Finding], after: Vec) -> Vec { + let known: HashSet<(&Rule, &str)> = before + .iter() + .map(|b| (&b.rule, b.masked.as_str())) + .collect(); after .into_iter() - .filter(|f| { - !before - .iter() - .any(|b| b.rule == f.rule && b.masked == f.masked) - }) + .filter(|f| !known.contains(&(&f.rule, f.masked.as_str()))) .collect() } -/// 找到的东西写成事件里的样子。 -pub fn items(found: &[Finding]) -> Vec { +/// 一个请求报出去的值最多几个。 +/// +/// 安全日志一个值一行([`tw_api::Event::SecretsFound`] 的一项就是一行)。一个请求里贴进来 +/// 几万个不同的密钥时(一份导出的凭据清单、一段日志),逐个报就是几万行:日志被一个请求 +/// 刷满,真正要看的那几条反而淹没了,一条事件也有几 MB。前 100 个(按在请求里第一次出现 +/// 的先后)足够说明这个请求带了什么。 +/// +/// **只少报,不少换**:替换按的是账本([`look`]),和报了几个无关 —— 超出的值照样换成 +/// 占位符、回答里照样换回来。 +pub const REPORTED_MAX: usize = 100; + +/// 找到的东西写成事件里的样子。`already` 是这个请求先前已经找到过几个(插件改写过的请求 +/// 只报插件写进来的那些,见 [`more_found`]):**一个请求加起来最多报 [`REPORTED_MAX`] 个**, +/// 先报先出现的。 +pub fn items(found: &[Finding], already: usize) -> Vec { found .iter() + .take(REPORTED_MAX.saturating_sub(already)) .map(|f| tw_api::SecretItem { rule: f.rule.id().to_string(), custom: f.rule.custom(), @@ -369,13 +386,83 @@ mod tests { #[test] fn the_event_items_name_the_rule_and_never_carry_the_value() { - let it = items(&find(Mode::Observe, &RuleSet::defaults(), &body())); + let it = items(&find(Mode::Observe, &RuleSet::defaults(), &body()), 0); assert_eq!(it[0].rule, "anthropic-api-key"); assert_eq!(it[0].kind, tw_api::SecretKind::ApiKeys); assert!(!it[0].custom); assert!(!it[0].masked.contains("AAAAAAAAAAAA"), "{}", it[0].masked); } + /// 第 `i` 个 AWS 访问密钥的样子,两两不同,打出来的码(头 5 尾 4)也两两不同 + fn aws_key(i: usize) -> String { + let mut n = i; + let d: Vec = (0..5) + .map(|_| { + let c = char::from(b'A' + (n % 26) as u8); + n /= 26; + c + }) + .collect(); + format!("AKIA{}QQQQQQQQQQQ{}{}{}{}", d[0], d[1], d[2], d[3], d[4]) + } + + #[test] + fn a_request_reports_its_first_hundred_values_and_still_replaces_every_one() { + let keys: Vec = (0..150).map(aws_key).collect(); + let body = serde_json::json!({"messages": [{"content": keys.join(" ")}]}).to_string(); + let (found, ledger) = look(Mode::Enforce, &RuleSet::defaults(), body.as_bytes()); + // 找到的一个不少:总共几个不同的值,`found.len()` 说得出来 + assert_eq!(found.len(), 150); + let it = items(&found, 0); + assert_eq!(it.len(), REPORTED_MAX); + // 报的是先出现的那些,按出现的先后 + let masked: Vec = keys[..REPORTED_MAX] + .iter() + .map(|k| tw_guard::redact::rules::masked(&found[0].rule, k)) + .collect(); + assert_eq!( + it.iter().map(|i| i.masked.clone()).collect::>(), + masked + ); + // 插件改写过的请求再报一次(只报插件写进来的):和开头那一条加起来不超过上限 + assert_eq!(items(&found, 30).len(), REPORTED_MAX - 30); + assert!(items(&found, 150).is_empty()); + // 没报的照样换掉 + let (out, ledger) = replace(Mode::Enforce, &RuleSet::defaults(), body.into(), &ledger); + let out = String::from_utf8(out.to_vec()).unwrap(); + assert!(!out.contains("AKIA"), "{out}"); + assert!(out.contains("<>"), "{out}"); + assert_eq!(ledger.len(), 150); + } + + /// 查表比和两两比,留下的一模一样:同一条规则下打码一样的算见过,换一条规则就不算 + #[test] + fn more_found_keeps_exactly_what_comparing_every_pair_kept() { + let text = |range: std::ops::Range| { + let keys: Vec = range.map(aws_key).collect(); + serde_json::json!({ "content": keys.join(" ") }).to_string() + }; + let mut before = find(Mode::Observe, &RuleSet::defaults(), text(0..60).as_bytes()); + // 前 10 个当成是另一条规则认出来的:打码一样、规则不同,不算见过 + for f in &mut before[..10] { + f.rule = Rule::Custom(std::sync::Arc::from("aws-again")); + } + let after_text = format!("{} {}", text(30..90), text(0..5)); + let after = find(Mode::Observe, &RuleSet::defaults(), after_text.as_bytes()); + let pairwise: Vec = after + .iter() + .filter(|f| { + !before + .iter() + .any(|b| b.rule == f.rule && b.masked == f.masked) + }) + .cloned() + .collect(); + let got = more_found(&before, after); + assert_eq!(got, pairwise); + assert_eq!(got.len(), 35, "60..90 是新的,0..5 换了规则"); + } + /// 每一跳接着原文那本账换:同一把密钥在每一跳都是同一个号,哪怕那一跳发出去的那份 /// 把字段换了顺序(转换过格式,或者改写参数时按键名重排过)。以前各起一本账,下面 /// 这一跳里 `system` 排到了 `messages` 后面,两把密钥的号就对调了 diff --git a/crates/tw-gateway/src/plugin/bridge.rs b/crates/tw-gateway/src/plugin/bridge.rs index 92d4ea6f..02bb7504 100644 --- a/crates/tw-gateway/src/plugin/bridge.rs +++ b/crates/tw-gateway/src/plugin/bridge.rs @@ -17,6 +17,8 @@ //! 那个原值。危险的工具调用归工具调用审查管:它看的是换回之后、客户端要执行的那一个调用, //! 把凭据发往陌生主机的,内置规则 `secret-to-unknown-host` 在拦截档下切断。 +use std::collections::{BTreeMap, HashMap}; +use std::hash::{BuildHasherDefault, Hasher}; use std::sync::Arc; use serde_json::Value; @@ -28,8 +30,14 @@ use tw_guard::redact::rules::RuleSet; pub struct Bridge { rules: Arc, ledger: Ledger, - /// 账里的值,长的在前:换的时候长的先换,一个值是另一个的一部分时不会只换半截 + /// 账里的值和它的占位符,长的在前、一样长的按字典序:换的时候长的先换,一个值是另一个 + /// 的一部分时不会只换半截(见 [`Bridge::hide`]) values: Vec<(String, String)>, + /// 一段文字里账里的值都在哪儿(见 [`Lengths`]) + lengths: Lengths, + /// `values` 的下标,按值的字节序排:一段尾巴是不是哪个值的开头,二分查一次(见 + /// [`Bridge::hold_from`]) + by_bytes: Vec, } impl Bridge { @@ -38,6 +46,8 @@ impl Bridge { rules, ledger: Ledger::new(Scheme::SECRET), values: Vec::new(), + lengths: Lengths::default(), + by_bytes: Vec::new(), } } @@ -82,35 +92,26 @@ impl Bridge { let mut values: Vec<(String, String)> = self .ledger .replacements() + .filter(|(o, _)| !o.is_empty()) .map(|(o, p)| (o.to_string(), p.to_string())) .collect(); values.sort_by(|a, b| b.0.len().cmp(&a.0.len()).then_with(|| a.0.cmp(&b.0))); + let mut by_bytes: Vec = (0..values.len()).collect(); + by_bytes.sort_unstable_by(|&a, &b| values[a].0.cmp(&values[b].0)); + self.lengths = Lengths::of(&values); + self.by_bytes = by_bytes; self.values = values; } /// 一段文字里的密钥换成占位符:账里的值,加上按规则新找到的。 pub fn hide(&mut self, s: &str) -> String { - let mut out = None::; - for (original, placeholder) in &self.values { - let cur = out.as_deref().unwrap_or(s); - if cur.contains(original.as_str()) { - out = Some(cur.replace(original.as_str(), placeholder)); - } - } - let cur = out.unwrap_or_else(|| s.to_string()); + let cur = self.hide_known(s); if self.rules.is_empty() { return cur; } let mut hits = tw_guard::redact::rules::scan_text(&cur, &self.rules); // 压在一个占位符上的不算(连接串规则会把 `app:<>@` 当成口令) - if !hits.is_empty() && cur.contains(Scheme::SECRET.open) { - let ours = Scheme::SECRET.find_in(&cur); - hits.retain(|h| { - !ours - .iter() - .any(|(at, _, _)| at.start < h.bytes.end && h.bytes.start < at.end) - }); - } + tw_guard::redact::flow::off_placeholders(&cur, &mut hits); if hits.is_empty() { return cur; } @@ -121,6 +122,46 @@ impl Bridge { r.text } + /// 账里的值换成占位符。 + /// + /// **换出来的和「长的先换、一样长的按字典序,一个值一个值地整段 `replace`」一样**,只是 + /// 不拿每个值去整段里各找一遍 —— 账里几万个值时,每段文字都要找几万遍。先一遍找出 + /// 每个值出现的每一处(互相重叠的也算),再照那个先后挑:一处被先换的值占了的地方, + /// 后换的值在那儿就不在了;同一个值的几处互相重叠时,`replace` 从左往右取不重叠的, + /// 这里也是。 + /// + /// 和一个个 `replace` 只差在一种情况上:后换的值碰巧是**刚换上去的那个占位符**里的一截 + /// (一个口令就是 `1`),一个个换会把占位符换坏,这里不会 —— 找的是原文。 + fn hide_known(&self, s: &str) -> String { + let mut found = self.lengths.find(&self.values, s); + if found.is_empty() { + return s.to_string(); + } + // 下标就是先后(`values` 照换的先后排好了),同一个值从左往右 + found.sort_unstable(); + // 挑中的:起点 → (终点, 哪个值)。互不重叠,所以只看起点在它前面的最后一处 + let mut taken: BTreeMap = BTreeMap::new(); + for (v, start) in found { + let end = start + self.values[v].0.len(); + let clash = taken + .range(..end) + .next_back() + .is_some_and(|(_, &(e, _))| e > start); + if !clash { + taken.insert(start, (end, v)); + } + } + let mut out = String::with_capacity(s.len()); + let mut at = 0; + for (start, (end, v)) in taken { + out.push_str(&s[at..start]); + out.push_str(&self.values[v].1); + at = end; + } + out.push_str(&s[at..]); + out + } + /// 占位符换回原值 pub fn reveal(&self, s: &str) -> String { tw_guard::redact::replace::restore(s, &self.ledger) @@ -189,24 +230,132 @@ impl Bridge { /// 可能把它补全 —— 半截的值送进去,插件就看到了真值的一部分,换也换不掉。 /// /// 返回 `buf.len()` 是全都能给。切点总在字符边界上。 + /// + /// 从长到短试 `buf` 的每一截尾巴,是哪个值的真前缀就从那儿扣:按字节序排好的值里 + /// 二分查一次,不拿每个值的每个前缀去比 —— 后者在流式的每一段上都要比「账里有几个 + /// 值 × 值有多长」次。 pub fn hold_from(&self, buf: &str) -> usize { - let mut cut = buf.len(); - for (original, _) in &self.values { - // 从长到短试这个值的每一个真前缀 - let mut ends: Vec = original - .char_indices() - .map(|(i, _)| i) - .filter(|i| *i > 0) - .collect(); - ends.reverse(); - for k in ends { - if k <= buf.len() && buf.ends_with(&original[..k]) { - cut = cut.min(buf.len() - k); - break; + let longest = self.values.first().map_or(0, |(o, _)| o.len()); + for k in (1..longest.min(buf.len() + 1)).rev() { + let at = buf.len() - k; + if buf.is_char_boundary(at) && self.starts_a_value(&buf.as_bytes()[at..]) { + return at; + } + } + buf.len() + } + + /// `tail` 是不是账里某个值的**真**前缀。以它开头的值在字节序里连成一段,打头的是第一 + /// 个不小于它的;那个正好等于它的话(不是真前缀),这一段里还有的就是紧跟着的那个 + fn starts_a_value(&self, tail: &[u8]) -> bool { + let first = self + .by_bytes + .partition_point(|&v| self.values[v].0.as_bytes() < tail); + self.by_bytes[first..].iter().take(2).any(|&v| { + let o = self.values[v].0.as_bytes(); + o.len() > tail.len() && o.starts_with(tail) + }) + } +} + +/// 账里的值按长度分组,在一段文字里一遍找出它们出现的每一处(见 [`Bridge::hide_known`])。 +/// +/// 每一种长度在文字上滚一遍哈希(Rabin–Karp),哈希对上了再逐字节核对:一段文字的 +/// 代价是「文字长度 × 值有几种长度」,和账里有几个值无关 —— 同一种密钥一样长,几万把 +/// AWS 访问密钥只是一种长度。哈希撞了只多核对一次,不会换错。 +#[derive(Clone, Default)] +struct Lengths { + groups: Vec, +} + +/// 一样长的那些值 +#[derive(Clone)] +struct Group { + len: usize, + /// `BASE` 的 `len - 1` 次方:滚动时减掉移出窗口的那个字节用 + top: u64, + /// 哈希 → 这么长的值(`values` 的下标) + by_hash: HashMap, BuildHasherDefault>, +} + +/// 滚动哈希的底 +const BASE: u64 = 0x0100_0000_01b3; + +/// 一段字节的滚动哈希。字节加一,免得开头的 `\0` 不算数 +fn rolling(bytes: &[u8]) -> u64 { + bytes.iter().fold(0, |h: u64, &c| { + h.wrapping_mul(BASE).wrapping_add(u64::from(c) + 1) + }) +} + +impl Lengths { + /// `values` 已经按长度从长到短排好:一样长的连在一起 + fn of(values: &[(String, String)]) -> Self { + let mut groups: Vec = Vec::new(); + for (v, (original, _)) in values.iter().enumerate() { + let len = original.len(); + if groups.last().is_none_or(|g| g.len != len) { + groups.push(Group { + len, + top: (1..len).fold(1, |p: u64, _| p.wrapping_mul(BASE)), + by_hash: HashMap::default(), + }); + } + if let Some(g) = groups.last_mut() { + g.by_hash + .entry(rolling(original.as_bytes())) + .or_default() + .push(v); + } + } + Self { groups } + } + + /// `s` 里每一处账里的值:(`values` 的下标, 起点)。互相重叠的都在 + fn find(&self, values: &[(String, String)], s: &str) -> Vec<(usize, usize)> { + let b = s.as_bytes(); + let mut found = Vec::new(); + for g in self.groups.iter().filter(|g| g.len <= b.len()) { + let mut h = rolling(&b[..g.len]); + for at in 0..=b.len() - g.len { + if at > 0 { + let (out, inn) = (b[at - 1], b[at + g.len - 1]); + h = h + .wrapping_sub((u64::from(out) + 1).wrapping_mul(g.top)) + .wrapping_mul(BASE) + .wrapping_add(u64::from(inn) + 1); + } + // 值是完整的 UTF-8,字节对得上的地方自然落在字符边界上 + for &v in g.by_hash.get(&h).into_iter().flatten() { + if values[v].0.as_bytes() == &b[at..at + g.len] { + found.push((v, at)); + } } } } - cut + found + } +} + +/// 滚动哈希的低位分布得不匀(底是奇数,最低一位只看字节和的奇偶),进哈希表之前再 +/// 搅一下(splitmix64 的收尾) +#[derive(Default)] +struct Spread(u64); + +impl Hasher for Spread { + fn write(&mut self, bytes: &[u8]) { + for &b in bytes { + self.0 = self.0.rotate_left(8) ^ u64::from(b); + } + } + fn write_u64(&mut self, x: u64) { + self.0 = x; + } + fn finish(&self) -> u64 { + let mut z = self.0; + z = (z ^ (z >> 30)).wrapping_mul(0xbf58_476d_1ce4_e5b9); + z = (z ^ (z >> 27)).wrapping_mul(0x94d0_49bb_1331_11eb); + z ^ (z >> 31) } } @@ -273,4 +422,139 @@ mod tests { assert_eq!(b.hold_from("别的 sk"), "别的 sk".len() - 2); assert_eq!(b.hold_from("什么都不像"), "什么都不像".len()); } + + /// 一本账里装着这些值(按出现的先后发号) + fn bridge_with(values: &[&str]) -> Bridge { + use tw_guard::redact::rules::{Hit, Rule}; + let text = values.join("\u{1}"); + let mut hits = Vec::new(); + let mut at = 0; + for v in values { + hits.push(Hit { + bytes: at..at + v.len(), + rule: Rule::Builtin("aws-access-key-id"), + label: None, + }); + at += v.len() + 1; + } + let ledger = + tw_guard::redact::replace::apply(&text, &hits, Ledger::new(Scheme::SECRET)).ledger; + Bridge::new(Arc::new(RuleSet::none())).with_ledger(ledger) + } + + fn rng(mut seed: u64) -> impl FnMut(usize) -> usize { + move |n| { + seed ^= seed << 13; + seed ^= seed >> 7; + seed ^= seed << 17; + (seed % n as u64) as usize + } + } + + /// 一遍找出每一处再挑,和原来一个值一个值地整段 `replace`,换出来的一模一样:值互为 + /// 前缀、后缀、首尾相叠、自己和自己相叠、多字节的。值里没有占位符里会有的字(大写、 + /// 数字、`<>`),一个个换时不会在刚换上的占位符里再找到东西 + #[test] + fn hiding_in_one_pass_writes_what_replacing_value_by_value_wrote() { + let values = [ + "a", "ab", "ba", "aba", "abab", "bab", "abc", "ca", "中", "中文", "文中", "文a", "c中", + ]; + let b = bridge_with(&values); + let one_by_one = |s: &str| { + let mut out = None::; + for (original, placeholder) in &b.values { + let cur = out.as_deref().unwrap_or(s); + if cur.contains(original.as_str()) { + out = Some(cur.replace(original.as_str(), placeholder)); + } + } + out.unwrap_or_else(|| s.to_string()) + }; + let alphabet = ["a", "b", "c", "中", "文", " ", "xyz"]; + let mut next = rng(0x2545_f491_4f6c_dd1d); + for _ in 0..3000 { + let s: String = (0..next(30)) + .map(|_| alphabet[next(alphabet.len())]) + .collect(); + assert_eq!(b.hide_known(&s), one_by_one(&s), "{s:?}"); + } + } + + /// 二分查排好序的值,和拿每个值的每个真前缀去比,扣住的位置一模一样 + #[test] + fn holding_back_agrees_with_trying_every_prefix_of_every_value() { + let values = ["abab", "abc", "中文字", "a中", "bca", "sk-ant-api03-x"]; + let b = bridge_with(&values); + let every = |buf: &str| { + let mut cut = buf.len(); + for (original, _) in &b.values { + let mut ends: Vec = original + .char_indices() + .map(|(i, _)| i) + .filter(|i| *i > 0) + .collect(); + ends.reverse(); + for k in ends { + if k <= buf.len() && buf.ends_with(&original[..k]) { + cut = cut.min(buf.len() - k); + break; + } + } + } + cut + }; + let alphabet = [ + "a", "b", "c", "中", "文", "字", " ", "sk-", "ant-", "api03-", + ]; + let mut next = rng(0x9e37_79b9_7f4a_7c15); + for _ in 0..3000 { + let buf: String = (0..next(12)) + .map(|_| alphabet[next(alphabet.len())]) + .collect(); + assert_eq!(b.hold_from(&buf), every(&buf), "{buf:?}"); + } + assert_eq!(Bridge::new(rules()).hold_from("abc"), 3, "空账什么都不扣"); + } + + /// 账里几万个值时,插件看到的请求照样换得过来:花的时间跟着请求的大小走,不跟着 + /// 「值的个数 × 字符串的个数」走 + #[test] + fn a_request_with_tens_of_thousands_of_keys_is_hidden_in_linear_time() { + let key = |i: usize| { + let mut n = i; + let tail: String = (0..16) + .map(|_| { + let c = char::from(b'A' + (n % 26) as u8); + n /= 26; + c + }) + .collect(); + format!("AKIA{tail}") + }; + let round = |n: usize| { + let messages: Vec = (0..n) + .map(|i| serde_json::json!({"role": "user", "content": format!("key {}", key(i))})) + .collect(); + let body = serde_json::json!({ "messages": messages }); + let started = std::time::Instant::now(); + let mut b = Bridge::new(rules()); + b.learn(body.to_string().as_bytes()); + let mut shown = body.clone(); + b.hide_value(&mut shown); + let text = shown.to_string(); + assert!(!text.contains("AKIA"), "插件看到了真值"); + assert!(text.contains(&format!("<>"))); + b.reveal_value(&mut shown); + assert_eq!(shown, body); + started.elapsed() + }; + let fastest = |n| (0..2).map(|_| round(n)).min().unwrap(); + let (small, large) = (fastest(5_000), fastest(20_000)); + let ratio = large.as_secs_f64() / small.as_secs_f64(); + // 线性的是 4 倍上下;一个值一个值地找是 16 倍 + assert!( + ratio < 10.0, + "2 万个值花了 5 千个值的 {ratio:.1} 倍({small:?} → {large:?})" + ); + } } diff --git a/crates/tw-gateway/src/server/pipeline.rs b/crates/tw-gateway/src/server/pipeline.rs index d1b6cdd0..2866217a 100644 --- a/crates/tw-gateway/src/server/pipeline.rs +++ b/crates/tw-gateway/src/server/pipeline.rs @@ -657,7 +657,7 @@ fn start( id, provider: alive.first().cloned().unwrap_or_default(), replaced: redact_mode.acts(), - items: crate::guard::items(&found), + items: crate::guard::items(&found, 0), at_ms: now_ms(), }); } diff --git a/crates/tw-gateway/src/server/pipeline/plug.rs b/crates/tw-gateway/src/server/pipeline/plug.rs index c1e1ba6f..e74a3856 100644 --- a/crates/tw-gateway/src/server/pipeline/plug.rs +++ b/crates/tw-gateway/src/server/pipeline/plug.rs @@ -110,12 +110,14 @@ pub(super) async fn attempt( .map_or_else(|| started.ledger.clone(), |b| b.ledger().clone()); let (found, ledger) = crate::guard::look_from(mode, &rt.redact, &body, seed); let more = crate::guard::more_found(&started.found, found); - if !more.is_empty() { + // 和开头那一条加起来,一个请求报的有上限(见 `crate::guard::REPORTED_MAX`) + let items = crate::guard::items(&more, started.found.len()); + if !items.is_empty() { state.bus.emit(tw_api::Event::SecretsFound { id: started.id, provider: provider.name.clone(), replaced: mode.acts(), - items: crate::guard::items(&more), + items, at_ms: crate::server::now_ms(), }); } diff --git a/crates/tw-gateway/src/ws.rs b/crates/tw-gateway/src/ws.rs index 1d8c418a..82facdf3 100644 --- a/crates/tw-gateway/src/ws.rs +++ b/crates/tw-gateway/src/ws.rs @@ -401,11 +401,13 @@ async fn pump( if found.is_empty() { UpMsg::Text(text.into()) } else { + // 客户端发来的一帧是一次请求,各报各的(一次最多报几个见 + // `crate::guard::REPORTED_MAX`) state.bus.emit(tw_api::Event::SecretsFound { id: p.id, provider: p.provider.clone(), replaced: mode.acts(), - items: crate::guard::items(&found), + items: crate::guard::items(&found, 0), at_ms: crate::server::now_ms(), }); if mode.acts() { diff --git a/crates/tw-gateway/tests/m5_redact.rs b/crates/tw-gateway/tests/m5_redact.rs index 678c9c57..54b1bc29 100644 --- a/crates/tw-gateway/tests/m5_redact.rs +++ b/crates/tw-gateway/tests/m5_redact.rs @@ -339,6 +339,79 @@ async fn the_ui_is_told_what_was_replaced_without_being_told_the_value() { assert!(!dump.contains("USERSOWNKEY"), "事件里带出了原值:{dump}"); } +#[tokio::test] +async fn a_request_full_of_keys_is_reported_in_part_and_replaced_in_full() { + // 一份导出的凭据清单贴进了对话:150 把不同的 key。安全日志只记前 100 个(一个值一行, + // 不让一个请求刷满日志),可每一把都要换掉、回显里每一把都要还原 + let (up, seen) = start_upstream(false).await; + let cfg = Config { + version: 1, + listen: Listen::default(), + clients: vec![Client { + name: "claude-code".into(), + key: "tw-testkey".into(), + ..Default::default() + }], + providers: vec![provider("relay", up)], + security: Security { + redact: policy(SecurityMode::Enforce), + ..Default::default() + }, + ..Default::default() + }; + let state = tw_gateway::AppState::new(cfg).unwrap(); + let mut rx = state.bus.subscribe(); + let gw = tw_gateway::serve(state, ([127, 0, 0, 1], 0).into()) + .await + .unwrap(); + tokio::time::sleep(Duration::from_millis(50)).await; + + // 第 i 把:AKIA 后面是 i 的 26 进制,两两不同 + let keys: Vec = (0..150) + .map(|i| { + let mut n = i; + let tail: String = (0..16) + .map(|_| { + let c = char::from(b'A' + (n % 26) as u8); + n /= 26; + c + }) + .collect(); + format!("AKIA{tail}") + }) + .collect(); + let body = serde_json::json!({ + "model": "claude-sonnet-4-5", + "max_tokens": 64, + "messages": [{"role": "user", "content": keys.join(" ")}] + }) + .to_string(); + let got = ask(gw, &body, false).await; + + let sent = String::from_utf8(seen.lock().unwrap().clone()).unwrap(); + assert!(!sent.contains("AKIA"), "中转站看见了真 key:{sent}"); + assert!(sent.contains("<>"), "{sent}"); + let echoed: serde_json::Value = serde_json::from_str(&got).unwrap(); + assert_eq!(echoed["text"], keys.join(" "), "回显没全部还原"); + + let mut items = None; + while let Ok(Ok(ev)) = tokio::time::timeout(Duration::from_secs(3), rx.recv()).await { + if let tw_api::Event::SecretsFound { items: it, .. } = ev { + items = Some(it); + break; + } + } + let items = items.expect("没发脱敏事件"); + assert_eq!(items.len(), tw_gateway::guard::REPORTED_MAX); + // 报的是先出现的那些 + assert_eq!(items[0].masked, "AKIAA…AAAA"); + assert!( + items + .iter() + .all(|i| i.count == 1 && i.rule == "aws-access-key-id") + ); +} + #[tokio::test] async fn id_and_card_numbers_leave_as_named_placeholders_and_come_back_whole() { // 身份证号和卡号出厂就开着。占位符写明是哪一种;回显一个字符一帧,拼回来 diff --git a/crates/tw-guard/src/redact/flow.rs b/crates/tw-guard/src/redact/flow.rs index 684c2453..42e8a2eb 100644 --- a/crates/tw-guard/src/redact/flow.rs +++ b/crates/tw-guard/src/redact/flow.rs @@ -40,28 +40,51 @@ pub fn hits(text: &str, rules: &RuleSet) -> Vec { if hits.is_empty() { return hits; } - if text.contains(Scheme::SECRET.open) { - let ours = Scheme::SECRET.find_in(text); - hits.retain(|h| !ours.iter().any(|(at, _, _)| overlaps(at, &h.bytes))); - } + off_placeholders(text, &mut hits); // 大多数请求一处都不命中:载荷在哪儿等有了命中再找 if !hits.is_empty() { - let payloads = base64_payloads(text); - if !payloads.is_empty() { - hits.retain(|h| !payloads.iter().any(|p| overlaps(p, &h.bytes))); - } + outside(&mut hits, base64_payloads(text)); } hits } +/// 去掉压在我们自己的占位符上的命中(见 [`hits`])。在解码过的正文上找的调用方(桌面版的 +/// 插件那一层)也用它。 +pub fn off_placeholders(text: &str, hits: &mut Vec) { + if !hits.is_empty() && text.contains(Scheme::SECRET.open) { + let ours = Scheme::SECRET.find_in(text); + outside(hits, ours.into_iter().map(|(at, _, _)| at).collect()); + } +} + /// 一段**纯文本**按它出现在请求体里时的样子找(同 [`hits`]),区间是原文里的。管理界面 /// 上的「测试」用它:用户贴进来的是一段正文,网关找的是装着它的 JSON —— 结论得一致。 pub fn hits_plain(text: &str, rules: &RuleSet) -> Vec { crate::redact::rules::plain_with(text, |encoded| hits(encoded, rules)) } -fn overlaps(a: &Range, b: &Range) -> bool { - a.start < b.end && b.start < a.end +/// 只留和 `spans` 里哪一段都不重叠的命中。 +/// +/// **不拿每个命中去和每一段比**:一个请求里命中几万处、占位符或载荷又有几万段时(重放 +/// 一份换过的请求),那是平方级的。`spans` 按起点排好,起点在命中终点之前的是一个前缀, +/// 前缀里最远的终点越过了命中的起点,就是有一段和它重叠 —— 每个命中二分查一次。 +fn outside(hits: &mut Vec, mut spans: Vec>) { + if spans.is_empty() { + return; + } + spans.sort_unstable_by_key(|s| s.start); + // reach[i]:前 i + 1 段里最远的终点 + let reach: Vec = spans + .iter() + .scan(0, |far, s| { + *far = s.end.max(*far); + Some(*far) + }) + .collect(); + hits.retain(|h| { + let before = spans.partition_point(|s| s.start < h.bytes.end); + before == 0 || reach[before - 1] <= h.bytes.start + }); } /// 找一遍。**观察档和替换档都找**,关闭时不找。 @@ -91,6 +114,9 @@ pub fn ledger_for(body: &[u8]) -> Ledger { /// 看一遍客户端发来的原文:报出去的记录(同 [`find`]),和这个请求的账本。 /// +/// 记录是**一个不同的值一条、一条不少**(见 [`crate::redact::rules::findings`]):报几条、 +/// 怎么聚合由调用方定,总共有几个不同的值就是它的长度。 +/// /// **替换档下账本在这里就编好号**:原文里找到的每个值按出现的先后发号,让开原文里本来 /// 就写着的占位符。之后每一跳都接着这本账换([`replace`]),存下来的那份请求也照它换。 /// 不在替换档时账本是空的。 @@ -471,6 +497,50 @@ mod tests { assert_eq!(hits(&body, &RuleSet::defaults()).len(), 1); } + /// 排序加二分,和拿每个命中去和每一段比,留下的一模一样:段可以乱序、套着、挨着、 + /// 是空的 + #[test] + fn keeping_hits_outside_the_spans_agrees_with_checking_every_pair() { + use crate::redact::rules::Rule; + let mut seed = 0x9e37_79b9_7f4a_7c15_u64; + let mut next = |n: usize| { + seed ^= seed << 13; + seed ^= seed >> 7; + seed ^= seed << 17; + (seed % n as u64) as usize + }; + for _ in 0..300 { + // 命中:排好序、互不重叠(scan 给的就是这样) + let mut hits = Vec::new(); + let mut at = 0; + for _ in 0..next(40) { + let start = at + next(5); + let end = start + 1 + next(6); + hits.push(Hit { + bytes: start..end, + rule: Rule::Builtin("aws-access-key-id"), + label: None, + }); + at = end; + } + let spans: Vec> = (0..next(30)) + .map(|_| { + let start = next(at + 5); + start..start + next(12) + }) + .collect(); + let mut pairwise = hits.clone(); + pairwise.retain(|h| { + !spans + .iter() + .any(|s| s.start < h.bytes.end && h.bytes.start < s.end) + }); + let mut kept = hits; + outside(&mut kept, spans.clone()); + assert_eq!(kept, pairwise, "{spans:?}"); + } + } + #[test] fn payloads_are_found_where_json_puts_them() { let long = "QUJD".repeat(80); diff --git a/crates/tw-guard/src/redact/replace.rs b/crates/tw-guard/src/redact/replace.rs index 49c245c8..af29e3d9 100644 --- a/crates/tw-guard/src/redact/replace.rs +++ b/crates/tw-guard/src/redact/replace.rs @@ -88,6 +88,9 @@ pub struct Ledger { seen: HashMap, /// 每个标签发到几号了 issued: HashMap, + /// 发出去的占位符有哪几种长度,从短到长。还原时在每个开头处按它们各查一次表(见 + /// [`restore`]);标签就那几个、号的位数也就几种,所以这里只有寥寥几个数 + lens: Vec, } impl Ledger { @@ -97,6 +100,7 @@ impl Ledger { back: HashMap::new(), seen: HashMap::new(), issued: HashMap::new(), + lens: Vec::new(), } } pub fn scheme(&self) -> Scheme { @@ -112,6 +116,10 @@ impl Ledger { pub fn table(&self) -> &HashMap { &self.back } + /// 发出去的占位符有哪几种长度,从短到长 + pub(crate) fn lens(&self) -> &[usize] { + &self.lens + } /// 原值 → 占位符。在一处找到的值要换到别处去时用它(企业版在解码后的 /// 请求上找,再换进原样转发的那一份)。 pub fn replacements(&self) -> impl Iterator { @@ -142,6 +150,9 @@ impl Ledger { let n = self.issued.entry(label.to_string()).or_insert(0); *n += 1; let p = self.scheme.placeholder(label, *n); + if let Err(at) = self.lens.binary_search(&p.len()) { + self.lens.insert(at, p.len()); + } self.back.insert(p.clone(), original.to_string()); self.seen.insert(original.to_string(), p.clone()); p @@ -163,9 +174,11 @@ pub struct Redacted { /// 同一个占位符,还原时必然给错一个。**那不是会不会发生的问题,是第二段只要 /// 命中一次就一定发生。** /// -/// **从后往前替换。**从前往后的话,第一次替换就会让后面所有区间的偏移 -/// 失效 —— 而那种错不会立刻炸,它会安静地切错一个字节,然后你拿到一份 -/// 坏掉的 JSON。 +/// `hits` 要**按起点排好、互不重叠**([`crate::redact::rules::scan`] 给的就是):区间指的 +/// 都是 `text` 原文里的位置。 +/// +/// **一遍从前往后抄出新的一份**:原文的一段、一个占位符、原文的下一段……不在原文上就地 +/// 换 —— 就地换一处,后面整段都要挪一次,一个请求里换几万处时是平方级的。 pub fn apply(text: &str, hits: &[Hit], mut ledger: Ledger) -> Redacted { if hits.is_empty() { return Redacted { @@ -173,21 +186,24 @@ pub fn apply(text: &str, hits: &[Hit], mut ledger: Ledger) -> Redacted { ledger, }; } - // 编号按出现的先后发:从后往前换,但先从前往后把号发完 + debug_assert!( + hits.windows(2).all(|w| w[0].bytes.end <= w[1].bytes.start), + "hits must be sorted and disjoint" + ); + // 编号按出现的先后发 let default = ledger.scheme.label; - let placeholders: Vec = hits - .iter() - .map(|h| { - ledger.issue( - &text[h.bytes.clone()], - h.label.as_deref().unwrap_or(default), - ) - }) - .collect(); - let mut out = text.to_string(); - for (h, ph) in hits.iter().zip(&placeholders).rev() { - out.replace_range(h.bytes.clone(), ph); + let mut out = String::with_capacity(text.len()); + let mut at = 0; + for h in hits { + let ph = ledger.issue( + &text[h.bytes.clone()], + h.label.as_deref().unwrap_or(default), + ); + out.push_str(&text[at..h.bytes.start]); + out.push_str(&ph); + at = h.bytes.end; } + out.push_str(&text[at..]); Redacted { text: out, ledger } } @@ -222,15 +238,55 @@ pub fn restore_json(text: &str, ledger: &Ledger) -> String { } fn swap(text: &str, ledger: &Ledger, put: impl Fn(&str) -> String) -> String { - if ledger.is_empty() || !text.contains(ledger.scheme.open) { + if ledger.is_empty() { return text.to_string(); } - let mut out = text.to_string(); - for (ph, original) in ledger.table() { - if out.contains(ph.as_str()) { - out = out.replace(ph.as_str(), &put(original)); + let lookup = |ph: &str| ledger.back.get(ph).map(String::as_str); + swap_with(text, ledger.scheme.open, &ledger.lens, lookup, put) +} + +/// 一遍扫过去,把 `text` 里写着的占位符换回原值。`lens` 是占位符有哪几种长度(从短到长), +/// `lookup` 按占位符查原值。 +/// +/// **在每个占位符开头处按这几种长度各查一次表**,不拿账里的每个占位符去整段文字里找一遍: +/// 后者是「占位符个数 × 文字长度」,一个请求换了几万个值、回答又把它们念了一遍时,光 +/// 还原就要几秒。换回来的原值不再参与匹配。 +pub(crate) fn swap_with<'a>( + text: &str, + open: &str, + lens: &[usize], + lookup: impl Fn(&str) -> Option<&'a str>, + put: impl Fn(&str) -> String, +) -> String { + let mut out = String::new(); + // 抄到了哪儿、从哪儿接着找下一个开头 + let (mut done, mut from) = (0, 0); + while let Some(i) = text[from..].find(open) { + let at = from + i; + let found = lens.iter().find_map(|&n| { + let ph = text.get(at..at + n)?; + lookup(ph).map(|original| (n, original)) + }); + match found { + Some((n, original)) => { + // 头一处才备下整段的地方:满是 `<<` 却一个占位符都没有的(C++ 代码)原样抄一份 + if done == 0 { + out.reserve(text.len()); + } + out.push_str(&text[done..at]); + out.push_str(&put(original)); + done = at + n; + from = done; + } + // 开头那段是 ASCII,加一还落在字的边界上。`<<>` 里的占位符从第二个 + // `<` 起 + None => from = at + 1, } } + if done == 0 { + return text.to_string(); + } + out.push_str(&text[done..]); out } @@ -478,6 +534,147 @@ mod tests { assert_eq!(taken("<>"), "<>"); } + /// 一个确定的伪随机序列 + fn rng(mut seed: u64) -> impl FnMut(usize) -> usize { + move |n| { + seed ^= seed << 13; + seed ^= seed >> 7; + seed ^= seed << 17; + (seed % n as u64) as usize + } + } + + /// 原来的做法:先从前往后发号,再从后往前就地换 + fn apply_in_place(text: &str, hits: &[Hit], mut ledger: Ledger) -> Redacted { + let default = ledger.scheme.label; + let placeholders: Vec = hits + .iter() + .map(|h| { + ledger.issue( + &text[h.bytes.clone()], + h.label.as_deref().unwrap_or(default), + ) + }) + .collect(); + let mut out = text.to_string(); + for (h, ph) in hits.iter().zip(&placeholders).rev() { + out.replace_range(h.bytes.clone(), ph); + } + Redacted { text: out, ledger } + } + + /// 原来的做法:账里的每个占位符在整段里找一遍、换一遍 + fn swap_each(text: &str, ledger: &Ledger) -> String { + if ledger.is_empty() || !text.contains(ledger.scheme.open) { + return text.to_string(); + } + let mut out = text.to_string(); + for (ph, original) in ledger.table() { + if out.contains(ph.as_str()) { + out = out.replace(ph.as_str(), original); + } + } + out + } + + fn same_ledger(a: &Ledger, b: &Ledger) { + assert_eq!(a.back, b.back); + assert_eq!(a.seen, b.seen); + assert_eq!(a.issued, b.issued); + } + + /// 一遍抄出新的一份,和原来就地换,换出来的文字、发的号一模一样:有重复的值、带 + /// 自己标签的、接着一本已经发过号的账、多字节的字 + #[test] + fn copying_once_writes_what_replacing_in_place_wrote() { + let rules = all() + .with_labeled("email", r"[a-z]+@[a-z]+\.com", Some("TW_EMAIL")) + .unwrap(); + let pool = [ + KEY, + "ghp_AAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAA", + "AKIAQRSTUVWXYZ234567", + "10.0.0.7", + "a@b.com", + "c@d.com", + "postgres://u:pw@h/db", + "很长的中文", + "plain", + ]; + let mut next = rng(0x2545_f491_4f6c_dd1d); + for round in 0..50 { + let mut text = String::new(); + for _ in 0..next(60) { + text.push_str(pool[next(pool.len())]); + text.push_str(["", " ", ",", "\n"][next(4)]); + } + let hits = crate::redact::rules::scan(&text, &rules); + let seed = if round % 2 == 0 { + l() + } else { + // 接着一本发过号的账:同一个值复用旧号,新值往后编 + redact(&format!("{KEY} 10.0.0.9"), &all(), l()).ledger + }; + let want = apply_in_place(&text, &hits, seed.clone()); + let got = apply(&text, &hits, seed); + assert_eq!(got.text, want.text); + same_ledger(&got.ledger, &want.ledger); + } + } + + /// 一遍扫过去,和账里的每个占位符各找一遍,换回来的一模一样:占位符挨着、夹在 + /// `<` 里、只写了半截、模型自己编的、几种长度和标签混着 + #[test] + fn restoring_in_one_pass_puts_back_what_swapping_each_placeholder_did() { + let mut ledger = l(); + for i in 0..120 { + let label = ["TW_SECRET", "TW_ID_NUMBER", "TW_EMAIL"][i % 3]; + ledger.issue(&format!("value-{i}-密"), label); + } + let known: Vec = ledger.table().keys().cloned().collect(); + let noise = [ + "", + " ", + "中文", + "<", + "<<", + "<<<", + ">>", + "<>", + "<>", + "std::cout << x << y;", + "{{TW_SECRET_1}}", + ]; + let mut next = rng(0x9e37_79b9_7f4a_7c15); + for _ in 0..300 { + let mut text = String::new(); + for _ in 0..next(40) { + if next(2) == 0 { + text.push_str(&known[next(known.len())]); + } else { + text.push_str(noise[next(noise.len())]); + } + } + assert_eq!(restore(&text, &ledger), swap_each(&text, &ledger), "{text}"); + // 整包 JSON 的还原走同一遍 + let body = serde_json::json!({ "text": text }).to_string(); + let escaped = |s: &str| { + let q = serde_json::to_string(s).unwrap(); + q[1..q.len() - 1].to_string() + }; + let mut want = body.clone(); + for (ph, original) in ledger.table() { + want = want.replace(ph.as_str(), &escaped(original)); + } + assert_eq!(restore_json(&body, &ledger), want); + } + // 账是空的、文字里没有开头的:原样 + assert_eq!(restore("<>", &l()), "<>"); + assert_eq!(restore("没有占位符", &ledger), "没有占位符"); + } + #[test] fn decoded_text_is_matched_as_written_and_restored_into_json_escaped() { // 正文上的 `password="x"`:截在引号处的话,换下来的只是 `password=`, diff --git a/crates/tw-guard/src/redact/rules.rs b/crates/tw-guard/src/redact/rules.rs index f669aea4..a4ef6ec5 100644 --- a/crates/tw-guard/src/redact/rules.rs +++ b/crates/tw-guard/src/redact/rules.rs @@ -27,7 +27,8 @@ //! 用户还可以写自己的规则(正则)。它们和内置规则在同一遍里找、同一本账 //! 里换,于是「同一个值只占一个编号」这类纪律对它们同样成立。 -use std::collections::HashSet; +use std::collections::hash_map::Entry; +use std::collections::{HashMap, HashSet}; use std::ops::Range; use std::sync::Arc; @@ -1468,19 +1469,30 @@ pub struct Finding { /// 把命中合并成「哪条规则 × 哪个值 × 几次」。**同一个值只报一次**: /// 一把 key 在一个请求里出现三次,是一把 key,不是三把。 +/// +/// 按第一次出现的先后排,**一个不同的值一条、一条不少**:`len()` 就是这段文本里有几个 +/// 不同的值。逐条报出去的调用方自己定报几条(桌面网关一个请求报前 100 个,见 +/// `tw_gateway::guard::items`);按规则聚合的(企业版的审计)拿得到全部。 +/// +/// **合并按「规则 × 值」查表**,不在已经合过的里面一条条找:一个请求里贴进来几万个不同 +/// 的值时(一份导出的凭据清单),后者是平方级的,几十万个值要几分钟。 pub fn findings(text: &str, hits: &[Hit]) -> Vec { - let mut out: Vec<(Rule, &str, u64)> = Vec::new(); + let mut out: Vec<(&Rule, &str, u64)> = Vec::new(); + let mut seen: HashMap<(&Rule, &str), usize> = HashMap::new(); for h in hits { let value = &text[h.bytes.clone()]; - match out.iter_mut().find(|(r, v, _)| *r == h.rule && *v == value) { - Some((_, _, n)) => *n += 1, - None => out.push((h.rule.clone(), value, 1)), + match seen.entry((&h.rule, value)) { + Entry::Occupied(at) => out[*at.get()].2 += 1, + Entry::Vacant(slot) => { + slot.insert(out.len()); + out.push((&h.rule, value, 1)); + } } } out.into_iter() .map(|(rule, value, count)| Finding { - masked: masked(&rule, value), - rule, + masked: masked(rule, value), + rule: rule.clone(), count, }) .collect() @@ -1822,6 +1834,64 @@ mod tests { assert_eq!(f[1].masked, "10.0.0.1"); } + /// 查表合并和原来在合过的里面一条条找,合出来的一模一样:同样的条目、同样的先后、 + /// 同样的次数。同一个值被两条规则认出来的,是两条 + #[test] + fn findings_merge_exactly_as_the_one_by_one_search_did() { + fn one_by_one(text: &str, hits: &[Hit]) -> Vec { + let mut out: Vec<(Rule, &str, u64)> = Vec::new(); + for h in hits { + let value = &text[h.bytes.clone()]; + match out.iter_mut().find(|(r, v, _)| *r == h.rule && *v == value) { + Some((_, _, n)) => *n += 1, + None => out.push((h.rule.clone(), value, 1)), + } + } + out.into_iter() + .map(|(rule, value, count)| Finding { + masked: masked(&rule, value), + rule, + count, + }) + .collect() + } + let pool = [ + "sk-ant-api03-AAAAAAAAAAAAAAAAAAAAAAAAAA", + "ghp_BBBBBBBBBBBBBBBBBBBBBBBBBBBBBBBBBBBB", + "AKIAQRSTUVWXYZ234567", + "10.0.0.1", + "10.0.0.2", + "db.internal", + "postgres://app:hunter2@db/x", + "PRJ-12", + "PRJ-7", + ]; + let rules = all().with_custom("project", r"PRJ-\d+").unwrap(); + // 一个确定的伪随机序列:值有重复、先后打乱 + let mut seed = 0x2545_f491_4f6c_dd1d_u64; + let mut text = String::new(); + for _ in 0..400 { + seed ^= seed << 13; + seed ^= seed >> 7; + seed ^= seed << 17; + text.push_str(pool[(seed % pool.len() as u64) as usize]); + text.push_str(if seed % 3 == 0 { "," } else { " " }); + } + let mut hits = scan(&text, &rules); + // 同一个值换一条规则认:按「规则 × 值」合,两条 + hits.extend( + scan(&text, &RuleSet::only(&["aws-access-key-id"])) + .into_iter() + .map(|h| Hit { + rule: Rule::Custom(Arc::from("aws-again")), + ..h + }), + ); + let got = findings(&text, &hits); + assert_eq!(got, one_by_one(&text, &hits)); + assert!(got.len() >= pool.len(), "{got:?}"); + } + #[test] fn every_builtin_rule_describes_what_it_matches() { // 界面上的「匹配」一栏来自 `matcher`;它说的必须就是扫描用的判据。 diff --git a/crates/tw-guard/src/redact/sse.rs b/crates/tw-guard/src/redact/sse.rs index 6acaaeef..b73290b2 100644 --- a/crates/tw-guard/src/redact/sse.rs +++ b/crates/tw-guard/src/redact/sse.rs @@ -138,7 +138,7 @@ struct Open { /// 一条流上、按帧的还原器。 pub struct FrameRestorer { dialect: Dialect, - ledger: Ledger, + /// 整帧的一次性还原。各路的还原器从它分出去,共用一本账(见 [`Restorer::fresh`]) oneshot: Restorer, lanes: BTreeMap, } @@ -149,7 +149,6 @@ impl FrameRestorer { pub fn new(ledger: &Ledger, dialect: Dialect) -> Self { Self { dialect, - ledger: ledger.clone(), oneshot: Restorer::new(ledger), lanes: BTreeMap::new(), } @@ -157,7 +156,7 @@ impl FrameRestorer { /// 没东西要还原。**调用方据此整条短路。** pub fn is_noop(&self) -> bool { - self.ledger.is_empty() + self.oneshot.is_noop() } /// 改写一帧。帧本身就地改,返回要先于它发出去的补帧。 @@ -173,7 +172,7 @@ impl FrameRestorer { let fields = fields(self.dialect, v); for f in &fields { let open = self.lanes.entry(f.lane).or_insert_with(|| Open { - restorer: Restorer::new(&self.ledger), + restorer: self.oneshot.fresh(), last: Value::Null, }); open.last = v.clone(); diff --git a/crates/tw-guard/src/redact/stream.rs b/crates/tw-guard/src/redact/stream.rs index b01dc2a6..378291e3 100644 --- a/crates/tw-guard/src/redact/stream.rs +++ b/crates/tw-guard/src/redact/stream.rs @@ -20,6 +20,8 @@ //! 编出来的 `<>` 也原样过去。扣住的长度因此天然封顶在最长的那个 //! 占位符上 —— 几十个字节。 +use std::sync::Arc; + use crate::redact::replace::Ledger; /// 一条流上的还原器。 @@ -27,23 +29,84 @@ use crate::redact::replace::Ledger; /// **`process()` 的输出拼起来再接上 `flush()`,等于把整段内容一次性 /// 还原的结果。**顺序和内容都不变,唯一的差别是有些字节晚几毫秒发出去。 pub struct Restorer { - table: std::collections::HashMap, + book: Arc, /// 占位符开头的第一个字节。先按它筛,再比前缀 lead: u8, open: &'static str, + buffer: String, +} + +/// 还原要查的那本账,建一次、一条流里的几路共用(见 [`Restorer::fresh`])。 +/// +/// **按占位符排好序**:一段尾巴是不是某个占位符的前缀、一个占位符的原值是什么,都是二分 +/// 查一次,不拿账里的每个占位符挨个比 —— 一个请求换了几万个值时,后者在每个 chunk 上 +/// 都要比几万次。 +struct Book { + /// (占位符, 原值),按占位符的字节序排 + entries: Vec<(String, String)>, + /// 占位符有哪几种长度,从短到长 + lens: Vec, /// 最长的占位符有多长:扣住的永远不会比它更多 longest: usize, - buffer: String, +} + +impl Book { + fn of(ledger: &Ledger) -> Self { + let mut entries: Vec<(String, String)> = ledger + .table() + .iter() + .map(|(p, o)| (p.clone(), o.clone())) + .collect(); + entries.sort_unstable_by(|a, b| a.0.cmp(&b.0)); + let lens = ledger.lens().to_vec(); + Self { + longest: lens.last().copied().unwrap_or(0), + entries, + lens, + } + } + + /// 第一个不小于 `s` 的占位符在哪儿 + fn seek(&self, s: &[u8]) -> usize { + self.entries.partition_point(|(p, _)| p.as_bytes() < s) + } + + fn original(&self, placeholder: &str) -> Option<&str> { + let (p, o) = self.entries.get(self.seek(placeholder.as_bytes()))?; + (p == placeholder).then_some(o.as_str()) + } + + /// `tail` 是不是某个占位符的**严格**前缀。 + /// + /// 以 `tail` 开头的占位符在排好的序里连成一段,打头的就是第一个不小于 `tail` 的; + /// 它要是正好等于 `tail`(不是严格前缀),这一段里还有的话就是紧跟着的那个。所以 + /// 看两个就够了。 + fn viable(&self, tail: &[u8]) -> bool { + self.entries[self.seek(tail)..] + .iter() + .take(2) + .any(|(p, _)| p.len() > tail.len() && p.as_bytes().starts_with(tail)) + } } impl Restorer { pub fn new(ledger: &Ledger) -> Self { let open = ledger.scheme().open; Self { - table: ledger.table().clone(), + book: Arc::new(Book::of(ledger)), lead: open.as_bytes()[0], open, - longest: ledger.table().keys().map(String::len).max().unwrap_or(0), + buffer: String::new(), + } + } + + /// 同一本账、缓冲是空的另一个还原器。一条流里有好几路时各路一个,共用一份账 + /// (不各抄一份、各排一次序)。 + pub fn fresh(&self) -> Self { + Self { + book: Arc::clone(&self.book), + lead: self.lead, + open: self.open, buffer: String::new(), } } @@ -51,7 +114,7 @@ impl Restorer { /// 没东西要还原。**调用方据此整条短路** —— 没脱敏的请求不该为这个 /// 功能付任何延迟。 pub fn is_noop(&self) -> bool { - self.table.is_empty() + self.book.entries.is_empty() } /// 喂下一段,返回现在可以安全发出去的部分。 @@ -95,17 +158,9 @@ impl Restorer { /// 切点总落在占位符开头那个 ASCII 字节上,所以一定是字符边界。 fn hold_from(&self) -> usize { let buf = self.buffer.as_bytes(); - let window = buf.len().saturating_sub(self.longest); + let window = buf.len().saturating_sub(self.book.longest); for p in window..buf.len() { - if buf[p] != self.lead { - continue; - } - let tail = &buf[p..]; - if self - .table - .keys() - .any(|k| k.len() > tail.len() && k.as_bytes().starts_with(tail)) - { + if buf[p] == self.lead && self.book.viable(&buf[p..]) { return p; } } @@ -113,16 +168,16 @@ impl Restorer { } fn restore(&self, s: &str) -> String { - if !s.contains(self.open) { + if self.is_noop() { return s.to_string(); } - let mut out = s.to_string(); - for (ph, original) in &self.table { - if out.contains(ph.as_str()) { - out = out.replace(ph.as_str(), original); - } - } - out + crate::redact::replace::swap_with( + s, + self.open, + &self.book.lens, + |ph| self.book.original(ph), + str::to_string, + ) } } @@ -355,6 +410,88 @@ mod tests { assert_eq!(r.process("1>> 后"), format!("{KEY} 后")); } + /// 二分查排好序的账,和拿账里每个占位符挨个比,扣住的位置一模一样;切成多少段喂, + /// 拼起来都和一次性还原一样 + #[test] + fn holding_back_agrees_with_checking_every_placeholder() { + use crate::redact::rules::{Hit, Rule}; + // 一本几种标签、几种号长都有的账 + let mut text = String::new(); + let mut hits = Vec::new(); + for i in 0..150 { + let value = format!("value-{i}"); + hits.push(Hit { + bytes: text.len()..text.len() + value.len(), + rule: Rule::Builtin("aws-access-key-id"), + label: [None, Some("TW_ID_NUMBER"), Some("TW_EMAIL")][i % 3].map(Arc::from), + }); + text.push_str(&value); + text.push(' '); + } + let l = crate::redact::replace::apply(&text, &hits, Ledger::new(Scheme::SECRET)).ledger; + let table = l.table(); + let longest = table.keys().map(String::len).max().unwrap(); + let every = |buf: &str| { + let buf = buf.as_bytes(); + for p in buf.len().saturating_sub(longest)..buf.len() { + if buf[p] == b'<' + && table + .keys() + .any(|k| k.len() > buf.len() - p && k.as_bytes().starts_with(&buf[p..])) + { + return p; + } + } + buf.len() + }; + let keys: Vec<&String> = table.keys().collect(); + let mut seed = 0x2545_f491_4f6c_dd1d_u64; + let mut next = |n: usize| { + seed ^= seed << 13; + seed ^= seed >> 7; + seed ^= seed << 17; + (seed % n as u64) as usize + }; + let mut r = Restorer::new(&l); + for _ in 0..2000 { + // 一段话,尾巴上是某个占位符(或者别的什么)的一截 + let k = keys[next(keys.len())]; + let mut buf = ["", "前面 ", "a << b ", "中文"][next(4)].to_string(); + match next(4) { + 0 => buf.push_str(&k[..1 + next(k.len())]), + 1 => buf.push_str(&k[..next(k.len())]), + 2 => { + buf.push_str(["< buf.push_str(&format!("{k}{}", &k[..next(k.len())])), + } + r.buffer = buf.clone(); + assert_eq!(r.hold_from(), every(&buf), "{buf:?}"); + } + // 整段切碎了喂,拼起来和一次性还原一样 + let mut whole = String::new(); + for _ in 0..300 { + whole.push_str(keys[next(keys.len())]); + whole.push_str(["", " ", "<<", "中", "< String { + let mut tail = String::new(); + let mut n = i; + for _ in 0..16 { + tail.push(char::from(b'A' + (n % 26) as u8)); + n /= 26; + } + format!("AKIA{tail}") +} + +/// 一条用户消息,里面是 `n` 个不同的密钥 +fn request(n: usize) -> String { + let keys: Vec = (0..n).map(key).collect(); + json!({"messages": [{"role": "user", "content": keys.join(" ")}]}).to_string() +} + +/// 看一遍、换一遍、整段还原、切成小块流式还原、按 SSE 帧还原,每一步都核对结果 +fn round_trip(n: usize) -> Duration { + let rules = RuleSet::defaults(); + let body = request(n); + let started = Instant::now(); + + let (found, ledger) = flow::look(Mode::Enforce, &rules, body.as_bytes()); + assert_eq!(found.len(), n, "一个不同的值一条"); + assert_eq!(ledger.len(), n); + let (sent, ledger) = flow::replace(Mode::Enforce, &rules, body.clone().into(), &ledger); + let sent = String::from_utf8(sent.to_vec()).unwrap(); + assert!(!sent.contains("AKIA"), "每一个都换掉了"); + assert!(sent.contains(&format!("<>"))); + + // 回答把占位符原样念了一遍 + assert_eq!(replace::restore(&sent, &ledger), body); + let mut streamed = Restorer::new(&ledger); + let mut out = String::new(); + for piece in sent.as_bytes().chunks(64) { + out.push_str(&streamed.process(std::str::from_utf8(piece).unwrap())); + } + out.push_str(&streamed.flush()); + assert_eq!(out, body); + + // 同一段话按 Anthropic 的 SSE 一帧几十个字流回来 + let text = serde_json::from_str::(&sent).unwrap()["messages"][0]["content"] + .as_str() + .unwrap() + .to_string(); + let mut sse = SseRestorer::new(&ledger, Dialect::Anthropic); + let mut raw = Vec::new(); + for piece in text.as_bytes().chunks(48) { + let delta = json!({ + "type": "content_block_delta", + "index": 0, + "delta": {"type": "text_delta", "text": std::str::from_utf8(piece).unwrap()} + }); + raw.extend( + sse.process(format!("event: content_block_delta\ndata: {delta}\n\n").as_bytes()), + ); + } + raw.extend(sse.flush()); + let restored: String = String::from_utf8(raw) + .unwrap() + .lines() + .filter_map(|l| l.strip_prefix("data: ")) + .filter_map(|d| serde_json::from_str::(d).ok()) + .filter_map(|v| v["delta"]["text"].as_str().map(str::to_string)) + .collect(); + assert_eq!(restored, (0..n).map(key).collect::>().join(" ")); + + started.elapsed() +} + +#[test] +fn tens_of_thousands_of_distinct_values_cost_linear_time() { + let fastest = |n| (0..3).map(|_| round_trip(n)).min().unwrap(); + let small = fastest(10_000); + let large = fastest(40_000); + let ratio = large.as_secs_f64() / small.as_secs_f64(); + // 线性的是 4 倍上下;平方级的是 16 倍。留足余量,只挡住平方级的 + assert!( + ratio < 10.0, + "4 万个值花了 1 万个值的 {ratio:.1} 倍({small:?} → {large:?}):有一步不再是线性的" + ); +} diff --git a/release-notes/0.59.0.md b/release-notes/0.59.0.md index 7604b6ee..f09c0c57 100644 --- a/release-notes/0.59.0.md +++ b/release-notes/0.59.0.md @@ -52,3 +52,5 @@ When a save writes the file, the plugin file, its approved copy and its `sha256` **Default plugins.** `reply-language` and `wsl-paths` write `on_error` and their settings' `value` in their manifests, in the style core writes, so changing a setting in the app changes one line of the file. The record of what core offered follows data-only saves and approvals, so a default whose settings, scope or `on_error` were changed still counts as unchanged code. A later version replaces it and keeps what was written in the file: `on_error`, `match`, and the values of the settings it still declares with the same type. A default whose code was changed is still left alone, and a new version that asks for more permissions or request kinds still comes back turned off. **Windows.** `twcore.exe` no longer needs the Visual C++ Redistributable; the C runtime is linked statically. The `twcore.exe` of 0.58.0 imported `VCRUNTIME140.dll`, which comes with that package and not with Windows, so the gateway did not start on a machine without it. + +**Requests full of secrets.** Outbound redaction no longer slows down when one request holds many distinct values that look like secrets: finding, replacing and restoring them, and hiding them from plugins, now take time in proportion to the request's size instead of growing with the square of the number of values. The security log records at most the first 100 distinct values of a request, and every value is still replaced.