From 46f3d2ce02ad0544a2b4b885665803fe340d98ca Mon Sep 17 00:00:00 2001 From: fylorn <249551762+fylorn@users.noreply.github.com> Date: Mon, 5 Oct 2026 16:04:40 +0800 Subject: [PATCH 01/22] feat(routing): weights on load-balance group members A load-balance group could only take turns one by one, so an upstream with twice the capacity (or half the price) got the same share as the rest. A member can now be written as `{ name, weight }` (1-100); a plain name is weight 1 and is written back as a plain name, so existing configurations round-trip byte for byte, through the control plane's writer too. Weights on other group types, or out of range, are refused at load time (engine.group_weight_*) and when saving a group (control.group.weight_*). Distribution is nginx's smooth weighted round-robin, which interleaves 7:3 instead of sending seven in a row. It replaces the `seq`-based rotation (and with it EventBus::peek_id, its only user). The per-group state lives in the gateway (tw_gateway::balance) and reaches the engine through Facts, so ordering stays a pure function and the dry-run reads the same state without advancing it. The state is charged at decision time, under one lock from ordering to charging, so a burst of simultaneous requests spreads out. It charges the member that leads after conversation stickiness: a conversation that stays on A counts toward A, and new conversations make up the difference, so the long-run ratio holds. Members that cannot serve the request or are cooling down sit the round out and the rest share by weight; otherwise their turn would fall to whoever follows them in the group. A group whose members or weights change starts over. The overview gives every member's weight for load-balance groups (GroupView.weights), saving takes them (GroupInput.weights), and the dry-run shows each candidate's weight (DryRunCandidate.weight). Co-Authored-By: Claude Opus 5.5 --- crates/tw-api/msg-codes.txt | 5 + crates/tw-api/src/lib.rs | 12 +- crates/tw-config/src/refs.rs | 61 +++- crates/tw-config/src/validate.rs | 52 ++++ crates/tw-config/tests/manual.rs | 23 +- crates/tw-config/tests/manual/schema.rs | 36 ++- crates/tw-control/src/dryrun.rs | 42 +-- crates/tw-control/src/lib.rs | 3 +- crates/tw-control/src/routes.rs | 96 +++++- crates/tw-control/tests/dryrun.rs | 62 +++- crates/tw-control/tests/resources.rs | 2 +- crates/tw-control/tests/routes.rs | 81 ++++- crates/tw-engine/src/engine.rs | 201 +++++++++--- crates/tw-engine/src/lib.rs | 2 + crates/tw-engine/src/weighted.rs | 376 +++++++++++++++++++++++ crates/tw-gateway/src/balance.rs | 162 ++++++++++ crates/tw-gateway/src/latency.rs | 2 +- crates/tw-gateway/src/lib.rs | 1 + crates/tw-gateway/src/server/pipeline.rs | 55 ++-- crates/tw-gateway/src/state.rs | 41 +++ crates/tw-gateway/tests/affinity.rs | 43 ++- crates/tw-observe/src/bus.rs | 9 - docs/config.md | 35 ++- docs/config.zh-CN.md | 26 +- 24 files changed, 1303 insertions(+), 125 deletions(-) create mode 100644 crates/tw-engine/src/weighted.rs create mode 100644 crates/tw-gateway/src/balance.rs diff --git a/crates/tw-api/msg-codes.txt b/crates/tw-api/msg-codes.txt index 1531c479..3dc13378 100644 --- a/crates/tw-api/msg-codes.txt +++ b/crates/tw-api/msg-codes.txt @@ -131,6 +131,9 @@ control.group.name_is_upstream control.group.no_such_upstream control.group.preferred_not_member control.group.upstream_twice +control.group.weight_not_load_balance +control.group.weight_not_member +control.group.weight_out_of_range control.group_in_use control.header_no_value control.internal_error @@ -230,6 +233,8 @@ engine.duplicate_route engine.empty_group engine.group_unknown_upstream engine.group_upstream_twice +engine.group_weight_not_load_balance +engine.group_weight_out_of_range engine.no_action engine.no_match engine.phase_two_with_to diff --git a/crates/tw-api/src/lib.rs b/crates/tw-api/src/lib.rs index 73b7fc4c..a8f30b5a 100644 --- a/crates/tw-api/src/lib.rs +++ b/crates/tw-api/src/lib.rs @@ -2141,6 +2141,9 @@ pub struct GroupView { #[serde(default, skip_serializing_if = "Option::is_none")] pub selected: Option, pub providers: Vec, + /// `load-balance` 组每个成员的权重,**每个成员都在**,没写权重的是 1:新对话按这个 + /// 比例分。别的类型不用权重,是空的 + pub weights: std::collections::BTreeMap, } #[derive(Debug, Clone, Serialize, Deserialize)] @@ -3198,7 +3201,7 @@ slug_enum! { Fallback = "fallback", /// 用选中的那一个,不可用时按顺序 Select = "select", - /// 轮流 + /// 按成员的权重轮流([`GroupView::weights`]) LoadBalance = "load-balance", /// 选最快的 UrlTest = "url-test", @@ -3218,6 +3221,10 @@ pub struct GroupInput { /// `select` 组优先使用的成员 #[serde(default, skip_serializing_if = "Option::is_none")] pub selected: Option, + /// `load-balance` 组成员的权重,1 到 100。不给 = 都是 1;给了的话没写到的成员是 1。 + /// 别的类型只能不给、或者都是 1 + #[serde(default, skip_serializing_if = "Option::is_none")] + pub weights: Option>, } #[derive(Debug, Clone, Serialize, Deserialize)] @@ -4748,6 +4755,9 @@ pub struct DryRunCandidate { /// 改写了模型)、`pinned`(规则指定了这一家发什么模型)。一样时没有 #[serde(default, skip_serializing_if = "Option::is_none")] pub model_via: Option, + /// 经过的是 `load-balance` 组时,它在组里的权重(没写权重的是 1)。别的时候没有 + #[serde(default, skip_serializing_if = "Option::is_none")] + pub weight: Option, } /// 一个要转换格式的候选上游。 diff --git a/crates/tw-config/src/refs.rs b/crates/tw-config/src/refs.rs index 0e47e9f3..88601475 100644 --- a/crates/tw-config/src/refs.rs +++ b/crates/tw-config/src/refs.rs @@ -50,7 +50,7 @@ pub fn provider_refs(cfg: &Config, name: &str) -> Vec { } } for g in &cfg.groups { - if g.providers.iter().any(|p| p == name) || g.selected.as_deref() == Some(name) { + if g.has(name) || g.selected.as_deref() == Some(name) { out.push(ProviderRef::Group { group: g.name.clone(), }); @@ -145,13 +145,18 @@ pub fn rename_provider( } for (g, group) in cfg.groups.iter().enumerate() { let base = [Step::key("groups"), Step::Index(g)]; - if group.providers.iter().any(|p| p == old) { + if group.has(old) { + // 整个列表照写回的规矩重写:权重是 1 的写成名字,带权重的写成 `{name, weight}` + // —— 权重跟着成员走,不因改名丢掉 let value = Value::Sequence( group .providers .iter() - .map(|x| s(if x == old { new } else { x })) - .collect(), + .map(|m| { + let name = if m.name == old { new } else { &m.name }; + member_value(name, m.weight) + }) + .collect::>()?, ); let mut p = base.to_vec(); p.push(Step::key("providers")); @@ -166,6 +171,19 @@ pub fn rename_provider( Ok(out) } +/// 策略组的一个成员写进配置的样子:权重是 1 的是名字,否则是 `{name, weight}` +/// (和 [`tw_engine::Group`] 写回时一样) +fn member_value(name: &str, weight: u32) -> Result { + if weight == 1 { + return Ok(Value::String(name.to_string())); + } + serde_yaml_ng::to_value(tw_engine::Member { + name: name.to_string(), + weight, + }) + .map_err(|e| EditError::Unwritable(e.to_string())) +} + /// 一条规则:它在哪条路由里、叫什么。 #[derive(Debug, Clone, PartialEq)] pub struct RuleRef { @@ -545,7 +563,40 @@ routes: let after = cfg(&out); assert!(provider_refs(&after, "relay").is_empty()); assert_eq!(provider_refs(&after, "relay-hk").len(), 3); - assert_eq!(after.groups[0].providers, ["relay-hk", "官方"]); + assert_eq!(after.groups[0].names(), ["relay-hk", "官方"]); + } + + /// 改名的那一家带着权重:权重跟着走,别的成员照旧写成名字 + #[test] + fn a_renamed_member_keeps_its_weight() { + let text = CFG.replace( + " type: fallback\n providers: [relay, 官方]\n", + " type: load-balance\n providers: [{ name: relay, weight: 3 }, 官方]\n", + ); + let c = cfg(&text); + assert_eq!(c.groups[0].weight("relay"), 3); + let renamed = edit::upsert( + &text, + edit::PROVIDERS, + Some("relay"), + &serde_yaml_ng::from_str( + "name: relay-hk\nbase_url: https://relay.example\nkey: sk-a\nproxy: hk\n", + ) + .unwrap(), + ) + .unwrap(); + let out = rename_provider(&renamed, &c, "relay", "relay-hk").unwrap(); + let after = cfg(&out); + assert_eq!( + after.groups[0].providers, + [ + tw_engine::Member { + name: "relay-hk".into(), + weight: 3 + }, + tw_engine::Member::named("官方"), + ] + ); } /// 指定模型里写着这一家,也是引用它:删之前要说,改名时只改那一项的 `provider` diff --git a/crates/tw-config/src/validate.rs b/crates/tw-config/src/validate.rs index a4a05ffc..fc068b8f 100644 --- a/crates/tw-config/src/validate.rs +++ b/crates/tw-config/src/validate.rs @@ -814,6 +814,58 @@ groups: assert_eq!((m.arg("group"), m.arg("upstream")), ("pool", "typo")); } + /// 手写的权重:`load-balance` 收 1 到 100,别的类型写了不是 1 的权重、或者超出范围, + /// 加载时就拒绝 + #[test] + fn a_weight_is_checked_at_load_time() { + let text = |kind: &str, b: &str| { + format!( + "version: 1 +listen: + control: + key: c0ffee00c0ffee00c0ffee00c0ffee00c0ffee00c0ffee00c0ffee00c0ffee00 +clients: + - name: c + key: tw-k +providers: + - name: a + base_url: https://a.example + key: sk-a + - name: b + base_url: https://b.example + key: sk-b +groups: + - name: pool + type: {kind} + providers: + - a + - {b} +" + ) + }; + let ok = crate::try_parse(&text("load-balance", "{ name: b, weight: 7 }")).unwrap(); + assert_eq!(ok.groups[0].weight("b"), 7); + let code = |kind: &str, b: &str| crate::try_parse(&text(kind, b)).unwrap_err().msg().code; + assert_eq!( + code("fallback", "{ name: b, weight: 7 }"), + "engine.group_weight_not_load_balance" + ); + assert_eq!( + code("load-balance", "{ name: b, weight: 0 }"), + "engine.group_weight_out_of_range" + ); + assert_eq!( + code("load-balance", "{ name: b, weight: 101 }"), + "engine.group_weight_out_of_range" + ); + // 拼错的字段照常说是哪一个 + let m = crate::try_parse(&text("load-balance", "{ name: b, wieght: 7 }")) + .unwrap_err() + .msg(); + assert_eq!(m.code, "config.unknown_field", "{m:?}"); + assert_eq!(m.arg("field"), "groups[0].providers[1].wieght", "{m:?}"); + } + #[test] fn error_messages_say_what_to_do_next() { // 错误信息是降低使用难度最有效的杠杆。判据不是「说清 diff --git a/crates/tw-config/tests/manual.rs b/crates/tw-config/tests/manual.rs index a2966496..f6b65f4f 100644 --- a/crates/tw-config/tests/manual.rs +++ b/crates/tw-config/tests/manual.rs @@ -180,6 +180,8 @@ pub enum Kind { OneOrManyMap(T2), /// 一个字符串,或者一组对象(见那一节) StrOrObjs(&'static str), + /// 一个列表,每一项是字符串或者对象(见那一节) + StrsOrObjs(&'static str), } #[derive(Clone, Copy)] @@ -313,6 +315,11 @@ fn kind(k: &Kind, l: Lang) -> String { pick("string, or list of", "字符串,或对象列表,见"), link(p) ), + Kind::StrsOrObjs(p) => format!( + "{} {}", + pick("list of strings or", "列表,每项是字符串或对象,对象见"), + link(p) + ), } } @@ -357,14 +364,20 @@ pub fn render_table(s: &Section, l: Lang) -> String { fn check_section(s: &Section, all: &[Section], errs: &mut Vec) { for r in &s.rows { match r.kind { - Kind::Obj(p) | Kind::Objs(p) | Kind::ObjMap(_, p) | Kind::StrOrObjs(p) => { + Kind::Obj(p) + | Kind::Objs(p) + | Kind::ObjMap(_, p) + | Kind::StrOrObjs(p) + | Kind::StrsOrObjs(p) => { if !all.iter().any(|x| x.path == p) { errs.push(format!( "{}.{} points at section `{p}`, which is not declared", s.path, r.name )); } - if !matches!(r.def, Def::Section | Def::Is(_) | Def::Unset) { + // 一个列表可以是必填的(策略组的成员):它自己是值,不是一个有默认值的对象 + let list = matches!(r.kind, Kind::StrsOrObjs(_)) && matches!(r.def, Def::Required); + if !list && !matches!(r.def, Def::Section | Def::Is(_) | Def::Unset) { errs.push(format!( "{}.{} is an object; its default is the section's own", s.path, r.name @@ -389,7 +402,11 @@ fn check_section(s: &Section, all: &[Section], errs: &mut Vec) { // 指向一节还没进代码的对象的那一行,同样还没进代码:它由那一节的 // `Ty::Pending` 看着,这里不数它 let pending = |r: &Row| match r.kind { - Kind::Obj(p) | Kind::Objs(p) | Kind::ObjMap(_, p) | Kind::StrOrObjs(p) => all + Kind::Obj(p) + | Kind::Objs(p) + | Kind::ObjMap(_, p) + | Kind::StrOrObjs(p) + | Kind::StrsOrObjs(p) => all .iter() .any(|x| x.path == p && matches!(x.ty, Ty::Pending { .. })), _ => false, diff --git a/crates/tw-config/tests/manual/schema.rs b/crates/tw-config/tests/manual/schema.rs index 51169a99..9f3cd162 100644 --- a/crates/tw-config/tests/manual/schema.rs +++ b/crates/tw-config/tests/manual/schema.rs @@ -11,7 +11,7 @@ use super::{Def, Kind, Lang, Row, Section, T2}; use tw_config::proxy::ProxyAuth; use tw_config::*; use tw_engine::rule::When; -use tw_engine::{Group, GroupType, Pinned, RouteSet, Rule, SetAction}; +use tw_engine::{Group, GroupType, Member, Pinned, RouteSet, Rule, SetAction}; use tw_pricing::{PerMillion, PricingConfig, SheetDef}; const fn t(en: &'static str, zh: &'static str) -> T2 { @@ -1228,17 +1228,17 @@ pub fn sections() -> Vec
{ Kind::Enum(group_types), Def::Is("fallback"), t( - "`fallback`: the first healthy member, in order. `select`: the member named in `selected`. `load-balance`: take turns between new conversations. `url-test`: the fastest by measured time to first byte. `cheapest`: the lowest input price.", - "`fallback`:按顺序取第一个健康的。`select`:取 `selected` 指定的那个。`load-balance`:新对话轮流。`url-test`:按实测首字节时间取最快的。`cheapest`:取输入单价最低的。", + "`fallback`: the first healthy member, in order. `select`: the member named in `selected`. `load-balance`: new conversations take turns, in proportion to the members' weights. `url-test`: the fastest by measured time to first byte. `cheapest`: the lowest input price.", + "`fallback`:按顺序取第一个健康的。`select`:取 `selected` 指定的那个。`load-balance`:新对话按成员的权重轮流。`url-test`:按实测首字节时间取最快的。`cheapest`:取输入单价最低的。", ), ), row( "providers", - Kind::Strs, + Kind::StrsOrObjs("groups[].providers[]"), Def::Required, t( - "Member upstreams, by name; not groups. Each upstream appears once in a group.", - "成员上游的名字,不能是策略组。同一个上游在一个策略组中只出现一次。", + "Member upstreams, by name; not groups. Each upstream appears once in a group. In a `load-balance` group, a member can be written as `{name, weight}`.", + "成员上游的名字,不能是策略组。同一个上游在一个策略组中只出现一次。`load-balance` 组的成员可以写成 `{name, weight}`。", ), ), row( @@ -1252,6 +1252,30 @@ pub fn sections() -> Vec
{ ), ], }, + Section { + path: "groups[].providers[]", + ty: checked!(Member, "{name: a}"), + rows: vec![ + row( + "name", + Kind::Str, + Def::Required, + t( + "The upstream, by name. A member written as just its name has weight 1.", + "上游的名字。只写名字的成员权重为 1。", + ), + ), + row( + "weight", + Kind::Int, + Def::Is("1"), + t( + "The member's share of a `load-balance` group's requests, in proportion to the other members' weights. From 1 to 100. Other group types take no weight other than 1.", + "成员在 `load-balance` 组中分到的请求份额,与其他成员的权重成比例。取值 1 到 100。其他类型的策略组只能写 1。", + ), + ), + ], + }, Section { path: "routes[]", ty: checked!(RouteSet, "{name: r}"), diff --git a/crates/tw-control/src/dryrun.rs b/crates/tw-control/src/dryrun.rs index f86213b6..476f2a4d 100644 --- a/crates/tw-control/src/dryrun.rs +++ b/crates/tw-control/src/dryrun.rs @@ -13,33 +13,32 @@ use axum::{Json, extract::State, http::StatusCode}; /// 试算页存在的全部意义是「告诉你这条请求会走哪儿」,所以它**必须**用 /// 同一个函数、同一份数字 —— 各算各的话,两边迟早会不一样,而那时 /// 试算比没有更糟。`sent` 是每一家和发给它的模型名:比价按它算,和数据面一样。 +/// +/// `load-balance` 读的是数据面记着的那一份轮询状态,**只读不记** +/// ([`tw_gateway::balance::Balance::peek`]):试算说的是下一个新对话会排给谁, +/// 试算了几次都不该改变这个答案,也不该让数据面跳过谁。 fn order_like_the_data_plane( s: &crate::ControlState, engine: &tw_engine::Engine, d: &tw_engine::Decision, sent: &[(String, String)], ) -> Vec { - let Some(gname) = d.via_group.clone() else { - return d.candidates.clone(); - }; - let Some(kind) = engine - .groups() - .iter() - .find(|g| g.name == gname) - .map(|g| g.kind) + let Some(g) = d + .via_group + .as_deref() + .and_then(|n| engine.groups().iter().find(|g| g.name == n)) else { return d.candidates.clone(); }; - if !kind.needs_runtime() { + if !g.kind.needs_runtime() { return d.candidates.clone(); } let cfg = s.config(); - let facts = tw_engine::Facts { - seq: s.gateway.bus.peek_id(), - ttfb_ms: s.gateway.latency.snapshot(&d.candidates), - price: s.gateway.unit_prices(&cfg.providers, sent, &d.candidates), - }; - engine.order(Some(&gname), &d.candidates, &facts) + let current = (g.kind == tw_engine::GroupType::LoadBalance).then(|| s.gateway.balance.peek(g)); + let facts = s + .gateway + .group_facts(&cfg.providers, g.kind, &d.candidates, sent, current); + engine.order(Some(&g.name), &d.candidates, &facts) } use tw_engine::{Outcome, RequestFacts, RouteError}; @@ -264,14 +263,20 @@ pub async fn dry_run( .collect(); // **顺序要和数据面一样,否则试算就是在撒谎。**`load-balance` // / `url-test` / `cheapest` 的次序由运行时的数字定, - // 这里走的是同一个 `order`,喂的是同一份延迟表和价目表。 + // 这里走的是同一个 `order`,喂的是同一份延迟表、价目表和轮询状态。 // // 会话那一维**故意留空**:试算是「假设现在来一个请求」, // 而它属于哪次会话取决于请求正文,试算没有那个东西。 - // 于是它显示的是轮转序列里的当前位置 —— 而那正是一个没有 - // 会话指纹的请求真的会走的路。 + // 于是它显示的是轮询此刻轮到的位置 —— 而那正是一个新对话 + // 真的会走的路。 out.candidates = order_like_the_data_plane(&s, engine, &d, &tw_gateway::sent::pairs(&sent)); + // `load-balance` 的候选带上各自的权重:排头的为什么是它,一半在这个数里 + let balanced = d + .via_group + .as_deref() + .and_then(|g| engine.groups().iter().find(|x| x.name == g)) + .filter(|g| g.kind == tw_engine::GroupType::LoadBalance); // 每一家收到的模型名,和为什么不是请求里写的那个 out.candidate_models = out .candidates @@ -284,6 +289,7 @@ pub async fn dry_run( .and_then(|x| x.model.clone().ok()) .filter(|m| !m.is_empty()), model_via: one.and_then(|x| x.via).map(|v| v.slug().to_string()), + weight: balanced.map(|g| g.weight(name)), } }) .collect(); diff --git a/crates/tw-control/src/lib.rs b/crates/tw-control/src/lib.rs index 168740d4..d23365e5 100644 --- a/crates/tw-control/src/lib.rs +++ b/crates/tw-control/src/lib.rs @@ -367,7 +367,8 @@ async fn overview(State(s): State) -> Json { builtin: tw_engine::is_builtin_group(&g.name), kind: routes::group_kind(g.kind), selected: g.selected.clone(), - providers: g.providers.clone(), + providers: g.names(), + weights: routes::group_weights(g), }) .collect(), // 和 `GET /keys` 同一份视图:概览里少一个字段的话,两处会各自 diff --git a/crates/tw-control/src/routes.rs b/crates/tw-control/src/routes.rs index 1c7655ef..769d193a 100644 --- a/crates/tw-control/src/routes.rs +++ b/crates/tw-control/src/routes.rs @@ -575,6 +575,18 @@ fn group_type(k: tw_api::GroupKind) -> GroupType { } } +/// 界面看到的权重:`load-balance` 组每个成员都在(没写的是 1),别的类型是空的 +pub(crate) fn group_weights(g: &Group) -> std::collections::BTreeMap { + match g.kind { + GroupType::LoadBalance => g + .providers + .iter() + .map(|m| (m.name.clone(), m.weight)) + .collect(), + _ => Default::default(), + } +} + fn to_group(input: &tw_api::GroupInput, cfg: &tw_config::Config) -> Result { let name = checked_name(&input.name, "group")?; reserved(&name, "group")?; @@ -609,6 +621,33 @@ fn to_group(input: &tw_api::GroupInput, cfg: &tw_config::Config) -> Result + "a weight is given for `{upstream}`, which is not a member" + )); + } + if *w != 1 && kind != GroupType::LoadBalance { + return Err(msg!( + "control.group.weight_not_load_balance", upstream = p, weight = w => + "upstream `{upstream}` has a weight of {weight}; only a load-balance group uses \ + weights" + )); + } + if !(tw_engine::weighted::WEIGHT_MIN..=tw_engine::weighted::WEIGHT_MAX).contains(w) { + return Err(msg!( + "control.group.weight_out_of_range", upstream = p, weight = w => + "upstream `{upstream}` has a weight of {weight}; a weight is a whole number from \ + 1 to 100" + )); + } + weights.insert(p.to_string(), *w); + } // 手动选择要有一个优先使用的成员;没选就是第一个。其余策略不写这一项 let selected = match kind { GroupType::Select => { @@ -632,7 +671,13 @@ fn to_group(input: &tw_api::GroupInput, cfg: &tw_config::Config) -> Result| tw_api::GroupInput { + name: "g".into(), + kind: tw_api::GroupKind::from_slug(kind).unwrap(), + providers: vec!["a".into(), "b".into()], + selected: None, + weights: w.map(|w| w.iter().map(|(p, w)| (p.to_string(), *w)).collect()), + }; + let g = to_group(&input("load-balance", Some(&[("b", 7)])), &c).unwrap(); + assert_eq!((g.weight("a"), g.weight("b")), (1, 7)); + assert_eq!( + group_weights(&g), + [("a".to_string(), 1), ("b".to_string(), 7)] + .into_iter() + .collect() + ); + let g = to_group(&input("load-balance", None), &c).unwrap(); + assert!(g.providers.iter().all(|m| m.weight == 1)); + let g = to_group(&input("fallback", Some(&[("a", 1)])), &c).unwrap(); + assert!(g.providers.iter().all(|m| m.weight == 1)); + assert!(group_weights(&g).is_empty(), "别的类型不给权重"); } } diff --git a/crates/tw-control/tests/dryrun.rs b/crates/tw-control/tests/dryrun.rs index cba3300e..19df3efa 100644 --- a/crates/tw-control/tests/dryrun.rs +++ b/crates/tw-control/tests/dryrun.rs @@ -49,11 +49,18 @@ fn app() -> (tempfile::TempDir, axum::Router) { } fn app_with(text: &str) -> (tempfile::TempDir, axum::Router) { + let (d, app, _) = app_and_gateway(text); + (d, app) +} + +/// [`app_with`],连同它的数据面:要看试算和数据面是不是读的同一份状态 +fn app_and_gateway(text: &str) -> (tempfile::TempDir, axum::Router, tw_gateway::AppState) { let d = tempfile::tempdir().unwrap(); let p = d.path().join("config.yaml"); std::fs::write(&p, text).unwrap(); let cfg: tw_config::Config = serde_yaml_ng::from_str(text).unwrap(); let gw = tw_gateway::AppState::new(cfg).unwrap(); + let gw_for_test = gw.clone(); let bus = gw.bus.clone(); let state = ControlState { shutdown: Default::default(), @@ -66,7 +73,7 @@ fn app_with(text: &str) -> (tempfile::TempDir, axum::Router) { chatgpt: Default::default(), zai: Default::default(), }; - (d, tw_control::router(state)) + (d, tw_control::router(state), gw_for_test) } /// 试算要说清按哪条路由算。没指定密钥、路由、草稿时,用唯一那把密钥 @@ -290,6 +297,59 @@ async fn the_dry_run_changes_nothing_and_sends_nothing() { ); } +/// `load-balance` 的试算读数据面记着的那一份轮询状态,**只读不记**:试算几次都还是同一家 +/// 排头;数据面排过之后,试算跟着变。每个候选带着它的权重 +#[tokio::test] +async fn a_weighted_group_is_tried_from_the_data_planes_state_without_moving_it() { + let text = CFG.replace( + " providers: [官方, 中转]\n", + " providers: [{ name: 官方, weight: 3 }, 中转]\n", + ); + let (_d, app, gw) = app_and_gateway(&text); + let weights = |r: &tw_api::DryRunResult| { + r.candidate_models + .iter() + .map(|c| (c.provider.clone(), c.weight)) + .collect::>() + }; + for _ in 0..5 { + let r = run(&app, r#"{"model":"claude-sonnet-4-5"}"#).await; + assert_eq!(r.candidates, ["官方", "中转"]); + assert_eq!( + weights(&r), + [("官方".to_string(), Some(3)), ("中转".to_string(), Some(1))] + ); + } + // 数据面排了两次,都排给了官方:3:1 的下一个是中转 + let rt = gw.runtime(); + let g = rt + .engine + .groups() + .iter() + .find(|g| g.name == "都试试") + .unwrap(); + let members = g.names(); + for _ in 0..2 { + let turn = gw.balance.turn(g); + let f = tw_engine::Facts { + current_weight: turn.current(), + ..Default::default() + }; + assert_eq!(rt.engine.order(Some(&g.name), &members, &f)[0], "官方"); + turn.charge(&members, &f, "官方"); + } + let before = gw.balance.peek(g); + for _ in 0..3 { + let r = run(&app, r#"{"model":"claude-sonnet-4-5"}"#).await; + assert_eq!(r.candidates, ["中转", "官方"], "排头的之后其余按组里的顺序"); + } + assert_eq!(gw.balance.peek(g), before, "试算不记账"); + + // 不经过 `load-balance` 的候选不带权重 + let r = run(&app, r#"{"model":"claude-sonnet-4-5","cache":true}"#).await; + assert_eq!(weights(&r), [("官方".to_string(), None)]); +} + #[tokio::test] async fn the_defaults_describe_an_ordinary_request() { // 用户只该改他关心的那一两个字段。 diff --git a/crates/tw-control/tests/resources.rs b/crates/tw-control/tests/resources.rs index 4da817a2..a0aeb332 100644 --- a/crates/tw-control/tests/resources.rs +++ b/crates/tw-control/tests/resources.rs @@ -303,7 +303,7 @@ async fn renaming_an_upstream_moves_its_references_in_the_same_version() { assert_eq!(st, StatusCode::OK, "{body}"); let cfg = b.parsed(); assert_eq!(cfg.providers[0].name, "anthropic"); - assert_eq!(cfg.groups[0].providers, ["anthropic"]); + assert_eq!(cfg.groups[0].names(), ["anthropic"]); // 一次保存一个版本:历史里只多了改之前的那一份 let history = tw_config::history::list(&b.dir.path().join("config.yaml")).unwrap(); assert_eq!( diff --git a/crates/tw-control/tests/routes.rs b/crates/tw-control/tests/routes.rs index 261dd3ac..16cee4a9 100644 --- a/crates/tw-control/tests/routes.rs +++ b/crates/tw-control/tests/routes.rs @@ -645,7 +645,86 @@ async fn a_new_group_writes_no_defaults() { .into_iter() .find(|g| g.name == "便宜优先") .unwrap(); - assert_eq!(g.providers, ["中转", "官方"]); + assert_eq!(g.names(), ["中转", "官方"]); + assert!(!file.contains("weight"), "{file}"); +} + +/// 概览给的组原样存回去:文件一个字节都不变 —— 没写过权重的组不会因为多了权重这一项 +/// 被改写 +#[tokio::test] +async fn saving_a_group_as_the_overview_gave_it_changes_nothing() { + for kind in ["select", "load-balance"] { + let yaml = BASE.replace(" type: select\n", &format!(" type: {kind}\n")); + let yaml = if kind == "select" { + yaml + } else { + yaml.replace(" selected: 官方\n", "") + }; + let b = bed(&yaml); + let group = find(&b.overview().await["groups"], "主力").clone(); + if kind == "load-balance" { + assert_eq!(group["weights"], json!({ "官方": 1, "中转": 1 }), "{group}"); + } else { + assert_eq!(group["weights"], json!({}), "{group}"); + } + let (st, v) = call( + &b.app, + "PUT", + &format!("/groups/{}", enc("主力")), + json!({ "group": group }), + ) + .await; + assert_eq!(st, StatusCode::OK, "{v}"); + assert_eq!(b.file(), yaml, "{kind}"); + } +} + +/// 给 `load-balance` 的成员设权重:写成 `{name, weight}`,权重 1 的照旧是名字;概览把 +/// 每个成员的权重都给回来,原样存回去不再改文件 +#[tokio::test] +async fn weights_are_written_on_the_members_that_have_them() { + let b = bed(BASE); + let (st, v) = call( + &b.app, + "PUT", + &format!("/groups/{}", enc("主力")), + json!({ "group": { "name": "主力", "kind": "load-balance", + "providers": ["官方", "中转"], "weights": { "中转": 3 } } }), + ) + .await; + assert_eq!(st, StatusCode::OK, "{v}"); + let g = b.parsed().groups.remove(0); + assert_eq!((g.weight("官方"), g.weight("中转")), (1, 3)); + let file = b.file(); + assert!(file.contains("- 官方\n"), "{file}"); + assert!( + file.contains("name: 中转") && file.contains("weight: 3"), + "{file}" + ); + let group = find(&b.overview().await["groups"], "主力").clone(); + assert_eq!(group["weights"], json!({ "官方": 1, "中转": 3 })); + let (st, v) = call( + &b.app, + "PUT", + &format!("/groups/{}", enc("主力")), + json!({ "group": group }), + ) + .await; + assert_eq!(st, StatusCode::OK, "{v}"); + assert_eq!(b.file(), file); + + // 别的类型不收权重,写错了说出是哪一条 + let (st, v) = call( + &b.app, + "PUT", + &format!("/groups/{}", enc("主力")), + json!({ "group": { "name": "主力", "kind": "fallback", + "providers": ["官方", "中转"], "weights": { "中转": 3 } } }), + ) + .await; + assert_eq!(st, StatusCode::BAD_REQUEST, "{v}"); + assert_eq!(v["code"], "control.group.weight_not_load_balance", "{v}"); + assert_eq!(b.file(), file, "拒绝了就什么都没写"); } #[tokio::test] diff --git a/crates/tw-engine/src/engine.rs b/crates/tw-engine/src/engine.rs index 1bab7be6..50321b84 100644 --- a/crates/tw-engine/src/engine.rs +++ b/crates/tw-engine/src/engine.rs @@ -9,6 +9,7 @@ use tw_types::{Msg, msg}; use crate::facts::RequestFacts; use crate::rule::{MatchError, When}; +use crate::weighted::{Member, WEIGHT_MAX, WEIGHT_MIN}; /// 策略组类型。 /// @@ -25,7 +26,8 @@ pub enum GroupType { Fallback, /// 手动指定一个。缓存友好 Select, - /// 轮流。**轮的是新对话**:已经有人回答过、缓存还热着的对话留在那一家 + /// 按成员的权重轮流(平滑加权轮询,见 [`crate::weighted`]),权重都是 1 就是挨个轮。 + /// **轮的是新对话**:已经有人回答过、缓存还热着的对话留在那一家 LoadBalance, /// 选最快的。判据是**真实流量测出来的 TTFB**,样本不够时用启动时 /// 那次零成本的 L1 握手计时补。 @@ -69,8 +71,12 @@ impl GroupType { /// 和真实转发不一样的结果」。 #[derive(Debug, Clone, Default)] pub struct Facts { - /// 轮转的种子 —— 通常是请求序号 - pub seq: u64, + /// `load-balance` 每个成员此刻的「当前权重」:平滑加权轮询攒下的那个数(见 + /// [`crate::weighted`])。网关按组记着,排一次记一次账;试算只读不记。**缺席 = 0** + pub current_weight: std::collections::HashMap, + /// 此刻停着的上游:熔断着、失败之后冷却着。`load-balance` 这一轮不算它们 —— + /// 网关反正会跳过它们,轮到它们的那一次会落到组里排在后面的那一家头上 + pub paused: std::collections::HashSet, /// 每家的典型 TTFB(毫秒)。**缺席 = 样本不够**,不是「很快」 pub ttfb_ms: std::collections::HashMap, /// 每家跑这个模型的单价,(输入, 输出),微分/百万 token。 @@ -90,19 +96,41 @@ pub struct Group { pub name: String, #[serde(default, rename = "type")] pub kind: GroupType, - pub providers: Vec, + /// 成员:上游的名字,`load-balance` 还可以带权重(见 [`Member`])。权重是 1 的写成名字 + #[serde(with = "crate::weighted::members")] + pub providers: Vec, /// `select` 用:当前选中的那个 #[serde(default, skip_serializing_if = "Option::is_none")] pub selected: Option, } +impl Group { + /// 成员的名字,按组里的顺序 + pub fn names(&self) -> Vec { + self.providers.iter().map(|m| m.name.clone()).collect() + } + + /// `name` 是不是这个组的成员 + pub fn has(&self, name: &str) -> bool { + self.providers.iter().any(|m| m.name == name) + } + + /// 成员 `name` 的权重。不是成员的按 1 算 + pub fn weight(&self, name: &str) -> u32 { + self.providers + .iter() + .find(|m| m.name == name) + .map_or(1, |m| m.weight) + } +} + /// 按策略排序。**纯函数** —— 同样的输入永远给同样的顺序,试算页因此 /// 能如实预告数据面会怎么走。 pub fn order_by(g: &Group, members: &[String], f: &Facts) -> Vec { match g.kind { // 这两种的顺序在 `expand_group` 里就定好了 GroupType::Fallback | GroupType::Select => members.to_vec(), - GroupType::LoadBalance => rotate(members, f), + GroupType::LoadBalance => balance(g, members, f), GroupType::UrlTest => { // 有样本的按 TTFB 升序;没样本的保持原有相对次序排在后面。 // **`sort_by_key` 是稳定排序**,所以同速的两家不会每次换位 @@ -124,20 +152,18 @@ pub fn order_by(g: &Group, members: &[String], f: &Facts) -> Vec { } } -/// `load-balance` 的轮转:按 `seq` 轮到谁,谁排头。 +/// `load-balance`:按权重轮到的那一家排头(平滑加权轮询,见 [`crate::weighted`])。 /// -/// **选中的那一家排头,其余顺次跟上**,一个都不少 —— 故障转移还要用它们。 +/// **排头的之后,其余的按组里的顺序跟上**,一个都不少 —— 故障转移还要用它们。 /// 同一段对话不在这里粘:留在上次回答它的那一家由网关按对话记着,这里只管 /// 还没人回答过的。 -fn rotate(members: &[String], f: &Facts) -> Vec { - if members.is_empty() { - return Vec::new(); - } - let start = (f.seq % members.len() as u64) as usize; - let mut out = Vec::with_capacity(members.len()); - out.extend_from_slice(&members[start..]); - out.extend_from_slice(&members[..start]); - out +fn balance(g: &Group, members: &[String], f: &Facts) -> Vec { + let Some(first) = crate::weighted::lead(g, members, f) else { + return members.to_vec(); + }; + std::iter::once(first.to_string()) + .chain(members.iter().filter(|m| *m != first).cloned()) + .collect() } /// 改写请求参数。 @@ -559,6 +585,18 @@ pub enum RouteError { #[error("{}", self.msg())] GroupUnknownUpstream { group: String, provider: String }, #[error("{}", self.msg())] + GroupWeightNotLoadBalance { + group: String, + provider: String, + weight: u32, + }, + #[error("{}", self.msg())] + GroupWeightOutOfRange { + group: String, + provider: String, + weight: u32, + }, + #[error("{}", self.msg())] DuplicateRoute(String), #[error("{}", self.msg())] UnknownDefaultRoute(String), @@ -620,6 +658,26 @@ impl RouteError { "group `{group}` lists `{upstream}`, which is not an upstream. A group's members are \ upstreams, by name" ), + RouteError::GroupWeightNotLoadBalance { + group, + provider, + weight, + } => msg!( + "engine.group_weight_not_load_balance", group = group, upstream = provider, + weight = weight => + "group `{group}` gives upstream `{upstream}` a weight of {weight}, and only a \ + load-balance group uses weights. Remove the weight, or make the group load-balance" + ), + RouteError::GroupWeightOutOfRange { + group, + provider, + weight, + } => msg!( + "engine.group_weight_out_of_range", group = group, upstream = provider, + weight = weight => + "group `{group}` gives upstream `{upstream}` a weight of {weight}. A weight is a \ + whole number from 1 to 100" + ), RouteError::DuplicateRoute(route) => msg!( "engine.duplicate_route", route = route => "there is more than one route named `{route}`. A gateway key binds to a route by \ @@ -752,7 +810,7 @@ impl Engine { groups.push(Group { name: ALL_UPSTREAMS.to_string(), kind: GroupType::Fallback, - providers: providers.clone(), + providers: providers.iter().map(Member::named).collect(), selected: None, }); } @@ -894,7 +952,8 @@ impl Engine { // 还会多轮到它几次 —— 一个没人写下、也看不出来的权重。控制面保存时就拦着 // (`control.group.upstream_twice`),手写的配置在这里拦 let mut members = std::collections::HashSet::new(); - for p in &g.providers { + for m in &g.providers { + let p = &m.name; // 不认识的名字(拼错了、或者写了另一个组):候选里它对不上任何上游, // 这一位就静默地没了。控制面保存时拦着(`control.group.no_such_upstream`) if !self.providers.contains(p) { @@ -909,6 +968,22 @@ impl Engine { provider: p.clone(), }); } + // 权重只有 `load-balance` 用:别的类型写了也不起作用,而写的人以为它起了。 + // 控制面保存时同样拦着(`control.group.weight_*`) + if m.weight != 1 && g.kind != GroupType::LoadBalance { + return Err(RouteError::GroupWeightNotLoadBalance { + group: g.name.clone(), + provider: p.clone(), + weight: m.weight, + }); + } + if !(WEIGHT_MIN..=WEIGHT_MAX).contains(&m.weight) { + return Err(RouteError::GroupWeightOutOfRange { + group: g.name.clone(), + provider: p.clone(), + weight: m.weight, + }); + } } } for set in &self.sets { @@ -1194,29 +1269,27 @@ impl Engine { fn expand_group(&self, g: &Group) -> Vec { match g.kind { // 顺序就是优先级。第一个健康的就用,缓存持续命中。 - GroupType::Fallback => g.providers.clone(), + GroupType::Fallback => g.names(), GroupType::Select => { // 选中的排头,其余仍然留着做故障转移 —— **手动选一家不 // 等于放弃容错**,那家挂了照样该切。 let mut out = Vec::with_capacity(g.providers.len()); if let Some(sel) = &g.selected - && g.providers.contains(sel) + && g.has(sel) { out.push(sel.clone()); } out.extend( g.providers .iter() - .filter(|p| Some(*p) != g.selected.as_ref()) - .cloned(), + .filter(|p| Some(&p.name) != g.selected.as_ref()) + .map(|p| p.name.clone()), ); out } // 这三种要运行时的数字才排得出来。**这里只给集合, // 顺序由 `order` 定** —— 它是纯函数,数据面和试算页都调它。 - GroupType::LoadBalance | GroupType::UrlTest | GroupType::Cheapest => { - g.providers.clone() - } + GroupType::LoadBalance | GroupType::UrlTest | GroupType::Cheapest => g.names(), } } @@ -1989,13 +2062,13 @@ mod tests { // 一样:引擎给出集合,而没有任何一层去转它,于是 6 个请求 6 次 // 落在第一家 —— 一个宣称做完了、实际什么都没做的功能。 let g = grp(GroupType::LoadBalance); + let members = g.names(); + let mut f = Facts::default(); let firsts: Vec = (0..6) - .map(|seq| { - let f = Facts { - seq, - ..Default::default() - }; - order_by(&g, &g.providers, &f)[0].clone() + .map(|_| { + let first = order_by(&g, &members, &f)[0].clone(); + f.current_weight = crate::weighted::advance(&g, &members, &f, &first); + first }) .collect(); assert_eq!(firsts, vec!["甲", "乙", "丙", "甲", "乙", "丙"]); @@ -2006,10 +2079,10 @@ mod tests { // 「轮到乙」不等于「甲和丙不要了」——那一家挂了还要能切 let g = grp(GroupType::LoadBalance); let f = Facts { - seq: 1, + current_weight: [("乙".to_string(), 1)].into_iter().collect(), ..Default::default() }; - assert_eq!(order_by(&g, &g.providers, &f), vec!["乙", "丙", "甲"]); + assert_eq!(order_by(&g, &g.names(), &f), vec!["乙", "甲", "丙"]); } #[test] @@ -2024,7 +2097,7 @@ mod tests { ttfb_ms: ttfb, ..Default::default() }; - assert_eq!(order_by(&g, &g.providers, &f), vec!["丙", "甲", "乙"]); + assert_eq!(order_by(&g, &g.names(), &f), vec!["丙", "甲", "乙"]); } #[test] @@ -2043,7 +2116,7 @@ mod tests { ..Default::default() }; for _ in 0..5 { - assert_eq!(order_by(&g, &g.providers, &f), vec!["甲", "乙", "丙"]); + assert_eq!(order_by(&g, &g.names(), &f), vec!["甲", "乙", "丙"]); } } @@ -2062,7 +2135,7 @@ mod tests { price, ..Default::default() }; - assert_eq!(order_by(&g, &g.providers, &f), vec!["丙", "甲", "乙"]); + assert_eq!(order_by(&g, &g.names(), &f), vec!["丙", "甲", "乙"]); } #[test] @@ -2081,7 +2154,7 @@ mod tests { price, ..Default::default() }; - assert_eq!(order_by(&g, &g.providers, &f), vec!["乙", "甲"]); + assert_eq!(order_by(&g, &g.names(), &f), vec!["乙", "甲"]); } #[test] @@ -2089,15 +2162,11 @@ mod tests { for kind in [GroupType::Fallback, GroupType::Select] { let g = grp(kind); let f = Facts { - seq: 7, + current_weight: [("丙".to_string(), 7)].into_iter().collect(), ttfb_ms: [("丙".to_string(), 1u32)].into_iter().collect(), ..Default::default() }; - assert_eq!( - order_by(&g, &g.providers, &f), - g.providers, - "{kind:?} 被重排了" - ); + assert_eq!(order_by(&g, &g.names(), &f), g.names(), "{kind:?} 被重排了"); assert!(!kind.needs_runtime()); } } @@ -2327,6 +2396,52 @@ mod builtin_tests { assert_eq!(m.code, "engine.group_unknown_upstream"); assert_eq!((m.arg("group"), m.arg("upstream")), ("pool", "typo")); } + + /// 权重只给 `load-balance`、只能是 1 到 100:写错了说出是哪个组、哪一家、写了多少 + #[test] + fn a_weight_outside_load_balance_or_out_of_range_is_rejected() { + let check = |yaml: &str| { + let g: Group = serde_yaml_ng::from_str(yaml).unwrap(); + Engine::with_default_rules( + vec!["a".into(), "b".into()], + vec![g], + vec![rule("兜底", "{}", "pool")], + ) + .validate() + }; + let lb = |a: u32| { + format!("name: pool\ntype: load-balance\nproviders: [{{name: a, weight: {a}}}, b]\n") + }; + assert_eq!(check(&lb(1)), Ok(())); + assert_eq!(check(&lb(100)), Ok(())); + for bad in [0, 101] { + let err = check(&lb(bad)).unwrap_err(); + assert_eq!( + err, + RouteError::GroupWeightOutOfRange { + group: "pool".into(), + provider: "a".into(), + weight: bad, + } + ); + let m = err.msg(); + assert_eq!(m.code, "engine.group_weight_out_of_range"); + assert_eq!(m.arg("weight"), bad.to_string()); + } + for kind in ["fallback", "select", "url-test", "cheapest"] { + let yaml = + format!("name: pool\ntype: {kind}\nproviders: [a, {{name: b, weight: 3}}]\n"); + let m = check(&yaml).unwrap_err().msg(); + assert_eq!(m.code, "engine.group_weight_not_load_balance", "{kind}"); + assert_eq!( + (m.arg("group"), m.arg("upstream"), m.arg("weight")), + ("pool", "b", "3") + ); + // 写成 1 的不算写了权重 + let one = format!("name: pool\ntype: {kind}\nproviders: [a, {{name: b, weight: 1}}]\n"); + assert_eq!(check(&one), Ok(()), "{kind}"); + } + } } #[cfg(test)] diff --git a/crates/tw-engine/src/lib.rs b/crates/tw-engine/src/lib.rs index 80eb1d09..baf67c5d 100644 --- a/crates/tw-engine/src/lib.rs +++ b/crates/tw-engine/src/lib.rs @@ -3,6 +3,7 @@ pub mod engine; pub mod facts; pub mod num; pub mod rule; +pub mod weighted; pub use catalog::{Catalog, ProviderModels}; pub use engine::{ @@ -11,3 +12,4 @@ pub use engine::{ SetAction, Target, has_catch_all, is_builtin_group, notes, order_by, scalar_name, }; pub use facts::{RequestFacts, estimate_strings, estimate_tokens}; +pub use weighted::Member; diff --git a/crates/tw-engine/src/weighted.rs b/crates/tw-engine/src/weighted.rs new file mode 100644 index 00000000..2e993bae --- /dev/null +++ b/crates/tw-engine/src/weighted.rs @@ -0,0 +1,376 @@ +//! 策略组的成员和 `load-balance` 的权重:成员怎么写,平滑加权轮询怎么轮。 +//! +//! 轮法是 nginx 的平滑加权轮询:每排一次,这一轮的每个成员给自己的「当前权重」加上 +//! 自己的权重,加完最大的那个排头,排头的再减去这一轮的总权重。7:3 排出来是 +//! 甲乙甲甲甲乙甲甲乙甲 —— 穿插着来,而不是先连着七次甲、再连着三次乙;权重都一样 +//! 就是挨个轮。 +//! +//! **当前权重不在这里**:它由网关记着(`tw_gateway::balance`),排序时经 +//! [`Facts::current_weight`] 传进来,排完、会话粘性也定了之后,网关按 [`advance`] +//! 记账。引擎因此还是纯函数:试算拿同一份状态算、不记账,说的就是数据面下一个新对话 +//! 会排给谁。 + +use std::collections::HashMap; + +use serde::{Deserialize, Serialize}; + +use crate::engine::{Facts, Group}; + +/// 权重最小是 1:0 等于把它从组里拿掉,那就该真的拿掉 +pub const WEIGHT_MIN: u32 = 1; + +/// 权重最大是 100:比例上够用,当前权重也不会大到溢出 +pub const WEIGHT_MAX: u32 = 100; + +/// 策略组的一个成员:上游的名字,和它在 `load-balance` 里的权重。 +/// +/// 配置里写成名字(权重 1),或者写成 `{ name, weight }`。**权重是 1 的写回去还是 +/// 名字**(见 [`members`]):没写过权重的配置,读进来再写回去一个字节都不变。 +/// +/// 只有 `load-balance` 用得上权重,别的类型写了不是 1 的权重,校验拒绝 +/// (`engine.group_weight_not_load_balance`)。 +#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)] +#[serde(deny_unknown_fields)] +pub struct Member { + /// 上游的名字 + pub name: String, + /// 新对话按权重的比例分给各个成员。不写是 1 + #[serde(default = "one")] + pub weight: u32, +} + +fn one() -> u32 { + 1 +} + +impl Member { + /// 只写了名字的成员:权重 1 + pub fn named(name: impl Into) -> Self { + Self { + name: name.into(), + weight: 1, + } + } +} + +impl From<&str> for Member { + fn from(name: &str) -> Self { + Self::named(name) + } +} + +impl From for Member { + fn from(name: String) -> Self { + Self::named(name) + } +} + +/// `providers` 列表的读写:一项是名字,或者 `{ name, weight }`。 +/// +/// 读的时候一个字符串是名字、一个映射是带权重的成员(和 serde 的 untagged 同一种写法), +/// 只是手写了:untagged 对写错的地方只会说「哪一种都不像」,而 `wieght` 这种拼错要说出 +/// 是哪个字段。写的时候权重 1 写成名字。 +pub(crate) mod members { + use serde::de::{self, Deserializer, MapAccess, Visitor}; + use serde::ser::{SerializeSeq, Serializer}; + use serde::{Deserialize, Serialize}; + + use super::Member; + + pub fn serialize(v: &[Member], s: S) -> Result { + let mut seq = s.serialize_seq(Some(v.len()))?; + for m in v { + if m.weight == 1 { + seq.serialize_element(&m.name)?; + } else { + seq.serialize_element(m)?; + } + } + seq.end() + } + + pub fn deserialize<'de, D: Deserializer<'de>>(d: D) -> Result, D::Error> { + Ok(Vec::::deserialize(d)? + .into_iter() + .map(|i| i.0) + .collect()) + } + + /// 列表里的一项 + struct Item(Member); + + impl Serialize for Item { + fn serialize(&self, s: S) -> Result { + self.0.serialize(s) + } + } + + impl<'de> Deserialize<'de> for Item { + fn deserialize>(d: D) -> Result { + struct V; + impl<'de> Visitor<'de> for V { + type Value = Item; + fn expecting(&self, f: &mut std::fmt::Formatter) -> std::fmt::Result { + f.write_str("the name of an upstream, or {name, weight}") + } + fn visit_str(self, v: &str) -> Result { + Ok(Item(Member::named(v))) + } + fn visit_string(self, v: String) -> Result { + Ok(Item(Member::named(v))) + } + // 没加引号、YAML 读成数或真假的名字(`[2024, b]`):照名字收下,和规则的 + // `to` 一样(见 [`crate::engine::scalar_name`])。这一项以前只读字符串, + // serde_yaml 给的是原文 —— 不收下的话,一份原来读得进的配置就读不进了 + fn visit_bool(self, v: bool) -> Result { + Ok(Item(Member::named(v.to_string()))) + } + fn visit_i64(self, v: i64) -> Result { + Ok(Item(Member::named(v.to_string()))) + } + fn visit_u64(self, v: u64) -> Result { + Ok(Item(Member::named(v.to_string()))) + } + fn visit_i128(self, v: i128) -> Result { + Ok(Item(Member::named(v.to_string()))) + } + fn visit_u128(self, v: u128) -> Result { + Ok(Item(Member::named(v.to_string()))) + } + fn visit_f64(self, v: f64) -> Result { + Ok(Item(Member::named(crate::engine::scalar_name(v)))) + } + fn visit_map>(self, map: A) -> Result { + // 照 `Member` 自己的规矩读:写错的字段名由 serde 说出来 + Member::deserialize(de::value::MapAccessDeserializer::new(map)).map(Item) + } + } + d.deserialize_any(V) + } + } +} + +/// 这一轮参加轮询的成员和它们的权重,按 `members` 的顺序。 +/// +/// **停着的不算**([`Facts::paused`]):熔断、冷却着的那一家这一次本来就会被跳过, +/// 轮到它的那一次落到组里排在它后面的那一家,那一家就平白多拿一份。全都停着时都算 +/// —— 那时网关照样一家家试(fail-open),排头的还是按权重来。 +fn round<'a>(g: &Group, members: &'a [String], f: &Facts) -> Vec<(&'a str, i64)> { + let up: Vec<&'a String> = members.iter().filter(|m| !f.paused.contains(*m)).collect(); + let pool = if up.is_empty() { + members.iter().collect() + } else { + up + }; + pool.into_iter() + .map(|m| (m.as_str(), i64::from(g.weight(m)))) + .collect() +} + +/// 下一个排头:当前权重加上自己的权重,最大的那个;一样大取组里靠前的。 +/// +/// `members` 是这次的候选(服务不了这个请求的已经去掉了):不在里面的成员这一轮不参加, +/// 剩下的按各自的权重分。 +pub fn lead<'a>(g: &Group, members: &'a [String], f: &Facts) -> Option<&'a str> { + let mut best: Option<(&str, i64)> = None; + for (m, w) in round(g, members, f) { + let v = f.current_weight.get(m).copied().unwrap_or(0) + w; + if best.is_none_or(|(_, b)| v > b) { + best = Some((m, v)); + } + } + best.map(|(m, _)| m) +} + +/// 记一次账,交回记过之后的当前权重:这一轮的每个成员加上自己的权重,`leader` 再减去 +/// 这一轮的总权重。 +/// +/// `leader` 是**会话粘性之后实际排头的那一家**,不一定是 [`lead`] 挑的那个:一段对话 +/// 留在了上次回答它的那一家,这一次就记在那一家头上,之后的新对话把差的补回去。 +/// `members` 和 `f` 要和排序时的一样(同一轮)。`leader` 不在这一轮里时什么都不记。 +pub fn advance(g: &Group, members: &[String], f: &Facts, leader: &str) -> HashMap { + let round = round(g, members, f); + let mut out = f.current_weight.clone(); + if !round.iter().any(|(m, _)| *m == leader) { + return out; + } + let total: i64 = round.iter().map(|(_, w)| w).sum(); + for (m, w) in &round { + *out.entry((*m).to_string()).or_insert(0) += w; + } + *out.entry(leader.to_string()).or_insert(0) -= total; + out +} + +#[cfg(test)] +mod tests { + use super::*; + use crate::engine::{GroupType, order_by}; + + fn group(members: &[(&str, u32)]) -> Group { + Group { + name: "池子".into(), + kind: GroupType::LoadBalance, + providers: members + .iter() + .map(|(n, w)| Member { + name: n.to_string(), + weight: *w, + }) + .collect(), + selected: None, + } + } + + fn names(g: &Group) -> Vec { + g.providers.iter().map(|m| m.name.clone()).collect() + } + + /// 像网关那样排 `n` 次、每次记在排头的那一家,交回每次的排头 + fn run(g: &Group, members: &[String], f: &mut Facts, n: usize) -> Vec { + (0..n) + .map(|_| { + let first = order_by(g, members, f)[0].clone(); + f.current_weight = advance(g, members, f, &first); + first + }) + .collect() + } + + /// 7:3 穿插着来,十次里正好七次甲、三次乙,然后从头再来一遍 + #[test] + fn seven_to_three_interleaves() { + let g = group(&[("甲", 7), ("乙", 3)]); + let mut f = Facts::default(); + let got = run(&g, &names(&g), &mut f, 20); + let want = ["甲", "乙", "甲", "甲", "甲", "乙", "甲", "甲", "乙", "甲"]; + assert_eq!(got[..10], want); + assert_eq!(got[10..], want, "十次之后回到起点"); + assert!(f.current_weight.values().all(|v| *v == 0), "{f:?}"); + } + + /// 权重都一样就是挨个轮 + #[test] + fn equal_weights_take_turns() { + for w in [1, 5] { + let g = group(&[("甲", w), ("乙", w), ("丙", w)]); + let mut f = Facts::default(); + assert_eq!( + run(&g, &names(&g), &mut f, 6), + ["甲", "乙", "丙", "甲", "乙", "丙"] + ); + } + } + + /// 排头的之后,其余的按组里的顺序跟着,一个都不少 —— 故障转移还要用它们 + #[test] + fn the_rest_follow_in_the_groups_order() { + let g = group(&[("甲", 1), ("乙", 1), ("丙", 1)]); + let f = Facts { + current_weight: [("乙".to_string(), 5)].into_iter().collect(), + ..Default::default() + }; + assert_eq!(order_by(&g, &names(&g), &f), ["乙", "甲", "丙"]); + } + + /// 不在候选里的成员(停用了、没有这个模型)这一轮不参加,当前权重也不动;剩下的 + /// 照各自的权重分 + #[test] + fn a_member_left_out_of_the_candidates_sits_the_round_out() { + let g = group(&[("甲", 3), ("乙", 1), ("丙", 2)]); + let without: Vec = vec!["甲".into(), "丙".into()]; + let mut f = Facts::default(); + let got = run(&g, &without, &mut f, 50); + assert_eq!(got.iter().filter(|m| *m == "甲").count(), 30); + assert_eq!(got.iter().filter(|m| *m == "丙").count(), 20); + assert_eq!(f.current_weight.get("乙"), None, "没参加的不记账"); + + // 已经轮了一阵之后才缺席:比例照样是剩下的那几家的 + let mut f = Facts::default(); + run(&g, &names(&g), &mut f, 7); + let before = f.current_weight.get("乙").copied(); + let got = run(&g, &without, &mut f, 100); + let a = got.iter().filter(|m| *m == "甲").count() as i64; + assert!((a - 60).abs() <= 2, "甲排头 {a} 次"); + assert_eq!(f.current_weight.get("乙").copied(), before); + } + + /// 停着的(熔断、冷却)同样不参加;全都停着时都参加 + #[test] + fn paused_members_sit_out_unless_every_member_is_paused() { + let g = group(&[("甲", 1), ("乙", 1), ("丙", 1)]); + let mut f = Facts { + paused: ["甲".to_string()].into_iter().collect(), + ..Default::default() + }; + assert_eq!( + run(&g, &names(&g), &mut f, 4), + ["乙", "丙", "乙", "丙"], + "停着的那一家轮不到,它那一份不会落到它后面那一家头上" + ); + let mut f = Facts { + paused: names(&g).into_iter().collect(), + ..Default::default() + }; + assert_eq!(run(&g, &names(&g), &mut f, 3), ["甲", "乙", "丙"]); + } + + /// 记在实际排头的那一家头上:粘性把甲留在了前面,下一个新对话轮到乙 + #[test] + fn the_member_that_actually_led_is_charged() { + let g = group(&[("甲", 1), ("乙", 1)]); + let members = names(&g); + let mut f = Facts::default(); + assert_eq!(run(&g, &members, &mut f, 1), ["甲"]); + // 按轮到的该是乙,一段对话留在了甲 + assert_eq!(order_by(&g, &members, &f)[0], "乙"); + f.current_weight = advance(&g, &members, &f, "甲"); + // 新对话:两次都是乙,把差的补回来 + assert_eq!(run(&g, &members, &mut f, 2), ["乙", "乙"]); + assert_eq!(run(&g, &members, &mut f, 2), ["甲", "乙"]); + // 不在这一轮里的名字什么都不记 + let before = f.current_weight.clone(); + assert_eq!(advance(&g, &members, &f, "别家"), before); + } + + /// 权重 1 写回去是名字,写了别的权重写回去是映射;数、真假照名字读 + #[test] + fn a_member_is_a_name_or_a_name_with_a_weight() { + let text = "name: 池子\ntype: load-balance\nproviders:\n- 甲\n- name: 乙\n weight: 3\n- name: 丙\n weight: 1\n- 2024\n"; + let g: Group = serde_yaml_ng::from_str(text).unwrap(); + assert_eq!( + g.providers + .iter() + .map(|m| (m.name.as_str(), m.weight)) + .collect::>(), + [("甲", 1), ("乙", 3), ("丙", 1), ("2024", 1)] + ); + assert_eq!(g.weight("乙"), 3); + assert_eq!(g.weight("没有这家"), 1); + let back = serde_yaml_ng::to_string(&g).unwrap(); + assert_eq!( + back, + "name: 池子\ntype: load-balance\nproviders:\n- 甲\n- name: 乙\n weight: 3\n- 丙\n- '2024'\n" + ); + // 只写名字的组,写回去和原来一样 + let plain = "name: 池子\ntype: fallback\nproviders:\n- 甲\n- 乙\n"; + let g: Group = serde_yaml_ng::from_str(plain).unwrap(); + assert_eq!(serde_yaml_ng::to_string(&g).unwrap(), plain); + } + + /// 拼错的字段说出是哪一个,不是一句「哪一种都不像」 + #[test] + fn a_misspelled_member_field_is_named() { + let e = serde_yaml_ng::from_str::( + "name: g\ntype: load-balance\nproviders: [a, {name: b, wieght: 3}]\n", + ) + .unwrap_err() + .to_string(); + assert!(e.contains("wieght"), "{e}"); + assert!(e.contains("providers"), "{e}"); + let e = serde_yaml_ng::from_str::("name: g\nproviders: [a, {weight: 3}]\n") + .unwrap_err() + .to_string(); + assert!(e.contains("name"), "{e}"); + } +} diff --git a/crates/tw-gateway/src/balance.rs b/crates/tw-gateway/src/balance.rs new file mode 100644 index 00000000..356a202b --- /dev/null +++ b/crates/tw-gateway/src/balance.rs @@ -0,0 +1,162 @@ +//! `load-balance` 组轮到了谁:每个组、每个成员此刻的「当前权重」(平滑加权轮询, +//! 轮法见 [`tw_engine::weighted`])。 +//! +//! - **做决定的那一刻就记账**,不等回答成功:同时进来的一批请求一个接一个地排,后一个 +//! 看到的是前一个记过账的样子,于是分散到各家,而不是都排给同一家。从排序到记账一直 +//! 拿着锁([`Turn`])。 +//! - **记在会话粘性之后实际排头的那一家头上**:一段对话留在上次回答它的那一家(见 +//! [`crate::affinity`]),按权重本该轮到的那一家这一次没轮上 —— 账记给留下的那一家, +//! 之后的新对话把差的补回去,长期看各家拿到的还是配置的比例。 +//! - **组的成员或权重变了就从头轮**:旧的当前权重是按旧的比例攒下的。 +//! +//! 只在内存里:core 重启之后从头轮,差的最多是一轮。**跨配置重载存活**,没改的组接着轮。 +//! 试算只读不记([`Balance::peek`]):它说的是下一个新对话会排给谁,说完了不该改变答案。 + +use std::collections::HashMap; +use std::sync::{Mutex, MutexGuard}; + +use tw_engine::{Facts, Group, Member}; + +/// 一个组轮到哪儿了 +struct Round { + /// 这份状态是按哪一组成员、哪一组权重攒的。和配置对不上就作废 + members: Vec, + /// 成员 → 当前权重。没有的是 0 + current: HashMap, +} + +/// 每个 `load-balance` 组轮到哪儿了,按组名记。 +#[derive(Default)] +pub struct Balance { + rounds: Mutex>, +} + +impl Balance { + fn lock(&self) -> MutexGuard<'_, HashMap> { + // 锁中毒了照样拿里面的表:轮得偏一点,好过一个不转发的网关 + self.rounds.lock().unwrap_or_else(|p| p.into_inner()) + } + + /// 开始给组 `g` 排一次。**拿到它就拿着锁**,直到 [`Turn::charge`] 记完账(或者丢掉它)。 + /// + /// 组的成员或权重和记着的不一样(配置改了)就从头轮。 + pub fn turn<'a>(&'a self, g: &'a Group) -> Turn<'a> { + let mut rounds = self.lock(); + if rounds.get(&g.name).is_none_or(|r| r.members != g.providers) { + rounds.insert( + g.name.clone(), + Round { + members: g.providers.clone(), + current: HashMap::new(), + }, + ); + } + Turn { rounds, group: g } + } + + /// 组 `g` 此刻的当前权重,**只看不记**。试算用:它和数据面读同一份,说完了不改变答案 + pub fn peek(&self, g: &Group) -> HashMap { + self.lock() + .get(&g.name) + .filter(|r| r.members == g.providers) + .map(|r| r.current.clone()) + .unwrap_or_default() + } +} + +/// 正在给一个组排的这一次:排序读它的当前权重,排完、粘性也定了之后记账。 +pub struct Turn<'a> { + rounds: MutexGuard<'a, HashMap>, + group: &'a Group, +} + +impl Turn<'_> { + /// 这个组此刻的当前权重:排序要的那一份([`Facts::current_weight`]) + pub fn current(&self) -> HashMap { + self.rounds + .get(&self.group.name) + .map(|r| r.current.clone()) + .unwrap_or_default() + } + + /// 记账:这一次排头的是 `leader`。`members` 和 `f` 是排序时的那一份候选和事实 —— + /// 粘性只换了次序,没换集合,这一轮还是同一轮 + pub fn charge(mut self, members: &[String], f: &Facts, leader: &str) { + let next = tw_engine::weighted::advance(self.group, members, f, leader); + if let Some(r) = self.rounds.get_mut(&self.group.name) { + r.current = next; + } + } +} + +#[cfg(test)] +mod tests { + use super::*; + + fn group(weights: &[(&str, u32)]) -> Group { + Group { + name: "池子".into(), + kind: tw_engine::GroupType::LoadBalance, + providers: weights + .iter() + .map(|(n, w)| Member { + name: n.to_string(), + weight: *w, + }) + .collect(), + selected: None, + } + } + + /// 像管线那样排一次:读状态、排序、记在排头的那一家头上 + fn once(b: &Balance, g: &Group) -> String { + let members = g.names(); + let t = b.turn(g); + let f = Facts { + current_weight: t.current(), + ..Default::default() + }; + let first = tw_engine::order_by(g, &members, &f)[0].clone(); + t.charge(&members, &f, &first); + first + } + + #[test] + fn it_remembers_where_each_group_is() { + let b = Balance::default(); + let g = group(&[("甲", 2), ("乙", 1)]); + let got: Vec = (0..6).map(|_| once(&b, &g)).collect(); + assert_eq!(got, ["甲", "乙", "甲", "甲", "乙", "甲"]); + } + + /// 只看不记:看多少次,下一个都还是同一家 + #[test] + fn peeking_does_not_move_it() { + let b = Balance::default(); + let g = group(&[("甲", 1), ("乙", 1)]); + once(&b, &g); + let seen = b.peek(&g); + for _ in 0..3 { + assert_eq!(b.peek(&g), seen); + } + assert_eq!(once(&b, &g), "乙"); + } + + /// 成员或权重变了:从头轮。没变的接着轮 + #[test] + fn a_changed_group_starts_over() { + let b = Balance::default(); + let g = group(&[("甲", 1), ("乙", 1)]); + assert_eq!(once(&b, &g), "甲"); + // 同一份配置重载了一次:接着轮 + assert_eq!(once(&b, &g.clone()), "乙"); + assert_eq!(once(&b, &g), "甲"); + // 权重改了:攒下的作废,试算也看不到旧的 + let reweighted = group(&[("甲", 1), ("乙", 3)]); + assert!(b.peek(&reweighted).is_empty()); + assert_eq!(once(&b, &reweighted), "乙"); + // 多了一个成员:同样从头轮 + let grown = group(&[("甲", 1), ("乙", 3), ("丙", 1)]); + assert!(b.peek(&grown).is_empty()); + } +} diff --git a/crates/tw-gateway/src/latency.rs b/crates/tw-gateway/src/latency.rs index 294e8ef5..35145702 100644 --- a/crates/tw-gateway/src/latency.rs +++ b/crates/tw-gateway/src/latency.rs @@ -110,7 +110,7 @@ pub async fn seed_url_test(state: &crate::state::AppState) { let mut want: Vec = Vec::new(); for g in rt.engine.groups() { if g.kind == tw_engine::GroupType::UrlTest { - want.extend(g.providers.iter().cloned()); + want.extend(g.names()); } } want.sort(); diff --git a/crates/tw-gateway/src/lib.rs b/crates/tw-gateway/src/lib.rs index ed2f09e6..1f1cb398 100644 --- a/crates/tw-gateway/src/lib.rs +++ b/crates/tw-gateway/src/lib.rs @@ -9,6 +9,7 @@ pub mod access; pub mod affinity; pub mod answer_model; pub mod auth; +pub mod balance; pub mod bedrock; pub mod bodies; pub mod chatgpt; diff --git a/crates/tw-gateway/src/server/pipeline.rs b/crates/tw-gateway/src/server/pipeline.rs index d1e82006..cd010133 100644 --- a/crates/tw-gateway/src/server/pipeline.rs +++ b/crates/tw-gateway/src/server/pipeline.rs @@ -655,40 +655,34 @@ fn route( tracing::debug!(skipped = ?serving.skipped, %model, "skipping the candidates that cannot serve this request"); } decision.candidates = serving.usable; + let group = decision + .via_group + .as_deref() + .and_then(|n| rt.engine.groups().iter().find(|g| g.name == n)); + // `load-balance` 这一次轮到谁(见 `crate::balance`)。**从排序到记账一直拿着锁**: + // 同时进来的几个请求一个接一个地排,后一个看到的是前一个记过的账 + let turn = group + .filter(|g| g.kind == tw_engine::GroupType::LoadBalance) + .map(|g| state.balance.turn(g)); // 策略组排序。**引擎给的是集合,顺序在这儿定** —— // 因为 `load-balance` / `url-test` / `cheapest` 都要运行时的数字, // 而路由决策本身必须是纯的、可试算的。 // // `fallback` 和 `select` 走不到这里面 —— 那是绝大多数人的配置, // 它们连一个 HashMap 都不用建。 - if let Some(gname) = decision.via_group.clone() - && let Some(kind) = rt - .engine - .groups() - .iter() - .find(|g| g.name == gname) - .map(|g| g.kind) - && kind.needs_runtime() + let mut facts_rt = None; + if let Some(g) = group + && g.kind.needs_runtime() { - let facts_rt = tw_engine::Facts { - seq: state.bus.peek_id(), - ttfb_ms: match kind { - tw_engine::GroupType::UrlTest => state.latency.snapshot(&decision.candidates), - _ => Default::default(), - }, - price: match kind { - // 每一家按发给它的名字算价钱:同一个别名在各家是各家的模型名、各家的价目 - tw_engine::GroupType::Cheapest => state.unit_prices( - &rt.config.providers, - &crate::sent::pairs(&sent), - &decision.candidates, - ), - _ => Default::default(), - }, - }; - decision.candidates = rt - .engine - .order(Some(&gname), &decision.candidates, &facts_rt); + let f = state.group_facts( + &rt.config.providers, + g.kind, + &decision.candidates, + &crate::sent::pairs(&sent), + turn.as_ref().map(crate::balance::Turn::current), + ); + decision.candidates = rt.engine.order(Some(&g.name), &decision.candidates, &f); + facts_rt = Some(f); } // 留在上次回答这段对话的那一家:同一轮里一律留,跨轮看缓存值不值得留。**排在 // 策略组排序之后** —— 该留的时候盖过策略,放开的时候策略照常说了算 @@ -706,6 +700,13 @@ fn route( stayed: Some(why), }); } + // `load-balance` 记账:**记粘性之后排头的那一家**,不是按权重轮到的那一家。一段对话 + // 留在了上次回答它的那一家,这一次就算那一家的;之后的新对话把差的补回去 + if let (Some(turn), Some(f), Some(leader)) = + (turn, facts_rt.as_ref(), decision.candidates.first()) + { + turn.charge(&decision.candidates, f, leader); + } Ok(Routed::Go(choice, decision)) } diff --git a/crates/tw-gateway/src/state.rs b/crates/tw-gateway/src/state.rs index 62404b41..dc2bc341 100644 --- a/crates/tw-gateway/src/state.rs +++ b/crates/tw-gateway/src/state.rs @@ -220,6 +220,8 @@ pub struct AppState { pub sessions: Arc, /// 每段对话这一轮的路由决定、上次回答它的那一家(见 [`crate::affinity`])。**跨重载存活** pub affinity: Arc, + /// 每个 `load-balance` 组轮到了谁(见 [`crate::balance`])。**跨重载存活**,改了的组从头轮 + pub balance: Arc, /// 每段对话里、每一家上游拒过的别家封存的推理(见 [`crate::seal`])。**跨重载存活** pub seals: Arc, /// 脚本插件里跨重载存活的那一半:运行时、插件文件在哪儿、计数和日志、编译缓存 @@ -289,6 +291,7 @@ impl AppState { live: crate::live::Live::default(), sessions: Default::default(), affinity: Default::default(), + balance: Default::default(), seals: Default::default(), plugins, swap: Default::default(), @@ -360,6 +363,44 @@ impl AppState { self.pricing.rcu(|book| book.with_table(table.clone())); } + /// 策略组排序要的运行时数字([`tw_engine::Facts`]),只取 `kind` 用得上的那几样。 + /// + /// **数据面和试算共用这一个**,喂给同一个 `order`:试算说会排给谁,数据面就排给谁。 + /// `sent` 是每一家和发给它的模型名(比价按它算);`current_weight` 是 `load-balance` + /// 组此刻的轮询状态(见 [`crate::balance`]):数据面在 [`crate::balance::Turn`] 里读, + /// 试算用 [`crate::balance::Balance::peek`] 只读不记。 + pub fn group_facts( + &self, + providers: &[tw_config::Provider], + kind: tw_engine::GroupType, + candidates: &[String], + sent: &[(String, String)], + current_weight: Option>, + ) -> tw_engine::Facts { + use tw_engine::GroupType; + tw_engine::Facts { + current_weight: current_weight.unwrap_or_default(), + // 熔断着、冷却着的不参加这一轮:它们反正会被跳过 + paused: match kind { + GroupType::LoadBalance => candidates + .iter() + .filter(|p| !self.health.is_available(p)) + .cloned() + .collect(), + _ => Default::default(), + }, + ttfb_ms: match kind { + GroupType::UrlTest => self.latency.snapshot(candidates), + _ => Default::default(), + }, + price: match kind { + // 每一家按发给它的名字算价钱:同一个别名在各家是各家的模型名、各家的价目 + GroupType::Cheapest => self.unit_prices(providers, sent, candidates), + _ => Default::default(), + }, + } + } + /// `cheapest` 排序用的单价:每家跑这个模型的 (输入, 输出),微分/百万 token。 /// /// **路由和预演共用这一个** —— 各写一份的话,预演说会选 A,实际选的是 B。 diff --git a/crates/tw-gateway/tests/affinity.rs b/crates/tw-gateway/tests/affinity.rs index 9ae6669b..2716dce3 100644 --- a/crates/tw-gateway/tests/affinity.rs +++ b/crates/tw-gateway/tests/affinity.rs @@ -85,6 +85,16 @@ async fn ask( gw: SocketAddr, rx: &mut Receiver, messages: &str, +) -> (String, String, Option) { + ask_in(gw, rx, "会话-1", messages).await +} + +/// [`ask`],在会话 `session` 里 +async fn ask_in( + gw: SocketAddr, + rx: &mut Receiver, + session: &str, + messages: &str, ) -> (String, String, Option) { let body = format!( r#"{{"model":"claude-sonnet-4-5","max_tokens":16,"system":"你是一个助手","messages":{messages}}}"# @@ -95,7 +105,7 @@ async fn ask( .unwrap() .post(format!("http://{gw}/v1/messages")) .header("x-api-key", "tw-k") - .header("x-claude-code-session-id", "会话-1") + .header("x-claude-code-session-id", session) .body(body) .send() .await @@ -180,3 +190,34 @@ async fn a_turn_keeps_its_route_and_upstream_and_a_warm_cache_keeps_the_next_tur assert_eq!(rule, "大输入"); assert_eq!(to, "乙"); } + +/// `load-balance` 记在实际排头的那一家头上:一段对话的新一轮留在了甲(缓存热着),按 +/// 权重本该轮到乙 —— 这一次算甲的,之后的两个新对话都去乙,把差的补回来 +#[tokio::test] +async fn the_upstream_a_conversation_stays_on_is_charged_and_new_conversations_make_up_for_it() { + let up = upstream().await; + let (gw, mut rx) = serve(cfg(up)).await; + let first = user("帮我重构"); + // 会话一的第一轮:从头轮,甲 + let (_, to, _) = ask_in(gw, &mut rx, "会话-1", &format!("[{first}]")).await; + assert_eq!(to, "甲"); + // 会话一的第二轮:缓存热着,留在甲 + let history = format!("{first},{{\"role\":\"assistant\",\"content\":\"好了\"}}"); + let (_, to, affinity) = ask_in( + gw, + &mut rx, + "会话-1", + &format!("[{history},{}]", user("再加个测试")), + ) + .await; + assert_eq!(to, "甲"); + assert_eq!(affinity.and_then(|a| a.stayed), Some(Stay::Cache)); + // 新对话:都去乙 + for s in ["会话-2", "会话-3"] { + let (_, to, _) = ask_in(gw, &mut rx, s, &format!("[{}]", user(s))).await; + assert_eq!(to, "乙", "{s}"); + } + // 补齐了,接着挨个轮 + let (_, to, _) = ask_in(gw, &mut rx, "会话-4", &format!("[{}]", user("会话-4"))).await; + assert_eq!(to, "甲"); +} diff --git a/crates/tw-observe/src/bus.rs b/crates/tw-observe/src/bus.rs index 76cb0d5d..96097980 100644 --- a/crates/tw-observe/src/bus.rs +++ b/crates/tw-observe/src/bus.rs @@ -126,15 +126,6 @@ impl EventBus { } /// 拿一个请求 id。同一个请求的四个事件共用它。 - /// 现在发到第几号了,**不占号**。 - /// - /// 给 `load-balance` 当轮转的种子用:它要一个单调、便宜、 - /// 每个请求都不同的数,而事件序号正好是。**不能用 `next_id`** —— - /// 那会凭空占掉一个号,让事件流里出现一个不存在的 id。 - pub fn peek_id(&self) -> u64 { - self.next_id.load(std::sync::atomic::Ordering::Relaxed) - } - pub fn next_id(&self) -> u64 { self.next_id.fetch_add(1, Ordering::Relaxed) } diff --git a/docs/config.md b/docs/config.md index ff747cad..47bd55b6 100644 --- a/docs/config.md +++ b/docs/config.md @@ -952,14 +952,42 @@ group with `to`. | Field | Type | Default | Description | |---|---|---|---| | `name` | string | **required** | Name of the group; unique, and not the name of an upstream. | -| `type` | `fallback` \| `select` \| `load-balance` \| `url-test` \| `cheapest` | `fallback` | `fallback`: the first healthy member, in order. `select`: the member named in `selected`. `load-balance`: take turns between new conversations. `url-test`: the fastest by measured time to first byte. `cheapest`: the lowest input price. | -| `providers` | list of strings | **required** | Member upstreams, by name; not groups. Each upstream appears once in a group. | +| `type` | `fallback` \| `select` \| `load-balance` \| `url-test` \| `cheapest` | `fallback` | `fallback`: the first healthy member, in order. `select`: the member named in `selected`. `load-balance`: new conversations take turns, in proportion to the members' weights. `url-test`: the fastest by measured time to first byte. `cheapest`: the lowest input price. | +| `providers` | list of strings or [`groups[].providers[]`](#cfg-groups-providers) | **required** | Member upstreams, by name; not groups. Each upstream appears once in a group. In a `load-balance` group, a member can be written as `{name, weight}`. | | `selected` | string | — | For `select`: the chosen member. | `fallback` is the default because a single user's machine has no load to spread. +In a `load-balance` group, a member can carry a weight, from 1 to 100; a +member written as just its name has weight 1. Weights set how the group's +requests are shared out: with `{ name: anthropic, weight: 7 }` and `relay`, +the official API serves seven requests in ten. Conversations in progress +stay on the upstream that answers them (see below) and count toward its +share, so the balance is kept by where new conversations start. An upstream +that is cooling down after failures, or cannot serve a request, sits that +request out, and the others share it by their weights. Other group types take +no weights. + + + + +| Field | Type | Default | Description | +|---|---|---|---| +| `name` | string | **required** | The upstream, by name. A member written as just its name has weight 1. | +| `weight` | integer | `1` | The member's share of a `load-balance` group's requests, in proportion to the other members' weights. From 1 to 100. Other group types take no weight other than 1. | + + +```yaml +groups: + - name: pool + type: load-balance + providers: + - { name: anthropic, weight: 7 } + - relay +``` + Whatever the type, a conversation stays on the upstream that last answered it, so that what the upstream holds of it in its prompt cache is read again rather than paid for in full elsewhere. Within a turn (while the client sends @@ -970,7 +998,8 @@ the conversation, and whichever upstream answered after a failover is the one it stays on. The rule a turn matched at its start also holds for the rest of that turn: rules keyed on input size or images do not move a turn halfway, unless its input no longer fits the context window of a model the rule sends -it to. `load-balance` therefore takes turns between new conversations. +it to. `load-balance` therefore takes turns between new conversations, by +weight. ### `routes` diff --git a/docs/config.zh-CN.md b/docs/config.zh-CN.md index deac4828..d247202c 100644 --- a/docs/config.zh-CN.md +++ b/docs/config.zh-CN.md @@ -747,14 +747,34 @@ aliases: | 字段 | 类型 | 默认值 | 说明 | |---|---|---|---| | `name` | 字符串 | **必填** | 策略组的名字,不能重复,也不能和上游同名。 | -| `type` | `fallback` \| `select` \| `load-balance` \| `url-test` \| `cheapest` | `fallback` | `fallback`:按顺序取第一个健康的。`select`:取 `selected` 指定的那个。`load-balance`:新对话轮流。`url-test`:按实测首字节时间取最快的。`cheapest`:取输入单价最低的。 | -| `providers` | 字符串列表 | **必填** | 成员上游的名字,不能是策略组。同一个上游在一个策略组中只出现一次。 | +| `type` | `fallback` \| `select` \| `load-balance` \| `url-test` \| `cheapest` | `fallback` | `fallback`:按顺序取第一个健康的。`select`:取 `selected` 指定的那个。`load-balance`:新对话按成员的权重轮流。`url-test`:按实测首字节时间取最快的。`cheapest`:取输入单价最低的。 | +| `providers` | 列表,每项是字符串或对象,对象见 [`groups[].providers[]`](#cfg-groups-providers) | **必填** | 成员上游的名字,不能是策略组。同一个上游在一个策略组中只出现一次。`load-balance` 组的成员可以写成 `{name, weight}`。 | | `selected` | 字符串 | — | `select` 类型选中的成员。 | 默认类型为 `fallback`:单个使用者的机器上没有需要分散的负载。 -无论哪种类型,一段对话都留在上次回答它的那一家上游,让上游缓存着的那部分被再次读取,而不是换一家全价重算。同一轮之内(客户端正在回传工具结果)一律不换;跨轮时,上一次回答读或写了至少 1024 个 token 的 prompt cache、且距今不到五分钟,才继续留下。上游因失败进入冷却时,对话随之放开;故障转移之后接下回答的那一家,就是之后留下的那一家。一轮开始时命中的规则也沿用到这一轮结束:按输入大小或图片分流的规则不会让一轮半路换家,除非输入已经超出规则所指模型的上下文窗口。因此 `load-balance` 轮流的是新对话。 +`load-balance` 组的成员可以带权重,取值 1 到 100;只写名字的成员权重为 1。权重决定组内请求怎么分:写成 `{ name: anthropic, weight: 7 }` 和 `relay` 时,每十个请求有七个由官方 API 服务。进行中的对话留在回答它的那一家(见下文),也算进那一家的份额,因此份额靠新对话从哪一家开始来补齐。因失败处于冷却、或服务不了某个请求的上游不参与这一次分配,其余成员按各自的权重分。其他类型的策略组不用权重。 + + + + +| 字段 | 类型 | 默认值 | 说明 | +|---|---|---|---| +| `name` | 字符串 | **必填** | 上游的名字。只写名字的成员权重为 1。 | +| `weight` | 整数 | `1` | 成员在 `load-balance` 组中分到的请求份额,与其他成员的权重成比例。取值 1 到 100。其他类型的策略组只能写 1。 | + + +```yaml +groups: + - name: pool + type: load-balance + providers: + - { name: anthropic, weight: 7 } + - relay +``` + +无论哪种类型,一段对话都留在上次回答它的那一家上游,让上游缓存着的那部分被再次读取,而不是换一家全价重算。同一轮之内(客户端正在回传工具结果)一律不换;跨轮时,上一次回答读或写了至少 1024 个 token 的 prompt cache、且距今不到五分钟,才继续留下。上游因失败进入冷却时,对话随之放开;故障转移之后接下回答的那一家,就是之后留下的那一家。一轮开始时命中的规则也沿用到这一轮结束:按输入大小或图片分流的规则不会让一轮半路换家,除非输入已经超出规则所指模型的上下文窗口。因此 `load-balance` 按权重轮流的是新对话。 ### `routes` From d01b0dc833cc0c9773a28accdbaf79013e78b5a9 Mon Sep 17 00:00:00 2001 From: fylorn <249551762+fylorn@users.noreply.github.com> Date: Mon, 5 Oct 2026 16:02:59 +0800 Subject: [PATCH 02/22] Load-balance groups can weigh new conversations by speed and reliability A load-balance group spreads new conversations by its members' weights only, so a member that is slow or keeps failing gets as many new conversations as a fast, healthy one until the pauses catch it. groups[].balance_by (weights, latency, health, latency-health; default weights, omitted when default) multiplies each member's weight by a factor. The engine stays pure: balance_factors(by, members, facts) computes the factor per member from Facts. Latency uses the median TTFB that url-test already uses ((median / own)^2, clamped to 0.1..10); health uses a new Facts.success ((rate)^2, floored at 0.05 so a flaky upstream still gets an occasional request and its recovery is seen). A member without enough samples counts as 1.0, so a new upstream is tried but not flooded. The smooth weighted round-robin turns on effective weights: weight x factor x 1000, rounded, at least 1 (weighted::effective), computed in the one place the schedule reads weights (weighted::round), for picking the leader and for charging it alike. The x1000 keeps fractional factors meaningful with integer current weights; equal factors scale every member alike, so the interleaving is exactly what the weights alone give. Factors are taken over all candidates before the paused ones sit the round out, which is also what the dry-run shows. The success rate comes from Health itself: every record_* call also lands in a per-upstream window (last 50 outcomes within 30 minutes, at least 5 to report), so it counts exactly what the breaker counts and needs no new classification at the call sites. Client errors count as answered, a missing model and client cancellations are not counted. The data plane and the dry-run build their facts in one function (AppState::group_facts, now taking the group), which fills TTFB and success rates only for groups that use them; the dry-run reads the same snapshot and the same round-robin state without advancing it, and shows each candidate's weight, TTFB, success rate and factor. An end-to-end test checks that the dry-run names the upstream the next new conversation takes, request after request, while the factors move. balance_by on any other group type is refused at load time (engine.group_balance_not_load_balance) and by the control plane (control.group.balance_not_load_balance). Co-Authored-By: Claude Opus 5.5 --- crates/tw-api/msg-codes.txt | 2 + crates/tw-api/src/lib.rs | 39 ++ crates/tw-config/src/validate.rs | 49 +++ crates/tw-config/tests/manual/schema.rs | 14 +- crates/tw-control/src/dryrun.rs | 106 +++-- crates/tw-control/src/lib.rs | 1 + crates/tw-control/src/routes.rs | 48 ++- crates/tw-control/tests/dryrun.rs | 216 ++++++++++- crates/tw-control/tests/routes.rs | 51 +++ crates/tw-engine/src/engine.rs | 388 ++++++++++++++++++- crates/tw-engine/src/lib.rs | 7 +- crates/tw-engine/src/weighted.rs | 176 ++++++++- crates/tw-gateway/src/balance.rs | 1 + crates/tw-gateway/src/health.rs | 173 ++++++++- crates/tw-gateway/src/server/pipeline.rs | 2 +- crates/tw-gateway/src/state.rs | 30 +- crates/tw-gateway/tests/affinity.rs | 1 + crates/tw-gateway/tests/client_model_name.rs | 1 + crates/tw-gateway/tests/passthrough.rs | 3 + crates/tw-gateway/tests/routing_facts.rs | 1 + crates/tw-gateway/tests/success_rate.rs | 195 ++++++++++ crates/tw-gateway/tests/ws.rs | 1 + docs/config.md | 34 ++ docs/config.zh-CN.md | 20 + 24 files changed, 1497 insertions(+), 62 deletions(-) create mode 100644 crates/tw-gateway/tests/success_rate.rs diff --git a/crates/tw-api/msg-codes.txt b/crates/tw-api/msg-codes.txt index 3dc13378..a6ab51fc 100644 --- a/crates/tw-api/msg-codes.txt +++ b/crates/tw-api/msg-codes.txt @@ -126,6 +126,7 @@ control.device_code_unavailable control.device_code_unreadable control.disabled_key_cannot_be_default control.dryrun_needs_target +control.group.balance_not_load_balance control.group.empty control.group.name_is_upstream control.group.no_such_upstream @@ -231,6 +232,7 @@ engine.compare.no_operator engine.duplicate_group engine.duplicate_route engine.empty_group +engine.group_balance_not_load_balance engine.group_unknown_upstream engine.group_upstream_twice engine.group_weight_not_load_balance diff --git a/crates/tw-api/src/lib.rs b/crates/tw-api/src/lib.rs index a8f30b5a..9cd290da 100644 --- a/crates/tw-api/src/lib.rs +++ b/crates/tw-api/src/lib.rs @@ -2144,6 +2144,8 @@ pub struct GroupView { /// `load-balance` 组每个成员的权重,**每个成员都在**,没写权重的是 1:新对话按这个 /// 比例分。别的类型不用权重,是空的 pub weights: std::collections::BTreeMap, + /// `load-balance` 按什么分新对话。别的类型永远是 `weights` + pub balance_by: BalanceBy, } #[derive(Debug, Clone, Serialize, Deserialize)] @@ -3210,6 +3212,24 @@ slug_enum! { } } +slug_enum! { + /// `load-balance` 组按什么分新对话:配置里 `balance_by` 写的那个词。 + /// + /// 成员的权重永远是底数,快慢、成败算出的系数乘在上面 + /// ([`DryRunCandidate::balance_factor`]);进行中的对话照旧留在回答它的那一家。 + /// 没有测到的上游算中等。 + pub enum BalanceBy { + /// 只按成员的权重 + Weights = "weights", + /// 首字节越快,分得越多 + Latency = "latency", + /// 最近失败越少,分得越多 + Health = "health", + /// 两样一起看 + LatencyHealth = "latency-health", + } +} + /// 新建或修改一个策略组时交过来的定义。 #[derive(Debug, Clone, Serialize, Deserialize)] #[cfg_attr(feature = "ts", derive(ts_rs::TS))] @@ -3225,6 +3245,9 @@ pub struct GroupInput { /// 别的类型只能不给、或者都是 1 #[serde(default, skip_serializing_if = "Option::is_none")] pub weights: Option>, + /// `load-balance` 组按什么分新对话。不给 = `weights`;别的类型只能是 `weights` + #[serde(default, skip_serializing_if = "Option::is_none")] + pub balance_by: Option, } #[derive(Debug, Clone, Serialize, Deserialize)] @@ -4712,6 +4735,9 @@ pub struct DryRunResult { /// 写在第一个」。直指 provider 时是 None。 #[serde(default, skip_serializing_if = "Option::is_none")] pub strategy: Option, + /// 经过的是 `load-balance` 组时,它按什么分新对话。别的时候没有 + #[serde(default, skip_serializing_if = "Option::is_none")] + pub balance_by: Option, /// `route` | `deny` | `no_match` | `unavailable`(选中的上游都服务不了, /// 见 `skipped`)| `intercepted` /// @@ -4758,6 +4784,18 @@ pub struct DryRunCandidate { /// 经过的是 `load-balance` 组时,它在组里的权重(没写权重的是 1)。别的时候没有 #[serde(default, skip_serializing_if = "Option::is_none")] pub weight: Option, + /// 它的典型首字节时间(最近样本的中位数),毫秒。只在顺序看它时有:`url-test`, + /// 按快慢分的 `load-balance`。样本不够时没有 + #[serde(default, skip_serializing_if = "Option::is_none")] + pub ttfb_ms: Option, + /// 它最近的成功率,0 到 1(最近 50 次、30 分钟以内)。只在按成败分的 `load-balance` + /// 里有;不到 5 次时没有 + #[serde(default, skip_serializing_if = "Option::is_none")] + pub success_rate: Option, + /// 按快慢、成败算出的系数,乘在权重([`Self::weight`])上:大于 1 分得多,小于 1 + /// 分得少,没有样本的那一项算 1。只在 `balance_by` 不是 `weights` 的 `load-balance` 里有 + #[serde(default, skip_serializing_if = "Option::is_none")] + pub balance_factor: Option, } /// 一个要转换格式的候选上游。 @@ -5573,6 +5611,7 @@ mod tests { ModelListStatus::from_slug, ); check(GroupKind::ALL, GroupKind::slug, GroupKind::from_slug); + check(BalanceBy::ALL, BalanceBy::slug, BalanceBy::from_slug); check(RuleVerdict::ALL, RuleVerdict::slug, RuleVerdict::from_slug); check(RuleEffect::ALL, RuleEffect::slug, RuleEffect::from_slug); check( diff --git a/crates/tw-config/src/validate.rs b/crates/tw-config/src/validate.rs index fc068b8f..0a7af0ef 100644 --- a/crates/tw-config/src/validate.rs +++ b/crates/tw-config/src/validate.rs @@ -866,6 +866,55 @@ groups: assert_eq!(m.arg("field"), "groups[0].providers[1].wieght", "{m:?}"); } + /// `balance_by` 写在不是负载均衡的组上:加载时就拒绝,它在那里什么都不做。写在负载 + /// 均衡组上的照收,写回时原样;默认的不写进去 + #[test] + fn balance_by_belongs_to_load_balance_groups_and_round_trips() { + let text = "version: 1 +listen: + control: + key: c0ffee00c0ffee00c0ffee00c0ffee00c0ffee00c0ffee00c0ffee00c0ffee00 +clients: + - name: c + key: tw-k +providers: + - name: a + base_url: https://a.example + key: sk-a + - name: b + base_url: https://b.example + key: sk-b +groups: + - name: pool + type: load-balance + providers: [a, b] + balance_by: latency-health +"; + let cfg = crate::try_parse(text).unwrap(); + assert_eq!( + cfg.groups[0].balance_by, + tw_engine::BalanceBy::LatencyHealth + ); + let back = serde_yaml_ng::to_string(&cfg.groups).unwrap(); + assert!(back.contains("balance_by: latency-health"), "{back}"); + + let fallback = text.replace("type: load-balance", "type: fallback"); + let m = crate::try_parse(&fallback).unwrap_err().msg(); + assert_eq!(m.code, "engine.group_balance_not_load_balance", "{m:?}"); + assert_eq!( + (m.arg("group"), m.arg("balance_by")), + ("pool", "latency-health") + ); + // 写明默认值的照收:它本来就什么都不做 + let weights = fallback.replace("latency-health", "weights"); + assert!(crate::try_parse(&weights).is_ok(), "{weights}"); + + let plain = text.replace(" balance_by: latency-health\n", ""); + let cfg = crate::try_parse(&plain).unwrap(); + let back = serde_yaml_ng::to_string(&cfg.groups).unwrap(); + assert!(!back.contains("balance_by"), "{back}"); + } + #[test] fn error_messages_say_what_to_do_next() { // 错误信息是降低使用难度最有效的杠杆。判据不是「说清 diff --git a/crates/tw-config/tests/manual/schema.rs b/crates/tw-config/tests/manual/schema.rs index 9f3cd162..8916784e 100644 --- a/crates/tw-config/tests/manual/schema.rs +++ b/crates/tw-config/tests/manual/schema.rs @@ -11,7 +11,7 @@ use super::{Def, Kind, Lang, Row, Section, T2}; use tw_config::proxy::ProxyAuth; use tw_config::*; use tw_engine::rule::When; -use tw_engine::{Group, GroupType, Member, Pinned, RouteSet, Rule, SetAction}; +use tw_engine::{BalanceBy, Group, GroupType, Member, Pinned, RouteSet, Rule, SetAction}; use tw_pricing::{PerMillion, PricingConfig, SheetDef}; const fn t(en: &'static str, zh: &'static str) -> T2 { @@ -58,6 +58,9 @@ fn content_matches() -> Vec<&'static str> { fn group_types() -> Vec<&'static str> { super::fields::() } +fn balance_bys() -> Vec<&'static str> { + super::fields::() +} const RULE_ID: T2 = t("built-in rule id", "内置规则 id"); const MODE_DOC: T2 = t( @@ -1250,6 +1253,15 @@ pub fn sections() -> Vec
{ "`select` 类型选中的成员。", ), ), + row( + "balance_by", + Kind::Enum(balance_bys), + Def::Is("weights"), + t( + "For `load-balance`: what the members' weights are multiplied by. `weights`: nothing; the weights alone. `latency`: faster upstreams get more. `health`: upstreams that fail less get more. `latency-health`: both. Other group types take only `weights`.", + "`load-balance` 类型用:成员的权重再乘上什么。`weights`:不乘,只按权重。`latency`:越快的上游分得越多。`health`:越少失败的上游分得越多。`latency-health`:两者都看。其他类型只能是 `weights`。", + ), + ), ], }, Section { diff --git a/crates/tw-control/src/dryrun.rs b/crates/tw-control/src/dryrun.rs index 476f2a4d..3c0d1f75 100644 --- a/crates/tw-control/src/dryrun.rs +++ b/crates/tw-control/src/dryrun.rs @@ -8,37 +8,84 @@ use axum::{Json, extract::State, http::StatusCode}; -/// 和数据面同一段排序。 +/// 经过的组,和给它排序要的那些数字:和数据面同一份延迟表、成功率、价目表、轮询状态。 +/// 没经过组、或者组的顺序用不着运行时的数字(`fallback`、`select`)时是 None。 /// /// 试算页存在的全部意义是「告诉你这条请求会走哪儿」,所以它**必须**用 /// 同一个函数、同一份数字 —— 各算各的话,两边迟早会不一样,而那时 -/// 试算比没有更糟。`sent` 是每一家和发给它的模型名:比价按它算,和数据面一样。 +/// 试算比没有更糟。数字由 [`tw_gateway::AppState::group_facts`] 取,数据面也调它。 +/// `sent` 是每一家和发给它的模型名:比价按它算,和数据面一样。 /// /// `load-balance` 读的是数据面记着的那一份轮询状态,**只读不记** /// ([`tw_gateway::balance::Balance::peek`]):试算说的是下一个新对话会排给谁, /// 试算了几次都不该改变这个答案,也不该让数据面跳过谁。 -fn order_like_the_data_plane( +fn runtime_facts<'e>( s: &crate::ControlState, - engine: &tw_engine::Engine, + engine: &'e tw_engine::Engine, d: &tw_engine::Decision, sent: &[(String, String)], -) -> Vec { - let Some(g) = d - .via_group - .as_deref() - .and_then(|n| engine.groups().iter().find(|g| g.name == n)) - else { - return d.candidates.clone(); - }; +) -> Option<(&'e tw_engine::Group, tw_engine::Facts)> { + let gname = d.via_group.as_deref()?; + let g = engine.groups().iter().find(|g| g.name == gname)?; if !g.kind.needs_runtime() { - return d.candidates.clone(); + return None; } let cfg = s.config(); let current = (g.kind == tw_engine::GroupType::LoadBalance).then(|| s.gateway.balance.peek(g)); let facts = s .gateway - .group_facts(&cfg.providers, g.kind, &d.candidates, sent, current); - engine.order(Some(&g.name), &d.candidates, &facts) + .group_facts(&cfg.providers, g, &d.candidates, sent, current); + Some((g, facts)) +} + +/// 和数据面同一段排序(见 [`runtime_facts`])。 +fn order_like_the_data_plane( + engine: &tw_engine::Engine, + d: &tw_engine::Decision, + runtime: Option<&(&tw_engine::Group, tw_engine::Facts)>, +) -> Vec { + match runtime { + Some((g, facts)) => engine.order(Some(&g.name), &d.candidates, facts), + None => d.candidates.clone(), + } +} + +/// 一家候选的排序依据:典型首字节时间、最近的成功率、按它们算出的系数。**只给顺序 +/// 真用到的那几样**([`runtime_facts`] 只取用得上的)—— `url-test` 看首字节时间; +/// `load-balance` 按 `balance_by` 看快慢、成败,系数乘在权重上。用不着的、没有样本的 +/// 是 None。 +struct Basis { + ttfb_ms: Option, + success_rate: Option, + balance_factor: Option, +} + +/// 每一家候选的排序依据,按 `d.candidates` 的名字查。系数和数据面一样按这份候选、 +/// 这份数字算(`tw_engine::balance_factors`) +fn bases( + d: &tw_engine::Decision, + runtime: Option<&(&tw_engine::Group, tw_engine::Facts)>, +) -> std::collections::HashMap { + let Some((g, f)) = runtime else { + return Default::default(); + }; + let factors = if g.kind == tw_engine::GroupType::LoadBalance && !g.balance_by.is_weights() { + tw_engine::balance_factors(g.balance_by, &d.candidates, f) + } else { + Vec::new() + }; + d.candidates + .iter() + .enumerate() + .map(|(i, name)| { + let basis = Basis { + ttfb_ms: f.ttfb_ms.get(name).copied(), + success_rate: f.success.get(name).copied(), + balance_factor: factors.get(i).copied(), + }; + (name.clone(), basis) + }) + .collect() } use tw_engine::{Outcome, RequestFacts, RouteError}; @@ -198,6 +245,7 @@ pub async fn dry_run( route, outcome: short.unwrap_or(DryRunOutcome::NoMatch), strategy: None, + balance_by: None, rule: None, reason: None, candidates: Vec::new(), @@ -220,11 +268,14 @@ pub async fn dry_run( out.rule = Some(d.matched_rule.clone()); out.via_group = d.via_group.clone(); out.set = describe(&d.set); - out.strategy = d + let group = d .via_group .as_deref() - .and_then(|g| engine.groups().iter().find(|x| x.name == g)) - .map(|g| crate::routes::group_kind(g.kind)); + .and_then(|g| engine.groups().iter().find(|x| x.name == g)); + out.strategy = group.map(|g| crate::routes::group_kind(g.kind)); + out.balance_by = group + .filter(|g| g.kind == tw_engine::GroupType::LoadBalance) + .map(|g| crate::routes::balance_by_view(g.balance_by)); // 和数据面同一步:去掉服务不了这个请求的候选(停用的、范围外的、 // 清单里没有这个模型的)。**被跳过的要列出来** —— 「规则明明写的 // 是 A」正是用户会来试算的原因 @@ -269,20 +320,20 @@ pub async fn dry_run( // 而它属于哪次会话取决于请求正文,试算没有那个东西。 // 于是它显示的是轮询此刻轮到的位置 —— 而那正是一个新对话 // 真的会走的路。 - out.candidates = - order_like_the_data_plane(&s, engine, &d, &tw_gateway::sent::pairs(&sent)); + let runtime = runtime_facts(&s, engine, &d, &tw_gateway::sent::pairs(&sent)); + out.candidates = order_like_the_data_plane(engine, &d, runtime.as_ref()); // `load-balance` 的候选带上各自的权重:排头的为什么是它,一半在这个数里 - let balanced = d - .via_group - .as_deref() - .and_then(|g| engine.groups().iter().find(|x| x.name == g)) - .filter(|g| g.kind == tw_engine::GroupType::LoadBalance); - // 每一家收到的模型名,和为什么不是请求里写的那个 + let balanced = group.filter(|g| g.kind == tw_engine::GroupType::LoadBalance); + // 每一家收到的模型名,和为什么不是请求里写的那个;顺序看运行时数字的,再加上 + // 每一家的那几个数字(权重、首字节时间、成功率、系数)—— 「为什么轮到它」 + // 要能从这里看出来 + let mut basis_of = bases(&d, runtime.as_ref()); out.candidate_models = out .candidates .iter() .map(|name| { let one = sent.iter().find(|x| x.provider == *name); + let basis = basis_of.remove(name); tw_api::DryRunCandidate { provider: name.clone(), sent_model: one @@ -290,6 +341,9 @@ pub async fn dry_run( .filter(|m| !m.is_empty()), model_via: one.and_then(|x| x.via).map(|v| v.slug().to_string()), weight: balanced.map(|g| g.weight(name)), + ttfb_ms: basis.as_ref().and_then(|b| b.ttfb_ms), + success_rate: basis.as_ref().and_then(|b| b.success_rate), + balance_factor: basis.as_ref().and_then(|b| b.balance_factor), } }) .collect(); diff --git a/crates/tw-control/src/lib.rs b/crates/tw-control/src/lib.rs index d23365e5..abf70bd5 100644 --- a/crates/tw-control/src/lib.rs +++ b/crates/tw-control/src/lib.rs @@ -369,6 +369,7 @@ async fn overview(State(s): State) -> Json { selected: g.selected.clone(), providers: g.names(), weights: routes::group_weights(g), + balance_by: routes::balance_by_view(g.balance_by), }) .collect(), // 和 `GET /keys` 同一份视图:概览里少一个字段的话,两处会各自 diff --git a/crates/tw-control/src/routes.rs b/crates/tw-control/src/routes.rs index 769d193a..2b21eeb6 100644 --- a/crates/tw-control/src/routes.rs +++ b/crates/tw-control/src/routes.rs @@ -21,7 +21,7 @@ use tw_config::edit::{self, EditError}; use tw_config::history::Origin; use tw_config::refs; use tw_engine::rule::{OneOrMany, When}; -use tw_engine::{Group, GroupType, RouteSet, Rule, SetAction}; +use tw_engine::{BalanceBy, Group, GroupType, RouteSet, Rule, SetAction}; use tw_types::{Msg, msg}; use tw_yaml::Step; @@ -587,6 +587,25 @@ pub(crate) fn group_weights(g: &Group) -> std::collections::BTreeMap tw_api::BalanceBy { + match b { + BalanceBy::Weights => tw_api::BalanceBy::Weights, + BalanceBy::Latency => tw_api::BalanceBy::Latency, + BalanceBy::Health => tw_api::BalanceBy::Health, + BalanceBy::LatencyHealth => tw_api::BalanceBy::LatencyHealth, + } +} + +fn balance_by_of(b: tw_api::BalanceBy) -> BalanceBy { + match b { + tw_api::BalanceBy::Weights => BalanceBy::Weights, + tw_api::BalanceBy::Latency => BalanceBy::Latency, + tw_api::BalanceBy::Health => BalanceBy::Health, + tw_api::BalanceBy::LatencyHealth => BalanceBy::LatencyHealth, + } +} + fn to_group(input: &tw_api::GroupInput, cfg: &tw_config::Config) -> Result { let name = checked_name(&input.name, "group")?; reserved(&name, "group")?; @@ -668,6 +687,15 @@ fn to_group(input: &tw_api::GroupInput, cfg: &tw_config::Config) -> Result None, }; + // 按快慢、成败分只有负载均衡用得上:别的类型写了也不起作用(引擎校验也拦, + // `engine.group_balance_not_load_balance`) + let balance_by = input.balance_by.map(balance_by_of).unwrap_or_default(); + if kind != GroupType::LoadBalance && !balance_by.is_weights() { + return Err(msg!( + "control.group.balance_not_load_balance", balance_by = balance_by.slug() => + "balance_by `{balance_by}` applies only to a load-balance group" + )); + } Ok(Group { name, kind, @@ -679,6 +707,7 @@ fn to_group(input: &tw_api::GroupInput, cfg: &tw_config::Config) -> Result (tempfile::TempDir, axum::Router) { (d, app) } -/// [`app_with`],连同它的数据面:要看试算和数据面是不是读的同一份状态 +/// [`app_with`],连同它的数据面:要看试算和数据面是不是读的同一份状态,或者先往延迟表、 +/// 成功率里放数字 fn app_and_gateway(text: &str) -> (tempfile::TempDir, axum::Router, tw_gateway::AppState) { let d = tempfile::tempdir().unwrap(); let p = d.path().join("config.yaml"); @@ -661,3 +662,216 @@ async fn a_candidate_whose_model_the_key_may_not_use_is_skipped() { let r = run(&app, r#"{"model":"glm-5","route":"default"}"#).await; assert_eq!(sent(&r), [("智谱", Some("glm-air"), Some("rule"))]); } + +fn candidate<'r>(r: &'r tw_api::DryRunResult, name: &str) -> &'r tw_api::DryRunCandidate { + r.candidate_models + .iter() + .find(|c| c.provider == name) + .unwrap_or_else(|| panic!("{name} 不在候选里:{r:?}")) +} + +/// 按快慢、成败分的负载均衡:每一家的首字节时间、成功率和算出的系数都写在候选上, +/// 用的是数据面同一份数字。「为什么这家分得多」要能从试算里看出来 +#[tokio::test] +async fn a_balancing_group_shows_what_each_member_is_weighed_by() { + let cfg = CFG.replace( + " type: load-balance\n", + " type: load-balance\n balance_by: latency-health\n", + ); + let (_d, app, gw) = app_and_gateway(&cfg); + // 官方快、没有成败的样本;中转慢一倍,最近五次失败了一次 + for _ in 0..3 { + gw.latency.record("官方", 100); + gw.latency.record("中转", 200); + } + for _ in 0..4 { + gw.health.record_success("中转"); + } + gw.health.record_failure("中转"); + + let r = run(&app, r#"{"model":"claude-sonnet-4-5"}"#).await; + assert_eq!(r.balance_by, Some(tw_api::BalanceBy::LatencyHealth)); + let official = candidate(&r, "官方"); + assert_eq!(official.ttfb_ms, Some(100)); + assert_eq!(official.success_rate, None, "样本不够,不说成功率"); + // 中位数 150:(150 / 100)² = 2.25,没有成败的样本算 1 + assert!( + (official.balance_factor.unwrap() - 2.25).abs() < 1e-9, + "{official:?}" + ); + let relay = candidate(&r, "中转"); + assert_eq!(relay.ttfb_ms, Some(200)); + assert!( + (relay.success_rate.unwrap() - 0.8).abs() < 1e-9, + "{relay:?}" + ); + // (150 / 200)² × 0.8² = 0.5625 × 0.64 + assert!( + (relay.balance_factor.unwrap() - 0.36).abs() < 1e-9, + "{relay:?}" + ); + + // 只按快慢:成功率不说,系数里也没有它 + let (_d, app, gw) = app_and_gateway(&cfg.replace("latency-health", "latency")); + for _ in 0..3 { + gw.latency.record("官方", 100); + gw.latency.record("中转", 200); + } + for _ in 0..5 { + gw.health.record_failure("中转"); + } + let r = run(&app, r#"{"model":"claude-sonnet-4-5"}"#).await; + let relay = candidate(&r, "中转"); + assert_eq!((relay.ttfb_ms, relay.success_rate), (Some(200), None)); + assert!( + (relay.balance_factor.unwrap() - 0.5625).abs() < 1e-9, + "{relay:?}" + ); +} + +/// 只按权重分时没有系数可说;`url-test` 只说首字节时间 —— 它就按这个排 +#[tokio::test] +async fn only_the_numbers_the_order_uses_are_shown() { + let (_d, app, gw) = app_and_gateway(CFG); + for _ in 0..3 { + gw.latency.record("官方", 100); + } + for _ in 0..5 { + gw.health.record_success("官方"); + } + let r = run(&app, r#"{"model":"claude-sonnet-4-5"}"#).await; + assert_eq!(r.balance_by, Some(tw_api::BalanceBy::Weights)); + let official = candidate(&r, "官方"); + assert_eq!( + ( + official.ttfb_ms, + official.success_rate, + official.balance_factor + ), + (None, None, None) + ); + + let (_d, app, gw) = app_and_gateway(&CFG.replace("type: load-balance", "type: url-test")); + for _ in 0..3 { + gw.latency.record("中转", 80); + } + let r = run(&app, r#"{"model":"claude-sonnet-4-5"}"#).await; + assert_eq!(r.balance_by, None, "不是负载均衡"); + assert_eq!(r.candidates, ["中转", "官方"], "测到的排前面"); + let relay = candidate(&r, "中转"); + assert_eq!( + (relay.ttfb_ms, relay.success_rate, relay.balance_factor), + (Some(80), None, None) + ); +} + +/// 按快慢、成败分时,试算说的排头就是数据面下一个新对话真的去的那一家:同一份延迟表、 +/// 成功率、轮询状态,同一个有效权重。起真网关、真上游,每发一个新对话之前先试算一次, +/// 一次都不能对不上 —— 真请求会添新的首字节样本和成败,系数一直在变,对得上才说明两边 +/// 每一次读的都是同一份 +#[tokio::test] +async fn with_factors_the_dry_run_names_the_upstream_the_next_new_conversation_takes() { + let up = { + let app = axum::Router::new().fallback(axum::routing::any(|| async { + axum::response::Response::builder() + .header("content-type", "application/json") + .body(Body::from( + r#"{"type":"message","content":[],"usage":{"input_tokens":3,"output_tokens":1}}"#, + )) + .unwrap() + })); + let l = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap(); + let a = l.local_addr().unwrap(); + tokio::spawn(async move { axum::serve(l, app).await.unwrap() }); + a + }; + let text = format!( + r#"version: 1 +listen: + control: + key: c0ffee00c0ffee00c0ffee00c0ffee00c0ffee00c0ffee00c0ffee00c0ffee00 +clients: + - name: 我 + key: tw-k +providers: + - {{ name: 甲, base_url: "http://{up}", key: sk-a, protocol: anthropic }} + - {{ name: 乙, base_url: "http://{up}", key: sk-b, protocol: anthropic }} + - {{ name: 丙, base_url: "http://{up}", key: sk-c, protocol: anthropic }} +groups: + - name: 池 + type: load-balance + balance_by: latency-health + providers: [{{ name: 甲, weight: 2 }}, 乙, 丙] +routes: + - name: default + rules: + - name: 都去池子 + to: 池 +"# + ); + let (_d, app, gw) = app_and_gateway(&text); + let mut events = gw.bus.subscribe(); + // 甲慢、乙快但最近一半失败、丙居中:有效权重 2×0.25 : 1×4×0.25 : 1×1 = 500 : 1000 : 1000 + for _ in 0..32 { + gw.latency.record("甲", 400); + gw.latency.record("乙", 100); + gw.latency.record("丙", 200); + } + for _ in 0..5 { + gw.health.record_success("乙"); + gw.health.record_failure("乙"); + } + let addr = tw_gateway::serve(gw.clone(), ([127, 0, 0, 1], 0).into()) + .await + .unwrap(); + let http = reqwest::Client::builder().no_proxy().build().unwrap(); + + let first = run(&app, r#"{"model":"claude-sonnet-4-5"}"#).await; + let factor = |name: &str| candidate(&first, name).balance_factor.unwrap(); + assert!((factor("甲") - 0.25).abs() < 1e-9, "{first:?}"); + assert!((factor("乙") - 1.0).abs() < 1e-9, "{first:?}"); + assert!((factor("丙") - 1.0).abs() < 1e-9, "{first:?}"); + assert_eq!(candidate(&first, "甲").weight, Some(2)); + + let mut seen = std::collections::HashMap::::new(); + for i in 0..30 { + let predicted = run(&app, r#"{"model":"claude-sonnet-4-5"}"#) + .await + .candidates[0] + .clone(); + // 每次一段新对话:没有粘性,排头的就是轮到的那一家 + let st = http + .post(format!("http://{addr}/v1/messages")) + .header("x-api-key", "tw-k") + .header("x-claude-code-session-id", format!("新对话-{i}")) + .body(format!( + r#"{{"model":"claude-sonnet-4-5","max_tokens":16,"messages":[{{"role":"user","content":"第 {i} 个"}}]}}"# + )) + .send() + .await + .unwrap() + .status(); + assert_eq!(st, 200); + let mut went = None; + loop { + let ev = tokio::time::timeout(std::time::Duration::from_secs(5), events.recv()) + .await + .expect("5 秒内没等到结局") + .expect("事件流断了"); + match ev { + tw_api::Event::RequestRouted { attempts, .. } => { + went = attempts.first().map(|a| a.provider.clone()) + } + tw_api::Event::RequestFinished { .. } => break, + tw_api::Event::RequestFailed { message, .. } => panic!("失败了:{message:?}"), + _ => {} + } + } + let went = went.expect("没有路由事件"); + assert_eq!(went, predicted, "第 {i} 个新对话"); + *seen.entry(went).or_default() += 1; + } + // 三家都轮到过,慢的那一家(权重 2)分得最少 + assert_eq!(seen.len(), 3, "{seen:?}"); + assert!(seen["甲"] < seen["丙"], "{seen:?}"); +} diff --git a/crates/tw-control/tests/routes.rs b/crates/tw-control/tests/routes.rs index 16cee4a9..f2b22480 100644 --- a/crates/tw-control/tests/routes.rs +++ b/crates/tw-control/tests/routes.rs @@ -741,6 +741,57 @@ async fn the_built_in_group_and_reserved_names_are_refused() { assert_eq!(st, StatusCode::BAD_REQUEST, "和上游同名:{v}"); } +/// 负载均衡按快慢、成败分:写进配置、在概览里读得回来;默认的只按比例不写进文件; +/// 别的类型写了要拒 +#[tokio::test] +async fn a_load_balance_group_balances_by_what_it_is_told_and_others_refuse() { + let b = bed(BASE); + let (st, v) = call( + &b.app, + "POST", + "/groups", + json!({ "group": { "name": "均摊", "kind": "load-balance", + "providers": ["官方", "中转"], "balance_by": "latency-health" } }), + ) + .await; + assert_eq!(st, StatusCode::OK, "{v}"); + let file = b.file(); + assert!(file.contains("balance_by: latency-health"), "{file}"); + let g = find(&b.overview().await["groups"], "均摊").clone(); + assert_eq!(g["balance_by"], "latency-health", "{g}"); + + // 改回只按比例:这一项从文件里消失,概览照样说出来 + let (st, v) = call( + &b.app, + "PUT", + &format!("/groups/{}", enc("均摊")), + json!({ "group": { "name": "均摊", "kind": "load-balance", + "providers": ["官方", "中转"] } }), + ) + .await; + assert_eq!(st, StatusCode::OK, "{v}"); + assert!(!b.file().contains("balance_by"), "{}", b.file()); + assert_eq!( + find(&b.overview().await["groups"], "均摊")["balance_by"], + "weights" + ); + assert_eq!( + find(&b.overview().await["groups"], "主力")["balance_by"], + "weights" + ); + + let (st, v) = call( + &b.app, + "POST", + "/groups", + json!({ "group": { "name": "按顺序", "kind": "fallback", + "providers": ["官方"], "balance_by": "health" } }), + ) + .await; + assert_eq!(st, StatusCode::BAD_REQUEST, "{v}"); + assert_eq!(v["code"], "control.group.balance_not_load_balance", "{v}"); +} + // ─────────────────────────────────────────────────────────── 已知模型 #[tokio::test] diff --git a/crates/tw-engine/src/engine.rs b/crates/tw-engine/src/engine.rs index 50321b84..91c04561 100644 --- a/crates/tw-engine/src/engine.rs +++ b/crates/tw-engine/src/engine.rs @@ -26,7 +26,8 @@ pub enum GroupType { Fallback, /// 手动指定一个。缓存友好 Select, - /// 按成员的权重轮流(平滑加权轮询,见 [`crate::weighted`]),权重都是 1 就是挨个轮。 + /// 按成员的权重轮流(平滑加权轮询,见 [`crate::weighted`]),权重都是 1 就是挨个轮; + /// `balance_by` 还可以按快慢、成败给权重乘一个系数([`BalanceBy`])。 /// **轮的是新对话**:已经有人回答过、缓存还热着的对话留在那一家 LoadBalance, /// 选最快的。判据是**真实流量测出来的 TTFB**,样本不够时用启动时 @@ -66,6 +67,130 @@ impl GroupType { } } +/// `load-balance` 按什么分新对话([`Group::balance_by`])。 +/// +/// **成员的权重永远是底数**:快慢、成败算出一个系数([`balance_factors`]),乘在每一家 +/// 的权重上,平滑加权轮询按乘出来的数轮([`crate::weighted`])。分的只是新对话 —— +/// 进行中的对话照旧留在回答它的那一家(`tw_gateway::affinity`)。 +#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize, Default)] +#[serde(rename_all = "kebab-case")] +pub enum BalanceBy { + /// 只按成员的权重 + #[default] + Weights, + /// 越快的分得越多:看首字节时间,和 `url-test` 同一份样本 + Latency, + /// 越少失败的分得越多:看最近的成功率 + Health, + /// 两样一起看:两个系数相乘 + LatencyHealth, +} + +impl BalanceBy { + /// 配置里写的那个词,也是控制面发给界面的值 + pub fn slug(&self) -> &'static str { + match self { + BalanceBy::Weights => "weights", + BalanceBy::Latency => "latency", + BalanceBy::Health => "health", + BalanceBy::LatencyHealth => "latency-health", + } + } + + /// 是不是默认的「只按成员的权重」。**写回配置时默认值不写**,老配置原样往返 + pub fn is_weights(&self) -> bool { + *self == BalanceBy::Weights + } + + /// 要不要每家的首字节时间([`Facts::ttfb_ms`]) + pub fn uses_latency(&self) -> bool { + matches!(self, BalanceBy::Latency | BalanceBy::LatencyHealth) + } + + /// 要不要每家的成功率([`Facts::success`]) + pub fn uses_health(&self) -> bool { + matches!(self, BalanceBy::Health | BalanceBy::LatencyHealth) + } +} + +/// 快慢系数的范围。**两头都要夹**:快十倍的一家不该把别家饿死(它们还要被测到、还要 +/// 做故障转移的备选),慢的那家也还该偶尔被试到,否则它变快了没人知道 +const LATENCY_FACTOR_MIN: f64 = 0.1; +const LATENCY_FACTOR_MAX: f64 = 10.0; + +/// 成败系数的下限。**不是 0**:常失败的那家照样偶尔分到一个新对话,它恢复了才看得出来。 +/// 一直失败的由熔断停用(`tw_gateway::health`),不归这里管 +const HEALTH_FACTOR_FLOOR: f64 = 0.05; + +/// `load-balance` 每一家的系数,和 `members` 一一对应:**成员的权重乘上它**,就是这一家 +/// 这一轮的有效权重([`crate::weighted`] 按它轮)。`weights` 全是 1。 +/// +/// - 快慢:(测到的那几家首字节时间的中位数 ÷ 这一家的)²,夹在 0.1 到 10 之间。平方让差别 +/// 看得出来:快一倍的分到四倍。 +/// - 成败:成功率²,最低 0.05。 +/// - **没有样本的那一项算 1,就是「中等」**:新加的上游会被试到,但不会一上来就被灌满。 +/// +/// **纯函数**,和 [`order_by`] 一样:数据面和试算页喂的是同一份数字。 +pub fn balance_factors(by: BalanceBy, members: &[String], f: &Facts) -> Vec { + let typical = if by.uses_latency() { + median_ttfb(members, f) + } else { + None + }; + members + .iter() + .map(|m| { + let mut x = 1.0; + if let Some(mid) = typical + && let Some(t) = f.ttfb_ms.get(m) + { + x *= latency_factor(mid, *t); + } + if by.uses_health() + && let Some(s) = f.success.get(m) + { + x *= health_factor(*s); + } + x + }) + .collect() +} + +/// 测到了的那几家首字节时间的中位数,毫秒。一家都没测到是 `None`。 +/// +/// 偶数家时取中间两家的平均:系数围着它往两头夹,取其中一家的话,夹的那一刀会偏向一边 +fn median_ttfb(members: &[String], f: &Facts) -> Option { + let mut v: Vec = members + .iter() + .filter_map(|m| f.ttfb_ms.get(m).copied()) + .collect(); + if v.is_empty() { + return None; + } + v.sort_unstable(); + let n = v.len(); + Some(if n % 2 == 1 { + f64::from(v[n / 2]) + } else { + (f64::from(v[n / 2 - 1]) + f64::from(v[n / 2])) / 2.0 + }) +} + +/// 快慢系数。**不到 1 毫秒的按 1 毫秒算**:本机的假上游、或者垫底的握手计时可能是 0 +fn latency_factor(median_ms: f64, ttfb_ms: u32) -> f64 { + let ratio = median_ms.max(1.0) / f64::from(ttfb_ms.max(1)); + (ratio * ratio).clamp(LATENCY_FACTOR_MIN, LATENCY_FACTOR_MAX) +} + +/// 成败系数。不是一个数(NaN)的当作没有样本 +fn health_factor(success: f64) -> f64 { + if success.is_nan() { + return 1.0; + } + let s = success.clamp(0.0, 1.0); + (s * s).max(HEALTH_FACTOR_FLOOR) +} + /// 排顺序时才知道的那些数字。**引擎是纯函数,这些从外面传进来** —— /// 于是数据面和试算页走的是同一段逻辑,试算不会「算出一个 /// 和真实转发不一样的结果」。 @@ -82,6 +207,9 @@ pub struct Facts { /// 每家跑这个模型的单价,(输入, 输出),微分/百万 token。 /// **缺席 = 算不出价钱**,不是「免费」 pub price: std::collections::HashMap, + /// 每家最近的成功率,0 到 1(`load-balance` 按成败分时用)。**缺席 = 样本不够**, + /// 不是「从不失败」 + pub success: std::collections::HashMap, } /// 一组上游,以及从里面挑一个的策略。 @@ -102,6 +230,9 @@ pub struct Group { /// `select` 用:当前选中的那个 #[serde(default, skip_serializing_if = "Option::is_none")] pub selected: Option, + /// `load-balance` 用:按什么分新对话(见 [`BalanceBy`])。默认只按成员的权重,不写回配置 + #[serde(default, skip_serializing_if = "BalanceBy::is_weights")] + pub balance_by: BalanceBy, } impl Group { @@ -597,6 +728,8 @@ pub enum RouteError { weight: u32, }, #[error("{}", self.msg())] + GroupBalanceNotLoadBalance { group: String, by: BalanceBy }, + #[error("{}", self.msg())] DuplicateRoute(String), #[error("{}", self.msg())] UnknownDefaultRoute(String), @@ -678,6 +811,11 @@ impl RouteError { "group `{group}` gives upstream `{upstream}` a weight of {weight}. A weight is a \ whole number from 1 to 100" ), + RouteError::GroupBalanceNotLoadBalance { group, by } => msg!( + "engine.group_balance_not_load_balance", group = group, balance_by = by.slug() => + "group `{group}` sets balance_by `{balance_by}`, which only a load-balance group \ + uses. Remove it, or make the group load-balance" + ), RouteError::DuplicateRoute(route) => msg!( "engine.duplicate_route", route = route => "there is more than one route named `{route}`. A gateway key binds to a route by \ @@ -812,6 +950,7 @@ impl Engine { kind: GroupType::Fallback, providers: providers.iter().map(Member::named).collect(), selected: None, + balance_by: Default::default(), }); } // **默认路由必须永远存在。**判据是「有没有叫这个名字的路由」, @@ -985,6 +1124,14 @@ impl Engine { }); } } + // 按快慢、成败分只有 `load-balance` 用得上:别的类型写了也不起作用,而写的人 + // 以为它在起作用。控制面保存时就拦着(`control.group.balance_not_load_balance`) + if g.kind != GroupType::LoadBalance && !g.balance_by.is_weights() { + return Err(RouteError::GroupBalanceNotLoadBalance { + group: g.name.clone(), + by: g.balance_by, + }); + } } for set in &self.sets { self.check_rules(&set.rules)?; @@ -1626,6 +1773,7 @@ mod tests { kind: GroupType::Fallback, providers: vec!["official".into(), "relay".into()], selected: None, + balance_by: Default::default(), }; let e = Engine::with_default_rules( vec!["official".into(), "relay".into()], @@ -1646,6 +1794,7 @@ mod tests { kind: GroupType::Select, providers: vec!["a".into(), "b".into(), "c".into()], selected: Some("b".into()), + balance_by: Default::default(), }; let e = Engine::with_default_rules( vec!["a".into(), "b".into(), "c".into()], @@ -1683,6 +1832,7 @@ mod tests { kind: GroupType::Fallback, providers: vec![], selected: None, + balance_by: Default::default(), }; let e = Engine::with_default_rules(vec!["a".into()], vec![g], vec![route("x", "{}", "empty")]); @@ -2053,6 +2203,7 @@ mod tests { kind, providers: vec!["甲".into(), "乙".into(), "丙".into()], selected: None, + balance_by: Default::default(), } } @@ -2323,6 +2474,7 @@ mod builtin_tests { kind: GroupType::Fallback, providers: vec!["a".into()], selected: None, + balance_by: Default::default(), }; let e = Engine::with_default_rules( vec!["a".into()], @@ -2361,6 +2513,7 @@ mod builtin_tests { kind: GroupType::Fallback, providers: vec!["a".into(), "b".into(), "a".into()], selected: None, + balance_by: Default::default(), }; let e = Engine::with_default_rules( vec!["a".into(), "b".into()], @@ -2460,6 +2613,7 @@ mod pinned_tests { kind: GroupType::LoadBalance, providers: vec!["relay".into(), "anthropic".into()], selected: None, + balance_by: Default::default(), }], rules, ) @@ -2864,6 +3018,234 @@ to: } } +/// `load-balance` 按快慢、成败分:系数怎么算、写在哪种组上算错。 +#[cfg(test)] +mod balance_tests { + use super::*; + use std::collections::HashMap; + + fn names(v: &[&str]) -> Vec { + v.iter().map(|s| s.to_string()).collect() + } + + fn ttfb(v: &[(&str, u32)]) -> HashMap { + v.iter().map(|(n, t)| (n.to_string(), *t)).collect() + } + + fn success(v: &[(&str, f64)]) -> HashMap { + v.iter().map(|(n, s)| (n.to_string(), *s)).collect() + } + + fn close(got: &[f64], want: &[f64]) { + assert_eq!(got.len(), want.len(), "{got:?}"); + for (g, w) in got.iter().zip(want) { + assert!((g - w).abs() < 1e-9, "{got:?} ≠ {want:?}"); + } + } + + #[test] + fn weights_alone_leaves_every_member_at_one_whatever_was_measured() { + let f = Facts { + ttfb_ms: ttfb(&[("甲", 100), ("乙", 900)]), + success: success(&[("甲", 0.1)]), + ..Default::default() + }; + close( + &balance_factors(BalanceBy::Weights, &names(&["甲", "乙", "丙"]), &f), + &[1.0, 1.0, 1.0], + ); + } + + #[test] + fn latency_compares_each_member_with_the_median_squared() { + // 中位数 200:快一倍的分到四倍,慢一倍的四分之一 + let f = Facts { + ttfb_ms: ttfb(&[("甲", 100), ("乙", 200), ("丙", 400)]), + ..Default::default() + }; + close( + &balance_factors(BalanceBy::Latency, &names(&["甲", "乙", "丙"]), &f), + &[4.0, 1.0, 0.25], + ); + } + + #[test] + fn an_even_count_takes_the_middle_two_and_the_order_of_members_does_not_matter() { + // 中位数 (100 + 400) / 2 = 250 + let f = Facts { + ttfb_ms: ttfb(&[("甲", 400), ("乙", 100)]), + ..Default::default() + }; + close( + &balance_factors(BalanceBy::Latency, &names(&["甲", "乙"]), &f), + &[0.390625, 6.25], + ); + } + + #[test] + fn an_unmeasured_member_counts_as_average_and_does_not_move_the_median() { + // 新加的上游被试到,但不一上来就被灌满。**中位数只看测到了的成员**:组外那一家 + // 的数字也不算 + let f = Facts { + ttfb_ms: ttfb(&[("甲", 100), ("乙", 400), ("组外", 5)]), + ..Default::default() + }; + close( + &balance_factors(BalanceBy::Latency, &names(&["甲", "新来的", "乙"]), &f), + &[6.25, 1.0, 0.390625], + ); + // 一家都没测到:全是 1 + close( + &balance_factors(BalanceBy::Latency, &names(&["甲", "乙"]), &Facts::default()), + &[1.0, 1.0], + ); + } + + #[test] + fn the_latency_factor_is_clamped_between_a_tenth_and_ten() { + // 中位数 1000:快一百倍的本该是一万倍,慢一百倍的本该是万分之一 + let f = Facts { + ttfb_ms: ttfb(&[("快", 10), ("中", 1000), ("慢", 100_000)]), + ..Default::default() + }; + close( + &balance_factors(BalanceBy::Latency, &names(&["快", "中", "慢"]), &f), + &[10.0, 1.0, 0.1], + ); + // 0 毫秒按 1 毫秒算,不除以零 + let f = Facts { + ttfb_ms: ttfb(&[("零", 0), ("一", 1)]), + ..Default::default() + }; + close( + &balance_factors(BalanceBy::Latency, &names(&["零", "一"]), &f), + &[1.0, 1.0], + ); + } + + #[test] + fn health_squares_the_success_rate_and_keeps_a_floor() { + let f = Facts { + success: success(&[ + ("全成", 1.0), + ("一半", 0.5), + ("偶尔成", 0.2), + ("全败", 0.0), + ("坏数", f64::NAN), + ]), + ..Default::default() + }; + close( + &balance_factors( + BalanceBy::Health, + &names(&["全成", "一半", "偶尔成", "全败", "坏数", "没样本"]), + &f, + ), + // 0.2² = 0.04 也抬到下限:常失败的那家照样偶尔分到一个,恢复了才看得出来 + &[1.0, 0.25, 0.05, 0.05, 1.0, 1.0], + ); + } + + #[test] + fn latency_and_health_multiply() { + let f = Facts { + ttfb_ms: ttfb(&[("甲", 100), ("乙", 200), ("丙", 400)]), + success: success(&[("甲", 0.5), ("丙", 1.0)]), + ..Default::default() + }; + let members = names(&["甲", "乙", "丙"]); + close( + &balance_factors(BalanceBy::LatencyHealth, &members, &f), + &[1.0, 1.0, 0.25], + ); + // 单看一样时另一样不算进去 + close( + &balance_factors(BalanceBy::Health, &members, &f), + &[0.25, 1.0, 1.0], + ); + } + + #[test] + fn balance_by_reads_writes_and_is_left_out_when_it_is_the_default() { + let g: Group = serde_yaml_ng::from_str( + "name: g\ntype: load-balance\nproviders: [a, b]\nbalance_by: latency-health\n", + ) + .unwrap(); + assert_eq!(g.balance_by, BalanceBy::LatencyHealth); + let back = serde_yaml_ng::to_string(&g).unwrap(); + assert!(back.contains("balance_by: latency-health"), "{back}"); + for by in [ + BalanceBy::Weights, + BalanceBy::Latency, + BalanceBy::Health, + BalanceBy::LatencyHealth, + ] { + let text = format!( + "name: g\ntype: load-balance\nproviders: [a]\nbalance_by: {}\n", + by.slug() + ); + let g: Group = serde_yaml_ng::from_str(&text).unwrap(); + assert_eq!(g.balance_by, by); + } + + // 不写就是只按比例,写回时也不写:老配置原样往返 + let g: Group = + serde_yaml_ng::from_str("name: g\ntype: load-balance\nproviders: [a]\n").unwrap(); + assert_eq!(g.balance_by, BalanceBy::Weights); + let back = serde_yaml_ng::to_string(&g).unwrap(); + assert!(!back.contains("balance_by"), "{back}"); + + let e = serde_yaml_ng::from_str::( + "name: g\ntype: load-balance\nproviders: [a]\nbalance_by: speed\n", + ) + .unwrap_err(); + assert!(e.to_string().contains("speed"), "{e}"); + } + + fn engine_with(kind: GroupType, by: BalanceBy) -> Engine { + Engine::with_default_rules( + vec!["a".into(), "b".into()], + vec![Group { + name: "pool".into(), + kind, + providers: vec!["a".into(), "b".into()], + selected: None, + balance_by: by, + }], + vec![], + ) + } + + #[test] + fn only_a_load_balance_group_may_balance_by_latency_or_health() { + assert_eq!( + engine_with(GroupType::LoadBalance, BalanceBy::LatencyHealth).validate(), + Ok(()) + ); + for kind in [ + GroupType::Fallback, + GroupType::Select, + GroupType::UrlTest, + GroupType::Cheapest, + ] { + // 写明默认值不算错:它本来就什么都不做 + assert_eq!(engine_with(kind, BalanceBy::Weights).validate(), Ok(())); + let e = engine_with(kind, BalanceBy::Health).validate().unwrap_err(); + assert_eq!( + e, + RouteError::GroupBalanceNotLoadBalance { + group: "pool".into(), + by: BalanceBy::Health, + }, + "{kind:?}" + ); + let m = e.msg(); + assert_eq!(m.code, "engine.group_balance_not_load_balance"); + assert_eq!((m.arg("group"), m.arg("balance_by")), ("pool", "health")); + } + } +} + #[cfg(test)] mod msg_codes { use super::*; @@ -2895,6 +3277,10 @@ mod msg_codes { group: "g".into(), provider: "p".into(), }, + RouteError::GroupBalanceNotLoadBalance { + group: "g".into(), + by: BalanceBy::Latency, + }, RouteError::DuplicateRoute("x".into()), RouteError::UnknownDefaultRoute("x".into()), RouteError::UnknownRoute { diff --git a/crates/tw-engine/src/lib.rs b/crates/tw-engine/src/lib.rs index baf67c5d..38ab032a 100644 --- a/crates/tw-engine/src/lib.rs +++ b/crates/tw-engine/src/lib.rs @@ -7,9 +7,10 @@ pub mod weighted; pub use catalog::{Catalog, ProviderModels}; pub use engine::{ - ALL_UPSTREAMS, Asked, CATCH_ALL_RULE, DEFAULT_ROUTE, Decision, Engine, Facts, Group, GroupType, - Origin, Outcome, Outcome2, Pinned, RESERVED_PREFIX, RouteError, RouteSet, Rule, RuleNotes, - SetAction, Target, has_catch_all, is_builtin_group, notes, order_by, scalar_name, + ALL_UPSTREAMS, Asked, BalanceBy, CATCH_ALL_RULE, DEFAULT_ROUTE, Decision, Engine, Facts, Group, + GroupType, Origin, Outcome, Outcome2, Pinned, RESERVED_PREFIX, RouteError, RouteSet, Rule, + RuleNotes, SetAction, Target, balance_factors, has_catch_all, is_builtin_group, notes, + order_by, scalar_name, }; pub use facts::{RequestFacts, estimate_strings, estimate_tokens}; pub use weighted::Member; diff --git a/crates/tw-engine/src/weighted.rs b/crates/tw-engine/src/weighted.rs index 2e993bae..6c031753 100644 --- a/crates/tw-engine/src/weighted.rs +++ b/crates/tw-engine/src/weighted.rs @@ -5,6 +5,9 @@ //! 甲乙甲甲甲乙甲甲乙甲 —— 穿插着来,而不是先连着七次甲、再连着三次乙;权重都一样 //! 就是挨个轮。 //! +//! **轮的是有效权重**:成员的权重乘上 `balance_by` 按快慢、成败算出的系数 +//! ([`crate::engine::balance_factors`],只按权重分时全是 1),见 [`effective`]。 +//! //! **当前权重不在这里**:它由网关记着(`tw_gateway::balance`),排序时经 //! [`Facts::current_weight`] 传进来,排完、会话粘性也定了之后,网关按 [`advance`] //! 记账。引擎因此还是纯函数:试算拿同一份状态算、不记账,说的就是数据面下一个新对话 @@ -14,7 +17,7 @@ use std::collections::HashMap; use serde::{Deserialize, Serialize}; -use crate::engine::{Facts, Group}; +use crate::engine::{Facts, Group, balance_factors}; /// 权重最小是 1:0 等于把它从组里拿掉,那就该真的拿掉 pub const WEIGHT_MIN: u32 = 1; @@ -22,6 +25,19 @@ pub const WEIGHT_MIN: u32 = 1; /// 权重最大是 100:比例上够用,当前权重也不会大到溢出 pub const WEIGHT_MAX: u32 = 100; +/// 有效权重放大多少倍再取整。系数是小数(0.05、0.39…),当前权重是整数:直接取整的话, +/// 权重 1 乘 0.39 就成了 0。放大一千倍,小数点后三位都算数;系数都是 1 时各家同样 +/// 放大,轮出来的次序和不放大一模一样 +const SCALE: f64 = 1000.0; + +/// 一个成员这一轮的有效权重:权重 × 系数,放大 [`SCALE`] 倍取整,**最少是 1** —— +/// 系数再小,它也还在轮里,偶尔轮到一次(成败系数的下限也是这个意思)。 +/// +/// 最大是 100 × 10(快慢)× 1(成败)× 1000 = 一百万,当前权重攒得再多也远不到 i64 的边 +pub fn effective(weight: u32, factor: f64) -> i64 { + ((f64::from(weight) * factor * SCALE).round() as i64).max(1) +} + /// 策略组的一个成员:上游的名字,和它在 `load-balance` 里的权重。 /// /// 配置里写成名字(权重 1),或者写成 `{ name, weight }`。**权重是 1 的写回去还是 @@ -150,24 +166,30 @@ pub(crate) mod members { } } -/// 这一轮参加轮询的成员和它们的权重,按 `members` 的顺序。 +/// 这一轮参加轮询的成员和它们的有效权重([`effective`]),按 `members` 的顺序。 +/// +/// 系数按**全部候选**算([`balance_factors`],快慢比的是候选里测到的那几家的中位数), +/// 和试算列出来的那一份一样;算完再去掉停着的。 /// /// **停着的不算**([`Facts::paused`]):熔断、冷却着的那一家这一次本来就会被跳过, /// 轮到它的那一次落到组里排在它后面的那一家,那一家就平白多拿一份。全都停着时都算 /// —— 那时网关照样一家家试(fail-open),排头的还是按权重来。 fn round<'a>(g: &Group, members: &'a [String], f: &Facts) -> Vec<(&'a str, i64)> { - let up: Vec<&'a String> = members.iter().filter(|m| !f.paused.contains(*m)).collect(); - let pool = if up.is_empty() { - members.iter().collect() - } else { - up - }; - pool.into_iter() - .map(|m| (m.as_str(), i64::from(g.weight(m)))) - .collect() + let factors = balance_factors(g.balance_by, members, f); + let all: Vec<(&'a str, i64)> = members + .iter() + .zip(factors) + .map(|(m, x)| (m.as_str(), effective(g.weight(m), x))) + .collect(); + let up: Vec<(&'a str, i64)> = all + .iter() + .copied() + .filter(|(m, _)| !f.paused.contains(*m)) + .collect(); + if up.is_empty() { all } else { up } } -/// 下一个排头:当前权重加上自己的权重,最大的那个;一样大取组里靠前的。 +/// 下一个排头:当前权重加上自己的有效权重,最大的那个;一样大取组里靠前的。 /// /// `members` 是这次的候选(服务不了这个请求的已经去掉了):不在里面的成员这一轮不参加, /// 剩下的按各自的权重分。 @@ -182,12 +204,13 @@ pub fn lead<'a>(g: &Group, members: &'a [String], f: &Facts) -> Option<&'a str> best.map(|(m, _)| m) } -/// 记一次账,交回记过之后的当前权重:这一轮的每个成员加上自己的权重,`leader` 再减去 -/// 这一轮的总权重。 +/// 记一次账,交回记过之后的当前权重:这一轮的每个成员加上自己的有效权重,`leader` +/// 再减去这一轮的总和。 /// /// `leader` 是**会话粘性之后实际排头的那一家**,不一定是 [`lead`] 挑的那个:一段对话 /// 留在了上次回答它的那一家,这一次就记在那一家头上,之后的新对话把差的补回去。 -/// `members` 和 `f` 要和排序时的一样(同一轮)。`leader` 不在这一轮里时什么都不记。 +/// `members` 和 `f` 要和排序时的一样(同一轮;系数也就是排序时的那一份)。`leader` +/// 不在这一轮里时什么都不记。 pub fn advance(g: &Group, members: &[String], f: &Facts, leader: &str) -> HashMap { let round = round(g, members, f); let mut out = f.current_weight.clone(); @@ -205,7 +228,7 @@ pub fn advance(g: &Group, members: &[String], f: &Facts, leader: &str) -> HashMa #[cfg(test)] mod tests { use super::*; - use crate::engine::{GroupType, order_by}; + use crate::engine::{BalanceBy, GroupType, order_by}; fn group(members: &[(&str, u32)]) -> Group { Group { @@ -219,6 +242,7 @@ mod tests { }) .collect(), selected: None, + balance_by: Default::default(), } } @@ -373,4 +397,124 @@ mod tests { .to_string(); assert!(e.contains("name"), "{e}"); } + + /// 按 `by` 分的组:成员和权重同 [`group`] + fn balanced(members: &[(&str, u32)], by: BalanceBy) -> Group { + Group { + balance_by: by, + ..group(members) + } + } + + /// 排 `n` 次,数每一家排头了几次 + fn count(g: &Group, f: &mut Facts, n: usize) -> HashMap { + let mut out = HashMap::new(); + for m in run(g, &names(g), f, n) { + *out.entry(m).or_default() += 1; + } + out + } + + fn ttfb(v: &[(&str, u32)]) -> HashMap { + v.iter().map(|(n, t)| (n.to_string(), *t)).collect() + } + + /// 有效权重:权重 × 系数,放大一千倍取整,最少是 1 + #[test] + fn the_effective_weight_is_the_weight_times_the_factor() { + assert_eq!(effective(1, 1.0), 1000); + assert_eq!(effective(7, 1.0), 7000); + assert_eq!(effective(3, 0.390625), 1172, "1171.875 取整"); + assert_eq!(effective(100, 10.0), 1_000_000); + // 系数再小也还在轮里 + assert_eq!(effective(1, 0.0), 1); + assert_eq!(effective(1, 1e-9), 1); + } + + /// 按快慢分:首字节快的那一家分到的新对话多。一整圈(有效权重之和那么多次)下来, + /// 各家排头的次数正好是各自的有效权重,当前权重回到零 + #[test] + fn latency_gives_the_faster_member_more_new_conversations() { + let g = balanced(&[("甲", 1), ("乙", 1)], BalanceBy::Latency); + // 中位数 250:甲 (250/100)² = 6.25,乙 (250/400)² = 0.390625 + let mut f = Facts { + ttfb_ms: ttfb(&[("甲", 100), ("乙", 400)]), + ..Default::default() + }; + let got = count(&g, &mut f, 6250 + 391); + assert_eq!((got["甲"], got["乙"]), (6250, 391)); + assert!(f.current_weight.values().all(|v| *v == 0), "{f:?}"); + // 只按权重分时,同样的样本不起作用:挨个轮 + let g = balanced(&[("甲", 1), ("乙", 1)], BalanceBy::Weights); + assert_eq!(run(&g, &names(&g), &mut f, 4), ["甲", "乙", "甲", "乙"]); + } + + /// 按成败分:常失败的那一家分得少,但还在轮里 —— 成功率 0.1 的系数本该是 0.01, + /// 抬到下限 0.05,每 21 个新对话里有它一个,它恢复了才看得出来 + #[test] + fn health_demotes_a_flaky_member_but_keeps_it_in_rotation() { + let g = balanced(&[("稳", 1), ("抖", 1)], BalanceBy::Health); + let mut f = Facts { + success: [("稳".to_string(), 1.0), ("抖".to_string(), 0.1)] + .into_iter() + .collect(), + ..Default::default() + }; + let firsts = run(&g, &names(&g), &mut f, 1050); + assert_eq!(firsts.iter().filter(|m| *m == "抖").count(), 50); + for window in firsts.chunks(21) { + assert!(window.iter().any(|m| m == "抖"), "{window:?}"); + } + // 一次都没成功过的也一样:停用它是熔断的事,这里只是少分 + f.success.insert("抖".into(), 0.0); + f.current_weight.clear(); + let got = count(&g, &mut f, 1050); + assert_eq!(got["抖"], 50); + } + + /// 权重和系数相乘:3:1 的组,慢的那一家(权重 3)系数 0.39,快的(权重 1)系数 6.25 + /// —— 有效权重 1172 : 6250,快的那一家反过来分得多 + #[test] + fn weights_and_factors_multiply() { + let g = balanced(&[("慢", 3), ("快", 1)], BalanceBy::LatencyHealth); + let mut f = Facts { + ttfb_ms: ttfb(&[("慢", 400), ("快", 100)]), + ..Default::default() + }; + let got = count(&g, &mut f, 1172 + 6250); + assert_eq!((got["慢"], got["快"]), (1172, 6250)); + // 再加上成败:快的那一家成功率 0.5,系数再乘 0.25 —— 6.25 × 0.25 × 1000 = 1563 + let mut f = Facts { + ttfb_ms: ttfb(&[("慢", 400), ("快", 100)]), + success: [("快".to_string(), 0.5)].into_iter().collect(), + ..Default::default() + }; + let got = count(&g, &mut f, 1172 + 1563); + assert_eq!((got["慢"], got["快"]), (1172, 1563)); + } + + /// 各家系数一样时,次序和只按权重分一模一样:7:3 照样穿插着来,不是先七后三 + #[test] + fn equal_factors_keep_the_exact_interleaving() { + let want = ["甲", "乙", "甲", "甲", "甲", "乙", "甲", "甲", "乙", "甲"]; + let cases = [ + // 首字节一样:系数都是 1 + Facts { + ttfb_ms: ttfb(&[("甲", 200), ("乙", 200)]), + ..Default::default() + }, + // 成功率一样:系数都是 0.25,比例不变 + Facts { + success: [("甲".to_string(), 0.5), ("乙".to_string(), 0.5)] + .into_iter() + .collect(), + ..Default::default() + }, + ]; + for mut f in cases { + let g = balanced(&[("甲", 7), ("乙", 3)], BalanceBy::LatencyHealth); + assert_eq!(run(&g, &names(&g), &mut f, 10), want, "{f:?}"); + assert!(f.current_weight.values().all(|v| *v == 0), "{f:?}"); + } + } } diff --git a/crates/tw-gateway/src/balance.rs b/crates/tw-gateway/src/balance.rs index 356a202b..53b1727b 100644 --- a/crates/tw-gateway/src/balance.rs +++ b/crates/tw-gateway/src/balance.rs @@ -105,6 +105,7 @@ mod tests { }) .collect(), selected: None, + balance_by: Default::default(), } } diff --git a/crates/tw-gateway/src/health.rs b/crates/tw-gateway/src/health.rs index 5cdbfe45..098e839a 100644 --- a/crates/tw-gateway/src/health.rs +++ b/crates/tw-gateway/src/health.rs @@ -15,8 +15,11 @@ //! 看见真实错误,也不要返回一个我们自己编的「无可用上游」。 //! 2. **只有一个候选时完全旁路熔断器。**否则唯一的上游一旦被自己熔断, //! 就把用户锁死了,熔断纯粹是自伤。 +//! +//! 同一批成败还记成每家最近的成功率([`Health::success_rates`]),`load-balance` +//! 按成败分新对话时看它。 -use std::collections::HashMap; +use std::collections::{HashMap, VecDeque}; use std::sync::{Arc, Mutex}; use std::time::Duration; @@ -39,13 +42,58 @@ struct Entry { held_until_ms: Option, } +/// 成功率看最近多少次。 +const RECENT_MAX: usize = 50; +/// 也只看最近这么久(毫秒):**一家恢复了,半小时前的失败不该还压着它的份额**。 +/// 请求稀疏的时候,按次数的窗口会把很久以前的事一直留着 +const RECENT_MAX_AGE_MS: i64 = 30 * 60 * 1000; +/// 少于这么多次就说不出成功率。头一两次失败可能是碰巧,不该就此少分到新对话 +const RECENT_MIN_SAMPLES: u32 = 5; + +/// 每家最近的成败:最近 [`RECENT_MAX`] 次、[`RECENT_MAX_AGE_MS`] 以内。 +/// +/// **判据就是熔断的判据**:`record_*` 记一次,这里也记一次(见 [`Health::success_rates`])。 +#[derive(Default)] +struct Recent(HashMap>); + +impl Recent { + fn record(&mut self, name: &str, ok: bool, now: i64) { + let w = self.0.entry(name.to_string()).or_default(); + if w.len() == RECENT_MAX { + w.pop_front(); + } + w.push_back((now, ok)); + } + + /// 窗口里的统计。过了时的不算 + fn tally(&self, name: &str, now: i64) -> Tally { + self.0 + .get(name) + .into_iter() + .flatten() + .filter(|(at, _)| now.saturating_sub(*at) < RECENT_MAX_AGE_MS) + .fold(Tally::default(), |t, (_, ok)| Tally { + total: t.total + 1, + errors: t.errors + u32::from(!ok), + }) + } + + /// 成功率,0 到 1。样本不够是 `None` + fn rate(&self, name: &str, now: i64) -> Option { + let t = self.tally(name, now); + (t.total >= RECENT_MIN_SAMPLES).then(|| f64::from(t.total - t.errors) / f64::from(t.total)) + } +} + /// 每个 provider 的健康状态。 /// /// **不持久化**。桌面应用重启频繁,把「这家挂了」的判断带过重启 /// 意味着用户重启后第一个请求还在被上一次的故障惩罚 —— 而重启本身往往 -/// 就是他为了解决问题做的事。 +/// 就是他为了解决问题做的事。最近的成功率也一样:重启之后从头攒。 pub struct Health { map: Mutex>, + /// 每家最近的成败([`Self::success_rates`]) + recent: Mutex, /// 配置里的 `failover`。**跟着配置换**([`Self::configure`]),状态不丢 settings: Mutex, clock: Clock, @@ -65,6 +113,7 @@ impl Health { fn with_clock(clock: Clock) -> Self { Self { map: Mutex::new(HashMap::new()), + recent: Mutex::new(Recent::default()), settings: Mutex::new(tw_config::Failover::default()), clock, } @@ -141,6 +190,7 @@ impl Health { /// 调用方要拿它去发事件:熔断开合是界面上看得见的状态,而看得见的 /// 状态必须能被推出去 —— 否则界面只能轮询。 pub fn record_success(&self, name: &str) -> Option { + self.note(name, true); self.update(name, |e, s, now| { e.breaker .record(true, Tally::default(), &Self::policy(s, e.trips), now); @@ -151,6 +201,7 @@ impl Health { /// 记一次说不出原因的失败,同样返回状态变化。 pub fn record_failure(&self, name: &str) -> Option { + self.note(name, false); self.update(name, |e, s, now| { let policy = Self::policy(s, e.trips); let before = e.breaker.state_at(&policy, now); @@ -183,11 +234,42 @@ impl Health { return self.record_failure(name); } }; + self.note(name, false); self.update(name, |e, _, _| { e.held_until_ms = Some(e.held_until_ms.map_or(until, |t| t.max(until))); }) } + /// 记进最近的成败 + fn note(&self, name: &str, ok: bool) { + let now = (self.clock)(); + if let Ok(mut r) = self.recent.lock() { + r.record(name, ok, now); + } + } + + /// 这几家最近的成功率,0 到 1:最近 50 次、30 分钟以内,至少 5 次才算。**样本不够的 + /// 不在里面**,不是「从不失败」。`load-balance` 按成败分新对话时用它 + /// (`tw_engine::balance_factors`)。 + /// + /// 成败的判据就是熔断的判据,不另起一套:上面几个 `record_*` 记一次,这里就记一次。 + /// 于是 5xx、连不上、超时、限流、额度用完、没钱了、凭据被拒或取不到、流在第一段内容 + /// 之前断了都算失败;请求本身的问题(别的 4xx)算这家答上了;这家没有这个模型不记 —— + /// 它对别的模型照样好好的。客户端中途走了的不经过这里,也不记。 + /// + /// **换下一家、却不是这家的错的**(它只是慢,或者正忙)不该调 `record_*`:调了,它在 + /// 熔断和这里都会被算成一次失败。 + pub fn success_rates(&self, names: &[String]) -> HashMap { + let now = (self.clock)(); + let Ok(r) = self.recent.lock() else { + return HashMap::new(); + }; + names + .iter() + .filter_map(|n| r.rate(n, now).map(|s| (n.clone(), s))) + .collect() + } + /// 这家还要等多久才轮到探测。 /// /// `None`:没停用过(或者停用之后已经成功过)。`Some(0)`:到点了。 @@ -505,4 +587,91 @@ mod tests { h.record_cause("b", Cause::NoBalance); assert_eq!(h.cooldown_left("b"), Some(Duration::from_secs(7))); } + + fn rate(h: &Health, name: &str) -> Option { + h.success_rates(&names(&[name])).get(name).copied() + } + + #[test] + fn a_success_rate_needs_five_outcomes_and_looks_at_the_last_fifty() { + let (h, _) = clocked(); + for _ in 0..4 { + h.record_failure("a"); + } + assert_eq!(rate(&h, "a"), None, "四次还说不出成功率"); + h.record_success("a"); + assert_eq!(rate(&h, "a"), Some(0.2)); + // 攒满 50 次:四次失败还在窗口里 + for _ in 0..45 { + h.record_success("a"); + } + assert_eq!(rate(&h, "a"), Some(46.0 / 50.0)); + // 再来四次成功,最早那四次失败被挤出窗口 + for _ in 0..4 { + h.record_success("a"); + } + assert_eq!(rate(&h, "a"), Some(1.0)); + } + + #[test] + fn outcomes_older_than_half_an_hour_drop_out_so_a_recovered_upstream_regains_its_share() { + let (h, now) = clocked(); + for _ in 0..10 { + h.record_failure("a"); + } + assert_eq!(rate(&h, "a"), Some(0.0)); + advance(&now, 29 * 60 * 1000); + for _ in 0..5 { + h.record_success("a"); + } + assert_eq!(rate(&h, "a"), Some(5.0 / 15.0)); + // 那十次失败满半小时了:只剩后来的五次成功 + advance(&now, 60 * 1000); + assert_eq!(rate(&h, "a"), Some(1.0)); + advance(&now, 30 * 60 * 1000); + assert_eq!(rate(&h, "a"), None, "全都过了时,等于没有样本"); + } + + #[test] + fn the_success_rate_counts_what_the_breaker_counts() { + let (h, _) = clocked(); + // 上游的问题:说了原因的、没说原因的,都算失败 + for cause in [ + Cause::NoBalance, + Cause::QuotaUsedUp { resets_at_ms: None }, + Cause::RateLimited { + retry_after: Some(Duration::from_secs(30)), + }, + Cause::RateLimited { retry_after: None }, + Cause::AuthRejected, + Cause::Unexplained, + ] { + h.record_cause("a", cause); + } + h.record_failure("a"); + assert_eq!(rate(&h, "a"), Some(0.0), "七次失败"); + // 这家没有这个模型:不算它坏了,也不算它答上了 + for _ in 0..10 { + h.record_cause("b", Cause::ModelUnavailable); + } + assert_eq!(rate(&h, "b"), None); + for _ in 0..3 { + h.record_success("a"); + } + assert_eq!(rate(&h, "a"), Some(0.3)); + } + + #[test] + fn success_rates_lists_only_the_asked_ones_with_enough_samples() { + let (h, _) = clocked(); + for _ in 0..5 { + h.record_success("a"); + h.record_failure("b"); + h.record_success("组外"); + } + h.record_success("c"); + let got = h.success_rates(&names(&["a", "b", "c", "d"])); + assert_eq!(got.len(), 2, "{got:?}"); + assert_eq!((got["a"], got["b"]), (1.0, 0.0)); + } } diff --git a/crates/tw-gateway/src/server/pipeline.rs b/crates/tw-gateway/src/server/pipeline.rs index cd010133..cbd6bdc4 100644 --- a/crates/tw-gateway/src/server/pipeline.rs +++ b/crates/tw-gateway/src/server/pipeline.rs @@ -676,7 +676,7 @@ fn route( { let f = state.group_facts( &rt.config.providers, - g.kind, + g, &decision.candidates, &crate::sent::pairs(&sent), turn.as_ref().map(crate::balance::Turn::current), diff --git a/crates/tw-gateway/src/state.rs b/crates/tw-gateway/src/state.rs index dc2bc341..fc122d6d 100644 --- a/crates/tw-gateway/src/state.rs +++ b/crates/tw-gateway/src/state.rs @@ -363,7 +363,7 @@ impl AppState { self.pricing.rcu(|book| book.with_table(table.clone())); } - /// 策略组排序要的运行时数字([`tw_engine::Facts`]),只取 `kind` 用得上的那几样。 + /// 策略组排序要的运行时数字([`tw_engine::Facts`]),只取组 `g` 用得上的那几样。 /// /// **数据面和试算共用这一个**,喂给同一个 `order`:试算说会排给谁,数据面就排给谁。 /// `sent` 是每一家和发给它的模型名(比价按它算);`current_weight` 是 `load-balance` @@ -372,28 +372,38 @@ impl AppState { pub fn group_facts( &self, providers: &[tw_config::Provider], - kind: tw_engine::GroupType, + g: &tw_engine::Group, candidates: &[String], sent: &[(String, String)], current_weight: Option>, ) -> tw_engine::Facts { use tw_engine::GroupType; + let balanced = g.kind == GroupType::LoadBalance; tw_engine::Facts { current_weight: current_weight.unwrap_or_default(), // 熔断着、冷却着的不参加这一轮:它们反正会被跳过 - paused: match kind { - GroupType::LoadBalance => candidates + paused: if balanced { + candidates .iter() .filter(|p| !self.health.is_available(p)) .cloned() - .collect(), - _ => Default::default(), + .collect() + } else { + Default::default() }, - ttfb_ms: match kind { - GroupType::UrlTest => self.latency.snapshot(candidates), - _ => Default::default(), + // `url-test` 选最快的;`load-balance` 按快慢分的也看它 + ttfb_ms: if g.kind == GroupType::UrlTest || (balanced && g.balance_by.uses_latency()) { + self.latency.snapshot(candidates) + } else { + Default::default() + }, + // `load-balance` 按成败分时看每家最近的成功率 + success: if balanced && g.balance_by.uses_health() { + self.health.success_rates(candidates) + } else { + Default::default() }, - price: match kind { + price: match g.kind { // 每一家按发给它的名字算价钱:同一个别名在各家是各家的模型名、各家的价目 GroupType::Cheapest => self.unit_prices(providers, sent, candidates), _ => Default::default(), diff --git a/crates/tw-gateway/tests/affinity.rs b/crates/tw-gateway/tests/affinity.rs index 2716dce3..e06a6066 100644 --- a/crates/tw-gateway/tests/affinity.rs +++ b/crates/tw-gateway/tests/affinity.rs @@ -61,6 +61,7 @@ fn cfg(up: SocketAddr) -> Config { kind: tw_engine::GroupType::LoadBalance, providers: vec!["甲".into(), "乙".into()], selected: None, + balance_by: Default::default(), }], routes: vec![RouteSet::default_with(vec![ rule("大输入", "{ input_tokens: \">2000\" }", "乙"), diff --git a/crates/tw-gateway/tests/client_model_name.rs b/crates/tw-gateway/tests/client_model_name.rs index 29bab4d0..e21c0395 100644 --- a/crates/tw-gateway/tests/client_model_name.rs +++ b/crates/tw-gateway/tests/client_model_name.rs @@ -735,6 +735,7 @@ async fn after_failover_the_name_the_answering_hop_sent_decides() { kind: tw_engine::GroupType::Fallback, providers: vec!["official".into(), "relay".into()], selected: None, + balance_by: Default::default(), }], routes, ) diff --git a/crates/tw-gateway/tests/passthrough.rs b/crates/tw-gateway/tests/passthrough.rs index 5fc45305..b16b781e 100644 --- a/crates/tw-gateway/tests/passthrough.rs +++ b/crates/tw-gateway/tests/passthrough.rs @@ -1197,6 +1197,7 @@ async fn a_phase_two_rule_is_recomputed_after_failover() { kind: tw_engine::GroupType::Fallback, providers: vec!["official".into(), "relay".into()], selected: None, + balance_by: Default::default(), }]; cfg.routes[0].rules.last_mut().unwrap().to = Some("全部".into()); @@ -1334,6 +1335,7 @@ async fn a_phase_one_set_applies_on_every_attempt_including_after_failover() { kind: tw_engine::GroupType::Fallback, providers: vec!["dead".into(), "good".into()], selected: None, + balance_by: Default::default(), }]; let gw = serve_cfg(cfg).await; let r = reqwest::Client::new() @@ -2088,6 +2090,7 @@ async fn the_attempt_chain_records_every_hop_and_why_each_one_failed() { kind: tw_engine::GroupType::Fallback, providers: vec!["挂了的".into(), "限流的".into(), "好的".into()], selected: None, + balance_by: Default::default(), }]; cfg.routes = vec![tw_engine::RouteSet::default_with(vec![tw_engine::Rule { name: "都走这一组".into(), diff --git a/crates/tw-gateway/tests/routing_facts.rs b/crates/tw-gateway/tests/routing_facts.rs index 2b570a45..f300dc1a 100644 --- a/crates/tw-gateway/tests/routing_facts.rs +++ b/crates/tw-gateway/tests/routing_facts.rs @@ -79,6 +79,7 @@ fn cfg(providers: Vec, rules: Vec) -> Config { kind: tw_engine::GroupType::Fallback, providers: vec!["up".into()], selected: None, + balance_by: Default::default(), }], routes: vec![ RouteSet { diff --git a/crates/tw-gateway/tests/success_rate.rs b/crates/tw-gateway/tests/success_rate.rs new file mode 100644 index 00000000..9b9f60fe --- /dev/null +++ b/crates/tw-gateway/tests/success_rate.rs @@ -0,0 +1,195 @@ +//! 每家最近的成功率(`load-balance` 按成败分新对话时看它),端到端。 +//! +//! 盯的是**哪些结果算这家的失败**:上游自己的问题算(5xx、限流、连不上、流在第一段 +//! 内容之前报错),请求本身的问题不算、算它答上了,这家没有这个模型、客户端中途走了 +//! 都不记。判据和熔断是同一套(`tw_gateway::health`),这里从真请求一路看过去。 + +use std::net::SocketAddr; +use std::sync::Arc; +use std::time::Duration; + +use axum::Router; +use axum::routing::any; +use tw_config::{Client, Config, Failover, Provider}; +use tw_gateway::Health; + +mod common; + +/// 每次都回同一个状态码、类型和正文的上游。 +async fn answering(status: u16, content_type: &'static str, body: &'static str) -> SocketAddr { + listen(Router::new().fallback(any(move || async move { + axum::response::Response::builder() + .status(status) + .header("content-type", content_type) + .body(axum::body::Body::from(body)) + .unwrap() + }))) + .await +} + +/// 一直不回响应头的上游。 +async fn silent() -> SocketAddr { + listen(Router::new().fallback(any(|| async { + tokio::time::sleep(Duration::from_secs(60)).await; + "late" + }))) + .await +} + +async fn listen(app: Router) -> SocketAddr { + let l = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap(); + let a = l.local_addr().unwrap(); + tokio::spawn(async move { axum::serve(l, app).await.unwrap() }); + a +} + +fn provider(name: &str, at: SocketAddr) -> Provider { + Provider { + name: name.into(), + base_url: format!("http://{at}"), + key: Some("k".into()), + ..Default::default() + } +} + +/// 起一个网关,交回它的地址和健康状态。**停用的门槛调高**:这里要数每一次的成败, +/// 不想让熔断把失败的那家从候选里拿掉 +async fn gateway(providers: Vec) -> (SocketAddr, Arc) { + let cfg = Config { + clients: vec![Client { + name: "c".into(), + key: "tw-k".into(), + ..Default::default() + }], + providers, + failover: Failover { + failures_to_pause: 100, + ..Default::default() + }, + ..Default::default() + }; + let state = tw_gateway::AppState::new(cfg).unwrap(); + let health = state.health.clone(); + let gw = tw_gateway::serve(state, ([127, 0, 0, 1], 0).into()) + .await + .unwrap(); + tokio::time::sleep(Duration::from_millis(50)).await; + (gw, health) +} + +/// 发 `n` 个请求,每个最多等 `wait`:等不到就放弃(连接随之关闭) +async fn send(gw: SocketAddr, body: &str, n: usize, wait: Duration) { + for _ in 0..n { + let r = tokio::time::timeout( + wait, + reqwest::Client::new() + .post(format!("http://{gw}/v1/messages")) + .header("x-api-key", "tw-k") + .body(body.to_string()) + .send(), + ) + .await; + if let Ok(r) = r { + // 把正文读完:流式回答的那一家要等到这时才算交完 + let _ = r.unwrap().bytes().await; + } + } +} + +const PLAIN: &str = r#"{"model":"claude-sonnet-4-5"}"#; + +fn rate(h: &Health, name: &str) -> Option { + h.success_rates(&[name.to_string()]).get(name).copied() +} + +/// 只有这一家时,回 `status` 五次之后它的成功率。 +async fn after_five(status: u16, body: &'static str) -> Option { + let up = answering(status, "application/json", body).await; + let (gw, health) = gateway(vec![provider("only", up)]).await; + send(gw, PLAIN, 5, Duration::from_secs(10)).await; + rate(&health, "only") +} + +#[tokio::test] +async fn an_upstream_that_answers_counts_as_a_success() { + assert_eq!(after_five(200, r#"{"content":[]}"#).await, Some(1.0)); +} + +#[tokio::test] +async fn server_errors_and_rate_limits_count_against_the_upstream() { + assert_eq!(after_five(503, "{}").await, Some(0.0)); + assert_eq!(after_five(429, "{}").await, Some(0.0)); + // 凭据被拒:换一家有意义,那边是另一把密钥 + assert_eq!(after_five(401, "{}").await, Some(0.0)); +} + +#[tokio::test] +async fn a_bad_request_is_the_requests_fault_and_counts_as_answered() { + let got = after_five( + 400, + r#"{"type":"error","error":{"type":"invalid_request_error","message":"max_tokens: too large"}}"#, + ) + .await; + assert_eq!(got, Some(1.0)); +} + +#[tokio::test] +async fn a_missing_model_is_not_counted_either_way() { + // 它对别的模型照样好好的 + let got = after_five( + 404, + r#"{"type":"error","error":{"type":"not_found_error","message":"model not found"}}"#, + ) + .await; + assert_eq!(got, None); +} + +#[tokio::test] +async fn an_upstream_that_cannot_be_reached_counts_against_it() { + let nobody = SocketAddr::from(([127, 0, 0, 1], common::spare_port())); + let (gw, health) = gateway(vec![provider("gone", nobody)]).await; + send(gw, PLAIN, 5, Duration::from_secs(10)).await; + assert_eq!(rate(&health, "gone"), Some(0.0)); +} + +#[tokio::test] +async fn a_client_that_walks_away_leaves_no_mark_on_the_upstream() { + let (gw, health) = gateway(vec![provider("slow", silent().await)]).await; + send(gw, PLAIN, 5, Duration::from_millis(300)).await; + // 给网关一点时间:丢掉的 handler 要是还记了什么,这时也该记上了 + tokio::time::sleep(Duration::from_millis(200)).await; + assert_eq!(rate(&health, "slow"), None); +} + +#[tokio::test] +async fn a_stream_that_reports_an_error_before_any_content_counts_against_it() { + // 第一家开了流、第一个事件就是过载:换到第二家。前者算失败,后者算答上了 + let overloaded = answering( + 200, + "text/event-stream", + "event: message_start\ndata: {\"type\":\"message_start\",\"message\":{}}\n\n\ + event: error\ndata: {\"type\":\"error\",\"error\":{\"type\":\"overloaded_error\",\"message\":\"Overloaded\"}}\n\n", + ) + .await; + let good = answering( + 200, + "text/event-stream", + "event: message_start\ndata: {\"type\":\"message_start\",\"message\":{}}\n\n\ + event: content_block_delta\ndata: {\"type\":\"content_block_delta\",\"delta\":{\"text\":\"answer\"}}\n\n", + ) + .await; + let (gw, health) = gateway(vec![ + provider("first", overloaded), + provider("second", good), + ]) + .await; + send( + gw, + r#"{"model":"claude-sonnet-4-5","stream":true}"#, + 5, + Duration::from_secs(10), + ) + .await; + assert_eq!(rate(&health, "first"), Some(0.0)); + assert_eq!(rate(&health, "second"), Some(1.0)); +} diff --git a/crates/tw-gateway/tests/ws.rs b/crates/tw-gateway/tests/ws.rs index c6248066..8b348afc 100644 --- a/crates/tw-gateway/tests/ws.rs +++ b/crates/tw-gateway/tests/ws.rs @@ -292,6 +292,7 @@ fn routed_to_an_account(up: SocketAddr) -> Config { kind: tw_engine::GroupType::Fallback, providers: vec!["订阅账号".into()], selected: None, + balance_by: Default::default(), }], routes: vec![tw_engine::RouteSet::default_with(vec![tw_engine::Rule { name: "Codex 走账号".into(), diff --git a/docs/config.md b/docs/config.md index 47bd55b6..ff388656 100644 --- a/docs/config.md +++ b/docs/config.md @@ -955,6 +955,7 @@ group with `to`. | `type` | `fallback` \| `select` \| `load-balance` \| `url-test` \| `cheapest` | `fallback` | `fallback`: the first healthy member, in order. `select`: the member named in `selected`. `load-balance`: new conversations take turns, in proportion to the members' weights. `url-test`: the fastest by measured time to first byte. `cheapest`: the lowest input price. | | `providers` | list of strings or [`groups[].providers[]`](#cfg-groups-providers) | **required** | Member upstreams, by name; not groups. Each upstream appears once in a group. In a `load-balance` group, a member can be written as `{name, weight}`. | | `selected` | string | — | For `select`: the chosen member. | +| `balance_by` | `weights` \| `latency` \| `health` \| `latency-health` | `weights` | For `load-balance`: what the members' weights are multiplied by. `weights`: nothing; the weights alone. `latency`: faster upstreams get more. `health`: upstreams that fail less get more. `latency-health`: both. Other group types take only `weights`. | `fallback` is the default because a single user's machine has no load to @@ -1001,6 +1002,39 @@ unless its input no longer fits the context window of a model the rule sends it to. `load-balance` therefore takes turns between new conversations, by weight. +`balance_by` lets a `load-balance` group also look at how each upstream has +been doing lately. Each member's weight is multiplied by a factor, and the +group shares out requests by the result in the same way as above. + +- `weights` (the default): the weights alone. +- `latency`: faster upstreams get a larger share. Speed is the typical time + to first byte, the same measurement `url-test` uses. An upstream twice as + fast as the middle of the group has its weight multiplied by four, by at + most ten and at least a tenth. +- `health`: upstreams that fail less get a larger share. It looks at the + last 50 requests within the past 30 minutes. Server errors, rate limits, + used-up quota or balance, rejected credentials, timeouts and connection + errors count as failures; errors caused by the request itself do not, and + neither does a client that cancels. An upstream that keeps failing keeps a + twentieth of its weight, so it still gets the occasional new conversation + and its recovery is noticed; one that fails outright is set aside by + [`failover`](#cfg-failover) as before. +- `latency-health`: both factors, multiplied. + +An upstream without enough measurements yet counts as average. As with +weights alone, conversations in progress stay where they are, and new +conversations make up the difference. + +```yaml +groups: + - name: pool + type: load-balance + balance_by: latency-health + providers: + - { name: official, weight: 3 } + - relay +``` + ### `routes` A route is a list of rules evaluated top to bottom. Each key takes the route diff --git a/docs/config.zh-CN.md b/docs/config.zh-CN.md index d247202c..fdd50299 100644 --- a/docs/config.zh-CN.md +++ b/docs/config.zh-CN.md @@ -750,6 +750,7 @@ aliases: | `type` | `fallback` \| `select` \| `load-balance` \| `url-test` \| `cheapest` | `fallback` | `fallback`:按顺序取第一个健康的。`select`:取 `selected` 指定的那个。`load-balance`:新对话按成员的权重轮流。`url-test`:按实测首字节时间取最快的。`cheapest`:取输入单价最低的。 | | `providers` | 列表,每项是字符串或对象,对象见 [`groups[].providers[]`](#cfg-groups-providers) | **必填** | 成员上游的名字,不能是策略组。同一个上游在一个策略组中只出现一次。`load-balance` 组的成员可以写成 `{name, weight}`。 | | `selected` | 字符串 | — | `select` 类型选中的成员。 | +| `balance_by` | `weights` \| `latency` \| `health` \| `latency-health` | `weights` | `load-balance` 类型用:成员的权重再乘上什么。`weights`:不乘,只按权重。`latency`:越快的上游分得越多。`health`:越少失败的上游分得越多。`latency-health`:两者都看。其他类型只能是 `weights`。 | 默认类型为 `fallback`:单个使用者的机器上没有需要分散的负载。 @@ -776,6 +777,25 @@ groups: 无论哪种类型,一段对话都留在上次回答它的那一家上游,让上游缓存着的那部分被再次读取,而不是换一家全价重算。同一轮之内(客户端正在回传工具结果)一律不换;跨轮时,上一次回答读或写了至少 1024 个 token 的 prompt cache、且距今不到五分钟,才继续留下。上游因失败进入冷却时,对话随之放开;故障转移之后接下回答的那一家,就是之后留下的那一家。一轮开始时命中的规则也沿用到这一轮结束:按输入大小或图片分流的规则不会让一轮半路换家,除非输入已经超出规则所指模型的上下文窗口。因此 `load-balance` 按权重轮流的是新对话。 +`balance_by` 让 `load-balance` 组再看各上游最近的表现:每个成员的权重乘上一个系数,组内请求按乘出来的结果照上文的方式分。 + +- `weights`(默认):只按权重。 +- `latency`:越快的上游分得越多。快慢看典型的首字节时间,与 `url-test` 使用同一份测量。比组内居中者快一倍的上游,权重乘以四;最多乘以十,最少乘以十分之一。 +- `health`:越少失败的上游分得越多。依据是最近 30 分钟内的最近 50 次请求:服务器错误、限流、额度或余额用尽、凭据被拒、超时和连接失败算作失败;请求本身导致的错误不算,客户端取消也不算。经常失败的上游至少保留权重的二十分之一,仍会偶尔分到新对话,以便发现它已经恢复;完全失败的上游照旧由 [`failover`](#cfg-failover) 暂停。 +- `latency-health`:两个系数相乘。 + +测量还不够的上游按中等对待。与只按权重时一样,进行中的对话留在原来的上游,差额由新对话补齐。 + +```yaml +groups: + - name: 均摊 + type: load-balance + balance_by: latency-health + providers: + - { name: 官方, weight: 3 } + - 中转 +``` + ### `routes` 一条路由是一组自上而下求值的规则。每把密钥使用其 `route` 指定的路由;没有指定时用 `default_route`;再没有时用名为 `default` 的路由;一条路由都没有时,请求按上游的声明顺序故障转移。 From dabcf5465d1a1c3b339be1a859352e4ec2569df8 Mon Sep 17 00:00:00 2001 From: fylorn <249551762+fylorn@users.noreply.github.com> Date: Mon, 5 Oct 2026 16:03:45 +0800 Subject: [PATCH 03/22] Move a stream that is slow to start on to the next upstream Some upstreams accept a request and then send nothing for a long time: a relay queueing it, an overloaded provider that does not say so, response headers that do not come. The client sees a hung request while another candidate could already be answering. failover.next_on_slow_start (off by default) gives up on such an upstream when a streamed answer has no content stream_start_wait_secs after the request was sent, covering both the wait for response headers and the held start of the stream. What counts as content is what the stream-start hold already judges (text, thinking/reasoning, tool calls; not role or start frames). Giving up drops the pending request or the response, so the connection closes and the upstream stops generating. - The last candidate never switches. "Last" accounts for the candidates that could not take the request anyway (paused by then, or the hop could not be sent: phase-two denial, no servable name, conversion failure, forced server tool), so the effectively last one waits instead of trading a slow answer for a certain failure. - The slow upstream is neither paused nor counted as a failure, nor as an outcome in the success rate load-balance groups weigh by. - The abandoned try is an attempt with outcome slow_start, with the usage the upstream reported at the start (Anthropic message_start) or else the gateway's input estimate marked as estimated (AttemptView.usage). - Only requests the client asked to stream; a non-streamed request is unaffected, and with the setting off the hold behaves as before. - Validation refuses switching with a wait under 5 seconds (config.slow_start_too_short); FailoverView.next_on_slow_start shows it. No keepalive is sent while holding: response headers are only sent once an upstream is chosen, and committing a 200 early would turn every later 429, 5xx and the last upstream's own 4xx into in-stream errors that clients do not retry on; Google's Python SDK also fails on SSE comment lines. Co-Authored-By: Claude Opus 5.5 --- crates/tw-api/msg-codes.txt | 2 + crates/tw-api/src/lib.rs | 31 +- crates/tw-config/src/failover.rs | 18 +- crates/tw-config/src/lib.rs | 4 +- crates/tw-config/src/validate.rs | 30 + crates/tw-config/tests/manual/schema.rs | 9 + crates/tw-control/src/lib.rs | 1 + crates/tw-control/tests/in_flight.rs | 1 + crates/tw-control/tests/live_state.rs | 1 + crates/tw-control/tests/replay.rs | 1 + crates/tw-control/tests/resources.rs | 42 ++ crates/tw-gateway/src/server.rs | 2 + crates/tw-gateway/src/server/pipeline.rs | 1 + crates/tw-gateway/src/server/pipeline/hop.rs | 226 ++++++-- .../tw-gateway/src/server/pipeline/opening.rs | 99 +++- crates/tw-gateway/src/server/pipeline/slow.rs | 87 +++ crates/tw-gateway/tests/slow_start.rs | 542 ++++++++++++++++++ crates/tw-observe/src/bus.rs | 2 + crates/tw-store/src/db.rs | 1 + crates/tw-store/src/recorder.rs | 11 + docs/config.md | 21 +- docs/config.zh-CN.md | 14 +- 22 files changed, 1082 insertions(+), 64 deletions(-) create mode 100644 crates/tw-gateway/src/server/pipeline/slow.rs create mode 100644 crates/tw-gateway/tests/slow_start.rs diff --git a/crates/tw-api/msg-codes.txt b/crates/tw-api/msg-codes.txt index a6ab51fc..b296f7e4 100644 --- a/crates/tw-api/msg-codes.txt +++ b/crates/tw-api/msg-codes.txt @@ -90,6 +90,7 @@ config.schema_too_new config.secret.empty_name config.secret.env_missing config.secret.unterminated +config.slow_start_too_short config.store.conflict config.store.missing config.store.read_failed @@ -355,6 +356,7 @@ gw.route.protocol_mismatch gw.route.rule_failed gw.route.selected_upstream_missing gw.route.upstream_missing +gw.slow_start gw.toolcall.connection_cut gw.toolcall.response_cut gw.toolcall.response_withheld diff --git a/crates/tw-api/src/lib.rs b/crates/tw-api/src/lib.rs index 9cd290da..341eb504 100644 --- a/crates/tw-api/src/lib.rs +++ b/crates/tw-api/src/lib.rs @@ -1438,6 +1438,11 @@ slug_enum! { /// (`status` 是它回的那个)。尝试链到此为止,这一行记成网关自己答的 /// ([`HistoryRow::local`]),费用 0 Estimated = "estimated", + /// 流式回答等了 `failover.stream_start_wait_secs` 还没有内容,开着 + /// `failover.next_on_slow_start`,放弃这一家、换下一家(连接断开,上游不再生成)。 + /// 响应头到了的有 `status`,没到的没有。**这一家不停用、不算失败**。上游可能已经按 + /// 输入收了钱:知道多少的在 `usage` 里 + SlowStart = "slow_start", } } @@ -1462,10 +1467,32 @@ pub struct AttemptView { /// 上游返回的状态码。`error` 时没有 #[serde(default, skip_serializing_if = "Option::is_none")] pub status: Option, - /// `error` 时的说明。和这一跳报给客户端的那条错误是同一句 + /// `error` 时的说明。和这一跳报给客户端的那条错误是同一句。`slow_start` 时说等了多久 #[serde(default, skip_serializing_if = "Option::is_none")] pub error: Option, pub ms: u64, + /// 放弃了的这一跳(`slow_start`)上游可能已经收了钱的输入(见 [`AttemptUsage`])。估不 + /// 出来的(请求解不开)没有。别的结果都没有:接下请求的那一跳的用量在结局里 + #[serde(default, skip_serializing_if = "Option::is_none")] + pub usage: Option, +} + +/// 放弃了的一跳([`AttemptOutcome::SlowStart`])上游可能已经收了钱的输入。 +/// +/// 上游在流开头报了的(Anthropic 的 `message_start`)是它报的数;没报的只有 `input`,是网关 +/// 估的(`estimated`,和 [`Event::RequestStarted`] 的 `input_estimate` 同一个数)。**输出不知道**: +/// 先想好再输出的模型,放弃之前可能已经想了一阵,上游不说就看不到。 +/// +/// **不算进这个请求的费用**:上游收没收、收了多少,网关看不到 +#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)] +#[cfg_attr(feature = "ts", derive(ts_rs::TS))] +pub struct AttemptUsage { + /// 输入 token,不含缓存读写 + pub input: u64, + pub cache_read: u64, + pub cache_write: u64, + /// `input` 是网关估的,上游什么都没报 + pub estimated: bool, } /// 一次请求的路由决策。**详情抽屉的 Routing 那一页吃它。** @@ -1832,6 +1859,8 @@ pub struct FailoverView { pub rate_limit_max_pause_secs: u64, /// 流式回答的开头最多等多少秒 pub stream_start_wait_secs: u64, + /// 等过 `stream_start_wait_secs` 还没有内容就换下一家(最后一家照常等) + pub next_on_slow_start: bool, } /// 每项防护各在哪一档:`off` / `observe` / `enforce`。 diff --git a/crates/tw-config/src/failover.rs b/crates/tw-config/src/failover.rs index 1213f029..621af9e2 100644 --- a/crates/tw-config/src/failover.rs +++ b/crates/tw-config/src/failover.rs @@ -1,4 +1,4 @@ -//! 故障转移:一家上游失败之后停用多久、流开头最多等多久。 +//! 故障转移:一家上游失败之后停用多久、流开头最多等多久、等不到内容换不换下一家。 use serde::{Deserialize, Serialize}; @@ -37,6 +37,11 @@ pub struct Failover { /// 等过这么久还没有内容,就不再等,把已经收到的交给客户端 #[serde(default = "d_stream_start_wait_secs")] pub stream_start_wait_secs: u64, + /// 流式回答等过 [`Self::stream_start_wait_secs`] 还没有内容时,放弃这一家、换下一家。 + /// **最后一家不换**,照常等下去;这一家不停用,也不算一次失败。默认关:先想好再 + /// 输出的模型开头本来就慢,开着时要把等待调长 + #[serde(default)] + pub next_on_slow_start: bool, } fn d_failures_to_pause() -> u32 { @@ -71,6 +76,7 @@ impl Default for Failover { quota_pause_secs: d_quota_pause_secs(), rate_limit_max_pause_secs: d_rate_limit_max_pause_secs(), stream_start_wait_secs: d_stream_start_wait_secs(), + next_on_slow_start: false, } } } @@ -82,6 +88,10 @@ pub const MAX_PAUSE_SECS: u64 = 7 * 24 * 3600; /// 流开头最多等多少秒。再长的话,一家卡在半路的上游会让客户端先超时 pub const MAX_STREAM_START_WAIT_SECS: u64 = 120; +/// 开着「开头慢就换下一家」时,流开头至少等多少秒。再短的话,平常的请求还没开口就被 +/// 切掉了 +pub const MIN_SLOW_START_WAIT_SECS: u64 = 5; + impl Failover { /// 不在允许范围里的第一项:字段名、写的值、下限、上限。 pub(crate) fn out_of_range(&self) -> Option<(&'static str, u64, u64, u64)> { @@ -123,4 +133,10 @@ impl Failover { .into_iter() .find(|(_, v, min, max)| v < min || v > max) } + + /// 开着「开头慢就换下一家」、等待却短于 [`MIN_SLOW_START_WAIT_SECS`]:写的等待秒数。 + pub(crate) fn slow_start_too_short(&self) -> Option { + (self.next_on_slow_start && self.stream_start_wait_secs < MIN_SLOW_START_WAIT_SECS) + .then_some(self.stream_start_wait_secs) + } } diff --git a/crates/tw-config/src/lib.rs b/crates/tw-config/src/lib.rs index 171c3520..28490198 100644 --- a/crates/tw-config/src/lib.rs +++ b/crates/tw-config/src/lib.rs @@ -1168,7 +1168,9 @@ pub fn write(path: &Path, cfg: &Config) -> Result<(), WriteError> { Ok(()) } -pub use failover::{Failover, MAX_PAUSE_SECS, MAX_STREAM_START_WAIT_SECS}; +pub use failover::{ + Failover, MAX_PAUSE_SECS, MAX_STREAM_START_WAIT_SECS, MIN_SLOW_START_WAIT_SECS, +}; pub use probes::{ClientProbes, ProbeAction}; pub use reload::{Rejected, Stage, stand_in, try_parse}; pub use retention::Retention; diff --git a/crates/tw-config/src/validate.rs b/crates/tw-config/src/validate.rs index 0a7af0ef..c74b9581 100644 --- a/crates/tw-config/src/validate.rs +++ b/crates/tw-config/src/validate.rs @@ -64,6 +64,9 @@ pub enum ValidationError { min: u64, max: u64, }, + /// 开着「开头慢就换下一家」,流开头的等待却短于 [`crate::MIN_SLOW_START_WAIT_SECS`] + #[error("{}", self.msg())] + SlowStartTooShort { secs: u64 }, #[error("{}", self.msg())] ControlKeyMissing, #[error("{}", self.msg())] @@ -206,6 +209,12 @@ impl ValidationError { "config.failover_range", field = field, value = value, min = min, max = max => "failover.{field} is {value}; it has to be between {min} and {max}" ), + SlowStartTooShort { secs } => msg!( + "config.slow_start_too_short", + secs = secs, min = crate::MIN_SLOW_START_WAIT_SECS => + "failover.stream_start_wait_secs is {secs} while failover.next_on_slow_start is on; \ + it has to be at least {min}, or ordinary answers are cut off before they start" + ), ControlKeyMissing => msg!( "config.control_key_missing" => "the configuration has no listen.control.key, the key the desktop app connects \ @@ -475,6 +484,10 @@ pub fn validate(cfg: &Config) -> Result<(), ValidationError> { max, }); } + // 开头慢就换下一家:等得太短的话,平常的回答还没开口就被切到下一家 + if let Some(secs) = cfg.failover.slow_start_too_short() { + return Err(ValidationError::SlowStartTooShort { secs }); + } // 控制面的钥匙。**缺了、短了、不是十六进制,整份配置都不收**,旧的继续 // 服务:一份没有钥匙的配置换进来,下一条连接谁都进不来 —— 包括要把它 // 改回去的那个界面;一把好猜的短钥匙和没有差不多 @@ -1059,6 +1072,23 @@ groups: } } + /// 开头慢就换下一家:开着时等待至少 5 秒,关着时 1 秒也照收(只是交得早) + #[test] + fn switching_on_a_slow_start_needs_a_long_enough_wait() { + let mut x = with_rules(&[], &[]); + x.failover.stream_start_wait_secs = 3; + assert!(validate(&x).is_ok(), "关着时不管"); + x.failover.next_on_slow_start = true; + let e = validate(&x).unwrap_err(); + assert!( + matches!(e, ValidationError::SlowStartTooShort { secs: 3 }), + "{e:?}" + ); + assert_eq!(e.msg().code, "config.slow_start_too_short"); + x.failover.stream_start_wait_secs = crate::MIN_SLOW_START_WAIT_SECS; + assert!(validate(&x).is_ok()); + } + /// 别名表:名字一个一个、不带通配、不撞内置前缀,每个别名列着别的名称, /// 不指向别的别名。 #[test] diff --git a/crates/tw-config/tests/manual/schema.rs b/crates/tw-config/tests/manual/schema.rs index 8916784e..c220778d 100644 --- a/crates/tw-config/tests/manual/schema.rs +++ b/crates/tw-config/tests/manual/schema.rs @@ -1210,6 +1210,15 @@ pub fn sections() -> Vec
{ "流式回答在第一段内容到达前最多暂存的秒数。在此之前上游报错,请求换到下一家;超过这个时间,已收到的部分照常交给客户端。取值 1 到 120。", ), ), + row( + "next_on_slow_start", + Kind::Bool, + Def::Is("false"), + t( + "When a streamed answer still has no content `stream_start_wait_secs` after the request was sent, give up on that upstream and send the request to the next one. The last upstream always waits. The upstream given up on is not set aside. Needs `stream_start_wait_secs` of at least 5.", + "流式回答在请求发出 `stream_start_wait_secs` 秒后仍没有内容时,放弃这家上游,把请求交给下一家。最后一家总是等下去。被放弃的上游不会停用。开启时 `stream_start_wait_secs` 至少为 5。", + ), + ), ], }, // ── groups / routes ─────────────────────────────────── diff --git a/crates/tw-control/src/lib.rs b/crates/tw-control/src/lib.rs index abf70bd5..59668921 100644 --- a/crates/tw-control/src/lib.rs +++ b/crates/tw-control/src/lib.rs @@ -410,6 +410,7 @@ async fn overview(State(s): State) -> Json { quota_pause_secs: f.quota_pause_secs, rate_limit_max_pause_secs: f.rate_limit_max_pause_secs, stream_start_wait_secs: f.stream_start_wait_secs, + next_on_slow_start: f.next_on_slow_start, } }, listen: tw_api::ListenView { diff --git a/crates/tw-control/tests/in_flight.rs b/crates/tw-control/tests/in_flight.rs index bed48680..aa6f5e6a 100644 --- a/crates/tw-control/tests/in_flight.rs +++ b/crates/tw-control/tests/in_flight.rs @@ -97,6 +97,7 @@ async fn it_gives_every_running_request_with_what_has_happened_to_it_so_far() { status: Some(200), error: None, ms: 700, + usage: None, }], billing: tw_api::Billing::PerToken, }); diff --git a/crates/tw-control/tests/live_state.rs b/crates/tw-control/tests/live_state.rs index 7e91360c..ddcee347 100644 --- a/crates/tw-control/tests/live_state.rs +++ b/crates/tw-control/tests/live_state.rs @@ -154,6 +154,7 @@ async fn a_request_still_running_can_be_opened_and_becomes_whole_when_it_ends() status: Some(200), error: None, ms: 900, + usage: None, }], billing: tw_api::Billing::PerToken, }); diff --git a/crates/tw-control/tests/replay.rs b/crates/tw-control/tests/replay.rs index 98a20df9..f069212b 100644 --- a/crates/tw-control/tests/replay.rs +++ b/crates/tw-control/tests/replay.rs @@ -385,6 +385,7 @@ fn routed(hops: &[(&str, Option<&str>)]) -> Option { status: Some(200), error: None, ms: 100, + usage: None, }) .collect(); Some( diff --git a/crates/tw-control/tests/resources.rs b/crates/tw-control/tests/resources.rs index a0aeb332..333895c8 100644 --- a/crates/tw-control/tests/resources.rs +++ b/crates/tw-control/tests/resources.rs @@ -211,6 +211,48 @@ async fn failover_settings_show_their_defaults_and_take_an_edit() { assert!(body.contains("config.failover_range"), "{body}"); } +/// 开头慢就换下一家:默认关;打开要等得够久,等得太短的被拒 +#[tokio::test] +async fn switching_on_a_slow_start_is_shown_and_needs_a_long_enough_wait() { + let b = bed(BASE); + let (_, body) = call(&b.app, "GET", "/overview", serde_json::Value::Null).await; + assert_eq!( + json(&body)["failover"]["next_on_slow_start"], + false, + "{body}" + ); + + let (st, body) = call( + &b.app, + "PATCH", + "/config", + serde_json::json!({ + "ops": [{ "op": "replace", "path": "/failover/next_on_slow_start", "value": true }], + }), + ) + .await; + assert_eq!(st, StatusCode::OK, "{body}"); + assert!(b.parsed().failover.next_on_slow_start); + let (_, body) = call(&b.app, "GET", "/overview", serde_json::Value::Null).await; + assert_eq!( + json(&body)["failover"]["next_on_slow_start"], + true, + "{body}" + ); + + let (st, body) = call( + &b.app, + "PATCH", + "/config", + serde_json::json!({ + "ops": [{ "op": "replace", "path": "/failover/stream_start_wait_secs", "value": 3 }], + }), + ) + .await; + assert_eq!(st, StatusCode::BAD_REQUEST, "{body}"); + assert!(body.contains("config.slow_start_too_short"), "{body}"); +} + // ─────────────────────────────────────────────────────────── 上游 #[tokio::test] diff --git a/crates/tw-gateway/src/server.rs b/crates/tw-gateway/src/server.rs index 77a51303..dd05817c 100644 --- a/crates/tw-gateway/src/server.rs +++ b/crates/tw-gateway/src/server.rs @@ -346,6 +346,7 @@ pub(crate) fn hop( status: Some(status), error: None, ms: started.elapsed().as_millis() as u64, + usage: None, } } @@ -363,6 +364,7 @@ pub(crate) fn hop_failed( status: None, error: Some(error), ms: started.elapsed().as_millis() as u64, + usage: None, } } diff --git a/crates/tw-gateway/src/server/pipeline.rs b/crates/tw-gateway/src/server/pipeline.rs index cbd6bdc4..ca396e64 100644 --- a/crates/tw-gateway/src/server/pipeline.rs +++ b/crates/tw-gateway/src/server/pipeline.rs @@ -24,6 +24,7 @@ mod hop; mod opening; mod plug; mod relay; +mod slow; /// 256 MiB。大到能装下几张 4K 图的 base64(膨胀 33%),小到失控的 /// 客户端打不爆内存。 diff --git a/crates/tw-gateway/src/server/pipeline/hop.rs b/crates/tw-gateway/src/server/pipeline/hop.rs index dd4cdb1a..b80662ad 100644 --- a/crates/tw-gateway/src/server/pipeline/hop.rs +++ b/crates/tw-gateway/src/server/pipeline/hop.rs @@ -9,6 +9,8 @@ //! 每一跳先过插件的请求钩子([`super::plug`]):发往哪一家、发什么模型名这时都定了, //! 管这一跳的插件从客户端的原话起改,这一跳的转换、脱敏、发送用改过的那一份。换到下一 //! 家时从原话重来;同一家重发(OAuth 换 token、去封存)用这一跳定好的请求体,不重跑。 +//! +//! 开着 `failover.next_on_slow_start` 时,开头迟迟没有内容的一家也换掉(见 [`super::slow`])。 use bytes::Bytes; @@ -134,6 +136,8 @@ pub(super) async fn try_upstreams<'a>( let catalog = state.catalog.load(); // 这把密钥的模型范围:每一跳发出的名字都要过它(见 `crate::sent::name`) let allow = crate::models::key_allow(&rt.config, &req.client_name); + // 开头慢就换下一家:开着、客户端要的是流时,等多久(见 `super::slow`) + let slow_wait = super::slow::wait(rt, reading); for (i, name) in started.alive.iter().enumerate() { // 后面没有别的候选了 @@ -315,7 +319,15 @@ pub(super) async fn try_upstreams<'a>( let model = Some(plugged.model.clone().unwrap_or(sent)).filter(asked_other); let asked = Asked::of(req, reading, &plugged); - let out = match prepare(state, req, reading, &asked, provider, &effective_set, id) { + let out = match prepare( + state, + req, + reading, + &asked, + provider, + &effective_set, + Some(id), + ) { Ok(out) => out, Err(err) => { chain.push(hop_failed( @@ -413,31 +425,75 @@ pub(super) async fn try_upstreams<'a>( .clone() .unwrap_or_else(|| reading.facts.model.clone()); let (attempt, bridge) = (chain.len(), plugged.bridge); + // 后面还有没有接得下这个请求的(见 `successor`)。开头慢了才问 + let rest = &started.alive[i + 1..]; + let others = || successor(state, rt, req, reading, decision, &catalog, allow, rest); + // 开头慢就换下一家:等到什么时候,从这一刻(请求发出去)算起。最后一家不换。到点时 + // 问过、后面没有接得下的,清掉它:这一跳从此和不开时一样 + let mut slow_deadline = slow_wait + .filter(|_| !last) + .map(|w| tokio::time::Instant::now() + w); - let sent = send( - state, - req, - provider, - http, - &out, - body.clone(), - upstream_headers.clone(), - aws.as_ref(), - ) - .await; - // 上游拒绝了别家封存的推理:去掉它们,同一家再发一次 - let sent = match sent { - Ok(r) => { - let resend = Resend { - out: &out, - body: &body, - headers: &upstream_headers, - aws: aws.as_ref(), - conversation: started.conversation.as_deref(), - }; - resend_unsealed(state, req, provider, http, resend, r).await + // 发出去、等响应头。**等着的这个 future 只活在这一块里**:放弃这一家时它跟着丢掉, + // 连接随之断开 + let sent = { + let sending = async { + let sent = send( + state, + req, + provider, + http, + &out, + body.clone(), + upstream_headers.clone(), + aws.as_ref(), + ) + .await; + // 上游拒绝了别家封存的推理:去掉它们,同一家再发一次 + match sent { + Ok(r) => { + let resend = Resend { + out: &out, + body: &body, + headers: &upstream_headers, + aws: aws.as_ref(), + conversation: started.conversation.as_deref(), + }; + resend_unsealed(state, req, provider, http, resend, r).await + } + Err(e) => Err(e), + } + }; + let mut sending = std::pin::pin!(sending); + match slow_deadline { + None => sending.await, + Some(deadline) => match tokio::time::timeout_at(deadline, sending.as_mut()).await { + Ok(sent) => sent, + // 响应头都还没来,后面又有接得下的:放弃这一家。不停用、不算失败(见 + // `super::slow`) + Err(_) if others() => { + let waited = slow_wait.unwrap_or_default(); + chain.push(super::slow::abandoned( + &provider.name, + model.clone(), + None, + None, + reading, + waited, + hop_started, + )); + last_err = Some(GatewayError::upstream(super::slow::said( + &provider.name, + waited, + ))); + continue; + } + Err(_) => { + slow_deadline = None; + sending.await + } + }, } - Err(e) => Err(e), }; match sent { Ok(r) if !r.status().is_success() => { @@ -559,11 +615,42 @@ pub(super) async fn try_upstreams<'a>( let r = match opening_of(req, provider, &out, &r).filter(|_| !last) { None => r, Some((dialect, eventstream)) => { - let wait = std::time::Duration::from_secs( - rt.config.failover.stream_start_wait_secs, - ); - match super::opening::watch(r, dialect, eventstream, wait).await { + // 开头慢就换下一家的,等到发出请求之后的那一刻;别的从响应头到了算起 + let deadline = slow_deadline.unwrap_or_else(|| { + tokio::time::Instant::now() + + std::time::Duration::from_secs( + rt.config.failover.stream_start_wait_secs, + ) + }); + match super::opening::watch(r, dialect, eventstream, deadline).await { super::opening::Opening::Go(r) => r, + // 到点了还没有内容,后面又有接得下的:放弃这一家。**响应跟着这一轮 + // 循环丢掉**,和上游的连接随之断开,它不再接着生成。不停用、不算失败 + super::opening::Opening::Slow { response, usage } + if slow_deadline.is_some() && others() => + { + let status = response.status().as_u16(); + drop(response); + // 上游回了话,说明代理是通的 + state.note_proxy_ok(&provider.proxy); + let waited = slow_wait.unwrap_or_default(); + chain.push(super::slow::abandoned( + &provider.name, + model.clone(), + Some(status), + usage, + reading, + waited, + hop_started, + )); + last_err = Some(GatewayError::upstream(super::slow::said( + &provider.name, + waited, + ))); + continue; + } + // 不换(没开、或者后面没有接得下的):和内容来了一样交出去 + super::opening::Opening::Slow { response, .. } => response, super::opening::Opening::Failed { status, headers, @@ -798,6 +885,7 @@ fn estimated_hop( status, error: None, ms: started.elapsed().as_millis() as u64, + usage: None, } } @@ -879,9 +967,67 @@ fn unsendable_tool( Some(GatewayError::new(crate::error::Source::Request, msg)) } +/// 慢了的那一家后面,`rest` 里还有没有接得下这个请求的(见 [`super::slow`])。 +/// +/// 停用着的不算,这一跳发不出去的也不算:配置里没有了、格式对不上、阶段二的规则拒绝、 +/// 发给它的名字对不上(别名、清单、密钥范围)、转换不了、强制要用的工具发不过去。**一个都 +/// 没有的话,慢了的这一家就是最后一家**,照常等下去 —— 放弃了它,换来的是一个注定失败的 +/// 请求。 +/// +/// **只看不跑**:看的是客户端的原话,不跑插件、不取密钥(那两样在真发的那一跳才知道拒 +/// 不拒),不发转换事件 +#[allow(clippy::too_many_arguments)] +fn successor( + state: &AppState, + rt: &Runtime, + req: &Inbound, + reading: &crate::client_api::Reading, + decision: &tw_engine::Decision, + catalog: &tw_engine::Catalog, + allow: Option<&[String]>, + rest: &[String], +) -> bool { + let asked = Asked { + body: &req.body, + path: req.uri.path(), + decoded: reading.decoded.as_ref(), + }; + rest.iter().any(|name| { + let Some(provider) = rt.config.providers.iter().find(|p| &p.name == name) else { + return false; + }; + if !state.health.is_available(name) + || protocol_mismatch(req, reading.generates, provider).is_some() + { + return false; + } + let Ok(tw_engine::Outcome2::Proceed { + mut set, + model: renamed, + .. + }) = rt + .engine + .phase_two(&reading.facts, &provider.name, &decision.set) + else { + return false; + }; + let asked_model = + rt.engine + .asked_of(&reading.facts, decision, &provider.name, renamed.as_deref()); + let Ok(sent) = + crate::sent::name(&rt.config, catalog, decision, provider, &asked_model, allow) + else { + return false; + }; + set.model = (sent != reading.facts.model).then_some(sent); + prepare(state, req, reading, &asked, provider, &set, None).is_ok() + }) +} + /// 把这一跳的请求(客户端那种格式,插件改过的话是改过的,见 [`Asked`])改成要发的 /// 样子:同格式时只做参数改写,跨格式时转换。转换不了就换下一家:同格式的上游可能 -/// 还在后面。 +/// 还在后面。`id` 是这个请求的号,做了转换、丢了字段要报在它上面;只看发不发得出去时 +/// (见 [`successor`])是 None,什么都不报。 fn prepare( state: &AppState, req: &Inbound, @@ -889,7 +1035,7 @@ fn prepare( asked: &Asked<'_>, provider: &tw_config::Provider, effective_set: &tw_engine::SetAction, - id: u64, + id: Option, ) -> Result { let generates = reading.generates; // 方言互转。**同格式时是 None,这一整段零成本** @@ -949,7 +1095,7 @@ fn prepare( .and_then(|d| crate::egress::strip_body_identity(d, &out)) .unwrap_or(out) }; - if let Some(d) = client_dialect.filter(|_| !dropped.is_empty()) { + if let Some((d, id)) = client_dialect.filter(|_| !dropped.is_empty()).zip(id) { let same = crate::wire::dialect(d); state.bus.emit(tw_api::Event::Translated { id, @@ -1022,14 +1168,16 @@ fn prepare( // 却没生效」而完全不知道从哪儿查起 let mut dropped = p.dropped.clone(); dropped.extend(limit); - state.bus.emit(tw_api::Event::Translated { - id, - provider: provider.name.clone(), - from: crate::wire::dialect(d.client), - to: crate::wire::dialect(dialect), - dropped, - at_ms: crate::server::now_ms(), - }); + if let Some(id) = id { + state.bus.emit(tw_api::Event::Translated { + id, + provider: provider.name.clone(), + from: crate::wire::dialect(d.client), + to: crate::wire::dialect(dialect), + dropped, + at_ms: crate::server::now_ms(), + }); + } path = p.path.clone(); query = p.query.clone(); // Bedrock 上的 Claude:客户端 `anthropic-beta` 里 Bedrock 认的那几个放进请求体 diff --git a/crates/tw-gateway/src/server/pipeline/opening.rs b/crates/tw-gateway/src/server/pipeline/opening.rs index f8cd4ae3..7f74e6ab 100644 --- a/crates/tw-gateway/src/server/pipeline/opening.rs +++ b/crates/tw-gateway/src/server/pipeline/opening.rs @@ -10,9 +10,9 @@ //! 客户端收到的和直接转发一个字节都不差。 //! //! 等待有上限(配置的 `failover.stream_start_wait_secs`,以及 [`HOLD_LIMIT`]): -//! 上游迟迟不出内容时不能一直压着,那样客户端看到的就是一个卡住的请求。 - -use std::time::Duration; +//! 上游迟迟不出内容时不能一直压着,那样客户端看到的就是一个卡住的请求。等到点了是 +//! [`Opening::Slow`]:开着 `failover.next_on_slow_start` 时由调用方放弃这一家、换下一家, +//! 不开就和内容来了一样交出去。 use bytes::Bytes; use futures::StreamExt; @@ -25,8 +25,15 @@ pub(super) const HOLD_LIMIT: usize = 1024 * 1024; /// 流开头的结论。 pub(super) enum Opening { - /// 内容来了(或者等够了、流结束了):交给客户端。读过的字节已经接回去了 + /// 内容来了(或者流结束了、开头压得太多了):交给客户端。读过的字节已经接回去了 Go(reqwest::Response), + /// 等到点了还没有内容。`response` 和 [`Opening::Go`] 的一样,照旧交出去就是不换家; + /// **丢掉它就断开了和上游的连接**,上游不再接着生成。`usage` 是开头里上游报了的用量 + /// (Anthropic 的 `message_start` 带着输入),没报是 None + Slow { + response: reqwest::Response, + usage: Option, + }, /// 第一段内容之前上游报了错。`status` 是这个错误对应的状态码,`body` 是 /// 上游的原话,交给 [`crate::failure::classify`] 判断换不换。 /// @@ -44,13 +51,13 @@ pub(super) enum Opening { Broken(GatewayError), } -/// 读到第一段内容为止。`dialect` 是上游说的格式,`eventstream` 表示流是 Bedrock -/// 的二进制帧。 +/// 读到第一段内容为止,最多等到 `deadline`。`dialect` 是上游说的格式,`eventstream` +/// 表示流是 Bedrock 的二进制帧。 pub(super) async fn watch( r: reqwest::Response, dialect: Dialect, eventstream: bool, - wait: Duration, + deadline: tokio::time::Instant, ) -> Opening { let status = r.status(); let headers = r.headers().clone(); @@ -59,12 +66,17 @@ pub(super) async fn watch( let mut size = 0usize; let mut sse = Sse::default(); let mut unframe = eventstream.then(tw_bedrock::eventstream::Transcoder::new); - let deadline = tokio::time::Instant::now() + wait; + // 开头里上游报的用量:放弃这一家时,它可能已经按这些收了钱 + let mut usage = tw_dialect::usage::Sniffer::new(); + let mut slow = false; let verdict = loop { let next = match tokio::time::timeout_at(deadline, stream.next()).await { Ok(next) => next, - Err(_) => break Judge::Content, + Err(_) => { + slow = true; + break Judge::Content; + } }; let chunk = match next { None => break Judge::Content, @@ -86,6 +98,7 @@ pub(super) async fn watch( Err(_) => break Judge::Content, }, }; + usage.feed(&text); let judged = sse .feed(&text) .into_iter() @@ -119,6 +132,10 @@ pub(super) async fn watch( message, response, }, + Judge::Content | Judge::Preamble if slow => Opening::Slow { + response, + usage: usage.finish(), + }, Judge::Content | Judge::Preamble => Opening::Go(response), } } @@ -332,6 +349,8 @@ impl Sse { #[cfg(test)] mod tests { + use std::time::Duration; + use super::*; fn events(text: &str) -> Vec<(Option, String)> { @@ -446,7 +465,7 @@ mod tests { event: content_block_delta\ndata: {\"type\":\"content_block_delta\"}\n\n\ event: message_stop\ndata: {\"type\":\"message_stop\"}\n\n"; let r = served(body).await; - match watch(r, Dialect::Anthropic, false, Duration::from_secs(5)).await { + match watch(r, Dialect::Anthropic, false, soon(5)).await { Opening::Go(r) => { assert_eq!(r.headers()[http::header::CONTENT_TYPE], "text/event-stream"); assert_eq!(r.text().await.unwrap(), body); @@ -458,16 +477,60 @@ mod tests { #[tokio::test] async fn an_error_before_content_is_reported_not_handed_on() { let body = "event: error\ndata: {\"type\":\"error\",\"error\":{\"type\":\"rate_limit_error\",\"message\":\"slow\"}}\n\n"; - match watch( - served(body).await, - Dialect::Anthropic, - false, - Duration::from_secs(5), - ) - .await - { + match watch(served(body).await, Dialect::Anthropic, false, soon(5)).await { Opening::Failed { status, .. } => assert_eq!(status, 429), _ => panic!("该报失败"), } } + + fn soon(secs: u64) -> tokio::time::Instant { + tokio::time::Instant::now() + Duration::from_secs(secs) + } + + /// 发完开头就不说话的上游:响应头和 `head` 先到,之后流一直开着 + fn stalled(head: &'static str) -> reqwest::Response { + let first = + futures::stream::iter([Ok::<_, std::io::Error>(Bytes::from_static(head.as_bytes()))]); + let body = reqwest::Body::wrap_stream(first.chain(futures::stream::pending())); + let mut resp = http::Response::new(body); + resp.headers_mut().insert( + http::header::CONTENT_TYPE, + http::HeaderValue::from_static("text/event-stream"), + ); + reqwest::Response::from(resp) + } + + #[tokio::test] + async fn no_content_by_the_deadline_is_slow_and_keeps_what_the_upstream_reported() { + let head = "event: message_start\ndata: {\"type\":\"message_start\",\"message\":{\"usage\":{\"input_tokens\":1200,\"cache_read_input_tokens\":800,\"output_tokens\":1}}}\n\n\ + event: ping\ndata: {\"type\":\"ping\"}\n\n"; + let deadline = tokio::time::Instant::now() + Duration::from_millis(200); + match watch(stalled(head), Dialect::Anthropic, false, deadline).await { + Opening::Slow { usage, .. } => { + let u = usage.expect("message_start 报了输入"); + assert_eq!((u.input, u.cache_read), (1200, 800)); + } + _ => panic!("到点没有内容该是 Slow"), + } + } + + #[tokio::test] + async fn a_thinking_delta_is_content_not_a_slow_start() { + // Chat 格式的推理字(DeepSeek、Qwen 的 `reasoning_content`)也是模型开口了 + let head = "data: {\"choices\":[{\"delta\":{\"role\":\"assistant\",\"content\":\"\"}}]}\n\n\ + data: {\"choices\":[{\"delta\":{\"reasoning_content\":\"Let me think\"}}]}\n\n"; + let deadline = tokio::time::Instant::now() + Duration::from_millis(200); + assert!(matches!( + watch(stalled(head), Dialect::Chat, false, deadline).await, + Opening::Go(_) + )); + // 只来了角色的那一块不算 + let role = + "data: {\"choices\":[{\"delta\":{\"role\":\"assistant\",\"content\":\"\"}}]}\n\n"; + let deadline = tokio::time::Instant::now() + Duration::from_millis(200); + assert!(matches!( + watch(stalled(role), Dialect::Chat, false, deadline).await, + Opening::Slow { usage: None, .. } + )); + } } diff --git a/crates/tw-gateway/src/server/pipeline/slow.rs b/crates/tw-gateway/src/server/pipeline/slow.rs new file mode 100644 index 00000000..afa69400 --- /dev/null +++ b/crates/tw-gateway/src/server/pipeline/slow.rs @@ -0,0 +1,87 @@ +//! 开头慢就换下一家(配置的 `failover.next_on_slow_start`)。 +//! +//! 有的上游收下请求之后很久不出内容:中转站排着队,上游过载却不报错,响应头都迟迟不来。 +//! 开着这一项时,从请求发出去算起等 `failover.stream_start_wait_secs`,还没有内容就**断开 +//! 这一家**(丢掉响应或者还在等的请求,连接跟着断,上游不再接着生成),换下一家。客户端 +//! 这时一个字节都还没收到,换一家它无感。 +//! +//! 几条规矩: +//! +//! - **最后一家不换**,照常等下去。「最后」按后面还有没有接得下的算(见 +//! `hop::successor`):停用着的、这一跳发不出去的不算 —— 否则放弃了一个慢的,换来的是 +//! 一个注定失败的。 +//! - **这一家不停用、不算失败**:慢不是坏,下一个请求它可能就快了。 +//! - **尝试链上记一跳 `slow_start`**,带着上游可能已经收了钱的输入(见 +//! [`tw_api::AttemptUsage`])。 +//! - 只管客户端要流式的请求:整包的请求本来就要等全部生成完,开头慢说明不了什么。 +//! +//! **等的时候不给客户端发保活。**响应头要等选定了哪一家才发(见 `relay`),这期间客户端 +//! 那条连接上什么都没有;先发响应头再发 `: keepalive` 的话,状态码就定死成了 200 —— 之后 +//! 几家全都失败,429、5xx 和最后一家原样交出的 4xx 都给不出去,只能在流里报错,客户端按 +//! 状态码重试的逻辑就落空了;上游的响应头(请求号、额度)也带不过去。何况 Gemini 官方的 +//! Python SDK 会把注释行当成一段 JSON 去解析,直接报错。 + +use std::time::Duration; + +use crate::state::Runtime; +use tw_types::msg; + +/// 这个请求开头慢了换不换、换的话等多久:开着这一项、客户端要的是流时才有。 +pub(super) fn wait(rt: &Runtime, reading: &crate::client_api::Reading) -> Option { + let f = &rt.config.failover; + let streams = matches!(&reading.decoded, Some(Ok(d)) if d.request.stream); + (f.next_on_slow_start && streams).then(|| Duration::from_secs(f.stream_start_wait_secs)) +} + +/// 放弃了的那一跳:尝试链上的一行。`status` 是上游回的(响应头没到的没有),`seen` 是流 +/// 开头里上游报的用量。 +pub(super) fn abandoned( + provider: &str, + model: Option, + status: Option, + seen: Option, + reading: &crate::client_api::Reading, + waited: Duration, + started: std::time::Instant, +) -> tw_api::AttemptView { + tw_api::AttemptView { + provider: provider.to_string(), + model, + outcome: tw_api::AttemptOutcome::SlowStart, + status, + error: Some(said(provider, waited)), + ms: started.elapsed().as_millis() as u64, + usage: usage(seen, reading), + } +} + +/// 放弃的那一家可能已经收了钱的输入:上游报了的用它报的,没报的用网关估的(和开始事件的 +/// `input_estimate` 同一个数),估不出来(请求解不开)就没有。 +fn usage( + seen: Option, + reading: &crate::client_api::Reading, +) -> Option { + match seen.filter(|u| u.prompt_total() > 0) { + Some(u) => Some(tw_api::AttemptUsage { + input: u.input, + cache_read: u.cache_read, + cache_write: u.cache_write, + estimated: false, + }), + None => matches!(reading.decoded, Some(Ok(_))).then_some(tw_api::AttemptUsage { + input: reading.facts.input_tokens, + cache_read: 0, + cache_write: 0, + estimated: true, + }), + } +} + +/// 尝试链上那一跳的说明。后面几家也都不行时,它也是交给客户端的那条错误的退路 +pub(super) fn said(upstream: &str, waited: Duration) -> tw_types::Msg { + msg!( + "gw.slow_start", upstream = upstream, secs = waited.as_secs() => + "Upstream `{upstream}` sent no content within {secs} seconds, so the request moved on \ + to the next upstream." + ) +} diff --git a/crates/tw-gateway/tests/slow_start.rs b/crates/tw-gateway/tests/slow_start.rs new file mode 100644 index 00000000..41d415e7 --- /dev/null +++ b/crates/tw-gateway/tests/slow_start.rs @@ -0,0 +1,542 @@ +//! 开头慢就换下一家(`failover.next_on_slow_start`),端到端。 +//! +//! 假上游先回响应头和开头的例行帧,之后按脚本隔一阵才出内容,或者一直只发心跳不出内容; +//! 它记下自己的响应被丢掉了没有 —— 网关放弃它时,连接要真的断开,上游才会停下。 + +use std::net::SocketAddr; +use std::sync::Arc; +use std::sync::atomic::{AtomicBool, AtomicUsize, Ordering}; +use std::time::{Duration, Instant}; + +use axum::Router; +use bytes::Bytes; +use futures::StreamExt; +use serde_json::{Value, json}; +use tw_config::{Client, Config, Failover, Protocol, Provider}; + +/// 一个假上游的样子:响应头之前等多久,之后按顺序隔多少毫秒发哪一段,发完了挂不挂着 +#[derive(Clone)] +struct Script { + header_delay_ms: u64, + steps: Vec<(u64, &'static str)>, + /// 发完了不收尾,每 100 毫秒发一行注释(SSE 的心跳),直到连接断开 + hang: bool, +} + +/// 假上游被打了几次,以及它的响应(或者还没回的那个请求)被丢掉了没有 +struct Upstream { + addr: SocketAddr, + hits: Arc, + dropped: Arc, +} + +/// 被丢掉时举旗 +struct Flag(Arc); + +impl Drop for Flag { + fn drop(&mut self) { + self.0.store(true, Ordering::SeqCst); + } +} + +async fn upstream(script: Script) -> Upstream { + let hits = Arc::new(AtomicUsize::new(0)); + let dropped = Arc::new(AtomicBool::new(false)); + let (h, d) = (hits.clone(), dropped.clone()); + let app = Router::new().fallback(axum::routing::any(move || { + let (script, h, d) = (script.clone(), h.clone(), d.clone()); + async move { + h.fetch_add(1, Ordering::SeqCst); + // 响应头还没回时请求就被丢掉的,也要看得见 + let waiting = Flag(d.clone()); + tokio::time::sleep(Duration::from_millis(script.header_delay_ms)).await; + std::mem::forget(waiting); + let flag = Flag(d); + let steps = futures::stream::iter(script.steps).then(|(ms, chunk)| async move { + tokio::time::sleep(Duration::from_millis(ms)).await; + Ok::<_, std::convert::Infallible>(Bytes::from_static(chunk.as_bytes())) + }); + let hang = script.hang; + let tail = futures::stream::unfold((), move |_| async move { + if !hang { + return None; + } + tokio::time::sleep(Duration::from_millis(100)).await; + Some((Ok(Bytes::from_static(b": ping\n\n")), ())) + }); + let body = steps.chain(tail).map(move |x| { + let _ = &flag; + x + }); + axum::response::Response::builder() + .header("content-type", "text/event-stream") + .body(axum::body::Body::from_stream(body)) + .unwrap() + } + })); + let l = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap(); + let addr = l.local_addr().unwrap(); + tokio::spawn(async move { axum::serve(l, app).await.unwrap() }); + Upstream { + addr, + hits, + dropped, + } +} + +const MESSAGE_START: &str = "event: message_start\ndata: {\"type\":\"message_start\",\"message\":{\"usage\":{\"input_tokens\":1200,\"cache_read_input_tokens\":800,\"output_tokens\":1}}}\n\n"; +const PING: &str = "event: ping\ndata: {\"type\":\"ping\"}\n\n"; +const THINKING: &str = "event: content_block_start\ndata: {\"type\":\"content_block_start\",\"index\":0,\"content_block\":{\"type\":\"thinking\",\"thinking\":\"\"}}\n\n\ + event: content_block_delta\ndata: {\"type\":\"content_block_delta\",\"index\":0,\"delta\":{\"type\":\"thinking_delta\",\"thinking\":\"Let me think\"}}\n\n"; + +/// 一段完整的回答,`word` 是正文 +fn answer(word: &'static str) -> &'static str { + let text = format!( + "event: content_block_delta\ndata: {{\"type\":\"content_block_delta\",\"index\":0,\"delta\":{{\"type\":\"text_delta\",\"text\":\"{word}\"}}}}\n\n\ + event: message_stop\ndata: {{\"type\":\"message_stop\"}}\n\n" + ); + Box::leak(text.into_boxed_str()) +} + +/// 开了流、报了输入,之后只有心跳 +fn stalled() -> Script { + Script { + header_delay_ms: 0, + steps: vec![(0, MESSAGE_START), (0, PING)], + hang: true, + } +} + +/// 马上就答 +fn prompt(word: &'static str) -> Script { + Script { + header_delay_ms: 0, + steps: vec![(0, MESSAGE_START), (0, answer(word))], + hang: false, + } +} + +/// 开了流,隔 `ms` 毫秒才答 +fn late(ms: u64, word: &'static str) -> Script { + Script { + header_delay_ms: 0, + steps: vec![(0, MESSAGE_START), (ms, answer(word))], + hang: false, + } +} + +fn provider(name: &str, up: &Upstream, protocol: Protocol) -> Provider { + Provider { + name: name.into(), + base_url: format!("http://{}", up.addr), + key: Some("sk-upstream".into()), + protocol: Some(protocol), + ..Default::default() + } +} + +/// 起网关。等 1 秒(测试里图快;配置校验要求开着时至少 5 秒,网关自己不查) +async fn gateway( + providers: Vec, + switch: bool, +) -> ( + SocketAddr, + tokio::sync::broadcast::Receiver, + Arc, +) { + let cfg = Config { + version: 1, + clients: vec![Client { + name: "c".into(), + key: "tw-k".into(), + ..Default::default() + }], + providers, + failover: Failover { + stream_start_wait_secs: 1, + next_on_slow_start: switch, + ..Default::default() + }, + ..Default::default() + }; + let state = tw_gateway::AppState::new(cfg).unwrap(); + let rx = state.bus.subscribe(); + let health = state.health.clone(); + let addr = tw_gateway::serve(state, ([127, 0, 0, 1], 0).into()) + .await + .unwrap(); + tokio::time::sleep(Duration::from_millis(50)).await; + (addr, rx, health) +} + +fn messages(stream: bool) -> Value { + json!({"model": "claude-sonnet-5", "max_tokens": 1024, "stream": stream, + "messages": [{"role": "user", "content": "Say hello."}]}) +} + +async fn post(gw: SocketAddr, path: &str, body: &Value) -> (u16, String) { + let resp = reqwest::Client::new() + .post(format!("http://{gw}{path}")) + .header("content-type", "application/json") + .header("x-api-key", "tw-k") + .header("authorization", "Bearer tw-k") + .body(body.to_string()) + .send() + .await + .unwrap(); + let status = resp.status().as_u16(); + (status, resp.text().await.unwrap()) +} + +/// 这个请求的估算输入(开始事件带着)和尝试链 +async fn estimate_and_attempts( + rx: &mut tokio::sync::broadcast::Receiver, +) -> (Option, Vec) { + let mut estimate = None; + loop { + let e = tokio::time::timeout(Duration::from_secs(10), rx.recv()) + .await + .expect("no routing event") + .unwrap(); + match e { + tw_api::Event::RequestStarted { input_estimate, .. } => estimate = input_estimate, + tw_api::Event::RequestRouted { attempts, .. } => return (estimate, attempts), + _ => {} + } + } +} + +/// 等旗举起来,最多两秒 +async fn eventually(flag: &AtomicBool) -> bool { + for _ in 0..40 { + if flag.load(Ordering::SeqCst) { + return true; + } + tokio::time::sleep(Duration::from_millis(50)).await; + } + false +} + +fn outcomes(attempts: &[tw_api::AttemptView]) -> Vec<(&str, tw_api::AttemptOutcome)> { + attempts + .iter() + .map(|a| (a.provider.as_str(), a.outcome)) + .collect() +} + +#[tokio::test] +async fn a_slow_first_upstream_is_dropped_after_the_wait_and_the_next_one_answers() { + let slow = upstream(stalled()).await; + let good = upstream(prompt("hello")).await; + let (gw, mut rx, health) = gateway( + vec![ + provider("slow", &slow, Protocol::Anthropic), + provider("good", &good, Protocol::Anthropic), + ], + true, + ) + .await; + + let t = Instant::now(); + let (status, text) = post(gw, "/v1/messages", &messages(true)).await; + assert_eq!(status, 200, "{text}"); + assert!(text.contains("hello"), "{text}"); + assert!( + !text.contains("ping"), + "慢的那一家的开头不该到客户端:{text}" + ); + assert!(t.elapsed() >= Duration::from_secs(1), "等够了才换"); + assert!( + eventually(&slow.dropped).await, + "放弃的那一家连接要断开,它才会停下" + ); + + let (_, attempts) = estimate_and_attempts(&mut rx).await; + use tw_api::AttemptOutcome::{Served, SlowStart}; + assert_eq!(outcomes(&attempts), [("slow", SlowStart), ("good", Served)]); + let gave_up = &attempts[0]; + assert_eq!(gave_up.status, Some(200)); + assert_eq!( + gave_up.usage, + Some(tw_api::AttemptUsage { + input: 1200, + cache_read: 800, + cache_write: 0, + estimated: false, + }), + "上游在开头报了输入,记它报的" + ); + assert_eq!( + gave_up.error.as_ref().map(|m| m.code.as_str()), + Some("gw.slow_start") + ); + + // **慢不是坏**:不停用,几次之后也照样先试它 + for _ in 0..4 { + let (status, _) = post(gw, "/v1/messages", &messages(true)).await; + assert_eq!(status, 200); + } + assert!(health.is_available("slow"), "慢的那一家不该被停用"); + assert_eq!(slow.hits.load(Ordering::SeqCst), 5, "每次都还先试它"); + // 按成败分的负载均衡看的成功率也不记它:换走了五次,它一次都没失败过,也没答上过 + let rates = health.success_rates(&["slow".to_string(), "good".to_string()]); + assert_eq!(rates.get("slow"), None, "{rates:?}"); + assert_eq!(rates.get("good"), Some(&1.0), "{rates:?}"); +} + +#[tokio::test] +async fn the_last_upstream_is_not_given_up_on() { + // 第一家慢、被放弃;第二家也慢,但它是最后一家:照常等它 + let slow = upstream(stalled()).await; + let last = upstream(late(1_600, "finally")).await; + let (gw, mut rx, _) = gateway( + vec![ + provider("slow", &slow, Protocol::Anthropic), + provider("last", &last, Protocol::Anthropic), + ], + true, + ) + .await; + let (status, text) = post(gw, "/v1/messages", &messages(true)).await; + assert_eq!(status, 200, "{text}"); + assert!(text.contains("finally"), "{text}"); + let (_, attempts) = estimate_and_attempts(&mut rx).await; + use tw_api::AttemptOutcome::{Served, SlowStart}; + assert_eq!(outcomes(&attempts), [("slow", SlowStart), ("last", Served)]); +} + +#[tokio::test] +async fn a_next_upstream_that_is_paused_by_then_does_not_count() { + // 第一家慢,第二家在等的时候被停用了:后面没有接得下的,第一家就是最后一家,照常等它 + let slow = upstream(late(1_600, "patience")).await; + let other = upstream(prompt("other")).await; + let (gw, mut rx, health) = gateway( + vec![ + provider("slow", &slow, Protocol::Anthropic), + provider("other", &other, Protocol::Anthropic), + ], + true, + ) + .await; + let pause = tokio::spawn(async move { + tokio::time::sleep(Duration::from_millis(300)).await; + for _ in 0..3 { + health.record_failure("other"); + } + assert!(!health.is_available("other")); + }); + let (status, text) = post(gw, "/v1/messages", &messages(true)).await; + pause.await.unwrap(); + assert_eq!(status, 200, "{text}"); + assert!(text.contains("patience"), "{text}"); + assert_eq!(other.hits.load(Ordering::SeqCst), 0); + let (_, attempts) = estimate_and_attempts(&mut rx).await; + assert_eq!( + outcomes(&attempts), + [("slow", tw_api::AttemptOutcome::Served)] + ); +} + +#[tokio::test] +async fn a_next_upstream_that_cannot_take_the_request_does_not_count() { + // Claude Code 搜网页的请求强制要用服务端工具 `web_search`,发不到 Chat 格式的上游: + // 那一家接不下,第一家就是最后一家 + let slow = upstream(late(1_600, "searched")).await; + let chat = upstream(prompt("other")).await; + let (gw, mut rx, _) = gateway( + vec![ + provider("slow", &slow, Protocol::Anthropic), + provider("chat", &chat, Protocol::OpenaiChat), + ], + true, + ) + .await; + let body = json!({ + "model": "claude-sonnet-5", "max_tokens": 1024, "stream": true, + "messages": [{"role": "user", "content": "Perform a web search for the query: bedrock pricing"}], + "tools": [{"type": "web_search_20250305", "name": "web_search", "max_uses": 8}], + "tool_choice": {"type": "tool", "name": "web_search"} + }); + let (status, text) = post(gw, "/v1/messages", &body).await; + assert_eq!(status, 200, "{text}"); + assert!(text.contains("searched"), "{text}"); + assert_eq!(chat.hits.load(Ordering::SeqCst), 0); + let (_, attempts) = estimate_and_attempts(&mut rx).await; + assert_eq!( + outcomes(&attempts), + [("slow", tw_api::AttemptOutcome::Served)] + ); +} + +#[tokio::test] +async fn switched_off_the_slow_start_is_handed_on_as_before() { + let slow = upstream(late(1_600, "eventually")).await; + let other = upstream(prompt("other")).await; + let (gw, mut rx, _) = gateway( + vec![ + provider("slow", &slow, Protocol::Anthropic), + provider("other", &other, Protocol::Anthropic), + ], + false, + ) + .await; + let (status, text) = post(gw, "/v1/messages", &messages(true)).await; + assert_eq!(status, 200, "{text}"); + assert!(text.contains("eventually"), "{text}"); + assert_eq!(other.hits.load(Ordering::SeqCst), 0); + let (_, attempts) = estimate_and_attempts(&mut rx).await; + assert_eq!( + outcomes(&attempts), + [("slow", tw_api::AttemptOutcome::Served)] + ); +} + +#[tokio::test] +async fn thinking_counts_as_content() { + // 开头就在想(推理的字在流),正文要过一阵才来:这不是开头慢 + let thinker = upstream(Script { + header_delay_ms: 0, + steps: vec![ + (0, MESSAGE_START), + (200, THINKING), + (1_400, answer("thought")), + ], + hang: false, + }) + .await; + let other = upstream(prompt("other")).await; + let (gw, mut rx, _) = gateway( + vec![ + provider("thinker", &thinker, Protocol::Anthropic), + provider("other", &other, Protocol::Anthropic), + ], + true, + ) + .await; + let (status, text) = post(gw, "/v1/messages", &messages(true)).await; + assert_eq!(status, 200, "{text}"); + assert!( + text.contains("Let me think") && text.contains("thought"), + "{text}" + ); + assert_eq!(other.hits.load(Ordering::SeqCst), 0); + let (_, attempts) = estimate_and_attempts(&mut rx).await; + assert_eq!( + outcomes(&attempts), + [("thinker", tw_api::AttemptOutcome::Served)] + ); +} + +#[tokio::test] +async fn without_reported_usage_the_estimate_is_recorded_as_possibly_billed() { + // Chat 格式的上游开头只发一块角色、不报用量:记网关估的输入 + let silent = upstream(Script { + header_delay_ms: 0, + steps: vec![( + 0, + "data: {\"choices\":[{\"delta\":{\"role\":\"assistant\",\"content\":\"\"}}]}\n\n", + )], + hang: true, + }) + .await; + let good = upstream(Script { + header_delay_ms: 0, + steps: vec![( + 0, + "data: {\"choices\":[{\"delta\":{\"content\":\"hi there\"},\"finish_reason\":\"stop\"}]}\n\ndata: [DONE]\n\n", + )], + hang: false, + }) + .await; + let (gw, mut rx, _) = gateway( + vec![ + provider("silent", &silent, Protocol::OpenaiChat), + provider("good", &good, Protocol::OpenaiChat), + ], + true, + ) + .await; + let body = json!({"model": "gpt-5", "stream": true, + "messages": [{"role": "user", "content": "Say hello."}]}); + let (status, text) = post(gw, "/v1/chat/completions", &body).await; + assert_eq!(status, 200, "{text}"); + assert!(text.contains("hi there"), "{text}"); + assert!(eventually(&silent.dropped).await); + let (estimate, attempts) = estimate_and_attempts(&mut rx).await; + use tw_api::AttemptOutcome::{Served, SlowStart}; + assert_eq!( + outcomes(&attempts), + [("silent", SlowStart), ("good", Served)] + ); + let estimate = estimate.expect("解得开的请求有估算"); + assert!(estimate > 0); + assert_eq!( + attempts[0].usage, + Some(tw_api::AttemptUsage { + input: estimate, + cache_read: 0, + cache_write: 0, + estimated: true, + }) + ); +} + +#[tokio::test] +async fn response_headers_that_never_come_count_as_a_slow_start() { + // 响应头都迟迟不来(中转站排着队):一样从发出请求算起,到点放弃 + let queued = upstream(Script { + header_delay_ms: 5_000, + ..prompt("too late") + }) + .await; + let good = upstream(prompt("hello")).await; + let (gw, mut rx, _) = gateway( + vec![ + provider("queued", &queued, Protocol::Anthropic), + provider("good", &good, Protocol::Anthropic), + ], + true, + ) + .await; + let t = Instant::now(); + let (status, text) = post(gw, "/v1/messages", &messages(true)).await; + assert_eq!(status, 200, "{text}"); + assert!(text.contains("hello"), "{text}"); + assert!(t.elapsed() < Duration::from_secs(4), "没有等到响应头"); + assert!(eventually(&queued.dropped).await, "还在等的请求要断开"); + let (_, attempts) = estimate_and_attempts(&mut rx).await; + use tw_api::AttemptOutcome::{Served, SlowStart}; + assert_eq!( + outcomes(&attempts), + [("queued", SlowStart), ("good", Served)] + ); + assert_eq!(attempts[0].status, None, "响应头没到,没有状态码"); + assert!(attempts[0].usage.is_some_and(|u| u.estimated)); +} + +#[tokio::test] +async fn a_request_that_does_not_stream_is_not_switched() { + // 整包的请求本来就要等全部生成完:响应头来得晚也照常等 + let slow = upstream(Script { + header_delay_ms: 1_600, + ..prompt("whole") + }) + .await; + let other = upstream(prompt("other")).await; + let (gw, mut rx, _) = gateway( + vec![ + provider("slow", &slow, Protocol::Anthropic), + provider("other", &other, Protocol::Anthropic), + ], + true, + ) + .await; + let (status, text) = post(gw, "/v1/messages", &messages(false)).await; + assert_eq!(status, 200, "{text}"); + assert_eq!(other.hits.load(Ordering::SeqCst), 0); + let (_, attempts) = estimate_and_attempts(&mut rx).await; + assert_eq!( + outcomes(&attempts), + [("slow", tw_api::AttemptOutcome::Served)] + ); +} diff --git a/crates/tw-observe/src/bus.rs b/crates/tw-observe/src/bus.rs index 96097980..7ea989a2 100644 --- a/crates/tw-observe/src/bus.rs +++ b/crates/tw-observe/src/bus.rs @@ -424,6 +424,7 @@ mod tests { status: Some(503), error: None, ms: 10, + usage: None, }, tw_api::AttemptView { provider: served_by.into(), @@ -432,6 +433,7 @@ mod tests { status: Some(200), error: None, ms: 20, + usage: None, }, ], billing: tw_api::Billing::PerToken, diff --git a/crates/tw-store/src/db.rs b/crates/tw-store/src/db.rs index e97981b0..2476db7d 100644 --- a/crates/tw-store/src/db.rs +++ b/crates/tw-store/src/db.rs @@ -2730,6 +2730,7 @@ mod cost_state_tests { status: Some(200), error: None, ms: 1, + usage: None, }], ..Default::default() }) diff --git a/crates/tw-store/src/recorder.rs b/crates/tw-store/src/recorder.rs index 0f6e6279..5b60b8e0 100644 --- a/crates/tw-store/src/recorder.rs +++ b/crates/tw-store/src/recorder.rs @@ -847,6 +847,7 @@ mod tests { text: "上游响应超时".into(), }), ms: 10_003, + usage: None, }, tw_api::AttemptView { provider: "中转".into(), @@ -855,6 +856,7 @@ mod tests { status: Some(200), error: None, ms: 5_042, + usage: None, }, ], billing: tw_api::Billing::PerToken, @@ -899,6 +901,7 @@ mod tests { status: Some(200), error: None, ms: 300, + usage: None, }], billing: tw_api::Billing::PerToken, }); @@ -928,6 +931,7 @@ mod tests { status: Some(529), error: None, ms: 100, + usage: None, }, tw_api::AttemptView { provider: "中转".into(), @@ -936,6 +940,7 @@ mod tests { status: Some(200), error: None, ms: 300, + usage: None, }, ], billing: tw_api::Billing::PerToken, @@ -984,6 +989,7 @@ mod tests { status: Some(200), error: None, ms: 300, + usage: None, }], billing: tw_api::Billing::PerToken, }); @@ -1104,6 +1110,7 @@ mod tests { status: None, error: None, ms: 0, + usage: None, }], billing: tw_api::Billing::PerToken, }); @@ -1416,6 +1423,7 @@ mod tests { status: Some(404), error: None, ms: 80, + usage: None, }], billing: tw_api::Billing::Free, affinity: None, @@ -1489,6 +1497,7 @@ mod billing_tests { status: Some(200), error: None, ms: 5, + usage: None, }], billing, } @@ -1576,6 +1585,7 @@ mod billing_tests { status: Some(101), error: None, ms: 40, + usage: None, }], billing, } @@ -1964,6 +1974,7 @@ mod translation_tests { status: Some(200), error: None, ms: 1, + usage: None, }) .collect(), billing: tw_api::Billing::PerToken, diff --git a/docs/config.md b/docs/config.md index ff388656..f30d6155 100644 --- a/docs/config.md +++ b/docs/config.md @@ -893,6 +893,19 @@ Before the first content of a streamed answer reaches the client, an error the upstream sends in the stream moves the request to the next candidate, the same as an error status would. +An upstream can also be slow to start: it accepts the request and then sends +nothing for a long time. With `next_on_slow_start`, the request moves on to the +next candidate when no content has arrived `stream_start_wait_secs` after it +was sent. It is off by default, because models that think before they write +can take long to start; with it on, wait 30 seconds or more. The last +candidate always waits, and the upstream given up on is not set aside. + +```yaml +failover: + stream_start_wait_secs: 30 + next_on_slow_start: true +``` + @@ -905,6 +918,7 @@ the same as an error status would. | `quota_pause_secs` | integer | `3600` | Seconds to set aside an upstream whose quota is used up when it does not say when the quota resets. When it does, the upstream is set aside until then. | | `rate_limit_max_pause_secs` | integer | `3600` | A rate-limited upstream is set aside for the time its `Retry-After` gives, at most this many seconds. Without `Retry-After` it counts as a failure without a stated reason. | | `stream_start_wait_secs` | integer | `15` | Seconds to hold a streamed answer until its first content arrives. An error before then moves the request to the next upstream; after this long, what has arrived is passed on. From 1 to 120. | +| `next_on_slow_start` | bool | `false` | When a streamed answer still has no content `stream_start_wait_secs` after the request was sent, give up on that upstream and send the request to the next one. The last upstream always waits. The upstream given up on is not set aside. Needs `stream_start_wait_secs` of at least 5. | ### `aliases` @@ -1015,9 +1029,10 @@ group shares out requests by the result in the same way as above. last 50 requests within the past 30 minutes. Server errors, rate limits, used-up quota or balance, rejected credentials, timeouts and connection errors count as failures; errors caused by the request itself do not, and - neither does a client that cancels. An upstream that keeps failing keeps a - twentieth of its weight, so it still gets the occasional new conversation - and its recovery is noticed; one that fails outright is set aside by + neither does a client that cancels or a switch away from a stream that is + slow to start. An upstream that keeps failing keeps a twentieth of its + weight, so it still gets the occasional new conversation and its recovery + is noticed; one that fails outright is set aside by [`failover`](#cfg-failover) as before. - `latency-health`: both factors, multiplied. diff --git a/docs/config.zh-CN.md b/docs/config.zh-CN.md index fdd50299..7fee25d7 100644 --- a/docs/config.zh-CN.md +++ b/docs/config.zh-CN.md @@ -702,6 +702,17 @@ security: 流式回答在第一段内容交给客户端之前,上游在流里报的错误和错误状态码一样,会把 请求换到下一个候选。 +上游也可能开头很慢:收下请求之后很久都不发内容。开启 `next_on_slow_start` 后,请求 +发出 `stream_start_wait_secs` 秒仍没有内容,就换到下一个候选。默认关闭,因为先思考 +再输出的模型本来就可能很久才开始;开启时建议等 30 秒以上。最后一个候选总是等下去, +被放弃的上游不会停用。 + +```yaml +failover: + stream_start_wait_secs: 30 + next_on_slow_start: true +``` + @@ -714,6 +725,7 @@ security: | `quota_pause_secs` | 整数 | `3600` | 上游报告额度用完、但没有给出重置时间时停用的秒数。给出了重置时间的,停用到那一刻。 | | `rate_limit_max_pause_secs` | 整数 | `3600` | 被限流的上游按它给的 `Retry-After` 停用,最多这么多秒。没有 `Retry-After` 的按没有说明原因的失败计。 | | `stream_start_wait_secs` | 整数 | `15` | 流式回答在第一段内容到达前最多暂存的秒数。在此之前上游报错,请求换到下一家;超过这个时间,已收到的部分照常交给客户端。取值 1 到 120。 | +| `next_on_slow_start` | 布尔 | `false` | 流式回答在请求发出 `stream_start_wait_secs` 秒后仍没有内容时,放弃这家上游,把请求交给下一家。最后一家总是等下去。被放弃的上游不会停用。开启时 `stream_start_wait_secs` 至少为 5。 | ### `aliases` @@ -781,7 +793,7 @@ groups: - `weights`(默认):只按权重。 - `latency`:越快的上游分得越多。快慢看典型的首字节时间,与 `url-test` 使用同一份测量。比组内居中者快一倍的上游,权重乘以四;最多乘以十,最少乘以十分之一。 -- `health`:越少失败的上游分得越多。依据是最近 30 分钟内的最近 50 次请求:服务器错误、限流、额度或余额用尽、凭据被拒、超时和连接失败算作失败;请求本身导致的错误不算,客户端取消也不算。经常失败的上游至少保留权重的二十分之一,仍会偶尔分到新对话,以便发现它已经恢复;完全失败的上游照旧由 [`failover`](#cfg-failover) 暂停。 +- `health`:越少失败的上游分得越多。依据是最近 30 分钟内的最近 50 次请求:服务器错误、限流、额度或余额用尽、凭据被拒、超时和连接失败算作失败;请求本身导致的错误不算,客户端取消、因开头太慢而换走也不算。经常失败的上游至少保留权重的二十分之一,仍会偶尔分到新对话,以便发现它已经恢复;完全失败的上游照旧由 [`failover`](#cfg-failover) 暂停。 - `latency-health`:两个系数相乘。 测量还不够的上游按中等对待。与只按权重时一样,进行中的对话留在原来的上游,差额由新对话补齐。 From 800ffc2176bf117494dff0568492aa56c285ec2b Mon Sep 17 00:00:00 2001 From: fylorn <249551762+fylorn@users.noreply.github.com> Date: Mon, 5 Oct 2026 15:58:49 +0800 Subject: [PATCH 04/22] Model specs set by hand: an upstream's context window and output limit win over the price table A relay's own models are often missing from the price table, and the table is sometimes wrong, so clients were told no context window (or a wrong one) and the gateway could not tell whether a conversation still fits its model. `providers[].model_specs` maps an exact model id to `context_window` and/or `max_output_tokens`. One function in tw-config (`model_specs::resolve`, reached through `Provider::model_limits` / `Config::model_limits`) decides the precedence, field by field, and every reader goes through it: the `/v1/models` listing in all three shapes (a real model is now described by the first upstream offering it, so its spec applies; aliases as before), the upstream model rows, alias context windows, the pipeline's "input outgrew the held decision's context window" check, and the output limit filled in when a request is converted to Anthropic. `ModelRow` gains `context_window_source`, `max_output_tokens` and `max_output_tokens_source` so the UI can show which number is manual. `PUT /provider-model-spec` sets or (both null) removes one entry, editing only that upstream's `model_specs` so the rest of the entry stays byte for byte; the provider dialog keeps existing specs, and they follow renames and deletes because they live inside the entry. Validation: `config.model_spec_*`. Co-Authored-By: Claude Opus 5.5 --- crates/tw-api/msg-codes.txt | 4 + crates/tw-api/src/ep.rs | 3 + crates/tw-api/src/lib.rs | 42 ++- crates/tw-config/src/edit.rs | 217 ++++++++++++ crates/tw-config/src/lib.rs | 9 +- crates/tw-config/src/model_specs.rs | 233 +++++++++++++ crates/tw-config/src/validate.rs | 148 +++++++++ crates/tw-config/src/wire.rs | 11 +- crates/tw-config/tests/manual/schema.rs | 34 ++ crates/tw-control/src/aliases.rs | 5 +- crates/tw-control/src/resources.rs | 51 ++- crates/tw-control/tests/model_specs.rs | 333 +++++++++++++++++++ crates/tw-gateway/src/server/listing.rs | 84 +++-- crates/tw-gateway/src/server/pipeline.rs | 10 +- crates/tw-gateway/src/server/pipeline/hop.rs | 1 + crates/tw-gateway/src/translate.rs | 51 ++- crates/tw-gateway/tests/model_specs.rs | 244 ++++++++++++++ docs/config.md | 35 ++ docs/config.zh-CN.md | 26 ++ 19 files changed, 1495 insertions(+), 46 deletions(-) create mode 100644 crates/tw-config/src/model_specs.rs create mode 100644 crates/tw-control/tests/model_specs.rs create mode 100644 crates/tw-gateway/tests/model_specs.rs diff --git a/crates/tw-api/msg-codes.txt b/crates/tw-api/msg-codes.txt index b296f7e4..a8e71ea4 100644 --- a/crates/tw-api/msg-codes.txt +++ b/crates/tw-api/msg-codes.txt @@ -67,6 +67,10 @@ config.edit.unwritable config.empty_key config.empty_models_only config.failover_range +config.model_spec_blank_model +config.model_spec_empty +config.model_spec_wildcard +config.model_spec_zero config.name_collision config.no_clients config.plugin.bad_id diff --git a/crates/tw-api/src/ep.rs b/crates/tw-api/src/ep.rs index 7529a218..c4d6c3d6 100644 --- a/crates/tw-api/src/ep.rs +++ b/crates/tw-api/src/ep.rs @@ -95,6 +95,9 @@ endpoints! { UpdateProvider: PUT "/providers/{name}" [name], api::ProviderSave => api::ConfigWritten; DeleteProvider: DELETE "/providers/{name}" [name], api::BaseVersion => api::ConfigWritten; ProviderModels: GET "/providers/{name}/models" [name], () => api::ProviderModelsView; + /// 手写一家上游的一个模型的上下文窗口、输出上限,优先于价目表;两项都空就删掉。 + /// **不在 `/providers/` 底下**:写死的一段会盖住 `/providers/{name}` + SetModelSpec: PUT "/provider-model-spec", api::ModelSpecSave => api::ConfigWritten; RefreshProviderModels: POST "/providers/{name}/models/refresh" [name], () => api::ProviderModelsView; RefreshStaleModels: POST "/models/refresh", () => api::ModelsRefreshing; CreateProxy: POST "/proxies", api::ProxySave => api::ConfigWritten; diff --git a/crates/tw-api/src/lib.rs b/crates/tw-api/src/lib.rs index 341eb504..500026cb 100644 --- a/crates/tw-api/src/lib.rs +++ b/crates/tw-api/src/lib.rs @@ -165,6 +165,16 @@ slug_enum! { } } +slug_enum! { + /// 上下文窗口、输出上限这样的模型规格从哪儿来。 + pub enum SpecSource { + /// 价目表 + PriceTable = "price_table", + /// 这一家上游手写的(配置里的 `model_specs`),优先于价目表 + Manual = "manual", + } +} + slug_enum! { /// 一个上游现在能不能进候选链。 pub enum Health { @@ -2895,9 +2905,18 @@ pub struct ModelRow { pub id: String, /// 在启用范围里 pub enabled: bool, - /// 上下文窗口,来自默认价目表 + /// 上下文窗口:这一家手写的(`model_specs`),没写时来自价目表 #[serde(default, skip_serializing_if = "Option::is_none")] pub context_window: Option, + /// `context_window` 从哪儿来。不知道上下文窗口时没有 + #[serde(default, skip_serializing_if = "Option::is_none")] + pub context_window_source: Option, + /// 一次最多输出多少 token:这一家手写的,没写时来自价目表 + #[serde(default, skip_serializing_if = "Option::is_none")] + pub max_output_tokens: Option, + /// `max_output_tokens` 从哪儿来。不知道输出上限时没有 + #[serde(default, skip_serializing_if = "Option::is_none")] + pub max_output_tokens_source: Option, /// 按这个上游选的价目表查到的价格。空 = 无法计价 #[serde(default, skip_serializing_if = "Option::is_none")] pub price: Option, @@ -2961,6 +2980,25 @@ pub struct ProviderSave { pub base_version: Option, } +/// 设一家上游的一个模型的规格(`PUT /provider-model-spec`):价目表不认识这个模型、 +/// 或者写错了时手写。**两项都空就是删掉这一项**,回到价目表。 +#[derive(Debug, Clone, Default, Serialize, Deserialize)] +#[cfg_attr(feature = "ts", derive(ts_rs::TS))] +pub struct ModelSpecSave { + pub provider: String, + /// 模型 ID,和这家的清单里写的完全相等。去掉首尾空白 + pub model: String, + /// 上下文窗口(token)。空 = 用价目表的 + #[serde(default, skip_serializing_if = "Option::is_none")] + pub context_window: Option, + /// 输出上限(token)。空 = 用价目表的 + #[serde(default, skip_serializing_if = "Option::is_none")] + pub max_output_tokens: Option, + /// 你基于哪一版。**对不上就是 409** + #[serde(default, skip_serializing_if = "Option::is_none")] + pub base_version: Option, +} + /// 检测一个上游,**不保存**。 #[derive(Debug, Clone, Serialize, Deserialize)] #[cfg_attr(feature = "ts", derive(ts_rs::TS))] @@ -3059,7 +3097,7 @@ pub struct AliasView { /// 清单里有一个和别名同名的真模型、而别名的列表里没有这个名称的上游:**这个名称 /// 不会再发给它们**(别名优先) pub shadows: Vec, - /// 上下文窗口,来自默认价目表:第一家能服务它的上游发出的那个模型的 + /// 上下文窗口:第一家能服务它的上游发出的那个模型的,这一家手写的优先于价目表 #[serde(default, skip_serializing_if = "Option::is_none")] pub context_window: Option, /// 最近 24 小时里客户端用这个名称发来的请求 diff --git a/crates/tw-config/src/edit.rs b/crates/tw-config/src/edit.rs index 93e1da4d..22bf0508 100644 --- a/crates/tw-config/src/edit.rs +++ b/crates/tw-config/src/edit.rs @@ -620,6 +620,105 @@ pub fn remove_alias(text: &str, name: &str) -> Result { Ok(out) } +// ─────────────────────────────────────────────────────────── 上游的模型规格 + +/// 上游 `provider` 的 `model_specs` 里,设(`spec` 是 `Some`)或者删掉(`None`)`model` +/// 那一项。**这一家的其余字段一个字节都不动** —— 不走 [`upsert`]:那条路按读进来再写回去的 +/// 结构比,用户写成和默认值一样的字段(`proxy: direct`)会被顺手删掉。 +/// +/// - 删掉的是最后一项时,`model_specs` 整个删:默认值不写进文件 +/// - 要删的那一项本来就没有:原样返回 +/// - 那一家写成了行内(`- { name: a, … }`):整项换成块式,位置不变(和 [`upsert`] 一样); +/// `model_specs` 写成了行内:整张换掉 +/// +/// 两项都空的 `spec` 交过来是调用方的错 —— 那是「删掉」,传 `None`。 +pub fn set_model_spec( + text: &str, + provider: &str, + model: &str, + spec: Option<&crate::ModelSpec>, +) -> Result { + reject_multiline(&Value::String(model.to_string()))?; + let doc = parse(text)?; + let index = PROVIDERS + .index_of(&doc, provider) + .ok_or_else(|| EditError::NotFound { + what: PROVIDERS.what, + name: provider.to_string(), + })?; + let item = PROVIDERS.items(&doc)[index] + .as_mapping() + .cloned() + .unwrap_or_default(); + let mut specs = match item.get(MODEL_SPECS) { + Some(Value::Mapping(m)) => m.clone(), + _ => Mapping::new(), + }; + let key = Value::String(model.to_string()); + let had = specs.contains_key(&key); + let value = match spec { + Some(s) => { + Some(serde_yaml_ng::to_value(s).map_err(|e| EditError::Unwritable(e.to_string()))?) + } + None => None, + }; + match &value { + Some(v) if specs.get(&key) == Some(v) => return Ok(text.to_string()), + Some(v) => { + specs.insert(key.clone(), v.clone()); + } + None if !had => return Ok(text.to_string()), + None => { + specs.remove(&key); + } + } + let mut new_item = item.clone(); + if specs.is_empty() { + new_item.remove(MODEL_SPECS); + } else { + new_item.insert( + Value::String(MODEL_SPECS.into()), + Value::Mapping(specs.clone()), + ); + } + + let steps = PROVIDERS.steps(); + let mut at = steps.clone(); + at.push(Step::Index(index)); + let mut field = at.clone(); + field.push(Step::key(MODEL_SPECS)); + let mut entry = field.clone(); + entry.push(Step::Key(model.to_string())); + let out = if tw_yaml::is_flow_at(text, &at)? { + let block = render_block(&Value::Mapping(new_item.clone()))?; + tw_yaml::replace_item(text, &steps, index, &block)? + } else if specs.is_empty() { + tw_yaml::remove_key(text, &field)? + } else if tw_yaml::is_flow_at(text, &field)? { + put_value(text, &field, &Value::Mapping(specs))? + } else { + match &value { + Some(v) => put_value(text, &entry, v)?, + None => tw_yaml::remove_key(text, &entry)?, + } + }; + + // ── 语义核对 ───────────────────────────────────────────────────── + let mut expected = doc; + put_item(&mut expected, PROVIDERS, index, Value::Mapping(new_item)); + let got = parse(&out).map_err(|e| EditError::SelfCheck(e.to_string()))?; + if got != expected { + return Err(EditError::SelfCheck(format!( + "{} `{provider}` model_specs `{model}`", + PROVIDERS.what + ))); + } + Ok(out) +} + +/// 上游里手写的模型规格那一项的键 +const MODEL_SPECS: &str = "model_specs"; + /// 文件里的别名表,按书写顺序。键一律当字符串(配置读进来时就是这么读的) fn alias_table(doc: &Value) -> Vec<(String, Value)> { let Some(Value::Mapping(m)) = doc.get(ALIASES) else { @@ -1208,4 +1307,122 @@ routes: [] assert_eq!(aliases_of(&out)[2], alias(name, &["m"]), "{out}"); } } + + // ── 模型规格 ───────────────────────────────────────────────────── + + fn spec(context_window: Option, max_output_tokens: Option) -> crate::ModelSpec { + crate::ModelSpec { + context_window, + max_output_tokens, + } + } + + fn specs_of(text: &str, provider: &str) -> Vec<(String, crate::ModelSpec)> { + let cfg: crate::Config = serde_yaml_ng::from_str(text).unwrap(); + cfg.providers + .into_iter() + .find(|p| p.name == provider) + .unwrap() + .model_specs + .into_iter() + .collect() + } + + /// 设一项、改一项、删到一项不剩:这一家的其余字段和旁边的注释原样 + #[test] + fn a_model_spec_is_set_changed_and_removed_without_touching_the_rest() { + let out = set_model_spec(CFG, "官方", "glm-5", Some(&spec(Some(128_000), None))).unwrap(); + assert!( + out.contains("base_url: https://api.anthropic.com # 直连"), + "{out}" + ); + assert!(out.contains("# 两家上游"), "{out}"); + assert_eq!( + specs_of(&out, "官方"), + [("glm-5".to_string(), spec(Some(128_000), None))] + ); + // 要加引号的模型 ID 读回来还是它 + let odd = "us.anthropic.claude-fable-5-v1:0"; + let out = set_model_spec(&out, "官方", odd, Some(&spec(None, Some(32_000)))).unwrap(); + let out = set_model_spec( + &out, + "官方", + "glm-5", + Some(&spec(Some(200_000), Some(8_000))), + ) + .unwrap(); + assert_eq!( + specs_of(&out, "官方"), + [ + ("glm-5".to_string(), spec(Some(200_000), Some(8_000))), + (odd.to_string(), spec(None, Some(32_000))), + ] + ); + // 一样的值、删一项没有的:原样 + assert_eq!( + set_model_spec( + &out, + "官方", + "glm-5", + Some(&spec(Some(200_000), Some(8_000))) + ) + .unwrap(), + out + ); + assert_eq!(set_model_spec(&out, "官方", "ghost", None).unwrap(), out); + + let out = set_model_spec(&out, "官方", "glm-5", None).unwrap(); + assert_eq!( + specs_of(&out, "官方"), + [(odd.to_string(), spec(None, Some(32_000)))] + ); + // 最后一项删掉,`model_specs` 整个不写 + let out = set_model_spec(&out, "官方", odd, None).unwrap(); + assert_eq!(out, CFG); + } + + /// 写成行内的上游整项换成块式;行内的 `model_specs` 整张换掉 + #[test] + fn a_model_spec_goes_into_an_inline_upstream_or_table_too() { + let out = set_model_spec(CFG, "relay", "m", Some(&spec(Some(1_000), None))).unwrap(); + assert_eq!( + specs_of(&out, "relay"), + [("m".to_string(), spec(Some(1_000), None))] + ); + let cfg: crate::Config = serde_yaml_ng::from_str(&out).unwrap(); + assert_eq!(cfg.providers[1].base_url, "https://relay.example"); + assert_eq!(cfg.providers[0].name, "官方", "位置变了:{out}"); + + let inline = CFG.replace( + " key: sk-a\n", + " key: sk-a\n model_specs: { a: { context_window: 5 } }\n", + ); + let out = set_model_spec(&inline, "官方", "b", Some(&spec(None, Some(7)))).unwrap(); + assert_eq!( + specs_of(&out, "官方"), + [ + ("a".to_string(), spec(Some(5), None)), + ("b".to_string(), spec(None, Some(7))), + ] + ); + let out = set_model_spec(&inline, "官方", "a", None).unwrap(); + assert!(!out.contains("model_specs"), "{out}"); + } + + #[test] + fn a_model_spec_for_a_missing_upstream_is_refused() { + let e = set_model_spec(CFG, "ghost", "m", Some(&spec(Some(1), None))).unwrap_err(); + assert!( + matches!( + e, + EditError::NotFound { + what: "upstream", + .. + } + ), + "{e}" + ); + let e = set_model_spec(CFG, "官方", "a\nb", Some(&spec(Some(1), None))).unwrap_err(); + assert!(matches!(e, EditError::Multiline), "{e}"); + } } diff --git a/crates/tw-config/src/lib.rs b/crates/tw-config/src/lib.rs index 28490198..334702ad 100644 --- a/crates/tw-config/src/lib.rs +++ b/crates/tw-config/src/lib.rs @@ -15,6 +15,7 @@ pub mod edit; mod failover; pub mod history; mod init; +pub mod model_specs; pub mod nics; pub mod plugins; pub mod private_dir; @@ -34,6 +35,7 @@ mod wire; pub use aliases::{Alias, Aliases}; pub use credential::{CredentialError, Header, Headers, Secret, SecretResolveError, auth_header}; pub use init::{generate_control_key, generate_initial, generate_key}; +pub use model_specs::{ModelLimits, ModelSpec, Sourced, SpecSource}; pub use plugins::Plugin; pub use proxy::{DIRECT, OnProxyFail, Proxy, ProxyKind, SYSTEM}; pub use validate::ValidationError; @@ -186,6 +188,7 @@ impl Default for Provider { models_only: None, billing: Billing::PerToken, pricing: None, + model_specs: std::collections::BTreeMap::new(), disabled: false, } } @@ -815,6 +818,10 @@ pub struct Provider { /// 按哪张价目表计价。不写就是默认价目表。 #[serde(default, skip_serializing_if = "Option::is_none")] pub pricing: Option, + /// 手写的模型规格:模型 ID(完全相等,没有通配)→ 上下文窗口、输出上限。写了就 + /// 优先于价目表,见 [`model_specs`]。价目表不认识的中转站模型靠它说出上下文窗口 + #[serde(default, skip_serializing_if = "std::collections::BTreeMap::is_empty")] + pub model_specs: std::collections::BTreeMap, /// 停用。**配置原样留着**:不参与路由,它的模型也不出现在 /// `/v1/models` 里。要暂时不用一家上游时,比删掉再重新填一遍凭据好。 #[serde(default, skip_serializing_if = "is_default")] @@ -1181,7 +1188,7 @@ pub use security::{ }; // Billing 在本文件里定义,这里不必再导出 pub use store::{Fingerprint, Loaded, StoreError, version_of}; -pub use validate::{check_aliases, validate}; +pub use validate::{check_aliases, check_model_spec, validate}; pub fn default_path() -> PathBuf { tw_api::data::dir().join("config.yaml") diff --git a/crates/tw-config/src/model_specs.rs b/crates/tw-config/src/model_specs.rs new file mode 100644 index 00000000..6ec9335c --- /dev/null +++ b/crates/tw-config/src/model_specs.rs @@ -0,0 +1,233 @@ +//! 手写的模型规格:某家上游的某个模型的上下文窗口和输出上限。 +//! +//! ```yaml +//! providers: +//! - name: relay +//! base_url: https://relay.example.com/v1 +//! model_specs: +//! glm-5-air: { context_window: 128000, max_output_tokens: 16384 } +//! ``` +//! +//! 价目表里没有的模型(中转站自己的名字)说不出上下文窗口,价目表写错的也有。手写的 +//! **只管这一家的这一个模型**(名字要完全相等,没有通配),写了就优先于价目表。 +//! +//! **哪个数优先只在 [`resolve`] 里定一次。**列模型(`/v1/models` 的各种格式)、上游页的 +//! 模型清单、别名、管线里「输入超出了上下文」、转换到 Anthropic 时补的输出上限,全都 +//! 经它取数 —— 哪一处另起炉灶,哪一处就会忘了手写的那个数。 + +use serde::{Deserialize, Serialize}; + +use crate::{Config, Provider}; + +/// `providers[].model_specs` 的一项。两项都可选,但至少写一项(校验管)。 +#[derive(Debug, Clone, Copy, Default, PartialEq, Eq, Serialize, Deserialize)] +#[serde(deny_unknown_fields)] +pub struct ModelSpec { + /// 一次最多输入多少 token,也就是上下文窗口 + #[serde(default, skip_serializing_if = "Option::is_none")] + pub context_window: Option, + /// 一次最多输出多少 token + #[serde(default, skip_serializing_if = "Option::is_none")] + pub max_output_tokens: Option, +} + +impl ModelSpec { + /// 两项都没写。写进配置里的不许这样([`crate::ValidationError::ModelSpecEmpty`]); + /// 界面交过来两项都空,是要删掉这一项 + pub fn is_empty(&self) -> bool { + self.context_window.is_none() && self.max_output_tokens.is_none() + } +} + +/// 一个数是从哪儿来的。 +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub enum SpecSource { + /// 价目表 + PriceTable, + /// 这一家的 `model_specs` + Manual, +} + +/// 一个 token 数和它的来源。 +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub struct Sourced { + pub tokens: u64, + pub source: SpecSource, +} + +/// 一个模型的上下文窗口和输出上限。**不知道就是 `None`**:编出来的数客户端会照着截断。 +#[derive(Debug, Clone, Copy, Default, PartialEq, Eq)] +pub struct ModelLimits { + pub context_window: Option, + pub max_output_tokens: Option, +} + +impl ModelLimits { + pub fn context_window(&self) -> Option { + self.context_window.map(|s| s.tokens) + } + + pub fn max_output_tokens(&self) -> Option { + self.max_output_tokens.map(|s| s.tokens) + } + + /// 不知道哪家上游服务它时:只看默认价目表。 + pub fn priced(book: &tw_pricing::PriceBook, model: &str) -> Self { + resolve(book, None, None, model) + } +} + +impl Provider { + /// 这一家的这个模型的上下文窗口和输出上限:手写的优先,没写的取它选的价目表。 + pub fn model_limits(&self, book: &tw_pricing::PriceBook, model: &str) -> ModelLimits { + resolve(book, Some(&self.name), self.model_specs.get(model), model) + } +} + +impl Config { + /// 同 [`Provider::model_limits`],按上游的名字找。配置里没有这一家(刚删掉)时只看 + /// 价目表。 + pub fn model_limits( + &self, + book: &tw_pricing::PriceBook, + provider: &str, + model: &str, + ) -> ModelLimits { + let spec = self + .providers + .iter() + .find(|p| p.name == provider) + .and_then(|p| p.model_specs.get(model)); + resolve(book, Some(provider), spec, model) + } +} + +/// **唯一定先后的地方**:一项一项看,手写了就用手写的,没写的那一项取价目表。两项各管 +/// 各的 —— 只写了上下文窗口,输出上限照样来自价目表。 +/// +/// 价目表按这一家选的那张查(`provider`),和计价同一个查法;上下文窗口和输出上限在 +/// 每张价目表里都取自默认价目表,所以给不给上游通常是同一个数。 +fn resolve( + book: &tw_pricing::PriceBook, + provider: Option<&str>, + spec: Option<&ModelSpec>, + model: &str, +) -> ModelLimits { + let priced = match provider { + Some(p) => book.resolve_for(p, model), + None => book.resolve(None, model), + }; + let price = priced.as_ref().map(|r| &r.price); + let pick = |manual: Option, table: Option| match manual { + Some(n) => Some(Sourced { + tokens: n.into(), + source: SpecSource::Manual, + }), + None => table.map(|n| Sourced { + tokens: n, + source: SpecSource::PriceTable, + }), + }; + ModelLimits { + context_window: pick( + spec.and_then(|s| s.context_window), + price.and_then(|p| p.max_input_tokens), + ), + max_output_tokens: pick( + spec.and_then(|s| s.max_output_tokens), + price.and_then(|p| p.max_output_tokens), + ), + } +} + +#[cfg(test)] +mod tests { + use super::*; + + fn book() -> tw_pricing::PriceBook { + tw_pricing::PriceBook::builtin().unwrap() + } + + fn relay(specs: &[(&str, ModelSpec)]) -> Provider { + Provider { + name: "relay".into(), + base_url: "https://relay.example.com/v1".into(), + model_specs: specs.iter().map(|(m, s)| (m.to_string(), *s)).collect(), + ..Default::default() + } + } + + #[test] + fn a_spec_written_by_hand_wins_over_the_price_table_field_by_field() { + let b = book(); + let p = relay(&[( + "claude-sonnet-4-5", + ModelSpec { + context_window: Some(1_000_000), + max_output_tokens: None, + }, + )]); + let l = p.model_limits(&b, "claude-sonnet-4-5"); + assert_eq!( + l.context_window, + Some(Sourced { + tokens: 1_000_000, + source: SpecSource::Manual + }) + ); + // 没写的那一项照样取价目表 + assert_eq!( + l.max_output_tokens, + Some(Sourced { + tokens: 64_000, + source: SpecSource::PriceTable + }) + ); + // 别的模型、别的上游不受影响 + assert_eq!( + p.model_limits(&b, "claude-haiku-4-5") + .context_window + .map(|s| s.source), + Some(SpecSource::PriceTable) + ); + let cfg = Config { + providers: vec![p, relay(&[])], + ..Default::default() + }; + assert_eq!( + cfg.model_limits(&b, "relay", "claude-sonnet-4-5") + .context_window(), + Some(1_000_000) + ); + assert_eq!( + ModelLimits::priced(&b, "claude-sonnet-4-5").context_window(), + Some(200_000) + ); + } + + #[test] + fn a_model_the_price_table_does_not_know_has_only_what_was_written() { + let b = book(); + let p = relay(&[( + "中转自有模型", + ModelSpec { + context_window: None, + max_output_tokens: Some(8_000), + }, + )]); + let l = p.model_limits(&b, "中转自有模型"); + assert_eq!(l.context_window, None); + assert_eq!(l.max_output_tokens(), Some(8_000)); + // 名字要完全相等:带日期的另一个名字不算 + assert_eq!( + p.model_limits(&b, "中转自有模型-2026"), + ModelLimits::default() + ); + // 配置里没有这一家:只看价目表 + let cfg = Config::default(); + assert_eq!( + cfg.model_limits(&b, "relay", "中转自有模型"), + ModelLimits::default() + ); + } +} diff --git a/crates/tw-config/src/validate.rs b/crates/tw-config/src/validate.rs index c74b9581..7b9a52b1 100644 --- a/crates/tw-config/src/validate.rs +++ b/crates/tw-config/src/validate.rs @@ -101,6 +101,19 @@ pub enum ValidationError { AliasChained { alias: String, model: String }, #[error("{}", self.msg())] AliasOnlyItself { alias: String }, + #[error("{}", self.msg())] + ModelSpecBlankModel { upstream: String }, + #[error("{}", self.msg())] + ModelSpecWildcard { upstream: String, model: String }, + #[error("{}", self.msg())] + ModelSpecEmpty { upstream: String, model: String }, + /// `field` 是字段名本身(`context_window`、`max_output_tokens`),不翻 + #[error("{}", self.msg())] + ModelSpecZero { + upstream: String, + model: String, + field: &'static str, + }, } impl ValidationError { @@ -295,6 +308,29 @@ impl ValidationError { "alias `{alias}` lists only itself, so it changes nothing. List the names the \ upstreams use, or remove the alias" ), + ModelSpecBlankModel { upstream } => msg!( + "config.model_spec_blank_model", upstream = upstream => + "upstream `{upstream}` has a model spec (model_specs) for an empty model id" + ), + ModelSpecWildcard { upstream, model } => msg!( + "config.model_spec_wildcard", upstream = upstream, model = model => + "the model spec `{model}` of upstream `{upstream}` contains * or ?. A model spec \ + is for one exact model id" + ), + ModelSpecEmpty { upstream, model } => msg!( + "config.model_spec_empty", upstream = upstream, model = model => + "the model spec `{model}` of upstream `{upstream}` sets neither context_window \ + nor max_output_tokens. Set at least one, or remove it" + ), + ModelSpecZero { + upstream, + model, + field, + } => msg!( + "config.model_spec_zero", upstream = upstream, model = model, field = field => + "the model spec `{model}` of upstream `{upstream}` has {field}: 0; it has to be a \ + number of tokens above 0. Leave it out to use the price table" + ), } } } @@ -352,6 +388,9 @@ pub fn validate(cfg: &Config) -> Result<(), ValidationError> { }); } } + for (model, spec) in &p.model_specs { + check_model_spec(&p.name, model, Some(spec))?; + } } let mut names = std::collections::HashSet::new(); @@ -566,6 +605,49 @@ pub fn check_aliases(aliases: &crate::Aliases) -> Result<(), ValidationError> { Ok(()) } +/// 一家上游的一项手写模型规格写得对不对,和整份配置的校验是同一套(控制面保存一项 +/// 之前也用它)。`spec` 是 `None` 时只查模型 ID:界面要删掉这一项。**模型在不在这家的 +/// 清单里不查**:清单是运行时问来的。 +pub fn check_model_spec( + upstream: &str, + model: &str, + spec: Option<&crate::ModelSpec>, +) -> Result<(), ValidationError> { + if model.trim().is_empty() { + return Err(ValidationError::ModelSpecBlankModel { + upstream: upstream.to_string(), + }); + } + let named = || (upstream.to_string(), model.to_string()); + // 没有通配:`glm-*` 写在这里,读的人会以为一批模型都按它算 + if model.contains(['*', '?']) { + let (upstream, model) = named(); + return Err(ValidationError::ModelSpecWildcard { upstream, model }); + } + let Some(spec) = spec else { + return Ok(()); + }; + if spec.is_empty() { + let (upstream, model) = named(); + return Err(ValidationError::ModelSpecEmpty { upstream, model }); + } + // 0 不是「不知道」:上下文窗口是 0 的模型什么都装不下,输出上限是 0 的什么都答不出 + for (field, v) in [ + ("context_window", spec.context_window), + ("max_output_tokens", spec.max_output_tokens), + ] { + if v == Some(0) { + let (upstream, model) = named(); + return Err(ValidationError::ModelSpecZero { + upstream, + model, + field, + }); + } + } + Ok(()) +} + /// 最小的 CIDR 形状校验。**真正的匹配逻辑在 tw-gateway::access** —— /// 这里只是不想让 tw-config 依赖数据面,而「这条写法对不对」是配置层 /// 该回答的问题。 @@ -656,6 +738,56 @@ mod tests { assert!(validate(&cfg(vec![c("d", "tw-1")], vec![p("r", "https://x.com")])).is_ok()); } + #[test] + fn model_specs_name_one_exact_model_and_set_a_positive_number() { + let parse = |specs: &str| { + crate::try_parse(&format!( + "version: 1\nlisten:\n control:\n key: {}\nclients:\n - name: c\n key: tw-k\nproviders:\n - name: relay\n base_url: https://relay.example.com/v1\n model_specs:\n{specs}", + "c0".repeat(32) + )) + }; + let ok = parse( + " glm-5-air: { context_window: 128000, max_output_tokens: 16384 }\n \"us.anthropic.claude-fable-5-v1:0\": { max_output_tokens: 32000 }\n", + ) + .unwrap(); + let specs = &ok.providers[0].model_specs; + assert_eq!(specs["glm-5-air"].context_window, Some(128_000)); + assert_eq!( + specs["us.anthropic.claude-fable-5-v1:0"].max_output_tokens, + Some(32_000) + ); + for (specs, code) in [ + ( + " \"\": { context_window: 1000 }\n", + "config.model_spec_blank_model", + ), + ( + " glm-*: { context_window: 1000 }\n", + "config.model_spec_wildcard", + ), + (" glm-5-air: {}\n", "config.model_spec_empty"), + ( + " glm-5-air: { context_window: 0 }\n", + "config.model_spec_zero", + ), + ( + " glm-5-air: { context_window: 1000, max_output_tokens: 0 }\n", + "config.model_spec_zero", + ), + ] { + let m = parse(specs).unwrap_err().msg(); + assert_eq!(m.code, code, "{specs}: {m:?}"); + } + let m = parse(" glm-5-air: { max_output_tokens: 0 }\n") + .unwrap_err() + .msg(); + assert_eq!(m.arg("field"), "max_output_tokens"); + assert_eq!(m.arg("upstream"), "relay"); + // 字段名写错、写成负数,serde 自己说 + assert!(parse(" glm-5-air: { context: 1000 }\n").is_err()); + assert!(parse(" glm-5-air: { context_window: -1 }\n").is_err()); + } + #[test] fn a_newer_schema_is_refused_before_anything_else_is_judged() { // 顺序很重要:来自新版本的配置,我们对它的其他判断都不作数。 @@ -1373,6 +1505,22 @@ mod msg_codes { model: "b".into(), }, AliasOnlyItself { alias: "a".into() }, + ModelSpecBlankModel { + upstream: "a".into(), + }, + ModelSpecWildcard { + upstream: "a".into(), + model: "m*".into(), + }, + ModelSpecEmpty { + upstream: "a".into(), + model: "m".into(), + }, + ModelSpecZero { + upstream: "a".into(), + model: "m".into(), + field: "context_window", + }, ]; check( "config.", diff --git a/crates/tw-config/src/wire.rs b/crates/tw-config/src/wire.rs index 4cac48ac..1ee7917b 100644 --- a/crates/tw-config/src/wire.rs +++ b/crates/tw-config/src/wire.rs @@ -2,7 +2,7 @@ //! 对齐。多一个变体,这里编译不过。安全防护的档位、匹配方式两边是同一个类型 //! (`tw_guard::policy`),不用对齐。 -use crate::{Billing, OnProxyFail, ProbeAction, Protocol, ProxyKind, Stage}; +use crate::{Billing, OnProxyFail, ProbeAction, Protocol, ProxyKind, SpecSource, Stage}; impl From for tw_api::Billing { fn from(b: Billing) -> Self { @@ -22,6 +22,15 @@ impl From for Billing { } } +impl From for tw_api::SpecSource { + fn from(s: SpecSource) -> Self { + match s { + SpecSource::PriceTable => Self::PriceTable, + SpecSource::Manual => Self::Manual, + } + } +} + impl From for tw_api::Protocol { fn from(p: Protocol) -> Self { match p { diff --git a/crates/tw-config/tests/manual/schema.rs b/crates/tw-config/tests/manual/schema.rs index c220778d..653be6ad 100644 --- a/crates/tw-config/tests/manual/schema.rs +++ b/crates/tw-config/tests/manual/schema.rs @@ -561,6 +561,15 @@ pub fn sections() -> Vec
{ "`pricing.sheets` 中某张价目表的名字。不写:默认价目表。", ), ), + row( + "model_specs", + Kind::ObjMap(t("model id", "模型 ID"), "providers[].model_specs.*"), + Def::Is("{}"), + t( + "Context window and output limit of single models of this upstream, written by hand, by exact model id. They take precedence over the price table: for models it does not know, or gets wrong.", + "手写这家上游某些模型的上下文窗口和输出上限,按模型 ID 完全匹配。写了就优先于价目表,用于价目表里没有或写错的模型。", + ), + ), row( "disabled", Kind::Bool, @@ -572,6 +581,31 @@ pub fn sections() -> Vec
{ ), ], }, + Section { + path: "providers[].model_specs.*", + // 两项都可选,至少写一项由校验管 + ty: checked!(ModelSpec, "{}"), + rows: vec![ + row( + "context_window", + Kind::Int, + Def::Unset, + t( + "Context window: the most tokens a request can take in. Unset: the price table's.", + "上下文窗口,即一次请求最多输入多少 token。不写:取价目表的。", + ), + ), + row( + "max_output_tokens", + Kind::Int, + Def::Unset, + t( + "The most tokens an answer can have. Unset: the price table's.", + "一次回答最多输出多少 token。不写:取价目表的。", + ), + ), + ], + }, Section { path: "providers[].aws", // 字段都可选:写访问密钥还是写 profile,由写法检查管二选一 diff --git a/crates/tw-control/src/aliases.rs b/crates/tw-control/src/aliases.rs index d8429337..b3a834ca 100644 --- a/crates/tw-control/src/aliases.rs +++ b/crates/tw-control/src/aliases.rs @@ -68,9 +68,10 @@ async fn list(State(s): State) -> Json { }) .collect(), shadows: shadows(&lists, a), + // 和 `/v1/models` 列别名时同一个查法:那一家手写的优先 context_window: served_by.first().and_then(|first| { - book.resolve_for(&first.provider, &first.model) - .and_then(|r| r.price.max_input_tokens) + cfg.model_limits(&book, &first.provider, &first.model) + .context_window() }), served_by, requests_24h, diff --git a/crates/tw-control/src/resources.rs b/crates/tw-control/src/resources.rs index aaeb11ff..96b1faf8 100644 --- a/crates/tw-control/src/resources.rs +++ b/crates/tw-control/src/resources.rs @@ -35,6 +35,7 @@ pub fn router() -> axum::Router { .at(ep::UpdateProvider, update_provider) .at(ep::DeleteProvider, delete_provider) .at(ep::ProviderModels, provider_models) + .at(ep::SetModelSpec, set_model_spec) .at(ep::RefreshProviderModels, refresh_models) .at(ep::RefreshStaleModels, refresh_stale_models) .at(ep::CreateProxy, create_proxy) @@ -142,6 +143,42 @@ fn chatgpt_login(cfg: &tw_config::Config, name: &str, token_endpoint: &str) -> O (ours && !shared).then(|| o.refresh.clone()) } +/// 手写一家上游的一个模型的上下文窗口、输出上限;两项都空就删掉那一项,回到价目表。 +/// +/// **只动这一家 `model_specs` 里的这一项**([`edit::set_model_spec`]),不走编辑上游那条 +/// 路:那条路按读进来的结构把整项写回去,用户写成和默认值一样的字段会被顺手删掉。 +/// 模型 ID 和数值的毛病由整份配置的校验说([`tw_config::check_model_spec`]),先查一遍, +/// 好让它说的是这一项,而不是写进去之后被拒。 +async fn set_model_spec( + State(s): State, + Json(req): Json, +) -> Result, Fail> { + let model = req.model.trim(); + let spec = tw_config::ModelSpec { + context_window: req.context_window, + max_output_tokens: req.max_output_tokens, + }; + let spec = (!spec.is_empty()).then_some(spec); + let version = s + .cfg + .transform(req.base_version.as_deref(), Origin::Ui, |text, cfg| { + if !cfg.providers.iter().any(|p| p.name == req.provider) { + return Err(not_found("upstream", &req.provider)); + } + tw_config::check_model_spec(&req.provider, model, spec.as_ref()) + .map_err(|e| invalid(e.msg()))?; + Ok(edit::set_model_spec( + text, + &req.provider, + model, + spec.as_ref(), + )?) + }) + .await + .map_err(apply_fail)?; + Ok(Json(tw_api::ConfigWritten { version })) +} + /// 按接口地址自动识别会得到什么:协议、是不是官方端点、默认脱敏哪几类。 /// /// **不联网,只看地址。**编辑对话框在用户输入地址时调它,好让「自动识别」 @@ -246,8 +283,8 @@ async fn test_provider( })) } -/// 一个上游的模型清单:每个模型在不在启用范围里、上下文窗口多大、按它选的 -/// 价目表怎么计价。 +/// 一个上游的模型清单:每个模型在不在启用范围里、上下文窗口和输出上限多大(手写的 +/// 还是价目表的)、按它选的价目表怎么计价。 async fn provider_models( State(s): State, Path(name): Path, @@ -306,9 +343,14 @@ fn models_view( .into_iter() .map(|id| { let r = book.resolve_for(&p.name, &id); + // 上下文窗口、输出上限和 `/v1/models` 给客户端的是同一个查法 + let limits = p.model_limits(&book, &id); tw_api::ModelRow { enabled: p.uses_model(&id), - context_window: r.as_ref().and_then(|r| r.price.max_input_tokens), + context_window: limits.context_window(), + context_window_source: limits.context_window.map(|s| s.source.into()), + max_output_tokens: limits.max_output_tokens(), + max_output_tokens_source: limits.max_output_tokens.map(|s| s.source.into()), price: r.as_ref().map(|r| { crate::pricing::price_fields(&tw_pricing::PerMillion::of(&r.price)) }), @@ -426,6 +468,9 @@ fn to_provider( .as_ref() .map(|ms| ms.iter().map(|m| m.trim().to_string()).collect::>()), pricing: input.pricing.clone(), + // 手写的模型规格不在这个表单里(上游页的模型清单一行一行改,见 `set_model_spec`): + // 沿用原来的,改名时跟着这一项走 + model_specs: existing.map(|e| e.model_specs.clone()).unwrap_or_default(), disabled: input.disabled, }; // **保存和检测之前就说清楚凭据写法哪儿不对**,而不是等整份配置校验时 diff --git a/crates/tw-control/tests/model_specs.rs b/crates/tw-control/tests/model_specs.rs new file mode 100644 index 00000000..f55f9b6e --- /dev/null +++ b/crates/tw-control/tests/model_specs.rs @@ -0,0 +1,333 @@ +//! 手写的模型规格:上游页的模型清单说出每个数是手写的还是价目表的,`PUT +//! /provider-model-spec` 设一项、删一项,别名跟着服务它的模型走,改名和删上游时 +//! 规格跟着那一家走。 + +use std::sync::Arc; + +use axum::body::Body; +use axum::http::{Request, StatusCode}; +use serde_json::{Value, json}; +use tower::ServiceExt; +use tw_control::{ConfigManager, ControlState}; + +struct Bed { + dir: tempfile::TempDir, + app: axum::Router, +} + +impl Bed { + fn file(&self) -> String { + std::fs::read_to_string(self.dir.path().join("config.yaml")).unwrap() + } +} + +fn bed(yaml: &str) -> Bed { + let d = tempfile::tempdir().unwrap(); + let p = d.path().join("config.yaml"); + std::fs::write(&p, yaml).unwrap(); + let cfg = tw_config::try_parse(yaml).unwrap(); + let gw = tw_gateway::AppState::new(cfg).unwrap(); + let bus = gw.bus.clone(); + let state = ControlState { + shutdown: Default::default(), + remote: Default::default(), + cfg: Arc::new(ConfigManager::new(p, gw.clone(), bus)), + gateway: gw, + store: None, + started: std::time::Instant::now(), + price_updater: Default::default(), + chatgpt: Default::default(), + zai: Default::default(), + }; + Bed { + app: tw_control::router(state), + dir: d, + } +} + +async fn call(app: &axum::Router, method: &str, path: &str, body: Value) -> (StatusCode, Value) { + let r = app + .clone() + .oneshot( + Request::builder() + .method(method) + .uri(path) + .header("content-type", "application/json") + .body(Body::from(body.to_string())) + .unwrap(), + ) + .await + .unwrap(); + let st = r.status(); + let b = axum::body::to_bytes(r.into_body(), 1 << 20).await.unwrap(); + let text = String::from_utf8_lossy(&b).to_string(); + let v = serde_json::from_str(&text).unwrap_or(Value::String(text)); + (st, v) +} + +/// 一个假上游:列出三个模型,其中一个价目表不认识。 +async fn upstream() -> std::net::SocketAddr { + let app = axum::Router::new().route( + "/v1/models", + axum::routing::get(|| async { + r#"{"data":[{"id":"claude-sonnet-4-5"},{"id":"claude-haiku-4-5"},{"id":"中转自有模型"}]}"# + }), + ); + let l = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap(); + let a = l.local_addr().unwrap(); + tokio::spawn(async move { axum::serve(l, app).await.unwrap() }); + a +} + +/// `relay` 那一项后面接 `extra`(缩进四格的字段),整份配置后面接 `tail` +fn config(up: std::net::SocketAddr, extra: &str, tail: &str) -> String { + format!( + "version: 1 +listen: + control: + key: c0ffee00c0ffee00c0ffee00c0ffee00c0ffee00c0ffee00c0ffee00c0ffee00 +clients: + - name: c + key: tw-k +providers: + - name: relay + base_url: http://{up} + key: sk-good + protocol: anthropic +{extra}{tail}" + ) +} + +/// 问一遍清单,交回 `id` → 那一行 +async fn rows(b: &Bed) -> Vec { + let (st, v) = call( + &b.app, + "POST", + "/providers/relay/models/refresh", + Value::Null, + ) + .await; + assert_eq!(st, StatusCode::OK, "{v}"); + v["models"].as_array().unwrap().clone() +} + +fn row<'a>(rows: &'a [Value], id: &str) -> &'a Value { + rows.iter().find(|m| m["id"] == id).unwrap() +} + +#[tokio::test] +async fn each_row_says_whether_its_numbers_were_written_by_hand() { + let up = upstream().await; + let b = bed(&config( + up, + " model_specs: + claude-sonnet-4-5: { max_output_tokens: 8000 } + 中转自有模型: { context_window: 32000 } +", + "", + )); + let rows = rows(&b).await; + + let sonnet = row(&rows, "claude-sonnet-4-5"); + assert_eq!(sonnet["context_window"], 200_000, "{sonnet}"); + assert_eq!(sonnet["context_window_source"], "price_table"); + assert_eq!(sonnet["max_output_tokens"], 8_000); + assert_eq!(sonnet["max_output_tokens_source"], "manual"); + + let own = row(&rows, "中转自有模型"); + assert_eq!(own["context_window"], 32_000, "{own}"); + assert_eq!(own["context_window_source"], "manual"); + // 不知道就整个不出现,来源也没有 + assert!(own.get("max_output_tokens").is_none(), "{own}"); + assert!(own.get("max_output_tokens_source").is_none(), "{own}"); + + let haiku = row(&rows, "claude-haiku-4-5"); + assert_eq!(haiku["context_window_source"], "price_table", "{haiku}"); + assert_eq!(haiku["max_output_tokens_source"], "price_table", "{haiku}"); +} + +#[tokio::test] +async fn a_spec_is_set_changed_and_cleared_through_the_endpoint() { + let up = upstream().await; + let b = bed(&config(up, "", "")); + let before = b.file(); + let set = |cw: Value, out: Value| { + json!({ + "provider": "relay", + "model": " 中转自有模型 ", + "context_window": cw, + "max_output_tokens": out, + }) + }; + + let (st, v) = call( + &b.app, + "PUT", + "/provider-model-spec", + set(json!(64000), Value::Null), + ) + .await; + assert_eq!(st, StatusCode::OK, "{v}"); + assert!(v["version"].as_str().is_some(), "{v}"); + let cfg = tw_config::try_parse(&b.file()).unwrap(); + assert_eq!( + cfg.providers[0].model_specs["中转自有模型"], + tw_config::ModelSpec { + context_window: Some(64_000), + max_output_tokens: None, + }, + "模型 ID 去掉首尾空白:{}", + b.file() + ); + let own = row(&rows(&b).await, "中转自有模型").clone(); + assert_eq!(own["context_window"], 64_000); + assert_eq!(own["context_window_source"], "manual"); + + // 改:输出上限也写上 + let (st, v) = call( + &b.app, + "PUT", + "/provider-model-spec", + set(json!(64000), json!(4096)), + ) + .await; + assert_eq!(st, StatusCode::OK, "{v}"); + let own = row(&rows(&b).await, "中转自有模型").clone(); + assert_eq!(own["max_output_tokens"], 4_096); + assert_eq!(own["max_output_tokens_source"], "manual"); + + // 两项都空:删掉,回到价目表(它不认识这个模型),文件回到原样 + let (st, v) = call( + &b.app, + "PUT", + "/provider-model-spec", + set(Value::Null, Value::Null), + ) + .await; + assert_eq!(st, StatusCode::OK, "{v}"); + assert_eq!(b.file(), before); + let own = row(&rows(&b).await, "中转自有模型").clone(); + assert!(own.get("context_window").is_none(), "{own}"); + assert!(own.get("context_window_source").is_none(), "{own}"); +} + +#[tokio::test] +async fn a_bad_spec_is_refused_in_the_configurations_words() { + let up = upstream().await; + let b = bed(&config(up, "", "")); + let before = b.file(); + for (body, status, code) in [ + ( + json!({ "provider": "relay", "model": "m", "context_window": 0 }), + StatusCode::BAD_REQUEST, + "config.model_spec_zero", + ), + ( + json!({ "provider": "relay", "model": "m", "max_output_tokens": 0 }), + StatusCode::BAD_REQUEST, + "config.model_spec_zero", + ), + ( + json!({ "provider": "relay", "model": "glm-*", "context_window": 1000 }), + StatusCode::BAD_REQUEST, + "config.model_spec_wildcard", + ), + ( + json!({ "provider": "relay", "model": " ", "context_window": 1000 }), + StatusCode::BAD_REQUEST, + "config.model_spec_blank_model", + ), + // 删一项也要说清楚是哪个模型 + ( + json!({ "provider": "relay", "model": "" }), + StatusCode::BAD_REQUEST, + "config.model_spec_blank_model", + ), + ( + json!({ "provider": "ghost", "model": "m", "context_window": 1000 }), + StatusCode::NOT_FOUND, + "config.edit.not_found", + ), + ] { + let (st, v) = call(&b.app, "PUT", "/provider-model-spec", body.clone()).await; + assert_eq!(st, status, "{body}: {v}"); + assert_eq!(v["code"], code, "{body}: {v}"); + } + // 负数、写错字段名:请求体本身不对 + let (st, _) = call( + &b.app, + "PUT", + "/provider-model-spec", + json!({ "provider": "relay", "model": "m", "context_window": -1 }), + ) + .await; + assert!(st.is_client_error(), "{st}"); + assert_eq!(b.file(), before); +} + +#[tokio::test] +async fn an_alias_takes_the_spec_of_the_model_that_serves_it() { + let up = upstream().await; + let b = bed(&config( + up, + " model_specs: + 中转自有模型: { context_window: 32000 } +", + "aliases: + own: 中转自有模型 +", + )); + rows(&b).await; + let (st, v) = call(&b.app, "GET", "/aliases", Value::Null).await; + assert_eq!(st, StatusCode::OK, "{v}"); + let own = &v["aliases"][0]; + assert_eq!(own["name"], "own"); + assert_eq!(own["context_window"], 32_000, "{own}"); +} + +/// 规格在那一项里面:编辑对话框保存(不带规格)不丢,改名跟着走,删掉一起没 +#[tokio::test] +async fn specs_stay_with_their_upstream_through_edits_renames_and_deletes() { + let up = upstream().await; + let b = bed(&config( + up, + " model_specs: + 中转自有模型: { context_window: 32000 } +", + "", + )); + let (st, v) = call( + &b.app, + "PUT", + "/providers/relay", + json!({ "provider": { + "name": "中转", + "base_url": format!("http://{up}"), + "key": "sk-good", + "protocol": "anthropic", + }}), + ) + .await; + assert_eq!(st, StatusCode::OK, "{v}"); + let cfg = tw_config::try_parse(&b.file()).unwrap(); + assert_eq!(cfg.providers[0].name, "中转"); + assert_eq!( + cfg.providers[0].model_specs["中转自有模型"].context_window, + Some(32_000), + "{}", + b.file() + ); + // 规格按名字跟着这一家:旧名字上设不了 + let (st, _) = call( + &b.app, + "PUT", + "/provider-model-spec", + json!({ "provider": "relay", "model": "m", "context_window": 1000 }), + ) + .await; + assert_eq!(st, StatusCode::NOT_FOUND); + + let (st, v) = call(&b.app, "DELETE", "/providers/中转", json!({})).await; + assert_eq!(st, StatusCode::OK, "{v}"); + assert!(!b.file().contains("model_specs"), "{}", b.file()); +} diff --git a/crates/tw-gateway/src/server/listing.rs b/crates/tw-gateway/src/server/listing.rs index c2e9f6d5..c2f146d9 100644 --- a/crates/tw-gateway/src/server/listing.rs +++ b/crates/tw-gateway/src/server/listing.rs @@ -63,7 +63,8 @@ impl ListingShape { /// 列表里一个模型带给客户端的元数据。 /// -/// 来自价目表,和上游页模型一格的「上下文」(`ModelRow.context_window`)是同一个数。 +/// 这一家手写的(`model_specs`)优先,没写的来自价目表(见 [`tw_config::model_specs`]), +/// 和上游页模型一格的「上下文」(`ModelRow.context_window`)是同一个数。 /// **查不到就是 `None`,对应的字段整个不出现** —— 客户端读不到会用自己的默认值, /// 一个编出来的数它却会照着截断对话。 #[derive(Debug, Clone, Copy, Default, PartialEq, Eq)] @@ -76,31 +77,31 @@ pub(crate) struct ModelMeta { /// 查一个模型的元数据。 /// -/// `provider`:知道这个名称发给哪家上游时给上,按它选的价目表查,和上游页那一格是 -/// 同一个查法;不给就查默认价目表。自定义价目表只改单价,上下文窗口照样取自默认 -/// 价目表,所以两种查法给出的通常是同一个数。 +/// `provider`:知道这个名称发给哪家上游时给上,先看这一家手写的,再按它选的价目表查, +/// 和上游页那一格是同一个查法([`tw_config::Config::model_limits`]);不给就只查默认 +/// 价目表。自定义价目表只改单价,上下文窗口照样取自默认价目表。 pub(crate) fn model_meta( + cfg: &tw_config::Config, book: &tw_pricing::PriceBook, provider: Option<&str>, model: &str, ) -> ModelMeta { - let resolved = match provider { - Some(p) => book.resolve_for(p, model), - None => book.resolve(None, model), + let limits = match provider { + Some(p) => cfg.model_limits(book, p, model), + None => tw_config::ModelLimits::priced(book, model), }; - resolved - .map(|r| ModelMeta { - max_input_tokens: r.price.max_input_tokens, - max_output_tokens: r.price.max_output_tokens, - }) - .unwrap_or_default() + ModelMeta { + max_input_tokens: limits.context_window(), + max_output_tokens: limits.max_output_tokens(), + } } /// 列表里一个名称的元数据。 /// /// 别名用它第一个有上游提供的模型的(按列表顺序),按提供它的头一家查(见 -/// [`tw_engine::Catalog::first_served`])。目录空着、列表里谁都不提供时(这时单点 -/// 查询不拦),按它列表里的头一个查默认价目表。 +/// [`tw_engine::Catalog::first_served`]);真模型也按提供它的头一家(配置里的顺序)查 +/// —— 那一家手写的规格才对得上。目录空着、谁都不提供时(这时单点查询不拦),别名按 +/// 它列表里的头一个、真模型按它自己查默认价目表。 fn listed_meta( book: &tw_pricing::PriceBook, catalog: &tw_engine::Catalog, @@ -108,11 +109,16 @@ fn listed_meta( name: &str, ) -> ModelMeta { if let Some((provider, model)) = catalog.first_served(name) { - return model_meta(book, Some(provider), model); + return model_meta(cfg, book, Some(provider), model); + } + if !cfg.aliases.contains(name) + && let Some(provider) = catalog.providers_for(name).first() + { + return model_meta(cfg, book, Some(provider), name); } match cfg.aliases.find(name).and_then(|a| a.models.first()) { - Some(model) => model_meta(book, None, model), - None => model_meta(book, None, name), + Some(model) => model_meta(cfg, book, None, model), + None => model_meta(cfg, book, None, name), } } @@ -464,19 +470,53 @@ mod tests { #[test] fn metadata_comes_from_the_price_table() { let book = tw_pricing::PriceBook::builtin().unwrap(); - let sonnet = model_meta(&book, None, "claude-sonnet-4-5-20250929"); + let cfg = tw_config::Config::default(); + let sonnet = model_meta(&cfg, &book, None, "claude-sonnet-4-5-20250929"); assert_eq!(sonnet.max_input_tokens, Some(200_000)); assert_eq!(sonnet.max_output_tokens, Some(64_000)); // 给了上游就按它选的价目表查;没选价目表的上游和默认价目表一样 - assert_eq!(model_meta(&book, Some("up"), "claude-sonnet-4-5"), sonnet); + assert_eq!( + model_meta(&cfg, &book, Some("up"), "claude-sonnet-4-5"), + sonnet + ); // Bedrock 的名字也查得到 assert_eq!( - model_meta(&book, Some("bedrock"), "us.anthropic.claude-fable-5").max_input_tokens, + model_meta(&cfg, &book, Some("bedrock"), "us.anthropic.claude-fable-5") + .max_input_tokens, Some(1_000_000) ); assert_eq!( - model_meta(&book, None, "no-such-model-anywhere"), + model_meta(&cfg, &book, None, "no-such-model-anywhere"), ModelMeta::default() ); } + + #[test] + fn a_spec_written_for_the_upstream_wins_over_the_price_table() { + let book = tw_pricing::PriceBook::builtin().unwrap(); + let cfg = tw_config::Config { + providers: vec![tw_config::Provider { + name: "up".into(), + base_url: "https://relay.example.com".into(), + model_specs: [( + "claude-sonnet-4-5".to_string(), + tw_config::ModelSpec { + context_window: Some(1_000_000), + max_output_tokens: None, + }, + )] + .into(), + ..Default::default() + }], + ..Default::default() + }; + let m = model_meta(&cfg, &book, Some("up"), "claude-sonnet-4-5"); + assert_eq!(m.max_input_tokens, Some(1_000_000)); + assert_eq!(m.max_output_tokens, Some(64_000), "没写的照样取价目表"); + // 不知道是哪一家时没有手写的可看 + assert_eq!( + model_meta(&cfg, &book, None, "claude-sonnet-4-5").max_input_tokens, + Some(200_000) + ); + } } diff --git a/crates/tw-gateway/src/server/pipeline.rs b/crates/tw-gateway/src/server/pipeline.rs index ca396e64..8ef2b076 100644 --- a/crates/tw-gateway/src/server/pipeline.rs +++ b/crates/tw-gateway/src/server/pipeline.rs @@ -356,8 +356,8 @@ fn conversation( /// 输入超出了这个决定所选模型的上下文窗口:这一轮沿用的决定要重新求值。 /// /// 按候选里最小的那个窗口算,留 5% 的余量(输入是估的)。**知道窗口的才算**:价目表 -/// 里没写的模型,说不出它装不装得下,照常沿用。窗口按每一家发出去的名字查:别名在各家 -/// 是各家的名字(见 [`crate::sent`])。 +/// 里没写、那一家也没手写(`model_specs`)的模型,说不出它装不装得下,照常沿用。窗口按 +/// 每一家发出去的名字查:别名在各家是各家的名字(见 [`crate::sent`])。 fn outgrown( state: &AppState, rt: &Runtime, @@ -380,9 +380,9 @@ fn outgrown( ) .iter() .filter_map(|s| { - book.resolve_for(&s.provider, s.model.as_deref().ok()?)? - .price - .max_input_tokens + rt.config + .model_limits(&book, &s.provider, s.model.as_deref().ok()?) + .context_window() }) .min() .is_some_and(|limit| facts.input_tokens.saturating_mul(100) >= limit.saturating_mul(95)) diff --git a/crates/tw-gateway/src/server/pipeline/hop.rs b/crates/tw-gateway/src/server/pipeline/hop.rs index b80662ad..a213e9ef 100644 --- a/crates/tw-gateway/src/server/pipeline/hop.rs +++ b/crates/tw-gateway/src/server/pipeline/hop.rs @@ -1157,6 +1157,7 @@ fn prepare( official: tw_dialect::official::is_official_host(&provider.base_url), default_max_tokens: crate::translate::default_max_tokens( &state.pricing.load(), + provider, &d.request.model, ), }); diff --git a/crates/tw-gateway/src/translate.rs b/crates/tw-gateway/src/translate.rs index c714ac5f..f3c48a51 100644 --- a/crates/tw-gateway/src/translate.rs +++ b/crates/tw-gateway/src/translate.rs @@ -65,12 +65,17 @@ pub fn apply_set(r: &mut Request, set: &tw_engine::SetAction) { /// 客户端没写最大输出、目标格式又必须写(Anthropic)时用多少。 /// -/// **价目表里有这个模型的输出上限就用它**:写大了上游拒绝,写小了回答被截断。 -/// 查不到时按模型名兜底([`tw_dialect::official::fallback_max_output_tokens`])。 -pub fn default_max_tokens(book: &tw_pricing::PriceBook, model: &str) -> u64 { - book.table() - .get(model) - .and_then(|p| p.max_output_tokens) +/// **知道这一家的这个模型的输出上限就用它**(手写的优先,再看价目表,见 +/// [`tw_config::model_specs`]):写大了上游拒绝,写小了回答被截断。都查不到时按模型名 +/// 兜底([`tw_dialect::official::fallback_max_output_tokens`])。 +pub fn default_max_tokens( + book: &tw_pricing::PriceBook, + provider: &tw_config::Provider, + model: &str, +) -> u64 { + provider + .model_limits(book, model) + .max_output_tokens() .unwrap_or_else(|| tw_dialect::official::fallback_max_output_tokens(model)) } @@ -125,11 +130,37 @@ mod tests { Default::default(), [], ); - assert_eq!(default_max_tokens(&book, "claude-sonnet-4-5"), 64000); - assert_eq!(default_max_tokens(&book, "deepseek-chat"), 8000); + let up = tw_config::Provider { + name: "up".into(), + ..Default::default() + }; + assert_eq!(default_max_tokens(&book, &up, "claude-sonnet-4-5"), 64000); + assert_eq!(default_max_tokens(&book, &up, "deepseek-chat"), 8000); // 表里没有的按名字兜底 - assert_eq!(default_max_tokens(&book, "claude-mythos-5-1"), 32000); - assert_eq!(default_max_tokens(&book, "some-relay-model"), 8192); + assert_eq!(default_max_tokens(&book, &up, "claude-mythos-5-1"), 32000); + assert_eq!(default_max_tokens(&book, &up, "some-relay-model"), 8192); + // 这一家手写的优先 + let relay = tw_config::Provider { + model_specs: [("some-relay-model", 4096), ("claude-sonnet-4-5", 32000)] + .into_iter() + .map(|(m, n)| { + ( + m.to_string(), + tw_config::ModelSpec { + context_window: None, + max_output_tokens: Some(n), + }, + ) + }) + .collect(), + ..up + }; + assert_eq!(default_max_tokens(&book, &relay, "some-relay-model"), 4096); + assert_eq!( + default_max_tokens(&book, &relay, "claude-sonnet-4-5"), + 32000 + ); + assert_eq!(default_max_tokens(&book, &relay, "deepseek-chat"), 8000); } #[test] diff --git a/crates/tw-gateway/tests/model_specs.rs b/crates/tw-gateway/tests/model_specs.rs new file mode 100644 index 00000000..bac244df --- /dev/null +++ b/crates/tw-gateway/tests/model_specs.rs @@ -0,0 +1,244 @@ +//! 手写的模型规格(`providers[].model_specs`)在网关里处处优先于价目表:三种格式的 +//! `/v1/models`、单点查询、别名,以及一轮半路输入超出上下文时重新求值。 + +use std::net::SocketAddr; +use std::time::Duration; + +use axum::Router; +use axum::routing::{any, get}; +use serde_json::Value; +use tokio::sync::broadcast::Receiver; +use tw_api::Event; +use tw_config::Config; + +/// 一个假上游:列出 `models`,什么请求都答一个读了缓存的回答。 +async fn upstream(models: &'static [&'static str]) -> SocketAddr { + let app = Router::new() + .route( + "/v1/models", + get(move || async move { + axum::Json(serde_json::json!({ + "data": models.iter().map(|m| serde_json::json!({ "id": m })).collect::>() + })) + }), + ) + .fallback(any(|| async { + axum::response::Response::builder() + .header("content-type", "application/json") + .body(axum::body::Body::from( + r#"{"type":"message","content":[],"usage":{"input_tokens":3,"cache_read_input_tokens":5000,"output_tokens":1}}"#, + )) + .unwrap() + })); + let l = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap(); + let a = l.local_addr().unwrap(); + tokio::spawn(async move { axum::serve(l, app).await.unwrap() }); + a +} + +async fn serve(cfg: Config) -> (SocketAddr, Receiver) { + let state = tw_gateway::AppState::new(cfg).unwrap(); + tw_gateway::models::refresh_all(&state).await; + let events = state.bus.subscribe(); + let addr = tw_gateway::serve(state, ([127, 0, 0, 1], 0).into()) + .await + .unwrap(); + tokio::time::sleep(Duration::from_millis(50)).await; + (addr, events) +} + +/// 一家中转站:`claude-sonnet-4-5` 价目表里是 200k,这里手写成 1M、输出上限照旧; +/// `glm-5-air` 价目表里没有,两项都手写。别名 `air` 只列 `glm-5-air`。 +async fn relay() -> Config { + let up = upstream(&["claude-sonnet-4-5", "glm-5-air", "claude-haiku-4-5"]).await; + serde_yaml_ng::from_str(&format!( + "version: 1 +clients: + - name: me + key: tw-k +providers: + - name: relay + base_url: http://{up} + key: sk-x + protocol: anthropic + model_specs: + claude-sonnet-4-5: {{ context_window: 1000000 }} + glm-5-air: {{ context_window: 128000, max_output_tokens: 16384 }} +aliases: + air: glm-5-air +" + )) + .unwrap() +} + +/// `GET path`:`anthropic` 带 `x-api-key` 和版本头,`gemini` 用 Google 的位置,其余 Bearer。 +async fn get_json(gw: SocketAddr, path: &str, shape: &str) -> Value { + let r = reqwest::Client::new().get(format!("http://{gw}{path}")); + let r = match shape { + "anthropic" => r + .header("x-api-key", "tw-k") + .header("anthropic-version", "2023-06-01"), + "gemini" => r.header("x-goog-api-key", "tw-k"), + _ => r.header("authorization", "Bearer tw-k"), + }; + let r = r.send().await.unwrap(); + assert_eq!(r.status(), 200, "{path}"); + r.json().await.unwrap() +} + +fn find<'a>(list: &'a [Value], field: &str, id: &str) -> &'a Value { + list.iter() + .find(|m| m[field] == id) + .unwrap_or_else(|| panic!("{id} 不在列表里:{list:?}")) +} + +#[tokio::test] +async fn every_listing_shape_carries_the_spec_written_for_the_upstream() { + let (gw, _) = serve(relay().await).await; + + let openai = get_json(gw, "/v1/models", "openai").await; + let data = openai["data"].as_array().unwrap(); + let sonnet = find(data, "id", "claude-sonnet-4-5"); + for f in ["context_window", "context_length", "max_input_tokens"] { + assert_eq!(sonnet[f], 1_000_000, "{f}: {sonnet}"); + } + let air = find(data, "id", "glm-5-air"); + assert_eq!(air["context_window"], 128_000, "{air}"); + // 没手写的照旧取价目表 + assert_eq!( + find(data, "id", "claude-haiku-4-5")["context_window"], + 200_000 + ); + // 别名按服务它的那个模型:中转站手写的那个数 + assert_eq!(find(data, "id", "air")["context_window"], 128_000); + + let anthropic = get_json(gw, "/v1/models", "anthropic").await; + let data = anthropic["data"].as_array().unwrap(); + let sonnet = find(data, "id", "claude-sonnet-4-5"); + assert_eq!(sonnet["max_input_tokens"], 1_000_000, "{sonnet}"); + assert_eq!(sonnet["supports_1m"], true, "手写到一百万就是 1M:{sonnet}"); + let air = find(data, "id", "glm-5-air"); + assert_eq!(air["max_input_tokens"], 128_000); + assert_eq!(air["supports_1m"], false); + assert_eq!(find(data, "id", "air")["max_input_tokens"], 128_000); + let one = get_json(gw, "/v1/models/claude-sonnet-4-5", "anthropic").await; + assert_eq!(one, *sonnet, "单点查询和列表里的应该是同一个对象"); + + let gemini = get_json(gw, "/v1beta/models", "gemini").await; + let data = gemini["models"].as_array().unwrap(); + let sonnet = find(data, "name", "models/claude-sonnet-4-5"); + assert_eq!(sonnet["inputTokenLimit"], 1_000_000, "{sonnet}"); + // 输出上限没手写:价目表的 + assert_eq!(sonnet["outputTokenLimit"], 64_000, "{sonnet}"); + let air = find(data, "name", "models/glm-5-air"); + assert_eq!(air["inputTokenLimit"], 128_000, "{air}"); + assert_eq!(air["outputTokenLimit"], 16_384, "{air}"); + let alias = find(data, "name", "models/air"); + assert_eq!(alias["outputTokenLimit"], 16_384, "{alias}"); + let one = get_json(gw, "/v1beta/models/glm-5-air", "gemini").await; + assert_eq!(one, *air); +} + +// ─────────────────────────────────────────────────────────── 管线 + +/// 输入大的去乙,其余去甲。两家都在同一个假上游上,都列着价目表里没有的 `relay-model`; +/// `window` 给了就是甲手写的上下文窗口 +async fn two_upstreams(window: Option) -> Config { + let up = upstream(&["relay-model"]).await; + let specs = window + .map(|n| format!(" model_specs:\n relay-model: {{ context_window: {n} }}\n")) + .unwrap_or_default(); + serde_yaml_ng::from_str(&format!( + "version: 1 +clients: + - name: me + key: tw-k +providers: + - name: 甲 + base_url: http://{up} + key: sk-x + protocol: anthropic +{specs} - name: 乙 + base_url: http://{up} + key: sk-x + protocol: anthropic +routes: + - name: default + rules: + - name: 大输入 + when: {{ input_tokens: '>2000' }} + to: 乙 + - name: 其余 + to: 甲 +" + )) + .unwrap() +} + +/// 发一个请求,交回它的路由事件:规则、实际回答的那一家、这一轮的决定是不是沿用的 +async fn ask(gw: SocketAddr, rx: &mut Receiver, messages: &str) -> (String, String, bool) { + let body = format!( + r#"{{"model":"relay-model","max_tokens":16,"system":"你是一个助手","messages":{messages}}}"# + ); + let st = reqwest::Client::builder() + .no_proxy() + .build() + .unwrap() + .post(format!("http://{gw}/v1/messages")) + .header("x-api-key", "tw-k") + .header("x-claude-code-session-id", "会话-1") + .body(body) + .send() + .await + .unwrap() + .status(); + assert_eq!(st, 200); + let mut routed = None; + loop { + let ev = tokio::time::timeout(Duration::from_secs(5), rx.recv()) + .await + .expect("5 秒内没等到结局") + .expect("事件流断了"); + match ev { + Event::RequestRouted { + rule, + attempts, + affinity, + .. + } => { + routed = Some(( + rule, + attempts.last().unwrap().provider.clone(), + affinity.is_some_and(|a| a.held_route), + )) + } + Event::RequestFinished { .. } => return routed.expect("没有路由事件"), + Event::RequestFailed { message, .. } => panic!("失败了:{message:?}"), + _ => {} + } + } +} + +/// 同一轮里先小后大:第二个请求本该沿用这一轮开头的决定 +async fn a_turn_that_grows(cfg: Config) -> (String, String, bool) { + let (gw, mut rx) = serve(cfg).await; + let first = r#"{"role":"user","content":"帮我重构"}"#; + let (rule, to, _) = ask(gw, &mut rx, &format!("[{first}]")).await; + assert_eq!((rule.as_str(), to.as_str()), ("其余", "甲")); + let big = "读到的文件内容 ".repeat(1500); + let tool = format!( + r#"{{"role":"assistant","content":[{{"type":"tool_use","id":"t1","name":"Read","input":{{}}}}]}},{{"role":"user","content":[{{"type":"tool_result","tool_use_id":"t1","content":"{big}"}}]}}"# + ); + ask(gw, &mut rx, &format!("[{first},{tool}]")).await +} + +/// 价目表不认识 `relay-model`,说不出它装不装得下:一轮半路照常沿用。甲手写了 2000 的 +/// 上下文窗口,变大的输入装不下了:重新求值,去了乙 +#[tokio::test] +async fn a_turn_that_outgrows_the_window_written_by_hand_is_routed_again() { + let (rule, to, held) = a_turn_that_grows(two_upstreams(None).await).await; + assert_eq!((rule.as_str(), to.as_str(), held), ("其余", "甲", true)); + + let (rule, to, held) = a_turn_that_grows(two_upstreams(Some(2_000)).await).await; + assert_eq!((rule.as_str(), to.as_str(), held), ("大输入", "乙", false)); +} diff --git a/docs/config.md b/docs/config.md index f30d6155..7662d324 100644 --- a/docs/config.md +++ b/docs/config.md @@ -344,6 +344,7 @@ Upstreams: the APIs requests are forwarded to. | `models_only` | list of strings | — | Use only these of the upstream's models, as ids or globs. Others are not listed and are not routed here. Unset: all of them. Empty is refused; use `disabled`. | | `billing` | `per-token` \| `free` | `per-token` | `per-token`: cost is usage times the price in the upstream's price sheet, subscription accounts included. `free`: cost is recorded as 0. | | `pricing` | string | — | Name of a price sheet under `pricing.sheets`. Unset: the default price table. | +| `model_specs` | map of model id → [`providers[].model_specs.*`](#cfg-providers-model_specs) | `{}` | Context window and output limit of single models of this upstream, written by hand, by exact model id. They take precedence over the price table: for models it does not know, or gets wrong. | | `disabled` | bool | `false` | Take the upstream out of routing and out of the model list, and keep its configuration. | @@ -470,6 +471,40 @@ without them requests are still forwarded, and `models` can list the models by hand. For a VPC endpoint or a proxy, write its address in `base_url` and the region in `aws.region`; the model list is asked of that address too. +#### `providers[].model_specs` + +A model's context window and output limit come from the price table. A relay's +own models are often missing from it, and now and then it is wrong. Write the +numbers here, for this upstream and by the exact id in its model list. A value +written here takes precedence over the price table; one left out still comes +from it. At least one of the two is written, and neither can be 0. + +The same numbers are used everywhere: in `/v1/models` for every client format, +for an alias this upstream serves, when the gateway judges whether a +conversation still fits the model it is on, and as the output limit of a +request converted to Anthropic that does not set one. When several upstreams +offer the same model, `/v1/models` describes it by the first of them in +`providers`. + + + + +| Field | Type | Default | Description | +|---|---|---|---| +| `context_window` | integer | — | Context window: the most tokens a request can take in. Unset: the price table's. | +| `max_output_tokens` | integer | — | The most tokens an answer can have. Unset: the price table's. | + + +```yaml +providers: + - name: relay + base_url: https://relay.example.com/v1 + protocol: openai-chat + model_specs: + glm-5-air: { context_window: 128000, max_output_tokens: 16384 } + claude-sonnet-4-5: { context_window: 1000000 } +``` + ### `proxies` Outbound proxies. Different upstreams often need different ones, so there diff --git a/docs/config.zh-CN.md b/docs/config.zh-CN.md index 7fee25d7..f1731269 100644 --- a/docs/config.zh-CN.md +++ b/docs/config.zh-CN.md @@ -253,6 +253,7 @@ clients: | `models_only` | 字符串列表 | — | 只使用这家的这些模型,写 ID 或通配。范围外的模型不出现在模型列表里,也不会路由到这家。不写:全部。写空列表会被拒绝,暂停使用请用 `disabled`。 | | `billing` | `per-token` \| `free` | `per-token` | `per-token`:费用为用量乘以所选价目表中的单价,订阅账号同样如此。`free`:费用记为 0。 | | `pricing` | 字符串 | — | `pricing.sheets` 中某张价目表的名字。不写:默认价目表。 | +| `model_specs` | 映射: 模型 ID → [`providers[].model_specs.*`](#cfg-providers-model_specs) | `{}` | 手写这家上游某些模型的上下文窗口和输出上限,按模型 ID 完全匹配。写了就优先于价目表,用于价目表里没有或写错的模型。 | | `disabled` | 布尔 | `false` | 不参与路由,模型也不出现在模型列表里;配置原样保留。 | @@ -347,6 +348,31 @@ providers: 请求转换为 Converse 格式。模型清单取自所在区域的控制面:可按需调用的基础模型、AWS 预设的推理配置(`us.anthropic.claude-…`),以及账号自己创建的应用推理配置(按调用时使用的 ARN 列出)。列出清单需要 `bedrock:ListFoundationModels` 和 `bedrock:ListInferenceProfiles` 权限;没有这两项权限时请求照常转发,可以在 `models` 中手动列出模型。使用 VPC 端点或代理时,在 `base_url` 中写它的地址,在 `aws.region` 中写区域;模型清单也向该地址获取。 +#### `providers[].model_specs` + +模型的上下文窗口和输出上限取自价目表。中转站自有的模型常常不在价目表里,价目表偶尔也会写错。这时在这里按这家上游、按它模型清单里的 ID(完全匹配)手写。写了的一项优先于价目表,没写的一项仍取价目表。两项至少写一项,都不能是 0。 + +各处用的是同一个数:各种客户端格式的 `/v1/models`、由这家上游服务的别名、网关判断一段对话是否还装得下当前模型,以及请求转换为 Anthropic 格式且没有写输出上限时补上的值。同一个模型由几家上游提供时,`/v1/models` 按 `providers` 中排在最前的那一家给出。 + + + + +| 字段 | 类型 | 默认值 | 说明 | +|---|---|---|---| +| `context_window` | 整数 | — | 上下文窗口,即一次请求最多输入多少 token。不写:取价目表的。 | +| `max_output_tokens` | 整数 | — | 一次回答最多输出多少 token。不写:取价目表的。 | + + +```yaml +providers: + - name: relay + base_url: https://relay.example.com/v1 + protocol: openai-chat + model_specs: + glm-5-air: { context_window: 128000, max_output_tokens: 16384 } + claude-sonnet-4-5: { context_window: 1000000 } +``` + ### `proxies` 出站代理。不同上游需要的代理往往不同,因此没有全局开关:由每个上游用 `proxy` 选择。 From 4251aca2875244a807453a6d4c9622d0ae51373c Mon Sep 17 00:00:00 2001 From: fylorn <249551762+fylorn@users.noreply.github.com> Date: Mon, 5 Oct 2026 16:09:42 +0800 Subject: [PATCH 05/22] Cap concurrent requests per upstream Some relays and accounts accept only a few requests at a time and refuse the rest, so a busy upstream turned ordinary load into failed requests. `providers[].max_concurrent` (1-1000) makes the gateway respect that limit before a request is sent, instead of learning it from a refusal. A request takes a slot when its hop is about to be sent (before plugins and conversion run) and gives it back when the answer has been relayed to its end or the client is gone. The engine's order is unchanged: - the upstream that conversation stickiness kept (its cache, or the same turn) is waited for, since moving loses the cache; - any other full upstream is skipped at once (attempt skipped: busy); - when every remaining candidate is full, the request waits for the first to free, in candidate order, then answers 429 with Retry-After in the client's API shape (gw.busy_all, recorded as rate_limited). All waiting shares one budget per request, `failover.slot_wait_secs` (default 30, 0 = never wait), because the client receives nothing while it waits. Waiting is not a failure: no pause, no breaker count, and no outcome in the success rate load-balance groups weigh by. Token counting and WebSocket connections take no slot. Limits live in gateway state across reloads and are resized in place; a removed limit releases its waiters. With failover.next_on_slow_start, a full upstream still counts as one that can take the request when deciding whether the slow one is the last: it may free up within the slot wait, so the slow upstream is given up on and the request may then wait for the full one (or get the busy 429). Giving up on a slow upstream returns its slot at once. Traffic shows `AttemptView.queued_ms` and `skipped: busy`; the control plane exposes `ProviderView/ProviderInput.max_concurrent` and `FailoverView.slot_wait_secs`. Co-Authored-By: Claude Opus 5.5 --- crates/tw-api/msg-codes.txt | 3 + crates/tw-api/src/lib.rs | 19 + crates/tw-config/src/failover.rs | 18 +- crates/tw-config/src/lib.rs | 13 +- crates/tw-config/src/validate.rs | 48 +- crates/tw-config/tests/manual/schema.rs | 18 + crates/tw-control/src/lib.rs | 2 + crates/tw-control/src/resources.rs | 1 + crates/tw-control/tests/in_flight.rs | 2 + crates/tw-control/tests/live_state.rs | 2 + crates/tw-control/tests/replay.rs | 2 + crates/tw-control/tests/resources.rs | 56 ++ crates/tw-gateway/src/error.rs | 39 +- crates/tw-gateway/src/lib.rs | 1 + crates/tw-gateway/src/limits.rs | 9 +- crates/tw-gateway/src/server.rs | 34 + crates/tw-gateway/src/server/pipeline.rs | 4 + crates/tw-gateway/src/server/pipeline/hop.rs | 136 +++- .../tw-gateway/src/server/pipeline/relay.rs | 3 + crates/tw-gateway/src/server/pipeline/slow.rs | 3 + crates/tw-gateway/src/server/upgrade.rs | 3 + crates/tw-gateway/src/slots.rs | 395 ++++++++++++ crates/tw-gateway/src/state.rs | 7 + crates/tw-gateway/src/wire.rs | 2 +- crates/tw-gateway/tests/slow_start.rs | 64 ++ crates/tw-gateway/tests/upstream_slots.rs | 586 ++++++++++++++++++ crates/tw-observe/src/bus.rs | 4 + crates/tw-store/src/db.rs | 2 + crates/tw-store/src/recorder.rs | 22 + docs/config.md | 25 +- docs/config.zh-CN.md | 12 +- 31 files changed, 1511 insertions(+), 24 deletions(-) create mode 100644 crates/tw-gateway/src/slots.rs create mode 100644 crates/tw-gateway/tests/upstream_slots.rs diff --git a/crates/tw-api/msg-codes.txt b/crates/tw-api/msg-codes.txt index a8e71ea4..fd30156d 100644 --- a/crates/tw-api/msg-codes.txt +++ b/crates/tw-api/msg-codes.txt @@ -77,6 +77,7 @@ config.plugin.bad_id config.plugin.duplicate config.plugin.file config.plugin.sha256 +config.provider_concurrency_range config.rejected config.rejected_at config.remote_port_is_gateway @@ -259,6 +260,8 @@ gw.auth.key_invalid gw.auth.no_key gw.auth.source_not_allowed gw.auth.source_not_allowed_hint +gw.busy_all +gw.busy_upstream gw.chatgpt.token_missing gw.chatgpt.token_not_json gw.chatgpt.token_status diff --git a/crates/tw-api/src/lib.rs b/crates/tw-api/src/lib.rs index 500026cb..53aeb0d8 100644 --- a/crates/tw-api/src/lib.rs +++ b/crates/tw-api/src/lib.rs @@ -468,6 +468,9 @@ slug_enum! { /// 发给它的名字这把密钥不让用(`allow`):指定模型、阶段二改的名字一家一个, /// 只在试算给了密钥时出现 NotAllowed = "not_allowed", + /// 它的并发数满了(`max_concurrent`):这一跳没有发出去,换了下一家。只在尝试链里 + /// 出现([`AttemptView::skipped`]) + Busy = "busy", } } @@ -1485,6 +1488,14 @@ pub struct AttemptView { /// 出来的(请求解不开)没有。别的结果都没有:接下请求的那一跳的用量在结局里 #[serde(default, skip_serializing_if = "Option::is_none")] pub usage: Option, + /// 这一跳等了多少毫秒才轮到一个空位:这家设了 `max_concurrent` 而它满着。不算在 `ms` + /// 里。没等的没有 + #[serde(default, skip_serializing_if = "Option::is_none")] + pub queued_ms: Option, + /// 这一跳为什么没发出去:`busy`(这家满着,换了下一家;等过它的话 `queued_ms` 是等了 + /// 多久)。这时 `outcome` 是 `error`,`error` 是同一件事的那句话。发出去了的没有 + #[serde(default, skip_serializing_if = "Option::is_none")] + pub skipped: Option, } /// 放弃了的一跳([`AttemptOutcome::SlowStart`])上游可能已经收了钱的输入。 @@ -1871,6 +1882,8 @@ pub struct FailoverView { pub stream_start_wait_secs: u64, /// 等过 `stream_start_wait_secs` 还没有内容就换下一家(最后一家照常等) pub next_on_slow_start: bool, + /// 上游满着(`max_concurrent`)时,一个请求合计最多等多少秒空位。0 是不等 + pub slot_wait_secs: u64, } /// 每项防护各在哪一档:`off` / `observe` / `enforce`。 @@ -1963,6 +1976,9 @@ pub struct ProviderView { pub references: Vec, /// 选的价目表。空 = 默认价目表 pub pricing: Option, + /// 同时最多发给这家几个请求。不限是空 + #[serde(default, skip_serializing_if = "Option::is_none")] + pub max_concurrent: Option, } /// 一行请求头,配置里写的原样。 @@ -2826,6 +2842,9 @@ pub struct ProviderInput { /// 按哪张价目表计价。不给就是默认价目表 #[serde(default, skip_serializing_if = "Option::is_none")] pub pricing: Option, + /// 同时最多发给这家几个请求,1 到 1000。不给就是不限 + #[serde(default, skip_serializing_if = "Option::is_none")] + pub max_concurrent: Option, /// 停用 #[serde(default)] pub disabled: bool, diff --git a/crates/tw-config/src/failover.rs b/crates/tw-config/src/failover.rs index 621af9e2..48666c08 100644 --- a/crates/tw-config/src/failover.rs +++ b/crates/tw-config/src/failover.rs @@ -1,4 +1,5 @@ -//! 故障转移:一家上游失败之后停用多久、流开头最多等多久、等不到内容换不换下一家。 +//! 故障转移:一家上游失败之后停用多久、流开头最多等多久、等不到内容换不换下一家、上游 +//! 满着时最多等多久。 use serde::{Deserialize, Serialize}; @@ -42,6 +43,11 @@ pub struct Failover { /// 输出的模型开头本来就慢,开着时要把等待调长 #[serde(default)] pub next_on_slow_start: bool, + /// 上游的并发数满了(`providers[].max_concurrent`)时,一个请求最多等多少秒空位, + /// **整个请求合起来算**。留在那一家的对话等它空出来,候选都满了时等先空出来的那一家; + /// 等不到的换下一家,或者回 429。0 是不等 + #[serde(default = "d_slot_wait_secs")] + pub slot_wait_secs: u64, } fn d_failures_to_pause() -> u32 { @@ -65,6 +71,9 @@ fn d_rate_limit_max_pause_secs() -> u64 { fn d_stream_start_wait_secs() -> u64 { 15 } +fn d_slot_wait_secs() -> u64 { + 30 +} impl Default for Failover { fn default() -> Self { @@ -77,6 +86,7 @@ impl Default for Failover { rate_limit_max_pause_secs: d_rate_limit_max_pause_secs(), stream_start_wait_secs: d_stream_start_wait_secs(), next_on_slow_start: false, + slot_wait_secs: d_slot_wait_secs(), } } } @@ -92,10 +102,13 @@ pub const MAX_STREAM_START_WAIT_SECS: u64 = 120; /// 切掉了 pub const MIN_SLOW_START_WAIT_SECS: u64 = 5; +/// 等空位最多写多少秒。等的时候客户端一个字节都收不到,再长的话它先超时了 +pub const MAX_SLOT_WAIT_SECS: u64 = 300; + impl Failover { /// 不在允许范围里的第一项:字段名、写的值、下限、上限。 pub(crate) fn out_of_range(&self) -> Option<(&'static str, u64, u64, u64)> { - let fields: [(&'static str, u64, u64, u64); 7] = [ + let fields: [(&'static str, u64, u64, u64); 8] = [ ( "failures_to_pause", u64::from(self.failures_to_pause), @@ -128,6 +141,7 @@ impl Failover { 1, MAX_STREAM_START_WAIT_SECS, ), + ("slot_wait_secs", self.slot_wait_secs, 0, MAX_SLOT_WAIT_SECS), ]; fields .into_iter() diff --git a/crates/tw-config/src/lib.rs b/crates/tw-config/src/lib.rs index 334702ad..527aad57 100644 --- a/crates/tw-config/src/lib.rs +++ b/crates/tw-config/src/lib.rs @@ -189,6 +189,7 @@ impl Default for Provider { billing: Billing::PerToken, pricing: None, model_specs: std::collections::BTreeMap::new(), + max_concurrent: None, disabled: false, } } @@ -822,12 +823,21 @@ pub struct Provider { /// 优先于价目表,见 [`model_specs`]。价目表不认识的中转站模型靠它说出上下文窗口 #[serde(default, skip_serializing_if = "std::collections::BTreeMap::is_empty")] pub model_specs: std::collections::BTreeMap, + /// 同时最多发给这家几个请求,1 到 [`MAX_PROVIDER_CONCURRENCY`]。不写就是不限。 + /// + /// **给限制并发的中转站和账号用**:超出的那个请求到了上游只会被拒。满着的时候,留在 + /// 这家的对话等它空出来,别的请求换下一家(见 `tw_gateway::slots`) + #[serde(default, skip_serializing_if = "Option::is_none")] + pub max_concurrent: Option, /// 停用。**配置原样留着**:不参与路由,它的模型也不出现在 /// `/v1/models` 里。要暂时不用一家上游时,比删掉再重新填一遍凭据好。 #[serde(default, skip_serializing_if = "is_default")] pub disabled: bool, } +/// 一家上游的并发上限最多写多少 +pub const MAX_PROVIDER_CONCURRENCY: u32 = 1000; + impl Provider { /// 这家的这个模型在不在启用范围里(`models_only`)。**不管这家到底 /// 有没有这个模型** —— 那要看模型目录。 @@ -1176,7 +1186,8 @@ pub fn write(path: &Path, cfg: &Config) -> Result<(), WriteError> { } pub use failover::{ - Failover, MAX_PAUSE_SECS, MAX_STREAM_START_WAIT_SECS, MIN_SLOW_START_WAIT_SECS, + Failover, MAX_PAUSE_SECS, MAX_SLOT_WAIT_SECS, MAX_STREAM_START_WAIT_SECS, + MIN_SLOW_START_WAIT_SECS, }; pub use probes::{ClientProbes, ProbeAction}; pub use reload::{Rejected, Stage, stand_in, try_parse}; diff --git a/crates/tw-config/src/validate.rs b/crates/tw-config/src/validate.rs index 7b9a52b1..6474c953 100644 --- a/crates/tw-config/src/validate.rs +++ b/crates/tw-config/src/validate.rs @@ -33,6 +33,8 @@ pub enum ValidationError { #[error("{}", self.msg())] ZeroConcurrency { name: String }, #[error("{}", self.msg())] + ProviderConcurrency { name: String, value: u32 }, + #[error("{}", self.msg())] Routing(#[from] tw_engine::RouteError), #[error("{}", self.msg())] NameCollision(String), @@ -172,6 +174,12 @@ impl ValidationError { "gateway key `{key}` has max_concurrent: 0, so every request made with it would \ wait forever. Leave max_concurrent out for no limit" ), + ProviderConcurrency { name, value } => msg!( + "config.provider_concurrency_range", upstream = name, value = value, + max = crate::MAX_PROVIDER_CONCURRENCY => + "upstream `{upstream}` has max_concurrent: {value}; it has to be between 1 and \ + {max}. Leave max_concurrent out for no limit" + ), // 路由那几句本身就说清了是哪条规则、哪个组,前面不用再垫一句 Routing(e) => e.msg(), NameCollision(name) => msg!( @@ -374,6 +382,16 @@ pub fn validate(cfg: &Config) -> Result<(), ValidationError> { name: p.name.clone(), source, })?; + // 0 的话发给这家的请求一个都发不出去:留在它上面的对话每次都白等一场,别的请求 + // 每次都跳过它。不想用它是停用 + if let Some(n) = p.max_concurrent + && !(1..=crate::MAX_PROVIDER_CONCURRENCY).contains(&n) + { + return Err(ValidationError::ProviderConcurrency { + name: p.name.clone(), + value: n, + }); + } // **空范围不是「全部」,也不是一个合理的「停用」。**两种读法各有 // 人会当真,而停用有自己的开关 if let Some(only) = &p.models_only { @@ -891,6 +909,23 @@ mod tests { assert!(validate(&k).is_ok()); } + #[test] + fn an_upstreams_concurrency_limit_is_between_1_and_1000() { + let with = |n: Option| { + let mut prov = p("relay", "https://relay.example"); + prov.max_concurrent = n; + validate(&cfg(vec![c("a", "tw-a")], vec![prov])) + }; + assert!(with(None).is_ok(), "不写是不限"); + assert!(with(Some(1)).is_ok()); + assert!(with(Some(crate::MAX_PROVIDER_CONCURRENCY)).is_ok()); + for bad in [0, crate::MAX_PROVIDER_CONCURRENCY + 1] { + let e = with(Some(bad)).unwrap_err(); + assert_eq!(e.msg().code, "config.provider_concurrency_range", "{e}"); + assert!(e.to_string().contains("relay"), "{e}"); + } + } + #[test] fn an_empty_model_scope_is_refused_and_points_at_disabling_instead() { let mut prov = p("relay", "https://relay.example"); @@ -1183,7 +1218,7 @@ groups: let base = with_rules(&[], &[]); assert!(validate(&base).is_ok()); type Bend = fn(&mut crate::Failover); - let cases: [(&str, Bend); 4] = [ + let cases: [(&str, Bend); 5] = [ ("failures_to_pause", |f| f.failures_to_pause = 0), ("pause_secs", |f| f.pause_secs = 0), ("max_pause_secs", |f| { @@ -1193,6 +1228,9 @@ groups: ("stream_start_wait_secs", |f| { f.stream_start_wait_secs = crate::MAX_STREAM_START_WAIT_SECS + 1 }), + ("slot_wait_secs", |f| { + f.slot_wait_secs = crate::MAX_SLOT_WAIT_SECS + 1 + }), ]; for (want, bend) in cases { let mut x = base.clone(); @@ -1202,6 +1240,10 @@ groups: other => panic!("{want} 该被拒,实际 {other:?}"), } } + // 等空位写 0 是不等,不是写错 + let mut x = base.clone(); + x.failover.slot_wait_secs = 0; + assert!(validate(&x).is_ok()); } /// 开头慢就换下一家:开着时等待至少 5 秒,关着时 1 秒也照收(只是交得早) @@ -1454,6 +1496,10 @@ mod msg_codes { }, EmptyKey { name: "k".into() }, ZeroConcurrency { name: "k".into() }, + ProviderConcurrency { + name: "a".into(), + value: 0, + }, NameCollision("a".into()), BadCidr { entry: "x".into() }, UnknownPriceSheet { diff --git a/crates/tw-config/tests/manual/schema.rs b/crates/tw-config/tests/manual/schema.rs index 653be6ad..ed3b8482 100644 --- a/crates/tw-config/tests/manual/schema.rs +++ b/crates/tw-config/tests/manual/schema.rs @@ -570,6 +570,15 @@ pub fn sections() -> Vec
{ "手写这家上游某些模型的上下文窗口和输出上限,按模型 ID 完全匹配。写了就优先于价目表,用于价目表里没有或写错的模型。", ), ), + row( + "max_concurrent", + Kind::Int, + Def::Unset, + t( + "Most requests sent to this upstream at the same time, from 1 to 1000. When it is full, a conversation that stays on it waits for a free slot and other requests go to the next upstream; see `failover.slot_wait_secs`. Unset: no limit.", + "同时发给这家的请求最多几个,取值 1 到 1000。满了的时候,留在这家的对话等空位,别的请求换下一家;等多久见 `failover.slot_wait_secs`。不写:不限。", + ), + ), row( "disabled", Kind::Bool, @@ -1253,6 +1262,15 @@ pub fn sections() -> Vec
{ "流式回答在请求发出 `stream_start_wait_secs` 秒后仍没有内容时,放弃这家上游,把请求交给下一家。最后一家总是等下去。被放弃的上游不会停用。开启时 `stream_start_wait_secs` 至少为 5。", ), ), + row( + "slot_wait_secs", + Kind::Int, + Def::Is("30"), + t( + "Seconds a request waits in total for a free slot on upstreams that are at their `max_concurrent`. After that it goes to the next upstream, or, when every candidate is full, is answered with 429. `0`: never wait. From 0 to 300.", + "上游的并发数满了(`max_concurrent`)时,一个请求等空位合计最多等的秒数。等不到就换下一家;候选全满时回 429。`0`:不等。取值 0 到 300。", + ), + ), ], }, // ── groups / routes ─────────────────────────────────── diff --git a/crates/tw-control/src/lib.rs b/crates/tw-control/src/lib.rs index 59668921..cb7eb470 100644 --- a/crates/tw-control/src/lib.rs +++ b/crates/tw-control/src/lib.rs @@ -411,6 +411,7 @@ async fn overview(State(s): State) -> Json { rate_limit_max_pause_secs: f.rate_limit_max_pause_secs, stream_start_wait_secs: f.stream_start_wait_secs, next_on_slow_start: f.next_on_slow_start, + slot_wait_secs: f.slot_wait_secs, } }, listen: tw_api::ListenView { @@ -502,6 +503,7 @@ fn provider_view( .map(resources::reference_view) .collect(), pricing: p.pricing.clone(), + max_concurrent: p.max_concurrent, } } diff --git a/crates/tw-control/src/resources.rs b/crates/tw-control/src/resources.rs index 96b1faf8..40432825 100644 --- a/crates/tw-control/src/resources.rs +++ b/crates/tw-control/src/resources.rs @@ -471,6 +471,7 @@ fn to_provider( // 手写的模型规格不在这个表单里(上游页的模型清单一行一行改,见 `set_model_spec`): // 沿用原来的,改名时跟着这一项走 model_specs: existing.map(|e| e.model_specs.clone()).unwrap_or_default(), + max_concurrent: input.max_concurrent, disabled: input.disabled, }; // **保存和检测之前就说清楚凭据写法哪儿不对**,而不是等整份配置校验时 diff --git a/crates/tw-control/tests/in_flight.rs b/crates/tw-control/tests/in_flight.rs index aa6f5e6a..5adf2103 100644 --- a/crates/tw-control/tests/in_flight.rs +++ b/crates/tw-control/tests/in_flight.rs @@ -98,6 +98,8 @@ async fn it_gives_every_running_request_with_what_has_happened_to_it_so_far() { error: None, ms: 700, usage: None, + queued_ms: None, + skipped: None, }], billing: tw_api::Billing::PerToken, }); diff --git a/crates/tw-control/tests/live_state.rs b/crates/tw-control/tests/live_state.rs index ddcee347..c6e4261a 100644 --- a/crates/tw-control/tests/live_state.rs +++ b/crates/tw-control/tests/live_state.rs @@ -155,6 +155,8 @@ async fn a_request_still_running_can_be_opened_and_becomes_whole_when_it_ends() error: None, ms: 900, usage: None, + queued_ms: None, + skipped: None, }], billing: tw_api::Billing::PerToken, }); diff --git a/crates/tw-control/tests/replay.rs b/crates/tw-control/tests/replay.rs index f069212b..c79be1fd 100644 --- a/crates/tw-control/tests/replay.rs +++ b/crates/tw-control/tests/replay.rs @@ -386,6 +386,8 @@ fn routed(hops: &[(&str, Option<&str>)]) -> Option { error: None, ms: 100, usage: None, + queued_ms: None, + skipped: None, }) .collect(); Some( diff --git a/crates/tw-control/tests/resources.rs b/crates/tw-control/tests/resources.rs index 333895c8..8333eee7 100644 --- a/crates/tw-control/tests/resources.rs +++ b/crates/tw-control/tests/resources.rs @@ -180,6 +180,7 @@ async fn failover_settings_show_their_defaults_and_take_an_edit() { assert_eq!(f["failures_to_pause"], 3, "{body}"); assert_eq!(f["pause_secs"], 60); assert_eq!(f["stream_start_wait_secs"], 15); + assert_eq!(f["slot_wait_secs"], 30); let (st, body) = call( &b.app, @@ -209,6 +210,19 @@ async fn failover_settings_show_their_defaults_and_take_an_edit() { .await; assert_eq!(st, StatusCode::BAD_REQUEST, "{body}"); assert!(body.contains("config.failover_range"), "{body}"); + + // 等空位写 0 是不等,存得进去 + let (st, body) = call( + &b.app, + "PATCH", + "/config", + serde_json::json!({ + "ops": [{ "op": "replace", "path": "/failover/slot_wait_secs", "value": 0 }], + }), + ) + .await; + assert_eq!(st, StatusCode::OK, "{body}"); + assert_eq!(b.parsed().failover.slot_wait_secs, 0); } /// 开头慢就换下一家:默认关;打开要等得够久,等得太短的被拒 @@ -313,6 +327,48 @@ async fn saving_an_upstream_as_the_overview_shows_it_touches_only_what_changed() assert_eq!(p.billing, tw_config::Billing::Free); } +/// 并发上限:存进去、概览原样给回来;去掉就是不限。0 和超过 1000 的存不进去 +#[tokio::test] +async fn an_upstreams_concurrency_limit_is_saved_shown_and_checked() { + let b = bed(BASE); + let save = |n: serde_json::Value| { + let app = b.app.clone(); + async move { + call( + &app, + "PUT", + "/providers/官方", + serde_json::json!({ "provider": official(serde_json::json!({ "max_concurrent": n })) }), + ) + .await + } + }; + let (st, body) = save(serde_json::json!(4)).await; + assert_eq!(st, StatusCode::OK, "{body}"); + assert_eq!(b.parsed().providers[0].max_concurrent, Some(4)); + assert!(b.file().contains("max_concurrent: 4"), "{}", b.file()); + let (_, body) = call(&b.app, "GET", "/overview", serde_json::Value::Null).await; + assert_eq!(json(&body)["providers"][0]["max_concurrent"], 4, "{body}"); + + let before = b.file(); + for bad in [0, 1001] { + let (st, body) = save(serde_json::json!(bad)).await; + assert_eq!(st, StatusCode::BAD_REQUEST, "{bad}:{body}"); + assert!(body.contains("config.provider_concurrency_range"), "{body}"); + } + assert_eq!(b.file(), before, "拒绝了的保存不该动配置"); + + // 不交就是不限:这一行从文件里去掉,概览里也没有 + let (st, body) = save(serde_json::Value::Null).await; + assert_eq!(st, StatusCode::OK, "{body}"); + assert!(!b.file().contains("max_concurrent"), "{}", b.file()); + let (_, body) = call(&b.app, "GET", "/overview", serde_json::Value::Null).await; + assert!( + json(&body)["providers"][0].get("max_concurrent").is_none(), + "{body}" + ); +} + #[tokio::test] async fn billing_is_per_token_or_free_and_nothing_else() { // 订阅账号也按价目表算费用:「订阅」「未知」不再是计费方式,存不进去 diff --git a/crates/tw-gateway/src/error.rs b/crates/tw-gateway/src/error.rs index 98bca412..bde5a0c3 100644 --- a/crates/tw-gateway/src/error.rs +++ b/crates/tw-gateway/src/error.rs @@ -40,8 +40,18 @@ pub enum Source { /// 400 的话它当成请求写错了。对外的词表(`x-thinkwatch-error`、`RequestFailed.source`) /// 里算 `request`:服务不了的是这个请求要的接口 NotSupported, + /// 能服务这个请求的上游都满着(各自的 `max_concurrent`,见 [`crate::slots`]),等过了 + /// 也没空出来。**和上游限流一样回 429,另外带 `Retry-After`**:请求本身没问题,过一会儿 + /// 再来就能发出去 —— 客户端该退避再试,不是放弃。对外的词表里算 `rate_limited` + Busy, } +/// 上游都满着时告诉客户端过几秒再来(`Retry-After`)。 +/// +/// **不按 `slot_wait_secs` 算**:那么久已经在网关里等过了,再让客户端干等同样久没有意义。 +/// 重试进来照样排队等空位,所以这个数只管客户端别立刻打回来 +pub const BUSY_RETRY_AFTER_SECS: u64 = 5; + impl Source { /// `x-thinkwatch-error` 头和 `RequestFailed.source` 共用的词表。 pub fn slug(&self) -> &'static str { @@ -50,7 +60,7 @@ impl Source { Source::Config => "config", Source::Upstream => "upstream", Source::Request => "request", - Source::RateLimited => "rate_limited", + Source::RateLimited | Source::Busy => "rate_limited", Source::Denied => "denied", Source::NotSupported => "request", } @@ -62,7 +72,7 @@ impl Source { Source::Upstream => StatusCode::BAD_GATEWAY, Source::Request => StatusCode::BAD_REQUEST, // 429 而不是 503:客户端至少知道这是限流,可以退避。 - Source::RateLimited => StatusCode::TOO_MANY_REQUESTS, + Source::RateLimited | Source::Busy => StatusCode::TOO_MANY_REQUESTS, Source::Denied => StatusCode::FORBIDDEN, Source::NotSupported => StatusCode::NOT_IMPLEMENTED, } @@ -117,6 +127,9 @@ impl GatewayError { pub fn rate_limited(detail: Msg) -> Self { Self::new(Source::RateLimited, detail) } + pub fn busy(detail: Msg) -> Self { + Self::new(Source::Busy, detail) + } pub fn denied(detail: Msg) -> Self { Self::new(Source::Denied, detail) } @@ -181,6 +194,12 @@ impl IntoResponse for GatewayError { "x-thinkwatch-error", HeaderValue::from_static(self.source.slug()), ); + if self.source == Source::Busy { + h.insert( + header::RETRY_AFTER, + HeaderValue::from(BUSY_RETRY_AFTER_SECS), + ); + } resp } } @@ -267,6 +286,22 @@ mod tests { assert_eq!(json["error"]["status"], "UNAVAILABLE"); } + #[tokio::test] + async fn busy_upstreams_are_a_429_that_says_when_to_come_back() { + let r = GatewayError::busy(msg!("t.x" => "all busy")) + .in_dialect(Dialect::Chat) + .into_response(); + assert_eq!(r.status(), StatusCode::TOO_MANY_REQUESTS); + assert_eq!(r.headers()["retry-after"], "5"); + assert_eq!(r.headers()["x-thinkwatch-error"], "rate_limited"); + let b = to_bytes(r.into_body(), 64 * 1024).await.unwrap(); + let json: serde_json::Value = serde_json::from_slice(&b).unwrap(); + assert_eq!(json["error"]["type"], "rate_limit_error"); + // 上游自己的限流不带:要等多久由上游说,网关不替它编一个数 + let r = GatewayError::rate_limited(msg!("t.x" => "slow down")).into_response(); + assert!(!r.headers().contains_key("retry-after")); + } + #[test] fn a_responses_stream_is_told_with_response_failed() { // Chat 形状的 `{"error":…}` 在 Responses 的流里是一帧没人认的数据,客户端 diff --git a/crates/tw-gateway/src/lib.rs b/crates/tw-gateway/src/lib.rs index 1f1cb398..315fcb7c 100644 --- a/crates/tw-gateway/src/lib.rs +++ b/crates/tw-gateway/src/lib.rs @@ -43,6 +43,7 @@ pub mod seal; pub mod sent; pub mod server; pub mod session; +pub mod slots; pub mod state; pub mod translate; pub mod wire; diff --git a/crates/tw-gateway/src/limits.rs b/crates/tw-gateway/src/limits.rs index bb744bef..94b522dd 100644 --- a/crates/tw-gateway/src/limits.rs +++ b/crates/tw-gateway/src/limits.rs @@ -1,9 +1,10 @@ //! 每把密钥自己的并发上限。 //! -//! **网关本身不设上限** —— 没有全局的、没有单个上游的,也没有队列。本机 -//! 网关同时在跑的,就是这台电脑上几个客户端各自开着的会话;再压一道闸, -//! 挡住的只会是用户自己的并行任务。唯一的上限是用户给某一把密钥设的 -//! `max_concurrent`:那把密钥后面若是一个失控的脚本,它占不走别人的份。 +//! **网关本身不设上限** —— 没有全局的,也没有默认的。本机网关同时在跑的,就是这台 +//! 电脑上几个客户端各自开着的会话;再压一道闸,挡住的只会是用户自己的并行任务。上限 +//! 只有用户自己设的两种:一把密钥的 `max_concurrent`(这里),那把密钥后面若是一个 +//! 失控的脚本,它占不走别人的份;一家上游的 `max_concurrent`(见 [`crate::slots`]), +//! 那家同时只接得了这么多。 //! //! 超出上限的请求**等前面的结束**,不拒绝:客户端收到 429 通常不会优雅 //! 重试 —— Claude Code 把它当成硬失败,一个本来只需要多等两秒的请求会变成 diff --git a/crates/tw-gateway/src/server.rs b/crates/tw-gateway/src/server.rs index dd05817c..259cc6a7 100644 --- a/crates/tw-gateway/src/server.rs +++ b/crates/tw-gateway/src/server.rs @@ -308,6 +308,10 @@ pub(crate) struct Choice { pub(crate) rewritten_by: Vec, /// 这段对话之前的去向起的作用(见 [`crate::affinity`]) pub(crate) affinity: Option, + /// 这段对话留在的那一家:[`crate::affinity::Affinity::stay`] 把它挪到了候选的头上。 + /// **它满着时等它空出来**,不像别的候选那样当场跳过 —— 留下就是为了它的缓存(见 + /// [`crate::slots`])。没留的是 None + pub(crate) stayed_on: Option, } /// 规则做了决定、这个请求却一家上游都不会去时的路由事件:尝试链是空的。 @@ -347,6 +351,34 @@ pub(crate) fn hop( error: None, ms: started.elapsed().as_millis() as u64, usage: None, + queued_ms: None, + skipped: None, + } +} + +/// 这家满着、没发出去的一跳(见 [`crate::slots`])。`queued_ms`:等过它的话等了多久。 +/// +/// **不是失败**:上游什么都没说,不停用、不进熔断的账,这一行只说明请求为什么去了下一家。 +pub(crate) fn hop_busy( + provider: &str, + model: Option, + limit: usize, + queued_ms: Option, + started: std::time::Instant, +) -> tw_api::AttemptView { + tw_api::AttemptView { + provider: provider.to_string(), + model, + outcome: tw_api::AttemptOutcome::Error, + status: None, + error: Some(msg!( + "gw.busy_upstream", upstream = provider, limit = limit => + "Upstream `{upstream}` already has {limit} requests in progress, its max_concurrent." + )), + ms: started.elapsed().as_millis() as u64, + usage: None, + queued_ms, + skipped: Some(tw_api::ServeSkip::Busy), } } @@ -365,6 +397,8 @@ pub(crate) fn hop_failed( error: Some(error), ms: started.elapsed().as_millis() as u64, usage: None, + queued_ms: None, + skipped: None, } } diff --git a/crates/tw-gateway/src/server/pipeline.rs b/crates/tw-gateway/src/server/pipeline.rs index 8ef2b076..1e756dc1 100644 --- a/crates/tw-gateway/src/server/pipeline.rs +++ b/crates/tw-gateway/src/server/pipeline.rs @@ -613,6 +613,7 @@ fn route( group: None, rewritten_by: Vec::new(), affinity: None, + stayed_on: None, }; return Ok(Routed::Refused(choice, why)); } @@ -627,6 +628,7 @@ fn route( held_route, stayed: None, }), + stayed_on: None, }; // 每个候选实际要的模型和它的来历:规则改写过的按改写后的算;客户端写的、阶段一改写的 // 是客户端那一侧的名称(可能是别名),指定的、阶段二改的原样发出。准入看它 @@ -700,6 +702,8 @@ fn route( held_route, stayed: Some(why), }); + // 留下的那一家在头上。满着时等它,不当场跳过(见 `crate::slots`) + choice.stayed_on = decision.candidates.first().cloned(); } // `load-balance` 记账:**记粘性之后排头的那一家**,不是按权重轮到的那一家。一段对话 // 留在了上次回答它的那一家,这一次就算那一家的;之后的新对话把差的补回去 diff --git a/crates/tw-gateway/src/server/pipeline/hop.rs b/crates/tw-gateway/src/server/pipeline/hop.rs index a213e9ef..688d96b3 100644 --- a/crates/tw-gateway/src/server/pipeline/hop.rs +++ b/crates/tw-gateway/src/server/pipeline/hop.rs @@ -43,6 +43,9 @@ pub(super) struct Served<'a> { /// 交出去的是网关替上游说的一句话,不是它的原话(Bedrock 拒绝凭证,见 /// [`bedrock_refusal`]):这个请求失败的原因就是这一句。别的都是 None pub(super) refusal: Option, + /// 这一跳在这家占着的位置(见 [`crate::slots`])。**跟着回答走**:交完、或者客户端 + /// 走掉,响应体被丢掉时才还回去 + pub(super) slot: crate::slots::Slot, } /// 这个请求的着落。 @@ -138,10 +141,40 @@ pub(super) async fn try_upstreams<'a>( let allow = crate::models::key_allow(&rt.config, &req.client_name); // 开头慢就换下一家:开着、客户端要的是流时,等多久(见 `super::slow`) let slow_wait = super::slow::wait(rt, reading); + // 还没看的候选,按顺序 + let mut queue: std::collections::VecDeque<&String> = started.alive.iter().collect(); + // 满着没发的那几家(见 `crate::slots`),按候选的顺序。候选都看过一遍还没有着落时, + // 在它们里面等先空出来的那一家 + let mut busy: Vec<&String> = Vec::new(); + // 等空位等到什么时候:头一次要等时定下,**整个请求共用这一段** —— 等的时候客户端 + // 一个字节都收不到 + let wait = std::time::Duration::from_secs(rt.config.failover.slot_wait_secs); + let mut deadline: Option = None; + // 等过空位的那一跳在尝试链上的位置和等了多久。那一跳怎么收场都只进一行,进了之后补上 + let mut queued: Option<(usize, u64)> = None; + // 等到最后,剩下的候选还都满着 + let mut stalled = false; - for (i, name) in started.alive.iter().enumerate() { - // 后面没有别的候选了 - let last = i + 1 == started.alive.len(); + loop { + stamp_queued(&mut chain, queued.take()); + // 下一个候选。都看过了,就在满着的那几家里等先空出来的:等到的那一家带着占到的位置 + let (name, mut granted) = match queue.pop_front() { + Some(name) => (name, None), + None if busy.is_empty() => break, + None => { + let until = *deadline.get_or_insert_with(|| tokio::time::Instant::now() + wait); + let t = std::time::Instant::now(); + match state.slots.first_free(&busy, until).await { + Some((k, slot)) => (busy.remove(k), Some((slot, t.elapsed()))), + None => { + stalled = true; + break; + } + } + } + }; + // 后面没有别的候选了。满着、还等得到的那几家也算在后面 + let last = queue.is_empty() && busy.is_empty(); let Some(provider) = rt.config.providers.iter().find(|p| &p.name == name) else { // 校验时挡过一次,能到这儿说明配置在运行中被换过。 last_err = Some(GatewayError::config(msg!( @@ -157,7 +190,7 @@ pub(super) async fn try_upstreams<'a>( continue; } attempts.push(provider.name.clone()); - let hop_started = std::time::Instant::now(); + let mut hop_started = std::time::Instant::now(); // 数 token 选中的是别的格式的上游:不发,网关自己估 if counting && attempts.len() == 1 && estimates(req, provider) { @@ -274,6 +307,53 @@ pub(super) async fn try_upstreams<'a>( } } + // 并发上限:这一跳要发出去了,先占这家一个位置(见 `crate::slots`)。**在插件和 + // 转换之前**:满着没发的一跳不跑插件、不报转换。等不是失败:不停用、不进熔断的账。 + // + // 数 token 不占位置:它不跑模型,一眨眼就回来;为它排在几个长回答后面等上半分钟、 + // 最后回一个 429,客户端连上下文还剩多少都看不到 + let slot = match granted.take() { + Some((slot, waited)) => { + queued = Some((chain.len(), waited.as_millis() as u64)); + hop_started = std::time::Instant::now(); + slot + } + None if counting => crate::slots::Slot::free(), + None => { + let mut waited = None; + let mut slot = state.slots.try_take(&provider.name); + // 这段对话留在这家是为了它的缓存:等它空出来,等不到再换下一家(缓存就丢在 + // 这家了)。别的候选满着当场跳过 + if slot.is_none() && started.choice.stayed_on.as_ref() == Some(name) { + let until = *deadline.get_or_insert_with(|| tokio::time::Instant::now() + wait); + // 这个请求能等的已经等完了(或者配置的是不等):不再等 + if tokio::time::Instant::now() < until { + let t = std::time::Instant::now(); + slot = state.slots.take_by(&provider.name, until).await; + waited = Some(t.elapsed().as_millis() as u64); + hop_started = std::time::Instant::now(); + } + } + match slot { + Some(slot) => { + queued = waited.map(|ms| (chain.len(), ms)); + slot + } + None => { + chain.push(crate::server::hop_busy( + &provider.name, + Some(sent.clone()).filter(asked_other), + state.slots.limit(&provider.name).unwrap_or_default(), + waited, + hop_started, + )); + busy.push(name); + continue; + } + } + } + }; + // 插件的请求钩子:管这一跳的从客户端的原话起改。**拒绝的是整个请求**,不换下一家 let plugged = match super::plug::attempt( state, @@ -425,9 +505,12 @@ pub(super) async fn try_upstreams<'a>( .clone() .unwrap_or_else(|| reading.facts.model.clone()); let (attempt, bridge) = (chain.len(), plugged.bridge); - // 后面还有没有接得下这个请求的(见 `successor`)。开头慢了才问 - let rest = &started.alive[i + 1..]; - let others = || successor(state, rt, req, reading, decision, &catalog, allow, rest); + // 后面还有没有接得下这个请求的(见 `successor`)。开头慢了才问。后面的是还没看的 + // 候选,加上满着、跳过了的那几家:它们还可能空出来(见 `crate::slots`) + let others = || { + let rest = queue.iter().chain(busy.iter()).copied(); + successor(state, rt, req, reading, decision, &catalog, allow, rest) + }; // 开头慢就换下一家:等到什么时候,从这一刻(请求发出去)算起。最后一家不换。到点时 // 问过、后面没有接得下的,清掉它:这一跳从此和不开时一样 let mut slow_deadline = slow_wait @@ -603,6 +686,7 @@ pub(super) async fn try_upstreams<'a>( ledger, session: out.session, refusal, + slot, }); break; } @@ -734,6 +818,7 @@ pub(super) async fn try_upstreams<'a>( ledger, session: out.session, refusal: None, + slot, }); break; } @@ -769,6 +854,7 @@ pub(super) async fn try_upstreams<'a>( } } + stamp_queued(&mut chain, queued.take()); // 尝试链走完了,两条路都要发 —— 挂在 RequestFinished 上的话, // 失败那条路就没有尝试链,而那恰恰是最需要看它的时候。 // 最终服务的那家怎么收钱。**跟着请求走,不能事后查配置** —— @@ -826,6 +912,20 @@ pub(super) async fn try_upstreams<'a>( )))); } + // 剩下的候选等过了还都满着:**429,带 `Retry-After`**。请求本身没问题,过一会儿再来就 + // 发得出去;报成哪一家的失败都不对,它们一个字节都没收到 + if stalled { + let upstreams = busy + .iter() + .map(|n| format!("`{n}`")) + .collect::>() + .join(", "); + return Err(GatewayError::busy(msg!( + "gw.busy_all", upstreams = upstreams => + "Every upstream that can serve this request is at its concurrency limit \ + (max_concurrent): {upstreams}. None had a free slot in time; try again shortly." + ))); + } let Some(served) = served else { let mut err = last_err.unwrap_or_else(|| { GatewayError::config(msg!("gw.route.no_upstream_alive" => "No upstream is available.")) @@ -853,6 +953,16 @@ pub(super) async fn try_upstreams<'a>( Ok(Answer::Served(Box::new(served))) } +/// 等过空位的那一跳:在尝试链上补上它等了多久。`at` 是它那一行的位置 —— 等到之后, +/// 这一跳不管怎么收场都只进一行 +fn stamp_queued(chain: &mut [tw_api::AttemptView], queued: Option<(usize, u64)>) { + if let Some((at, ms)) = queued + && let Some(a) = chain.get_mut(at) + { + a.queued_ms = Some(ms); + } +} + /// 这一家和客户端是同一种格式(配置里没写格式的也算:照原样发过去) fn same_format(req: &Inbound, provider: &tw_config::Provider) -> bool { match (req.api, provider.effective_protocol()) { @@ -886,6 +996,8 @@ fn estimated_hop( error: None, ms: started.elapsed().as_millis() as u64, usage: None, + queued_ms: None, + skipped: None, } } @@ -974,10 +1086,14 @@ fn unsendable_tool( /// 没有的话,慢了的这一家就是最后一家**,照常等下去 —— 放弃了它,换来的是一个注定失败的 /// 请求。 /// +/// **满着的算接得下**(`max_concurrent`,见 [`crate::slots`]):满着只是此刻,等空位的 +/// 那一段(`failover.slot_wait_secs`)里它可能空出来。于是放弃了慢的这一家之后,请求可能 +/// 去等一家满着的、等不到时回 429 —— 不在这里猜它空不空得出来,「最后一家」只看接不接得下。 +/// /// **只看不跑**:看的是客户端的原话,不跑插件、不取密钥(那两样在真发的那一跳才知道拒 /// 不拒),不发转换事件 #[allow(clippy::too_many_arguments)] -fn successor( +fn successor<'r>( state: &AppState, rt: &Runtime, req: &Inbound, @@ -985,14 +1101,14 @@ fn successor( decision: &tw_engine::Decision, catalog: &tw_engine::Catalog, allow: Option<&[String]>, - rest: &[String], + mut rest: impl Iterator, ) -> bool { let asked = Asked { body: &req.body, path: req.uri.path(), decoded: reading.decoded.as_ref(), }; - rest.iter().any(|name| { + rest.any(|name| { let Some(provider) = rt.config.providers.iter().find(|p| &p.name == name) else { return false; }; diff --git a/crates/tw-gateway/src/server/pipeline/relay.rs b/crates/tw-gateway/src/server/pipeline/relay.rs index 6034e9f8..7b6c8a97 100644 --- a/crates/tw-gateway/src/server/pipeline/relay.rs +++ b/crates/tw-gateway/src/server/pipeline/relay.rs @@ -44,6 +44,7 @@ pub(super) fn respond( ledger, session, refusal, + slot, .. } = served; let status = @@ -160,6 +161,8 @@ pub(super) fn respond( // 发完是一种,客户端中途断开、hyper 丢掉响应体是另一种 —— 两种 // 都算这个请求结束了。 let _live = live; + // 在这家占着的位置也一样(见 `crate::slots`):回答交完、客户端走掉,才轮到下一个 + let _slot = slot; // 结局也一样:流被丢掉的时候,它替流报「客户端取消」。 let mut ending = ending; let mut chunks = std::pin::pin!(chunks); diff --git a/crates/tw-gateway/src/server/pipeline/slow.rs b/crates/tw-gateway/src/server/pipeline/slow.rs index afa69400..f9bb0c13 100644 --- a/crates/tw-gateway/src/server/pipeline/slow.rs +++ b/crates/tw-gateway/src/server/pipeline/slow.rs @@ -52,6 +52,9 @@ pub(super) fn abandoned( error: Some(said(provider, waited)), ms: started.elapsed().as_millis() as u64, usage: usage(seen, reading), + // 等过空位的话,等了多久由尝试链补上(`stamp_queued`) + queued_ms: None, + skipped: None, } } diff --git a/crates/tw-gateway/src/server/upgrade.rs b/crates/tw-gateway/src/server/upgrade.rs index c9b3d495..53b742e8 100644 --- a/crates/tw-gateway/src/server/upgrade.rs +++ b/crates/tw-gateway/src/server/upgrade.rs @@ -102,6 +102,7 @@ pub(super) async fn ws_upgrade( group: None, rewritten_by: Vec::new(), affinity: None, + stayed_on: None, }; let (id, ending) = open(&choice, "", tw_api::Billing::PerToken); state.bus.emit(super::routed_nowhere(id, choice)); @@ -156,6 +157,7 @@ pub(super) async fn ws_upgrade( group: decision.via_group.clone(), rewritten_by: decision.rewritten_by.clone(), affinity: None, + stayed_on: None, }; let Some(rule) = rule else { return Err(why) }; tracing::info!(%rule, provider = %c, "a phase-two rule denied the WebSocket upgrade"); @@ -210,6 +212,7 @@ pub(super) async fn ws_upgrade( rewritten_by, // WebSocket 那条路一条连接跑好几轮,不按对话记 affinity: None, + stayed_on: None, }; // **走代理的上游不代理 WS**,而且要明说。悄悄绕过用户配的代理, // 等于把他以为在代理后面的流量直接发出去 diff --git a/crates/tw-gateway/src/slots.rs b/crates/tw-gateway/src/slots.rs new file mode 100644 index 00000000..229c1473 --- /dev/null +++ b/crates/tw-gateway/src/slots.rs @@ -0,0 +1,395 @@ +//! 每家上游自己的并发上限(`providers[].max_concurrent`)。 +//! +//! 有的中转站、有的账号同时只接几个请求,多出来的那个到了上游直接被拒 —— 回一个 429, +//! 或者一句说不清的错误。上限记在网关这一侧,**一个请求在发出去之前就知道这家满了**, +//! 由管线决定等还是换(见 `server::pipeline::hop`): +//! +//! - 这段对话留在这家是为了它的缓存([`crate::affinity`]):等它空出一个位置,等不到 +//! 再换下一家,缓存就丢在这家了; +//! - 别的候选满着:当场跳过,试下一家; +//! - 候选都满着:在它们里面等先空出来的那一家,等不到就回 429。 +//! +//! 等多久是 `failover.slot_wait_secs`,**一个请求合起来算**:等的时候客户端一个字节都 +//! 收不到。**等不是失败**:满着的上游不停用、不进熔断的账。 +//! +//! 一个位置从发出请求占到回答交完、或者客户端走掉,由 [`Slot`] 的 Drop 还回去。和每把 +//! 密钥的闸([`crate::limits`])一样,**跨重载存活**,上限改了在原来那个信号量上加减: +//! 换一个新的,在跑的请求占着的位置就不算数了。上限去掉了的,等着的请求当场放行。 +//! +//! **不占位置的两种**:数 token 的请求,它不跑模型、一眨眼就回来;WebSocket 的连接(见 +//! [`crate::ws`]),一条连接跑好几轮、中间可以闲着很久,占着的话一条没在答话的连接就能 +//! 把这家堵死。 + +use std::collections::HashMap; +use std::sync::{Arc, Mutex}; + +use tokio::sync::{OwnedSemaphorePermit, Semaphore, TryAcquireError}; +use tokio::time::Instant; + +#[derive(Default)] +pub struct Slots { + per_upstream: Mutex>>, +} + +/// 一家上游的位置。 +struct Pool { + sem: Arc, + books: Mutex, +} + +struct Books { + limit: usize, + /// 调小上限时还没收回的位置。此刻空着的当场收,在跑的请求占着的那些等它们结束时收 + /// (见 [`Slot`] 的 Drop) + owed: usize, +} + +/// 占着的一个位置。**Drop 时自动归还** —— 手工释放迟早会在某条错误路径上漏掉,漏掉的 +/// 表现是这家的位置只减不增,最后发给它的请求全都在等。没设上限的上游给的是空的一个。 +pub struct Slot { + held: Option<(Arc, OwnedSemaphorePermit)>, +} + +impl Slot { + /// 这家不限并发:什么都不占 + pub fn free() -> Self { + Self { held: None } + } + + fn of(pool: Arc, permit: OwnedSemaphorePermit) -> Self { + Self { + held: Some((pool, permit)), + } + } +} + +impl Drop for Slot { + fn drop(&mut self) { + let Some((pool, permit)) = self.held.take() else { + return; + }; + let mut books = pool.books.lock().unwrap_or_else(|p| p.into_inner()); + if books.owed > 0 { + books.owed -= 1; + permit.forget(); + } + } +} + +impl Slots { + fn lock(&self) -> std::sync::MutexGuard<'_, HashMap>> { + self.per_upstream.lock().unwrap_or_else(|p| p.into_inner()) + } + + fn pool(&self, name: &str) -> Option> { + self.lock().get(name).cloned() + } + + /// 照这份配置定每家的上限。启动时和每次重载时调。 + /// + /// **改了上限的在原地加减**:调大的,等着的请求马上进;调小的,在跑的照样跑完,新的 + /// 要等它们降到新上限之下。去掉了上限的(或者这家上游没了)关掉它的信号量:等着的 + /// 请求当场放行,在跑的那些占着的位置随它们结束作废。 + pub fn configure(&self, providers: &[tw_config::Provider]) { + let limit_of = |name: &str| { + providers + .iter() + .find(|p| p.name == name) + .and_then(|p| p.max_concurrent) + }; + let mut per = self.lock(); + per.retain(|name, pool| { + let keep = limit_of(name).is_some(); + if !keep { + pool.sem.close(); + } + keep + }); + for p in providers { + let Some(limit) = p.max_concurrent.map(|n| n as usize) else { + continue; + }; + match per.get(&p.name) { + Some(pool) => pool.resize(limit), + None => { + per.insert( + p.name.clone(), + Arc::new(Pool { + sem: Arc::new(Semaphore::new(limit)), + books: Mutex::new(Books { limit, owed: 0 }), + }), + ); + } + } + } + } + + /// 这家此刻的上限。不限是 `None` + pub fn limit(&self, name: &str) -> Option { + let pool = self.pool(name)?; + let books = pool.books.lock().unwrap_or_else(|p| p.into_inner()); + Some(books.limit) + } + + /// 不等:有空位就占一个,这家满着是 `None`。不限并发的上游一律给一个空的。 + pub fn try_take(&self, name: &str) -> Option { + let Some(pool) = self.pool(name) else { + return Some(Slot::free()); + }; + match pool.sem.clone().try_acquire_owned() { + Ok(permit) => Some(Slot::of(pool, permit)), + // 上限刚去掉 + Err(TryAcquireError::Closed) => Some(Slot::free()), + Err(TryAcquireError::NoPermits) => None, + } + } + + /// 等这一家空出一个位置,最多等到 `until`。到点还满着是 `None`。 + pub async fn take_by(&self, name: &str, until: Instant) -> Option { + self.first_free(&[name], until).await.map(|(_, slot)| slot) + } + + /// 几家里先空出来的那一家:它在 `names` 里的位置和占到的位置。**同时空出来的按 + /// `names` 的顺序**,也就是候选的顺序。最多等到 `until`,到点都还满着是 `None`; + /// `until` 已经过了的话只看一眼。 + /// + /// 没等到的那几家不留下什么:等着的那一份被丢掉时,tokio 把已经分给它的位置还回去。 + pub async fn first_free>( + &self, + names: &[S], + until: Instant, + ) -> Option<(usize, Slot)> { + for (i, name) in names.iter().enumerate() { + if let Some(slot) = self.try_take(name.as_ref()) { + return Some((i, slot)); + } + } + if Instant::now() >= until { + return None; + } + let waits: Vec<_> = names + .iter() + .enumerate() + .filter_map(|(i, name)| { + let pool = self.pool(name.as_ref())?; + Some(Box::pin(async move { + match pool.sem.clone().acquire_owned().await { + Ok(permit) => (i, Slot::of(pool, permit)), + // 等的时候上限去掉了:不用再等 + Err(_) => (i, Slot::free()), + } + })) + }) + .collect(); + if waits.is_empty() { + // 看完那一眼之后上限都去掉了 + return Some((0, Slot::free())); + } + tokio::time::timeout_at(until, futures::future::select_all(waits)) + .await + .ok() + .map(|(got, _, _)| got) + } +} + +impl Pool { + /// 上限改了:**在原来那个信号量上加减,不换新的。**换新的话,在跑的请求占着的位置 + /// 还算在旧的那个上,新旧各放一份,那一阵发给这家的可以到两个上限之和。 + fn resize(&self, limit: usize) { + let mut books = self.books.lock().unwrap_or_else(|p| p.into_inner()); + if limit > books.limit { + // 先抵掉还欠着的,剩下的才是真要多放的 + let more = limit - books.limit; + let cancel = more.min(books.owed); + books.owed -= cancel; + self.sem.add_permits(more - cancel); + } else if limit < books.limit { + let less = books.limit - limit; + let now = self.sem.forget_permits(less); + books.owed += less - now; + } + books.limit = limit; + } +} + +#[cfg(test)] +mod tests { + use super::*; + use std::time::Duration; + + fn limited(name: &str, n: Option) -> tw_config::Provider { + tw_config::Provider { + name: name.into(), + base_url: "https://x".into(), + max_concurrent: n, + ..Default::default() + } + } + + fn slots(cfg: &[(&str, Option)]) -> Arc { + let s = Arc::new(Slots::default()); + s.configure(&cfg.iter().map(|(n, l)| limited(n, *l)).collect::>()); + s + } + + fn after(ms: u64) -> Instant { + Instant::now() + Duration::from_millis(ms) + } + + /// 让等着的请求有机会走一步 + async fn settle() { + tokio::time::sleep(Duration::from_millis(30)).await; + } + + #[test] + fn an_upstream_without_a_limit_never_fills() { + let s = slots(&[("甲", None)]); + let held: Vec<_> = (0..100).map(|_| s.try_take("甲").unwrap()).collect(); + assert_eq!(held.len(), 100); + // 配置里没有的也一样:不知道它的上限,就是不限 + assert!(s.try_take("不认识").is_some()); + } + + #[test] + fn a_full_upstream_says_so_at_once_and_a_dropped_slot_is_returned() { + let s = slots(&[("甲", Some(2)), ("乙", Some(1))]); + let a = s.try_take("甲").unwrap(); + let _b = s.try_take("甲").unwrap(); + assert!(s.try_take("甲").is_none(), "上限 2,第三个该被告知满了"); + // 一家满着不挡另一家 + assert!(s.try_take("乙").is_some()); + drop(a); + assert!(s.try_take("甲").is_some(), "还回去的位置没有回来"); + } + + #[tokio::test] + async fn a_wait_ends_with_the_slot_or_at_the_deadline() { + let s = slots(&[("甲", Some(1))]); + let first = s.try_take("甲").unwrap(); + let t = Instant::now(); + assert!(s.take_by("甲", after(80)).await.is_none()); + assert!(t.elapsed() >= Duration::from_millis(80), "没等到点就放弃了"); + + let s2 = s.clone(); + let waiter = tokio::spawn(async move { s2.take_by("甲", after(5_000)).await }); + settle().await; + assert!(!waiter.is_finished()); + drop(first); + let got = tokio::time::timeout(Duration::from_secs(1), waiter) + .await + .expect("空出来了却没轮到它") + .unwrap(); + assert!(got.is_some()); + } + + #[tokio::test] + async fn among_several_the_first_to_free_wins_and_ties_go_by_order() { + let s = slots(&[("甲", Some(1)), ("乙", Some(1)), ("丙", Some(1))]); + let (a, b, c) = ( + s.try_take("甲").unwrap(), + s.try_take("乙").unwrap(), + s.try_take("丙").unwrap(), + ); + let s2 = s.clone(); + let waiter = tokio::spawn(async move { + s2.first_free(&["甲", "乙", "丙"], after(5_000)) + .await + .map(|(i, _)| i) + }); + settle().await; + drop(c); + assert_eq!(waiter.await.unwrap(), Some(2), "丙先空出来"); + // 都空着的时候按顺序:头一个 + drop((a, b)); + let (i, _) = s.first_free(&["乙", "甲"], after(0)).await.unwrap(); + assert_eq!(i, 0); + // 等不到的那几家没有被占走位置 + assert!(s.try_take("甲").is_some()); + assert!(s.try_take("乙").is_some()); + } + + #[tokio::test] + async fn a_deadline_already_past_only_takes_a_look() { + let s = slots(&[("甲", Some(1))]); + let _held = s.try_take("甲").unwrap(); + let t = Instant::now(); + assert!(s.first_free(&["甲"], Instant::now()).await.is_none()); + assert!(t.elapsed() < Duration::from_millis(50)); + } + + #[tokio::test] + async fn a_waiter_dropped_just_as_it_is_granted_gives_the_slot_back() { + // 客户端在轮到它的那一刻走了:那个位置不能跟着丢 + let s = slots(&[("甲", Some(1))]); + let first = s.try_take("甲").unwrap(); + let s2 = s.clone(); + let waiter = tokio::spawn(async move { s2.take_by("甲", after(5_000)).await }); + settle().await; + drop(first); + waiter.abort(); + let _ = waiter.await; + settle().await; + assert!(s.try_take("甲").is_some(), "位置跟着走掉的请求丢了"); + } + + #[tokio::test] + async fn raising_a_limit_lets_a_waiter_in_at_once() { + let s = slots(&[("甲", Some(1))]); + let _a = s.try_take("甲").unwrap(); + let s2 = s.clone(); + let waiter = tokio::spawn(async move { s2.take_by("甲", after(5_000)).await }); + settle().await; + assert!(!waiter.is_finished()); + s.configure(&[limited("甲", Some(2))]); + let got = tokio::time::timeout(Duration::from_secs(1), waiter) + .await + .expect("调大之后等着的还在等") + .unwrap(); + assert!(got.is_some()); + } + + #[tokio::test] + async fn lowering_a_limit_counts_the_requests_already_running() { + // 从 2 调到 1:在跑的两个都结束之前,新来的不能进 + let s = slots(&[("甲", Some(2))]); + let a = s.try_take("甲").unwrap(); + let b = s.try_take("甲").unwrap(); + s.configure(&[limited("甲", Some(1))]); + drop(a); + assert!(s.try_take("甲").is_none(), "还有一个在跑,已经到新上限了"); + drop(b); + let c = s.try_take("甲"); + assert!(c.is_some()); + assert!(s.try_take("甲").is_none(), "新上限是 1"); + // 调回来:欠着的先抵掉,上限回到 2 + s.configure(&[limited("甲", Some(2))]); + let d = s.try_take("甲"); + assert!(d.is_some()); + assert!(s.try_take("甲").is_none(), "上限 2,两个都占着"); + } + + #[tokio::test] + async fn removing_a_limit_lets_every_waiter_go_and_nothing_breaks() { + let s = slots(&[("甲", Some(1))]); + let running = s.try_take("甲").unwrap(); + let s2 = s.clone(); + let waiter = tokio::spawn(async move { s2.take_by("甲", after(5_000)).await }); + settle().await; + // 上限去掉(上游删了也一样) + s.configure(&[limited("甲", None)]); + let got = tokio::time::timeout(Duration::from_secs(1), waiter) + .await + .expect("上限去掉了还在等") + .unwrap(); + assert!(got.is_some()); + assert!(s.try_take("甲").is_some(), "不限了"); + // 在跑的那个结束:它的位置属于一个已经关掉的信号量,还回去不出事 + drop(running); + // 再设回上限:重新数 + s.configure(&[limited("甲", Some(1))]); + let _one = s.try_take("甲").unwrap(); + assert!(s.try_take("甲").is_none()); + s.configure(&[]); + assert!(s.try_take("甲").is_some(), "上游没了,不再限它"); + } +} diff --git a/crates/tw-gateway/src/state.rs b/crates/tw-gateway/src/state.rs index fc122d6d..25fd0c3d 100644 --- a/crates/tw-gateway/src/state.rs +++ b/crates/tw-gateway/src/state.rs @@ -136,6 +136,9 @@ pub struct AppState { /// 换的话,每改一次配置,排着的请求就会失去位置,而已经在跑的那些的 /// 通行证会变成孤儿。上限改了由它自己在原地加减(见 `limits`)。 pub(crate) gate: Arc, + /// 每家上游自己的并发上限(见 [`crate::slots`])。**不在 Runtime 里**,理由和 `gate` + /// 一样:它握着在跑的请求占着的位置。配置换了由 [`Self::reload`] 在原地改上限 + pub slots: Arc, /// 一家上游不在当前运行时里时顶上的 Client(直连,不读系统代理) pub http: reqwest::Client, /// 观测事件往这里丢。没有订阅者时是零成本的 —— 数据面不该知道有 @@ -256,9 +259,12 @@ impl AppState { let rt = Runtime::build(config, None, &plugins)?; let health = Arc::new(Health::new()); health.configure(&rt.config.failover); + let slots = Arc::new(crate::slots::Slots::default()); + slots.configure(&rt.config.providers); let state = Self { rt: Arc::new(arc_swap::ArcSwap::from_pointee(rt)), gate: Default::default(), + slots, http, bus: tw_observe::EventBus::new(), health, @@ -476,6 +482,7 @@ impl AppState { self.pricing .rcu(|book| book.with_config(sheets.clone(), assign.clone())); self.health.configure(&next.config.failover); + self.slots.configure(&next.config.providers); self.rt.store(Arc::new(next)); self.announce_broken(broken); // 模型汇总马上按新配置重算:删掉、停用的上游的模型必须立刻消失(列表 diff --git a/crates/tw-gateway/src/wire.rs b/crates/tw-gateway/src/wire.rs index d957680a..897be607 100644 --- a/crates/tw-gateway/src/wire.rs +++ b/crates/tw-gateway/src/wire.rs @@ -11,7 +11,7 @@ impl From for FailureSource { Source::Config => Self::Config, Source::Upstream => Self::Upstream, Source::Request => Self::Request, - Source::RateLimited => Self::RateLimited, + Source::RateLimited | Source::Busy => Self::RateLimited, Source::Denied => Self::Denied, Source::NotSupported => Self::Request, } diff --git a/crates/tw-gateway/tests/slow_start.rs b/crates/tw-gateway/tests/slow_start.rs index 41d415e7..e06d71d1 100644 --- a/crates/tw-gateway/tests/slow_start.rs +++ b/crates/tw-gateway/tests/slow_start.rs @@ -540,3 +540,67 @@ async fn a_request_that_does_not_stream_is_not_switched() { [("slow", tw_api::AttemptOutcome::Served)] ); } + +/// 放弃的那一家占着的位置(`max_concurrent`,见 `tw_gateway::slots`)**当场**还回去:接下 +/// 请求的那一家还在答,慢的那一家已经空出来了,不等这个请求结束 +#[tokio::test] +async fn the_slot_of_an_upstream_given_up_on_is_free_at_once() { + const CONTENT: &str = "event: content_block_delta\ndata: {\"type\":\"content_block_delta\",\"index\":0,\"delta\":{\"type\":\"text_delta\",\"text\":\"hello\"}}\n\n"; + const STOP: &str = "event: message_stop\ndata: {\"type\":\"message_stop\"}\n\n"; + let slow = upstream(stalled()).await; + // 先答上,隔两秒才说完:这两秒里请求还没结束 + let good = upstream(Script { + header_delay_ms: 0, + steps: vec![(0, MESSAGE_START), (0, CONTENT), (2_000, STOP)], + hang: false, + }) + .await; + let cfg = Config { + version: 1, + clients: vec![Client { + name: "c".into(), + key: "tw-k".into(), + ..Default::default() + }], + providers: vec![ + Provider { + max_concurrent: Some(1), + ..provider("slow", &slow, Protocol::Anthropic) + }, + provider("good", &good, Protocol::Anthropic), + ], + failover: Failover { + stream_start_wait_secs: 1, + next_on_slow_start: true, + ..Default::default() + }, + ..Default::default() + }; + let state = tw_gateway::AppState::new(cfg).unwrap(); + let gw = tw_gateway::serve(state.clone(), ([127, 0, 0, 1], 0).into()) + .await + .unwrap(); + tokio::time::sleep(Duration::from_millis(50)).await; + + let asking = tokio::spawn(async move { post(gw, "/v1/messages", &messages(true)).await }); + // 还在等慢的那一家:它的位置占着 + tokio::time::sleep(Duration::from_millis(300)).await; + assert_eq!(slow.hits.load(Ordering::SeqCst), 1); + assert!(state.slots.try_take("slow").is_none(), "等着的时候占着位置"); + // 换到了好的那一家:慢的那一家的位置已经还回来了,这个请求还没结束 + for _ in 0..60 { + if good.hits.load(Ordering::SeqCst) == 1 { + break; + } + tokio::time::sleep(Duration::from_millis(50)).await; + } + assert_eq!(good.hits.load(Ordering::SeqCst), 1); + assert!(!asking.is_finished(), "测的是请求还在进行时"); + assert!( + state.slots.try_take("slow").is_some(), + "放弃的那一家的位置要当场还回去" + ); + let (status, text) = asking.await.unwrap(); + assert_eq!(status, 200, "{text}"); + assert!(text.contains("hello"), "{text}"); +} diff --git a/crates/tw-gateway/tests/upstream_slots.rs b/crates/tw-gateway/tests/upstream_slots.rs new file mode 100644 index 00000000..35ba725d --- /dev/null +++ b/crates/tw-gateway/tests/upstream_slots.rs @@ -0,0 +1,586 @@ +//! 上游的并发上限(`providers[].max_concurrent`):满着的上游,新的对话当场跳过,留在它 +//! 上面的对话等它空出来,都满着时等先空出来的那一家、等不到回 429 —— 走真实的管线,看 +//! 尝试链里说的和实际去的那一家。 + +use std::net::SocketAddr; +use std::sync::atomic::{AtomicUsize, Ordering}; +use std::sync::{Arc, Mutex}; +use std::time::Duration; + +use axum::Router; +use axum::routing::post; +use bytes::Bytes; +use futures::StreamExt; +use tokio::sync::Semaphore; +use tw_api::{AttemptOutcome, AttemptView, Event, FailureSource, ServeSkip, Stay}; + +/// 一家假的 Anthropic 上游。请求里写着 `HOLD` 的:先吐开头和一段内容,然后一直等到测试 +/// 放行([`Up::release`])才说完 —— 一个正在长篇作答的模型,占着一个位置。别的当场答完 +struct Up { + addr: SocketAddr, + /// 收到的生成请求数 + hits: Arc, + release: Arc, +} + +impl Up { + fn hits(&self) -> usize { + self.hits.load(Ordering::SeqCst) + } + /// 放走一个占着位置的请求 + fn release(&self) { + self.release.add_permits(1); + } +} + +const OPENING: &[u8] = b"event: message_start\ndata: {\"type\":\"message_start\",\"message\":{\"usage\":{\"input_tokens\":3,\"output_tokens\":1}}}\n\n\ +event: content_block_start\ndata: {\"type\":\"content_block_start\",\"index\":0,\"content_block\":{\"type\":\"text\",\"text\":\"\"}}\n\n\ +event: content_block_delta\ndata: {\"type\":\"content_block_delta\",\"index\":0,\"delta\":{\"type\":\"text_delta\",\"text\":\"hi\"}}\n\n"; + +const CLOSING: &[u8] = + b"event: content_block_stop\ndata: {\"type\":\"content_block_stop\",\"index\":0}\n\n\ +event: message_stop\ndata: {\"type\":\"message_stop\"}\n\n"; + +async fn upstream() -> Up { + let hits = Arc::new(AtomicUsize::new(0)); + let release = Arc::new(Semaphore::new(0)); + let (h, r) = (hits.clone(), release.clone()); + let app = Router::new().route( + "/v1/messages", + post(move |body: Bytes| { + let (hits, release) = (h.clone(), r.clone()); + async move { + hits.fetch_add(1, Ordering::SeqCst); + if String::from_utf8_lossy(&body).contains("HOLD") { + let stream = async_stream::stream! { + yield Ok::<_, std::io::Error>(Bytes::from_static(OPENING)); + if let Ok(p) = release.acquire().await { + p.forget(); + } + yield Ok(Bytes::from_static(CLOSING)); + }; + return axum::response::Response::builder() + .header("content-type", "text/event-stream") + .body(axum::body::Body::from_stream(stream)) + .unwrap(); + } + // 读了 5000 token 的缓存:这段对话跨轮也值得留下 + axum::response::Response::builder() + .header("content-type", "application/json") + .body(axum::body::Body::from( + r#"{"type":"message","role":"assistant","content":[{"type":"text","text":"ok"}],"usage":{"input_tokens":3,"cache_read_input_tokens":5000,"output_tokens":1}}"#, + )) + .unwrap() + } + }), + ); + let l = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap(); + let addr = l.local_addr().unwrap(); + tokio::spawn(async move { axum::serve(l, app).await.unwrap() }); + Up { + addr, + hits, + release, + } +} + +/// 甲、乙两家按顺序(`fallback`)。`limits` 是两家各自的上限 +fn config(a: &Up, b: &Up, limits: (Option, Option), wait_secs: u64) -> tw_config::Config { + let limit = |n: Option| { + n.map(|n| format!(" max_concurrent: {n}\n")) + .unwrap_or_default() + }; + let yaml = format!( + "version: 1 +clients: + - name: me + key: tw-k +providers: + - name: 甲 + base_url: http://{} + key: sk-x + protocol: anthropic +{} - name: 乙 + base_url: http://{} + key: sk-x + protocol: anthropic +{}groups: + - name: 池 + type: fallback + providers: [甲, 乙] +routes: + - name: 默认 + rules: + - name: 全部 + to: 池 +failover: + slot_wait_secs: {wait_secs} +", + a.addr, + limit(limits.0), + b.addr, + limit(limits.1), + ); + serde_yaml_ng::from_str(&yaml).unwrap() +} + +/// 总线上发生过的事件,后台一直收着 +#[derive(Clone)] +struct Log(Arc>>); + +impl Log { + async fn until(&self, what: &str, f: impl Fn(&[Event]) -> Option) -> T { + for _ in 0..500 { + if let Some(t) = f(&self.0.lock().unwrap()) { + return t; + } + tokio::time::sleep(Duration::from_millis(20)).await; + } + panic!("10 秒内没等到{what}:{:#?}", self.0.lock().unwrap()); + } + + /// 客户端写的模型是 `model` 的那个请求的号 + async fn id(&self, model: &str) -> u64 { + self.until(&format!(" {model} 开始"), |evs| { + evs.iter().find_map(|e| match e { + Event::RequestStarted { id, model: m, .. } if m == model => Some(*id), + _ => None, + }) + }) + .await + } + + /// 那个请求的尝试链和对话留在哪一家的说明 + async fn routed(&self, model: &str) -> (Vec, Option) { + let id = self.id(model).await; + self.until(&format!(" {model} 的路由"), |evs| { + evs.iter().find_map(|e| match e { + Event::RequestRouted { + id: i, + attempts, + affinity, + .. + } if *i == id => Some((attempts.clone(), affinity.as_ref().and_then(|a| a.stayed))), + _ => None, + }) + }) + .await + } + + async fn finished(&self, model: &str) { + let id = self.id(model).await; + self.until(&format!(" {model} 结束"), |evs| { + evs.iter() + .any(|e| matches!(e, Event::RequestFinished { id: i, .. } if *i == id)) + .then_some(()) + }) + .await + } + + fn health_changed(&self) -> bool { + self.0 + .lock() + .unwrap() + .iter() + .any(|e| matches!(e, Event::HealthChanged { .. })) + } +} + +async fn serve(cfg: tw_config::Config) -> (SocketAddr, tw_gateway::AppState, Log) { + let state = tw_gateway::AppState::new(cfg).unwrap(); + let mut rx = state.bus.subscribe(); + let log = Log(Arc::default()); + let into = log.0.clone(); + tokio::spawn(async move { + while let Ok(ev) = rx.recv().await { + into.lock().unwrap().push(ev); + } + }); + let addr = tw_gateway::serve(state.clone(), ([127, 0, 0, 1], 0).into()) + .await + .unwrap(); + (addr, state, log) +} + +fn client() -> reqwest::Client { + reqwest::Client::builder().no_proxy().build().unwrap() +} + +/// 一个请求。`conversation` 给了的话带着会话头,`messages` 是整段对话 +fn ask( + gw: SocketAddr, + model: &str, + conversation: Option<&str>, + messages: &str, + stream: bool, +) -> reqwest::RequestBuilder { + let body = format!( + r#"{{"model":"{model}","max_tokens":16,"stream":{stream},"system":"你是一个助手","messages":{messages}}}"# + ); + let mut r = client() + .post(format!("http://{gw}/v1/messages")) + .header("x-api-key", "tw-k") + .header("anthropic-version", "2023-06-01") + .body(body); + if let Some(c) = conversation { + r = r.header("x-claude-code-session-id", c); + } + r +} + +fn user(text: &str) -> String { + format!(r#"{{"role":"user","content":"{text}"}}"#) +} + +/// 同一轮里回传的工具结果:这段对话留在上次回答它的那一家 +fn tool_round() -> String { + r#"{"role":"assistant","content":[{"type":"tool_use","id":"t1","name":"Read","input":{}}]},{"role":"user","content":[{"type":"tool_result","tool_use_id":"t1","content":"文件内容"}]}"#.to_string() +} + +/// 起一个占着位置的长请求(模型名 `model`),等到它已经在那家答上了。交回读完它的那个任务 +async fn hold(gw: SocketAddr, log: &Log, model: &str) -> tokio::task::JoinHandle { + let r = ask( + gw, + model, + None, + &format!("[{}]", user(&format!("HOLD {model}"))), + true, + ); + let task = tokio::spawn(async move { r.send().await.unwrap().text().await.unwrap() }); + // 路由事件在开头的内容到了之后才发:那时它已经占着位置在答了 + log.routed(model).await; + task +} + +fn busy(a: &AttemptView) -> bool { + a.outcome == AttemptOutcome::Error && a.skipped == Some(ServeSkip::Busy) +} + +fn served(a: &AttemptView) -> bool { + a.outcome == AttemptOutcome::Served +} + +#[tokio::test] +async fn a_new_conversation_skips_a_full_upstream_and_waiting_is_not_a_failure() { + let (a, b) = (upstream().await, upstream().await); + let (gw, state, log) = serve(config(&a, &b, (Some(1), None), 5)).await; + let held = hold(gw, &log, "占着").await; + assert_eq!(a.hits(), 1); + + // 甲满着:新的对话不等它,当场换乙 + let st = ask(gw, "新的", None, &format!("[{}]", user("你好")), false) + .send() + .await + .unwrap() + .status(); + assert_eq!(st, 200); + let (chain, _) = log.routed("新的").await; + assert_eq!(chain.len(), 2, "{chain:#?}"); + assert_eq!(chain[0].provider, "甲"); + assert!(busy(&chain[0]), "{chain:#?}"); + assert_eq!(chain[0].queued_ms, None, "新的对话不等"); + assert_eq!( + chain[0].error.as_ref().map(|m| m.code.as_str()), + Some("gw.busy_upstream") + ); + assert_eq!(chain[1].provider, "乙"); + assert!(served(&chain[1])); + assert_eq!(a.hits(), 1, "满着的甲不该收到这个请求"); + + // 满着不是失败:甲没有停用,熔断也没动 + for _ in 0..5 { + let r = ask(gw, "又一个", None, &format!("[{}]", user("再来")), false); + assert_eq!(r.send().await.unwrap().status(), 200); + } + assert_eq!(state.health.state("甲"), tw_gateway::health::State::Closed); + assert!(!log.health_changed(), "满着跳过被算成了失败"); + // 按成败分的负载均衡看的成功率也不记它:甲只答过占着的那一个,跳过的六次不算样本 + assert_eq!( + state.health.success_rates(&["甲".to_string()]).get("甲"), + None + ); + + // 占着的那个答完,位置还回来:下一个请求照常去甲 + a.release(); + assert!(held.await.unwrap().contains("message_stop")); + log.finished("占着").await; + // 换一段对话:「你好」那段由乙答过、缓存热着,会留在乙 + let st = ask(gw, "之后", None, &format!("[{}]", user("另一件事")), false) + .send() + .await + .unwrap() + .status(); + assert_eq!(st, 200); + let (chain, _) = log.routed("之后").await; + assert_eq!(chain.len(), 1, "{chain:#?}"); + assert_eq!(chain[0].provider, "甲"); +} + +/// 这段对话第一轮由甲回答,交回第二个请求(同一轮里回传工具结果)的请求体 +async fn conversation_on_a(gw: SocketAddr, log: &Log, model: &str) -> String { + let first = user("帮我改一下"); + let st = ask(gw, model, Some("对话-1"), &format!("[{first}]"), false) + .send() + .await + .unwrap() + .status(); + assert_eq!(st, 200); + let (chain, _) = log.routed(model).await; + assert_eq!(chain.last().unwrap().provider, "甲"); + // 记下由谁回答是在结局里做的 + log.finished(model).await; + format!("[{first},{}]", tool_round()) +} + +#[tokio::test] +async fn a_conversation_that_stays_on_a_full_upstream_waits_for_its_slot() { + let (a, b) = (upstream().await, upstream().await); + let (gw, _state, log) = serve(config(&a, &b, (Some(1), None), 10)).await; + let next = conversation_on_a(gw, &log, "对话").await; + let held = hold(gw, &log, "占着").await; + + // 同一轮的下一个请求:留在甲是为了它的缓存,所以等甲空出来,不去乙 + let r = ask(gw, "对话", Some("对话-1"), &next, false); + let waiting = tokio::spawn(async move { r.send().await.unwrap().status() }); + tokio::time::sleep(Duration::from_millis(300)).await; + assert!(!waiting.is_finished(), "该在等甲的空位"); + assert_eq!(b.hits(), 0, "等的时候去了乙"); + + a.release(); + assert_eq!(waiting.await.unwrap(), 200); + held.await.unwrap(); + let (chain, stayed) = { + // 这段对话的两个请求模型名一样,取后一个的路由 + let id = log + .until("第二个请求", |evs| { + evs.iter() + .filter_map(|e| match e { + Event::RequestStarted { id, model, .. } if model == "对话" => Some(*id), + _ => None, + }) + .nth(1) + }) + .await; + log.until("它的路由", |evs| { + evs.iter().find_map(|e| match e { + Event::RequestRouted { + id: i, + attempts, + affinity, + .. + } if *i == id => Some((attempts.clone(), affinity.as_ref().and_then(|a| a.stayed))), + _ => None, + }) + }) + .await + }; + assert_eq!(stayed, Some(Stay::Turn)); + assert_eq!(chain.len(), 1, "{chain:#?}"); + assert_eq!(chain[0].provider, "甲"); + assert!(served(&chain[0])); + let queued = chain[0].queued_ms.expect("等过的那一跳没记等了多久"); + assert!(queued >= 250, "等了 {queued} 毫秒"); + assert!(chain[0].ms < queued, "等的时间不算在这一跳自己的耗时里"); + assert_eq!(b.hits(), 0); +} + +#[tokio::test] +async fn after_the_wait_the_conversation_moves_on_and_answers_elsewhere() { + let (a, b) = (upstream().await, upstream().await); + let (gw, _state, log) = serve(config(&a, &b, (Some(1), None), 1)).await; + let next = conversation_on_a(gw, &log, "对话").await; + let _held = hold(gw, &log, "占着").await; + + let t = std::time::Instant::now(); + let st = ask(gw, "对话", Some("对话-1"), &next, false) + .send() + .await + .unwrap() + .status(); + assert_eq!(st, 200); + assert!(t.elapsed() >= Duration::from_millis(900), "没等够就走了"); + let id = log + .until("第二个请求", |evs| { + evs.iter() + .filter_map(|e| match e { + Event::RequestStarted { id, model, .. } if model == "对话" => Some(*id), + _ => None, + }) + .nth(1) + }) + .await; + let chain = log + .until("它的路由", |evs| { + evs.iter().find_map(|e| match e { + Event::RequestRouted { + id: i, attempts, .. + } if *i == id => Some(attempts.clone()), + _ => None, + }) + }) + .await; + assert_eq!(chain.len(), 2, "{chain:#?}"); + assert_eq!(chain[0].provider, "甲"); + assert!(busy(&chain[0])); + let queued = chain[0].queued_ms.expect("等过甲,却没记等了多久"); + assert!(queued >= 900, "等了 {queued} 毫秒"); + assert_eq!(chain[1].provider, "乙"); + assert!(served(&chain[1])); + assert_eq!(chain[1].queued_ms, None, "乙有空位,没等"); +} + +#[tokio::test] +async fn when_every_upstream_is_full_the_request_waits_and_is_told_to_come_back() { + let (a, b) = (upstream().await, upstream().await); + let (gw, _state, log) = serve(config(&a, &b, (Some(1), Some(1)), 1)).await; + let _on_a = hold(gw, &log, "占着甲").await; + // 甲满了,这一个跳到乙,把乙也占满 + let _on_b = hold(gw, &log, "占着乙").await; + assert_eq!((a.hits(), b.hits()), (1, 1)); + + let t = std::time::Instant::now(); + let r = ask(gw, "挤不进", None, &format!("[{}]", user("你好")), false) + .send() + .await + .unwrap(); + assert!(t.elapsed() >= Duration::from_millis(900), "没等就拒了"); + assert_eq!(r.status(), 429); + assert_eq!(r.headers()["retry-after"], "5"); + assert_eq!(r.headers()["x-thinkwatch-error"], "rate_limited"); + let json: serde_json::Value = r.json().await.unwrap(); + // 客户端认得的形状:Anthropic 的限流错误 + assert_eq!(json["error"]["type"], "rate_limit_error", "{json}"); + let text = json["error"]["message"].as_str().unwrap(); + assert!(text.contains("`甲`") && text.contains("`乙`"), "{text}"); + + let id = log.id("挤不进").await; + let (source, code) = log + .until("失败", |evs| { + evs.iter().find_map(|e| match e { + Event::RequestFailed { + id: i, + source, + message, + .. + } if *i == id => Some((*source, message.code.clone())), + _ => None, + }) + }) + .await; + assert_eq!(source, FailureSource::RateLimited); + assert_eq!(code, "gw.busy_all"); + let (chain, _) = log.routed("挤不进").await; + assert_eq!(chain.len(), 2, "{chain:#?}"); + assert!(chain.iter().all(busy), "{chain:#?}"); + assert_eq!((a.hits(), b.hits()), (1, 1), "满着的上游一个都不该收到它"); +} + +#[tokio::test] +async fn when_every_upstream_is_full_the_first_to_free_takes_the_request() { + let (a, b) = (upstream().await, upstream().await); + let (gw, _state, log) = serve(config(&a, &b, (Some(1), Some(1)), 10)).await; + let _on_a = hold(gw, &log, "占着甲").await; + let on_b = hold(gw, &log, "占着乙").await; + + let r = ask(gw, "等着", None, &format!("[{}]", user("你好")), false); + let waiting = tokio::spawn(async move { r.send().await.unwrap().status() }); + tokio::time::sleep(Duration::from_millis(300)).await; + assert!(!waiting.is_finished()); + // 乙先空出来 + b.release(); + on_b.await.unwrap(); + assert_eq!(waiting.await.unwrap(), 200); + + let (chain, _) = log.routed("等着").await; + assert_eq!(chain.len(), 3, "{chain:#?}"); + assert!(busy(&chain[0]) && busy(&chain[1]), "{chain:#?}"); + assert_eq!(chain[2].provider, "乙"); + assert!(served(&chain[2])); + assert!(chain[2].queued_ms.is_some_and(|ms| ms >= 250), "{chain:#?}"); + assert_eq!(a.hits(), 1); +} + +#[tokio::test] +async fn a_client_that_leaves_mid_stream_gives_the_slot_back() { + let (a, b) = (upstream().await, upstream().await); + let (gw, _state, log) = serve(config(&a, &b, (Some(1), None), 5)).await; + + // 读到第一段就走(Claude Code 里按一下 Esc)。上游那边还在答 + let resp = ask(gw, "走了", None, &format!("[{}]", user("HOLD 走了")), true) + .send() + .await + .unwrap(); + let mut body = resp.bytes_stream(); + body.next().await.unwrap().unwrap(); + drop(body); + let id = log.id("走了").await; + log.until("取消", |evs| { + evs.iter() + .any(|e| matches!(e, Event::RequestCancelled { id: i, .. } if *i == id)) + .then_some(()) + }) + .await; + + let st = ask(gw, "下一个", None, &format!("[{}]", user("你好")), false) + .send() + .await + .unwrap() + .status(); + assert_eq!(st, 200); + let (chain, _) = log.routed("下一个").await; + assert_eq!(chain.len(), 1, "走掉的那个还占着甲的位置:{chain:#?}"); + assert_eq!(chain[0].provider, "甲"); + assert_eq!(b.hits(), 0); +} + +#[tokio::test] +async fn raising_the_limit_on_reload_lets_a_waiting_request_go() { + let (a, b) = (upstream().await, upstream().await); + let (gw, state, log) = serve(config(&a, &b, (Some(1), Some(1)), 10)).await; + let _on_a = hold(gw, &log, "占着甲").await; + let _on_b = hold(gw, &log, "占着乙").await; + + let r = ask(gw, "等着", None, &format!("[{}]", user("你好")), false); + let waiting = tokio::spawn(async move { r.send().await.unwrap().status() }); + tokio::time::sleep(Duration::from_millis(300)).await; + assert!(!waiting.is_finished()); + // 甲调到 2:等着的那个马上进去 + state + .reload(config(&a, &b, (Some(2), Some(1)), 10)) + .unwrap(); + let st = tokio::time::timeout(Duration::from_secs(2), waiting) + .await + .expect("调大了上限,等着的还在等") + .unwrap(); + assert_eq!(st, 200); + let (chain, _) = log.routed("等着").await; + assert_eq!(chain.last().unwrap().provider, "甲", "{chain:#?}"); +} + +#[tokio::test] +async fn counting_tokens_does_not_wait_for_a_slot() { + let (a, b) = (upstream().await, upstream().await); + let (gw, _state, log) = serve(config(&a, &b, (Some(1), None), 10)).await; + let _held = hold(gw, &log, "占着").await; + + // 甲满着,数 token 照样发给甲:它不跑模型。假的上游没有这个接口(404),由网关自己估 + let st = client() + .post(format!("http://{gw}/v1/messages/count_tokens")) + .header("x-api-key", "tw-k") + .header("anthropic-version", "2023-06-01") + .body(format!( + r#"{{"model":"数","system":"你是一个助手","messages":[{}]}}"#, + user("数一数") + )) + .send() + .await + .unwrap() + .status(); + assert_eq!(st, 200); + let (chain, _) = log.routed("数").await; + assert_eq!(chain.len(), 1, "{chain:#?}"); + assert_eq!(chain[0].provider, "甲"); + assert_eq!(chain[0].outcome, AttemptOutcome::Estimated); + assert_eq!(chain[0].status, Some(404), "该是发给了甲,而不是跳过"); + assert_eq!(chain[0].skipped, None); +} diff --git a/crates/tw-observe/src/bus.rs b/crates/tw-observe/src/bus.rs index 7ea989a2..60fcbbfc 100644 --- a/crates/tw-observe/src/bus.rs +++ b/crates/tw-observe/src/bus.rs @@ -425,6 +425,8 @@ mod tests { error: None, ms: 10, usage: None, + queued_ms: None, + skipped: None, }, tw_api::AttemptView { provider: served_by.into(), @@ -434,6 +436,8 @@ mod tests { error: None, ms: 20, usage: None, + queued_ms: None, + skipped: None, }, ], billing: tw_api::Billing::PerToken, diff --git a/crates/tw-store/src/db.rs b/crates/tw-store/src/db.rs index 2476db7d..ae06ae00 100644 --- a/crates/tw-store/src/db.rs +++ b/crates/tw-store/src/db.rs @@ -2731,6 +2731,8 @@ mod cost_state_tests { error: None, ms: 1, usage: None, + queued_ms: None, + skipped: None, }], ..Default::default() }) diff --git a/crates/tw-store/src/recorder.rs b/crates/tw-store/src/recorder.rs index 5b60b8e0..fa30a15d 100644 --- a/crates/tw-store/src/recorder.rs +++ b/crates/tw-store/src/recorder.rs @@ -848,6 +848,8 @@ mod tests { }), ms: 10_003, usage: None, + queued_ms: None, + skipped: None, }, tw_api::AttemptView { provider: "中转".into(), @@ -857,6 +859,8 @@ mod tests { error: None, ms: 5_042, usage: None, + queued_ms: None, + skipped: None, }, ], billing: tw_api::Billing::PerToken, @@ -902,6 +906,8 @@ mod tests { error: None, ms: 300, usage: None, + queued_ms: None, + skipped: None, }], billing: tw_api::Billing::PerToken, }); @@ -932,6 +938,8 @@ mod tests { error: None, ms: 100, usage: None, + queued_ms: None, + skipped: None, }, tw_api::AttemptView { provider: "中转".into(), @@ -941,6 +949,8 @@ mod tests { error: None, ms: 300, usage: None, + queued_ms: None, + skipped: None, }, ], billing: tw_api::Billing::PerToken, @@ -990,6 +1000,8 @@ mod tests { error: None, ms: 300, usage: None, + queued_ms: None, + skipped: None, }], billing: tw_api::Billing::PerToken, }); @@ -1111,6 +1123,8 @@ mod tests { error: None, ms: 0, usage: None, + queued_ms: None, + skipped: None, }], billing: tw_api::Billing::PerToken, }); @@ -1424,6 +1438,8 @@ mod tests { error: None, ms: 80, usage: None, + queued_ms: None, + skipped: None, }], billing: tw_api::Billing::Free, affinity: None, @@ -1498,6 +1514,8 @@ mod billing_tests { error: None, ms: 5, usage: None, + queued_ms: None, + skipped: None, }], billing, } @@ -1586,6 +1604,8 @@ mod billing_tests { error: None, ms: 40, usage: None, + queued_ms: None, + skipped: None, }], billing, } @@ -1975,6 +1995,8 @@ mod translation_tests { error: None, ms: 1, usage: None, + queued_ms: None, + skipped: None, }) .collect(), billing: tw_api::Billing::PerToken, diff --git a/docs/config.md b/docs/config.md index 7662d324..1114e853 100644 --- a/docs/config.md +++ b/docs/config.md @@ -345,6 +345,7 @@ Upstreams: the APIs requests are forwarded to. | `billing` | `per-token` \| `free` | `per-token` | `per-token`: cost is usage times the price in the upstream's price sheet, subscription accounts included. `free`: cost is recorded as 0. | | `pricing` | string | — | Name of a price sheet under `pricing.sheets`. Unset: the default price table. | | `model_specs` | map of model id → [`providers[].model_specs.*`](#cfg-providers-model_specs) | `{}` | Context window and output limit of single models of this upstream, written by hand, by exact model id. They take precedence over the price table: for models it does not know, or gets wrong. | +| `max_concurrent` | integer | — | Most requests sent to this upstream at the same time, from 1 to 1000. When it is full, a conversation that stays on it waits for a free slot and other requests go to the next upstream; see `failover.slot_wait_secs`. Unset: no limit. | | `disabled` | bool | `false` | Take the upstream out of routing and out of the model list, and keep its configuration. | @@ -368,6 +369,7 @@ providers: proxy: office models_only: [gpt-4.1*, o3] pricing: relay-discount + max_concurrent: 4 - name: local base_url: http://127.0.0.1:11434/v1 @@ -388,6 +390,15 @@ A ChatGPT account upstream (`protocol: chatgpt`) takes only the credential the desktop app obtains by signing in; it cannot be written by hand. Claude and Google subscription sign-ins are not supported; use an API key. +Some relays and accounts accept only a few requests at a time and refuse the +rest. `max_concurrent` keeps the gateway within that number: a request takes a +slot on the upstream when it is sent, and gives it back when the answer has +been passed on in full or the client has gone. When the upstream is full, a +conversation that stays on it to reuse its prompt cache waits for a slot; +any other request goes straight to the next upstream. How long a request +waits is `failover.slot_wait_secs`. Waiting is not a failure: the upstream is +not set aside. Requests that only count tokens do not take a slot. + #### `providers[].oauth` @@ -941,6 +952,14 @@ failover: next_on_slow_start: true ``` +When upstreams are at their `max_concurrent`, a request waits for a free slot +for at most `slot_wait_secs` in all. An ongoing conversation waits for the +upstream it stays on and, if no slot frees in time, moves on to the next one, +where its cache starts over. A new conversation skips a full upstream at once. +When every candidate is full, the request waits for whichever frees first; if +none does, the client gets a 429 with `Retry-After` saying the upstreams are +busy. + @@ -954,6 +973,7 @@ failover: | `rate_limit_max_pause_secs` | integer | `3600` | A rate-limited upstream is set aside for the time its `Retry-After` gives, at most this many seconds. Without `Retry-After` it counts as a failure without a stated reason. | | `stream_start_wait_secs` | integer | `15` | Seconds to hold a streamed answer until its first content arrives. An error before then moves the request to the next upstream; after this long, what has arrived is passed on. From 1 to 120. | | `next_on_slow_start` | bool | `false` | When a streamed answer still has no content `stream_start_wait_secs` after the request was sent, give up on that upstream and send the request to the next one. The last upstream always waits. The upstream given up on is not set aside. Needs `stream_start_wait_secs` of at least 5. | +| `slot_wait_secs` | integer | `30` | Seconds a request waits in total for a free slot on upstreams that are at their `max_concurrent`. After that it goes to the next upstream, or, when every candidate is full, is answered with 429. `0`: never wait. From 0 to 300. | ### `aliases` @@ -1064,8 +1084,9 @@ group shares out requests by the result in the same way as above. last 50 requests within the past 30 minutes. Server errors, rate limits, used-up quota or balance, rejected credentials, timeouts and connection errors count as failures; errors caused by the request itself do not, and - neither does a client that cancels or a switch away from a stream that is - slow to start. An upstream that keeps failing keeps a twentieth of its + neither does a client that cancels, a switch away from a stream that is + slow to start, or an upstream skipped because it is at its + `max_concurrent`. An upstream that keeps failing keeps a twentieth of its weight, so it still gets the occasional new conversation and its recovery is noticed; one that fails outright is set aside by [`failover`](#cfg-failover) as before. diff --git a/docs/config.zh-CN.md b/docs/config.zh-CN.md index f1731269..6cfeb41d 100644 --- a/docs/config.zh-CN.md +++ b/docs/config.zh-CN.md @@ -254,6 +254,7 @@ clients: | `billing` | `per-token` \| `free` | `per-token` | `per-token`:费用为用量乘以所选价目表中的单价,订阅账号同样如此。`free`:费用记为 0。 | | `pricing` | 字符串 | — | `pricing.sheets` 中某张价目表的名字。不写:默认价目表。 | | `model_specs` | 映射: 模型 ID → [`providers[].model_specs.*`](#cfg-providers-model_specs) | `{}` | 手写这家上游某些模型的上下文窗口和输出上限,按模型 ID 完全匹配。写了就优先于价目表,用于价目表里没有或写错的模型。 | +| `max_concurrent` | 整数 | — | 同时发给这家的请求最多几个,取值 1 到 1000。满了的时候,留在这家的对话等空位,别的请求换下一家;等多久见 `failover.slot_wait_secs`。不写:不限。 | | `disabled` | 布尔 | `false` | 不参与路由,模型也不出现在模型列表里;配置原样保留。 | @@ -273,6 +274,7 @@ providers: proxy: office models_only: [gpt-4.1*, o3] pricing: relay-discount + max_concurrent: 4 - name: local base_url: http://127.0.0.1:11434/v1 @@ -284,6 +286,8 @@ providers: ChatGPT 账号上游(`protocol: chatgpt`)只接受桌面应用登录得到的凭据,不能手写。不支持 Claude 和 Google 的订阅登录,请使用 API 密钥。 +有的中转站和账号同时只接受几个请求,多出来的直接拒绝。`max_concurrent` 让网关守住这个数:请求发出时占用这家的一个位置,回答完整交给客户端、或者客户端断开时归还。这家满了的时候,为复用提示缓存而留在这家的对话等空位,别的请求直接换下一家。最多等多久由 `failover.slot_wait_secs` 决定。等待不算失败,这家不会因此停用。只计算 token 数的请求不占位置。 + #### `providers[].oauth` @@ -739,6 +743,11 @@ failover: next_on_slow_start: true ``` +上游的并发数满了(`max_concurrent`)时,一个请求等空位合计最多 `slot_wait_secs` +秒。进行中的对话等它留在的那一家,到时还没有空位就换下一家,缓存在那边从头建; +新的对话遇到满着的上游直接跳过。候选全满时,请求等先空出来的那一家;都没有空出来, +客户端收到 429 和 `Retry-After`,说明上游都忙。 + @@ -752,6 +761,7 @@ failover: | `rate_limit_max_pause_secs` | 整数 | `3600` | 被限流的上游按它给的 `Retry-After` 停用,最多这么多秒。没有 `Retry-After` 的按没有说明原因的失败计。 | | `stream_start_wait_secs` | 整数 | `15` | 流式回答在第一段内容到达前最多暂存的秒数。在此之前上游报错,请求换到下一家;超过这个时间,已收到的部分照常交给客户端。取值 1 到 120。 | | `next_on_slow_start` | 布尔 | `false` | 流式回答在请求发出 `stream_start_wait_secs` 秒后仍没有内容时,放弃这家上游,把请求交给下一家。最后一家总是等下去。被放弃的上游不会停用。开启时 `stream_start_wait_secs` 至少为 5。 | +| `slot_wait_secs` | 整数 | `30` | 上游的并发数满了(`max_concurrent`)时,一个请求等空位合计最多等的秒数。等不到就换下一家;候选全满时回 429。`0`:不等。取值 0 到 300。 | ### `aliases` @@ -819,7 +829,7 @@ groups: - `weights`(默认):只按权重。 - `latency`:越快的上游分得越多。快慢看典型的首字节时间,与 `url-test` 使用同一份测量。比组内居中者快一倍的上游,权重乘以四;最多乘以十,最少乘以十分之一。 -- `health`:越少失败的上游分得越多。依据是最近 30 分钟内的最近 50 次请求:服务器错误、限流、额度或余额用尽、凭据被拒、超时和连接失败算作失败;请求本身导致的错误不算,客户端取消、因开头太慢而换走也不算。经常失败的上游至少保留权重的二十分之一,仍会偶尔分到新对话,以便发现它已经恢复;完全失败的上游照旧由 [`failover`](#cfg-failover) 暂停。 +- `health`:越少失败的上游分得越多。依据是最近 30 分钟内的最近 50 次请求:服务器错误、限流、额度或余额用尽、凭据被拒、超时和连接失败算作失败;请求本身导致的错误不算,客户端取消、因开头太慢而换走、因并发数满了(`max_concurrent`)而跳过也不算。经常失败的上游至少保留权重的二十分之一,仍会偶尔分到新对话,以便发现它已经恢复;完全失败的上游照旧由 [`failover`](#cfg-failover) 暂停。 - `latency-health`:两个系数相乘。 测量还不够的上游按中等对待。与只按权重时一样,进行中的对话留在原来的上游,差额由新对话补齐。 From 54452f4095508a17860f18e8872e754fd6caa2a9 Mon Sep 17 00:00:00 2001 From: fylorn <249551762+fylorn@users.noreply.github.com> Date: Mon, 5 Oct 2026 16:10:35 +0800 Subject: [PATCH 06/22] Hold a gateway key's permit until the answer has been relayed The permit for `clients[].max_concurrent` was a local in pipeline(), so it was returned the moment the response headers went out. A key limited to one request could run any number of streams at once; the limit covered only the wait for headers. The permit now moves into the response body stream next to the in-flight counter and the upstream slot, and is returned when the answer ends or the client goes away. Co-Authored-By: Claude Opus 5.5 --- crates/tw-gateway/src/server/pipeline.rs | 9 ++++-- .../tw-gateway/src/server/pipeline/relay.rs | 9 ++++-- crates/tw-gateway/tests/upstream_slots.rs | 32 +++++++++++++++++++ 3 files changed, 44 insertions(+), 6 deletions(-) diff --git a/crates/tw-gateway/src/server/pipeline.rs b/crates/tw-gateway/src/server/pipeline.rs index 1e756dc1..2424c32e 100644 --- a/crates/tw-gateway/src/server/pipeline.rs +++ b/crates/tw-gateway/src/server/pipeline.rs @@ -117,14 +117,17 @@ pub(super) async fn pipeline( }; // 管线第 3 步:这把密钥自己的并发上限。**等,不拒绝** —— 理由在 - // `crate::limits`。放在路由之后:被规则挡下的请求不用先等一轮 + // `crate::limits`。放在路由之后:被规则挡下的请求不用先等一轮。 + // + // **通行证交给回程,跟着响应体走**(见 `relay`):放在这里的话它在响应头交出去的那一刻 + // 就还了,一条还在流的回答不再算数,上限管的只是等响应头的那一段 let limit = rt .config .clients .iter() .find(|c| c.name == req.client_name) .and_then(|c| c.max_concurrent); - let _pass = state.gate.acquire(&req.client_name, limit).await; + let pass = state.gate.acquire(&req.client_name, limit).await; // 管线第 4 步:内容过滤先下结论,不发事件。删过的话,后面一律用删过的那一份 let screening = screen(&rt, &mut req, &mut reading); @@ -217,7 +220,7 @@ pub(super) async fn pipeline( &reading.facts.model, served, started.id, - live, + (live, pass), ending, reply_plugins, )) diff --git a/crates/tw-gateway/src/server/pipeline/relay.rs b/crates/tw-gateway/src/server/pipeline/relay.rs index 7b6c8a97..b36b6a45 100644 --- a/crates/tw-gateway/src/server/pipeline/relay.rs +++ b/crates/tw-gateway/src/server/pipeline/relay.rs @@ -23,7 +23,8 @@ use crate::state::{AppState, Runtime}; use tw_types::msg; /// `asked_model` 是客户端请求里写的模型名:发出去的不是它、上游答的又是发出去的那个 -/// 模型时,回答里的模型名写回它(见 [`crate::answer_model`])。 +/// 模型时,回答里的模型名写回它(见 [`crate::answer_model`])。`live` 和 `pass` 是这个 +/// 请求在服务中的那一笔和这把密钥的通行证:都跟着响应体走,回答交完或者客户端走掉才还。 #[allow(clippy::too_many_arguments)] pub(super) fn respond( state: &AppState, @@ -33,7 +34,7 @@ pub(super) fn respond( asked_model: &str, served: Served<'_>, id: u64, - live: crate::live::Pass, + (live, pass): (crate::live::Pass, crate::limits::Pass), mut ending: crate::ending::Ending, plugins: Option, ) -> Response { @@ -161,7 +162,9 @@ pub(super) fn respond( // 发完是一种,客户端中途断开、hyper 丢掉响应体是另一种 —— 两种 // 都算这个请求结束了。 let _live = live; - // 在这家占着的位置也一样(见 `crate::slots`):回答交完、客户端走掉,才轮到下一个 + // 这把密钥的通行证(见 `crate::limits`)、在这家占着的位置(见 `crate::slots`)也一样: + // 回答交完、客户端走掉,才轮到下一个 + let _pass = pass; let _slot = slot; // 结局也一样:流被丢掉的时候,它替流报「客户端取消」。 let mut ending = ending; diff --git a/crates/tw-gateway/tests/upstream_slots.rs b/crates/tw-gateway/tests/upstream_slots.rs index 35ba725d..02d73c8c 100644 --- a/crates/tw-gateway/tests/upstream_slots.rs +++ b/crates/tw-gateway/tests/upstream_slots.rs @@ -1,6 +1,9 @@ //! 上游的并发上限(`providers[].max_concurrent`):满着的上游,新的对话当场跳过,留在它 //! 上面的对话等它空出来,都满着时等先空出来的那一家、等不到回 429 —— 走真实的管线,看 //! 尝试链里说的和实际去的那一家。 +//! +//! 每把密钥的上限(`clients[].max_concurrent`)也在这里:它和上游的位置一样,占到回答 +//! 交完为止。 use std::net::SocketAddr; use std::sync::atomic::{AtomicUsize, Ordering}; @@ -584,3 +587,32 @@ async fn counting_tokens_does_not_wait_for_a_slot() { assert_eq!(chain[0].status, Some(404), "该是发给了甲,而不是跳过"); assert_eq!(chain[0].skipped, None); } + +/// 一把密钥的上限管的是整个回答:流还在走,它就还占着那一份。以前通行证在响应头交出去 +/// 时就还了,上限 1 的密钥照样能同时跑好几条流 +#[tokio::test] +async fn a_keys_limit_holds_until_the_streamed_answer_ends() { + let (a, b) = (upstream().await, upstream().await); + let mut cfg = config(&a, &b, (None, None), 10); + cfg.clients[0].max_concurrent = Some(1); + let (gw, _state, log) = serve(cfg).await; + let held = hold(gw, &log, "占着").await; + + // 同一把密钥的第二个请求:等前一条流走完 + let r = ask(gw, "第二个", None, &format!("[{}]", user("你好")), false); + let waiting = tokio::spawn(async move { r.send().await.unwrap().status() }); + tokio::time::sleep(Duration::from_millis(300)).await; + assert!( + !waiting.is_finished(), + "前一条流还在走,这把密钥已经到上限了" + ); + assert_eq!(a.hits() + b.hits(), 1, "第二个不该发出去"); + + a.release(); + assert!(held.await.unwrap().contains("message_stop")); + let st = tokio::time::timeout(Duration::from_secs(2), waiting) + .await + .expect("前一条流走完了,第二个还在等") + .unwrap(); + assert_eq!(st, 200); +} From 3e7841e2120134f883d90810f70aa9a1ad33a57e Mon Sep 17 00:00:00 2001 From: fylorn <249551762+fylorn@users.noreply.github.com> Date: Mon, 5 Oct 2026 16:40:23 +0800 Subject: [PATCH 07/22] Usage limits per gateway key: requests, tokens and cost per minute, hour, day, week or month A key on the LAN, or one handed to a script, can now be capped by what it spends rather than only by how many requests run at once. Limits live on the key (`clients[].limits`), each entry counts one measure, and every entry has to pass. - Day, week and month follow the calendar in the core machine's local time zone. Used up means refused until the next period: 429 with Retry-After to the reset and x-should-retry: false, and code insufficient_quota in the OpenAI shapes, so the official SDKs and Codex stop instead of retrying. - Minute and hour are rolling. Used up means waiting for the next free slot when it frees within failover.slot_wait_secs (the same setting as the wait for an upstream's free slot, counted separately), otherwise 429 with exact Retry-After and retry-after-ms. The 429 for upstreams that are all at their max_concurrent now sets its Retry-After through the same field, so it also carries retry-after-ms and x-should-retry: true. - Admission is one step after routing: calendar limits, then the existing max_concurrent gate (its permit still travels into the response body, so it is held until the answer ends), then rolling limits. A refused request is recorded like a routing refusal (start, empty attempt chain, failure with its own code). - Running requests hold their input estimate (and its input cost on the first candidate); the recorder settles each row to the usage and cost it writes, through a direct hook rather than the bus, so the in-memory totals equal what a restart adds back from the store. Token counts and local answers do not count; unpriced models and free upstreams count $0. - A restart rebuilds this day, week and month from the request records, so a monthly limit requires retention.row_days >= 31. - The key view carries each limit with used/max/resets_at/reached, and the models without a price for keys that have a cost limit; the key input takes the limits; renaming a key through the control plane carries its usage. - KeyLimitAlert on the event bus once per period at 80% and at the limit, for Lite's system notifications. - WebSocket connections are admitted once when they open: the store records one row per connection without usage. Co-Authored-By: Claude Opus 5.5 --- bin/twcore/src/main.rs | 12 +- crates/tw-api/msg-codes.txt | 12 + crates/tw-api/src/lib.rs | 91 ++ crates/tw-config/src/lib.rs | 7 + crates/tw-config/src/limits.rs | 193 ++++ crates/tw-config/src/validate.rs | 223 ++++ crates/tw-config/tests/manual/schema.rs | 63 ++ crates/tw-control/src/key_limits.rs | 65 ++ crates/tw-control/src/keys.rs | 32 + crates/tw-control/src/lib.rs | 1 + crates/tw-control/tests/key_limits.rs | 385 +++++++ crates/tw-gateway/src/error.rs | 116 +- crates/tw-gateway/src/key_limits/clock.rs | 189 ++++ crates/tw-gateway/src/key_limits/mod.rs | 988 ++++++++++++++++++ crates/tw-gateway/src/key_limits/tests.rs | 539 ++++++++++ crates/tw-gateway/src/lib.rs | 1 + crates/tw-gateway/src/server/pipeline.rs | 29 +- .../src/server/pipeline/admission.rs | 146 +++ crates/tw-gateway/src/server/upgrade.rs | 31 + crates/tw-gateway/src/state.rs | 21 +- crates/tw-gateway/tests/key_limits.rs | 260 +++++ crates/tw-observe/src/bus.rs | 15 + crates/tw-store/src/db.rs | 126 +++ crates/tw-store/src/lib.rs | 6 +- crates/tw-store/src/recorder.rs | 193 ++++ docs/config.md | 45 + docs/config.zh-CN.md | 39 + 27 files changed, 3804 insertions(+), 24 deletions(-) create mode 100644 crates/tw-config/src/limits.rs create mode 100644 crates/tw-control/src/key_limits.rs create mode 100644 crates/tw-control/tests/key_limits.rs create mode 100644 crates/tw-gateway/src/key_limits/clock.rs create mode 100644 crates/tw-gateway/src/key_limits/mod.rs create mode 100644 crates/tw-gateway/src/key_limits/tests.rs create mode 100644 crates/tw-gateway/src/server/pipeline/admission.rs create mode 100644 crates/tw-gateway/tests/key_limits.rs diff --git a/bin/twcore/src/main.rs b/bin/twcore/src/main.rs index eecb44f0..8f3e65bd 100644 --- a/bin/twcore/src/main.rs +++ b/bin/twcore/src/main.rs @@ -820,10 +820,14 @@ fn cmd_serve(path: &Path, port: Option, safe: bool, parent: Option) -> state.pricing.clone(), body_rx, run_rx, + tw_control::key_limits::settle_hook(&state), ); - if store.is_some() { + if let Some(rec) = &store { state.set_body_sink(body_tx); state.set_plugin_sink(run_tx); + // 密钥的用量上限:这一天、这一周、这个月已经用了多少,从记录里加回来。**在第一个 + // 请求之前**(网关还没开始听) + tw_control::key_limits::rebuild(&state, rec.lock().await.db()); } /* @@ -1004,6 +1008,8 @@ fn build_store( pricing: tw_pricing::Shared, bodies: tokio::sync::mpsc::Receiver, runs: tokio::sync::mpsc::Receiver, + // 每记下一行请求,交给网关的密钥用量上限结算(见 `tw_control::key_limits`) + settled: tw_store::SettleHook, ) -> Option>> { let events = bus.subscribe(); let (db, blobs) = match tw_store::open(dir) { @@ -1078,7 +1084,9 @@ fn build_store( let recorder = tw_store::task::spawn( // 算完价钱往回报一条 —— 见 `Event::RequestPriced`。这里是唯一 // 同时看得见总线和存储层的地方,所以接线在这儿完成。 - tw_store::Recorder::new(db, blobs, pricing).reporting_to(bus), + tw_store::Recorder::new(db, blobs, pricing) + .reporting_to(bus) + .settling_to(settled), events, rx, ); diff --git a/crates/tw-api/msg-codes.txt b/crates/tw-api/msg-codes.txt index fd30156d..19de9be7 100644 --- a/crates/tw-api/msg-codes.txt +++ b/crates/tw-api/msg-codes.txt @@ -67,6 +67,12 @@ config.edit.unwritable config.empty_key config.empty_models_only config.failover_range +config.key_limit_cache_reads +config.key_limit_duplicate +config.key_limit_empty +config.key_limit_month_retention +config.key_limit_not_positive +config.key_limit_two_measures config.model_spec_blank_model config.model_spec_empty config.model_spec_wildcard @@ -281,6 +287,12 @@ gw.convert.tools_unsendable gw.count_tokens.bedrock_upstream gw.files.unsupported gw.internal +gw.key_limit.cost_per_period +gw.key_limit.cost_rolling +gw.key_limit.requests_per_period +gw.key_limit.requests_rolling +gw.key_limit.tokens_per_period +gw.key_limit.tokens_rolling gw.listen.addr_unavailable gw.listen.bind_failed gw.listen.denied diff --git a/crates/tw-api/src/lib.rs b/crates/tw-api/src/lib.rs index 53aeb0d8..5b0d3caa 100644 --- a/crates/tw-api/src/lib.rs +++ b/crates/tw-api/src/lib.rs @@ -1332,6 +1332,28 @@ pub enum Event { resets_at_ms: Option, at_ms: u64, }, + /// 一把网关密钥这一期(天、周、月)的用量到了一条上限的八成,或者到了上限。 + /// + /// **每一期、每一档只报一次**:同一期里之后的请求照样被拒,同一句话说第二遍只会 + /// 让人学会忽略通知。下一期重新算;上限改了也重新算。重启之后从请求记录里加回来 + /// 时已经过了的档不再报。滚动的上限(分钟、小时)不报:它们几十秒就过去 + KeyLimitAlert { + id: u64, + /// 密钥的名字 + key: String, + per: LimitPer, + measure: LimitMeasure, + /// 上限,单位同 `KeyLimitView::max` + max: u64, + /// 报的时候用了多少,同上 + used: u64, + #[serde(default, skip_serializing_if = "std::ops::Not::not")] + cache_reads: bool, + /// `true` = 到了上限,之后的请求被拒到 `resets_at_ms`;`false` = 到了八成 + reached: bool, + resets_at_ms: u64, + at_ms: u64, + }, /// 数据面换了监听地址,或者没换成(旧的还在服务)。 /// /// **和 `ConfigReloaded` 是两件事。**配置换进去之后监听器才开始换, @@ -1632,6 +1654,7 @@ impl Event { | Event::ConfigRejected { id, .. } | Event::QuotaSeen { id, .. } | Event::QuotaExhausted { id, .. } + | Event::KeyLimitAlert { id, .. } | Event::SecretsFound { id, .. } | Event::ContentMatched { id, .. } | Event::RequestPriced { id, .. } @@ -2228,6 +2251,71 @@ pub struct ClientView { /// 那个可以伪造。从来没被用过时没有 #[serde(default, skip_serializing_if = "Option::is_none")] pub last_seen_ms: Option, + /// 用量上限,按配置里的顺序,各带此刻用了多少。没设的是空的 + pub limits: Vec, + /// 这把密钥用得到、却没有价格的模型。**只有设了费用上限的密钥才算**:这些模型的 + /// 请求费用记 0,费用上限管不住它们,对话框里要提醒一句。没有就不带 + #[serde(default, skip_serializing_if = "Vec::is_empty")] + pub unpriced_models: Vec, +} + +slug_enum! { + /// 用量上限按多长一段时间算。分钟、小时是**滚动的**(最近 60 秒、最近 60 分钟); + /// 天、周、月是**自然的**,按 core 所在机器的本地时区:零点、周一零点、一号零点重新算。 + pub enum LimitPer { + Minute = "minute", + Hour = "hour", + Day = "day", + Week = "week", + Month = "month", + } +} + +slug_enum! { + /// 一条用量上限数的是什么。 + pub enum LimitMeasure { + /// 请求数。数 token 的请求、网关自己答的不算 + Requests = "requests", + /// token:没走缓存的输入 + 写进缓存的 + 输出,`cache_reads` 时再加上从缓存读的 + Tokens = "tokens", + /// 费用,**微分**(百万分之一美元),和别处的费用同一个单位。没有价格的模型、 + /// 不计费的上游算 0 + Cost = "cost", + } +} + +/// 一条用量上限,和它此刻用了多少。 +#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)] +#[cfg_attr(feature = "ts", derive(ts_rs::TS))] +pub struct KeyLimitView { + pub per: LimitPer, + pub measure: LimitMeasure, + /// 上限:请求数、token 数,费用是微分 + pub max: u64, + /// token 上限把从缓存读的也算进去 + pub cache_reads: bool, + /// 用了多少,单位同 `max`。**在跑的请求也算**:按它们的输入估算占着,结束时换成 + /// 记下的实数 —— 准入看的就是这个数。滚动的是最近那一段时间里的,重启之后从空的 + /// 开始;自然的是这一期的,重启之后从请求记录里加回来 + pub used: u64, + /// 这一期什么时候结束、重新算。只有天、周、月有 + #[serde(default, skip_serializing_if = "Option::is_none")] + pub resets_at_ms: Option, + /// 到了:`used` 不小于 `max`,新的请求此刻会被拒(滚动的会先等一会儿) + pub reached: bool, +} + +/// 新建、保存密钥时的一条用量上限。 +#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)] +#[cfg_attr(feature = "ts", derive(ts_rs::TS))] +pub struct KeyLimitInput { + pub per: LimitPer, + pub measure: LimitMeasure, + /// 上限:请求数、token 数,费用是微分。要大于 0 + pub max: u64, + /// 只有 token 上限能开 + #[serde(default, skip_serializing_if = "std::ops::Not::not")] + pub cache_reads: bool, } /// 新建或保存一把网关密钥(`POST /keys`、`PUT /keys/{name}`)。 @@ -2255,6 +2343,9 @@ pub struct KeyInput { pub allow: Option>, #[serde(default, skip_serializing_if = "std::ops::Not::not")] pub disabled: bool, + /// 用量上限,整份替换。不带 = 一条都没有 + #[serde(default, skip_serializing_if = "Vec::is_empty")] + pub limits: Vec, } /// 换哪把密钥(`POST /keys/{name}/rotate`)。 diff --git a/crates/tw-config/src/lib.rs b/crates/tw-config/src/lib.rs index 527aad57..2b898f9f 100644 --- a/crates/tw-config/src/lib.rs +++ b/crates/tw-config/src/lib.rs @@ -15,6 +15,7 @@ pub mod edit; mod failover; pub mod history; mod init; +pub mod limits; pub mod model_specs; pub mod nics; pub mod plugins; @@ -35,6 +36,7 @@ mod wire; pub use aliases::{Alias, Aliases}; pub use credential::{CredentialError, Header, Headers, Secret, SecretResolveError, auth_header}; pub use init::{generate_control_key, generate_initial, generate_key}; +pub use limits::{KeyLimit, LimitMeasure, LimitPer}; pub use model_specs::{ModelLimits, ModelSpec, Sourced, SpecSource}; pub use plugins::Plugin; pub use proxy::{DIRECT, OnProxyFail, Proxy, ProxyKind, SYSTEM}; @@ -166,6 +168,7 @@ impl Default for Client { route: None, client: None, disabled: false, + limits: Vec::new(), } } } @@ -659,6 +662,10 @@ pub struct Client { /// 也能一键恢复的状态。 #[serde(default, skip_serializing_if = "std::ops::Not::not")] pub disabled: bool, + /// 用量上限:每分钟、每小时、每天、每周、每月最多多少个请求、多少 token、花多少钱。 + /// **每一条都要过**。不写就不限(见 [`limits`]) + #[serde(default, skip_serializing_if = "Vec::is_empty")] + pub limits: Vec, } /// OAuth 凭据(第 3 类)。 diff --git a/crates/tw-config/src/limits.rs b/crates/tw-config/src/limits.rs new file mode 100644 index 00000000..9ae5492e --- /dev/null +++ b/crates/tw-config/src/limits.rs @@ -0,0 +1,193 @@ +//! 一把网关密钥的用量上限:每分钟、每小时、每天、每周、每月最多多少个请求、多少 token、 +//! 花多少钱。 +//! +//! **上限挂在密钥上**,不挂在上游上:要管住的是「这台机器上的这个脚本」,不是「这家 +//! 上游」。一把密钥可以有好几条,**每一条都要过**。 +//! +//! 怎么数、怎么等、到了怎么拒,在 `tw_gateway::key_limits`;这里只有写法。 + +use serde::{Deserialize, Serialize}; + +/// 一段时间。分钟、小时是**滚动的**(最近 60 秒、最近 60 分钟);天、周、月是**自然的**, +/// 按 core 所在机器的本地时区:零点、周一零点、一号零点重新算。 +#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash, Serialize, Deserialize)] +#[serde(rename_all = "lowercase")] +pub enum LimitPer { + Minute, + Hour, + Day, + Week, + Month, +} + +impl LimitPer { + /// 写在配置里的那个词。报错、消息参数里用它 + pub fn word(&self) -> &'static str { + match self { + LimitPer::Minute => "minute", + LimitPer::Hour => "hour", + LimitPer::Day => "day", + LimitPer::Week => "week", + LimitPer::Month => "month", + } + } + + /// 滚动的那两种:窗口多长,毫秒。自然周期是 None + pub fn rolling_ms(&self) -> Option { + match self { + LimitPer::Minute => Some(60_000), + LimitPer::Hour => Some(3_600_000), + _ => None, + } + } +} + +/// 一条上限数的是什么。 +#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)] +pub enum LimitMeasure { + Requests, + Tokens, + Cost, +} + +impl LimitMeasure { + pub fn word(&self) -> &'static str { + match self { + LimitMeasure::Requests => "requests", + LimitMeasure::Tokens => "tokens", + LimitMeasure::Cost => "cost", + } + } +} + +/// 一条上限:`{ per: day, cost: 5 }`。**一条只写一种量**(请求数、token、费用三选一), +/// 校验时查。 +/// +/// 三种量写成三个可选字段而不是 `{ measure: cost, max: 5 }`:手写的人读 +/// `requests: 30` 不用再对一遍哪个数是什么单位。 +#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)] +#[serde(deny_unknown_fields)] +pub struct KeyLimit { + pub per: LimitPer, + /// 请求数。**写成有符号的**:写了负数时报的是「要大于 0」,不是一句解析器的话 + #[serde(default, skip_serializing_if = "Option::is_none")] + pub requests: Option, + /// token 数:没走缓存的输入 + 写进缓存的 + 输出,`cache_reads` 开着时再加上从缓存读的 + #[serde(default, skip_serializing_if = "Option::is_none")] + pub tokens: Option, + /// 费用,美元(core 记的费用都是美元)。没有价格的模型、不计费的上游算 0 + #[serde(default, skip_serializing_if = "Option::is_none")] + pub cost: Option, + /// token 上限把从缓存读的也算进去。**只有 token 上限能写**。缓存读的单价低、量大: + /// 一段几十万 token 的对话每一轮都整段读一遍,算进去的话上限很快就到 + #[serde(default, skip_serializing_if = "std::ops::Not::not")] + pub cache_reads: bool, +} + +impl KeyLimit { + /// 写了的那几种量。校验要求恰好一种 + pub fn measures(&self) -> Vec { + let mut out = Vec::new(); + if self.requests.is_some() { + out.push(LimitMeasure::Requests); + } + if self.tokens.is_some() { + out.push(LimitMeasure::Tokens); + } + if self.cost.is_some() { + out.push(LimitMeasure::Cost); + } + out + } + + /// 这一条数什么。**过了校验的才有意义**:一种都没写、写了两种时是写在前面的那种 + pub fn measure(&self) -> LimitMeasure { + self.measures() + .first() + .copied() + .unwrap_or(LimitMeasure::Requests) + } + + /// 上限,整数:请求数、token 数,费用是微分(百万分之一美元,和记账同一个单位) + pub fn max(&self) -> i64 { + match self.measure() { + LimitMeasure::Requests => self.requests.unwrap_or(0), + LimitMeasure::Tokens => self.tokens.unwrap_or(0), + LimitMeasure::Cost => self.cost.map(tw_pricing::to_micros).unwrap_or(0), + } + } + + /// 两条算不算同一条:同一段时间、同一种量、缓存读算不算进去也一样 + pub fn same_as(&self, other: &KeyLimit) -> bool { + self.per == other.per + && self.measure() == other.measure() + && self.cache_reads == other.cache_reads + } +} + +impl From for tw_api::LimitPer { + fn from(p: LimitPer) -> Self { + match p { + LimitPer::Minute => tw_api::LimitPer::Minute, + LimitPer::Hour => tw_api::LimitPer::Hour, + LimitPer::Day => tw_api::LimitPer::Day, + LimitPer::Week => tw_api::LimitPer::Week, + LimitPer::Month => tw_api::LimitPer::Month, + } + } +} + +impl From for LimitPer { + fn from(p: tw_api::LimitPer) -> Self { + match p { + tw_api::LimitPer::Minute => LimitPer::Minute, + tw_api::LimitPer::Hour => LimitPer::Hour, + tw_api::LimitPer::Day => LimitPer::Day, + tw_api::LimitPer::Week => LimitPer::Week, + tw_api::LimitPer::Month => LimitPer::Month, + } + } +} + +impl From for tw_api::LimitMeasure { + fn from(m: LimitMeasure) -> Self { + match m { + LimitMeasure::Requests => tw_api::LimitMeasure::Requests, + LimitMeasure::Tokens => tw_api::LimitMeasure::Tokens, + LimitMeasure::Cost => tw_api::LimitMeasure::Cost, + } + } +} + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn each_measure_reads_back_as_written_and_cost_is_in_micros() { + let l: KeyLimit = serde_yaml_ng::from_str("{per: day, cost: 5}").unwrap(); + assert_eq!((l.measure(), l.max()), (LimitMeasure::Cost, 5_000_000)); + let l: KeyLimit = serde_yaml_ng::from_str("{per: minute, requests: 30}").unwrap(); + assert_eq!((l.measure(), l.max()), (LimitMeasure::Requests, 30)); + let l: KeyLimit = + serde_yaml_ng::from_str("{per: week, tokens: 1000000, cache_reads: true}").unwrap(); + assert_eq!((l.measure(), l.max()), (LimitMeasure::Tokens, 1_000_000)); + assert!(l.cache_reads); + // 默认值不写回去 + let out = serde_yaml_ng::to_string(&KeyLimit { + per: LimitPer::Hour, + requests: Some(3), + tokens: None, + cost: None, + cache_reads: false, + }) + .unwrap(); + assert_eq!(out.trim(), "per: hour\nrequests: 3"); + } + + #[test] + fn a_typo_is_refused_rather_than_ignored() { + assert!(serde_yaml_ng::from_str::("{per: day, request: 3}").is_err()); + assert!(serde_yaml_ng::from_str::("{per: daily, requests: 3}").is_err()); + } +} diff --git a/crates/tw-config/src/validate.rs b/crates/tw-config/src/validate.rs index 6474c953..db69aa5c 100644 --- a/crates/tw-config/src/validate.rs +++ b/crates/tw-config/src/validate.rs @@ -116,6 +116,27 @@ pub enum ValidationError { model: String, field: &'static str, }, + #[error("{}", self.msg())] + KeyLimitEmpty { key: String }, + #[error("{}", self.msg())] + KeyLimitTwoMeasures { key: String, measures: String }, + #[error("{}", self.msg())] + KeyLimitNotPositive { + key: String, + per: &'static str, + measure: &'static str, + value: String, + }, + #[error("{}", self.msg())] + KeyLimitDuplicate { + key: String, + per: &'static str, + measure: &'static str, + }, + #[error("{}", self.msg())] + KeyLimitCacheReads { key: String, measure: &'static str }, + #[error("{}", self.msg())] + KeyLimitMonthRetention { key: String, days: u64 }, } impl ValidationError { @@ -339,6 +360,42 @@ impl ValidationError { "the model spec `{model}` of upstream `{upstream}` has {field}: 0; it has to be a \ number of tokens above 0. Leave it out to use the price table" ), + KeyLimitEmpty { key } => msg!( + "config.key_limit_empty", key = key => + "a limit of gateway key `{key}` names nothing to count. Each entry under limits \ + takes one of requests, tokens or cost" + ), + KeyLimitTwoMeasures { key, measures } => msg!( + "config.key_limit_two_measures", key = key, measures = measures => + "a limit of gateway key `{key}` names {measures} together. Each entry takes one \ + of requests, tokens or cost; write one entry for each" + ), + KeyLimitNotPositive { + key, + per, + measure, + value, + } => msg!( + "config.key_limit_not_positive", key = key, per = per, measure = measure, + value = value => + "the {measure} limit per {per} of gateway key `{key}` is {value}; it has to be \ + more than 0" + ), + KeyLimitDuplicate { key, per, measure } => msg!( + "config.key_limit_duplicate", key = key, per = per, measure = measure => + "gateway key `{key}` has two {measure} limits per {per}. Keep one of them" + ), + KeyLimitCacheReads { key, measure } => msg!( + "config.key_limit_cache_reads", key = key, measure = measure => + "the {measure} limit of gateway key `{key}` sets cache_reads, which only a \ + tokens limit takes" + ), + KeyLimitMonthRetention { key, days } => msg!( + "config.key_limit_month_retention", key = key, days = days => + "gateway key `{key}` has a limit per month, and retention.row_days is {days}. \ + After a restart the month's total is added up again from the request records, \ + so row_days has to be at least 31" + ), } } } @@ -430,6 +487,7 @@ pub fn validate(cfg: &Config) -> Result<(), ValidationError> { name: c.name.clone(), }); } + check_limits(c, cfg.retention.row_days)?; if let Some(prev) = keys.insert(&c.key, &c.name) { return Err(ValidationError::DuplicateKey( prev.to_string(), @@ -578,6 +636,75 @@ pub fn validate(cfg: &Config) -> Result<(), ValidationError> { Ok(()) } +/// 一个月的记录最少要留几天:月度上限的用量在重启之后从请求记录里重新加起来, +/// 留得比一个月短,月初那几天的就加不回来 +pub const MONTH_ROW_DAYS: u64 = 31; + +/// 一把密钥的用量上限写得对不对。 +/// +/// **每一条恰好数一种量**:一条里写两种,读的人分不清是「都要过」还是「过一个就行」。 +/// 同一段时间、同一种量写两条(缓存读算不算进去也一样)的,两条里总有一条不起作用, +/// 而用户以为它在起作用。 +fn check_limits(c: &crate::Client, row_days: u64) -> Result<(), ValidationError> { + let key = || c.name.clone(); + for (i, l) in c.limits.iter().enumerate() { + let measures = l.measures(); + let measure = match measures.as_slice() { + [] => return Err(ValidationError::KeyLimitEmpty { key: key() }), + [one] => *one, + many => { + return Err(ValidationError::KeyLimitTwoMeasures { + key: key(), + measures: many + .iter() + .map(|m| m.word()) + .collect::>() + .join(" and "), + }); + } + }; + // 写了 `.nan`、`.inf` 的费用也在这里挡:前者比什么都不大,后者换不成微分 + let positive = match measure { + crate::LimitMeasure::Requests => l.requests.is_some_and(|n| n > 0), + crate::LimitMeasure::Tokens => l.tokens.is_some_and(|n| n > 0), + crate::LimitMeasure::Cost => l.cost.is_some_and(|x| x.is_finite() && x > 0.0), + }; + if !positive { + let value = match measure { + crate::LimitMeasure::Requests => l.requests.unwrap_or_default().to_string(), + crate::LimitMeasure::Tokens => l.tokens.unwrap_or_default().to_string(), + crate::LimitMeasure::Cost => l.cost.unwrap_or_default().to_string(), + }; + return Err(ValidationError::KeyLimitNotPositive { + key: key(), + per: l.per.word(), + measure: measure.word(), + value, + }); + } + if l.cache_reads && measure != crate::LimitMeasure::Tokens { + return Err(ValidationError::KeyLimitCacheReads { + key: key(), + measure: measure.word(), + }); + } + if c.limits[..i].iter().any(|o| o.same_as(l)) { + return Err(ValidationError::KeyLimitDuplicate { + key: key(), + per: l.per.word(), + measure: measure.word(), + }); + } + if l.per == crate::LimitPer::Month && row_days < MONTH_ROW_DAYS { + return Err(ValidationError::KeyLimitMonthRetention { + key: key(), + days: row_days, + }); + } + } + Ok(()) +} + /// 别名表写得对不对,和整份配置的校验是同一套(界面预览一个别名时也用它)。 /// **模型名在不在上游清单里不在这里查**:清单是运行时问来的, /// 上游一时没列出来不该让整份配置不收。 @@ -1360,6 +1487,78 @@ groups: ); } } + + /// 用量上限:每一条恰好一种量、大于 0、`cache_reads` 只给 token、不重复;月度上限要 + /// 请求记录留够一个月。 + #[test] + fn key_limits_are_checked_one_entry_at_a_time() { + let with = |limits: &str, row_days: u64| { + let mut key = c("k", "tw-1"); + key.limits = serde_yaml_ng::from_str(limits).unwrap(); + let mut cfg = cfg(vec![key], vec![]); + cfg.retention.row_days = row_days; + validate(&cfg).map_err(|e| e.msg()) + }; + let code = |limits: &str| with(limits, 90).unwrap_err().code; + assert!( + with( + "[{per: minute, requests: 30}, {per: day, cost: 5.5}, \ + {per: day, tokens: 100000}, {per: day, tokens: 900000, cache_reads: true}, \ + {per: month, cost: 100}]", + 90 + ) + .is_ok(), + "缓存读算不算进去不一样,就是两条" + ); + assert_eq!(code("[{per: day}]"), "config.key_limit_empty"); + let m = with("[{per: day, requests: 3, cost: 1}]", 90).unwrap_err(); + assert_eq!(m.code, "config.key_limit_two_measures"); + assert_eq!(m.arg("measures"), "requests and cost"); + for bad in [ + "[{per: day, requests: 0}]", + "[{per: day, tokens: -5}]", + "[{per: day, cost: 0}]", + "[{per: day, cost: -1.5}]", + "[{per: day, cost: .nan}]", + "[{per: day, cost: .inf}]", + ] { + assert_eq!(code(bad), "config.key_limit_not_positive", "{bad}"); + } + let m = with("[{per: hour, cost: -1.5}]", 90).unwrap_err(); + assert_eq!( + (m.arg("per"), m.arg("measure"), m.arg("value")), + ("hour", "cost", "-1.5") + ); + assert_eq!( + code("[{per: day, requests: 3, cache_reads: true}]"), + "config.key_limit_cache_reads" + ); + assert_eq!( + code("[{per: day, cost: 3}, {per: day, cost: 5}]"), + "config.key_limit_duplicate" + ); + assert_eq!( + code( + "[{per: week, tokens: 3, cache_reads: true}, {per: week, tokens: 5, cache_reads: true}]" + ), + "config.key_limit_duplicate" + ); + // 月度上限:记录要留够 31 天,用量在重启之后从记录里加回来 + let m = with("[{per: month, requests: 3}]", 30).unwrap_err(); + assert_eq!(m.code, "config.key_limit_month_retention"); + assert_eq!(m.arg("days"), "30"); + assert!(with("[{per: month, requests: 3}]", 31).is_ok()); + assert!( + with("[{per: week, requests: 3}]", 7).is_ok(), + "周以内的不受影响" + ); + } + + #[test] + fn a_key_without_limits_writes_nothing_back() { + let out = serde_yaml_ng::to_string(&c("k", "tw-1")).unwrap(); + assert!(!out.contains("limits"), "{out}"); + } } #[cfg(test)] @@ -1567,6 +1766,30 @@ mod msg_codes { model: "m".into(), field: "context_window", }, + KeyLimitEmpty { key: "k".into() }, + KeyLimitTwoMeasures { + key: "k".into(), + measures: "requests and cost".into(), + }, + KeyLimitNotPositive { + key: "k".into(), + per: "day", + measure: "cost", + value: "0".into(), + }, + KeyLimitDuplicate { + key: "k".into(), + per: "day", + measure: "cost", + }, + KeyLimitCacheReads { + key: "k".into(), + measure: "cost", + }, + KeyLimitMonthRetention { + key: "k".into(), + days: 30, + }, ]; check( "config.", diff --git a/crates/tw-config/tests/manual/schema.rs b/crates/tw-config/tests/manual/schema.rs index ed3b8482..8a0ef785 100644 --- a/crates/tw-config/tests/manual/schema.rs +++ b/crates/tw-config/tests/manual/schema.rs @@ -61,6 +61,9 @@ fn group_types() -> Vec<&'static str> { fn balance_bys() -> Vec<&'static str> { super::fields::() } +fn limit_pers() -> Vec<&'static str> { + super::fields::() +} const RULE_ID: T2 = t("built-in rule id", "内置规则 id"); const MODE_DOC: T2 = t( @@ -428,6 +431,66 @@ pub fn sections() -> Vec
{ "拒绝使用这把密钥的所有请求,密钥本身保留。", ), ), + row( + "limits", + Kind::Objs("clients[].limits[]"), + Def::Unset, + t( + "Usage limits: requests, tokens or cost per minute, hour, day, week or month. A request has to pass every one. Unset: no limit.", + "用量上限:每分钟、每小时、每天、每周或每月的请求数、token 数或费用。请求要通过每一条。不写:不限。", + ), + ), + ], + }, + Section { + path: "clients[].limits[]", + ty: checked!(KeyLimit, "{per: day}"), + rows: vec![ + row( + "per", + Kind::Enum(limit_pers), + Def::Required, + t( + "The period. `minute` and `hour` are rolling (the last 60 seconds, the last 60 minutes); `day`, `week` and `month` start again at local midnight, on Monday and on the 1st.", + "按多长一段时间算。`minute`、`hour` 是滚动的(最近 60 秒、最近 60 分钟);`day`、`week`、`month` 在本地时间的零点、周一零点、每月一号零点重新算。", + ), + ), + row( + "requests", + Kind::Int, + Def::Unset, + t( + "At most this many requests. Token counts and answers the gateway gives itself do not count.", + "最多这么多个请求。数 token 的请求和网关自己答的不算。", + ), + ), + row( + "tokens", + Kind::Int, + Def::Unset, + t( + "At most this many tokens: uncached input, cache writes and output.", + "最多这么多 token:未命中缓存的输入、写入缓存的和输出。", + ), + ), + row( + "cost", + Kind::Num, + Def::Unset, + t( + "At most this much, in US dollars, as recorded for each request. Models without a price and upstreams with `billing: free` count as 0.", + "最多花这么多美元,按每个请求记下的费用算。没有价格的模型、`billing: free` 的上游算 0。", + ), + ), + row( + "cache_reads", + Kind::Bool, + Def::Is("false"), + t( + "Count cache reads too. Only for a `tokens` limit.", + "把读取缓存的 token 也算进去。只有 `tokens` 上限能写。", + ), + ), ], }, // ── providers ───────────────────────────────────────── diff --git a/crates/tw-control/src/key_limits.rs b/crates/tw-control/src/key_limits.rs new file mode 100644 index 00000000..7f88493f --- /dev/null +++ b/crates/tw-control/src/key_limits.rs @@ -0,0 +1,65 @@ +//! 密钥的用量上限和存储层之间的两根线:每记下一行请求就结算,启动时从请求记录里把这一天、 +//! 这一周、这个月用了多少加回来(见 `tw_gateway::key_limits`)。 +//! +//! **网关和存储层互不依赖**(两件平级的事),线在这里接:除了 twcore,控制面是唯一同时 +//! 看得见两边的地方。twcore 起来时接上,测试也从这里接。 + +use std::sync::Arc; + +use tw_gateway::key_limits::Recorded; + +/// 存储层每记下一行,交给网关结算。 +pub fn settle_hook(gw: &tw_gateway::AppState) -> tw_store::SettleHook { + let limits = gw.key_limits.clone(); + Arc::new(move |s: &tw_store::Settled| { + limits.settle( + s.id, + s.at_ms, + &Recorded { + client: s.client.clone(), + path: s.path.clone(), + local: s.local, + error_code: s.error_code.clone(), + attempted: s.attempted, + requests: 1, + input: s.input, + output: s.output, + cache_read: s.cache_read, + cache_write: s.cache_write, + // 算不出钱的(没有价格、没有用量)算 0:费用上限只管得住有价格的 + cost_micros: s.cost_micros.unwrap_or(0), + }, + ) + }) +} + +/// 从请求记录里把每把密钥这一期用了多少加回来。**在第一个请求之前调**。读不了库就 +/// 从 0 起,只记一行 —— 观测挂了,转发照常。 +pub fn rebuild(gw: &tw_gateway::AppState, db: &tw_store::Db) { + gw.key_limits.rebuild(|since| match db.key_usage_since(since) { + Ok(rows) => rows.into_iter().map(recorded).collect(), + Err(e) => { + tracing::warn!( + "the usage of the gateway keys could not be read back, so their limits count from zero: {e}" + ); + Vec::new() + } + }); +} + +fn recorded(u: tw_store::KeyUsage) -> Recorded { + let n = |v: i64| v.max(0) as u64; + Recorded { + client: u.client, + path: u.path, + local: false, + error_code: u.error_code, + attempted: u.attempted, + requests: n(u.requests), + input: n(u.input_tokens), + output: n(u.output_tokens), + cache_read: n(u.cache_read_tokens), + cache_write: n(u.cache_write_tokens), + cost_micros: u.cost_micros, + } +} diff --git a/crates/tw-control/src/keys.rs b/crates/tw-control/src/keys.rs index b4a7ed70..c2819ae1 100644 --- a/crates/tw-control/src/keys.rs +++ b/crates/tw-control/src/keys.rs @@ -81,6 +81,7 @@ pub async fn views(s: &ControlState, reveal: Reveal) -> Vec .unwrap_or_default(), None => Vec::new(), }; + let usage = &s.gateway.key_limits; cfg.clients .iter() .map(|c| tw_api::ClientView { @@ -99,6 +100,17 @@ pub async fn views(s: &ControlState, reveal: Reveal) -> Vec .iter() .find(|(k, _)| *k == c.name) .map(|(_, at)| *at as u64), + limits: usage.view(&c.name, &c.limits), + // 只有设了费用上限的才要提醒:这些模型的费用记 0,费用上限管不住它们 + unpriced_models: if c + .limits + .iter() + .any(|l| l.measure() == tw_config::LimitMeasure::Cost) + { + tw_gateway::key_limits::unpriced_models(&s.gateway, &cfg, &c.name) + } else { + Vec::new() + }, }) .collect() } @@ -181,6 +193,11 @@ async fn update( }) .await .map_err(apply_fail)?; + // 改了名:用量上限的账跟着走(今天、这周、这个月用了多少,在跑的请求的预留)。 + // 存储层那些行记的是旧名字,重启之后加回来的只认新名字下的 + if req.key.name != name { + s.gateway.key_limits.rename(&name, &req.key.name); + } Ok(Json(tw_api::ConfigWritten { version })) } @@ -326,5 +343,20 @@ fn to_client( // 绑定由接管写,用户改不了 —— 它记的是「这把钥匙是为谁生成的」 client: existing.and_then(|c| c.client.clone()), disabled: input.disabled, + // 写得对不对(大于 0、不重复、缓存读只给 token)由配置校验说,和手写的同一套 + limits: input.limits.iter().map(limit).collect(), }) } + +/// 界面上的一条上限 → 配置里的写法。费用在界面上是微分,配置里写美元 +fn limit(l: &tw_api::KeyLimitInput) -> tw_config::KeyLimit { + let max = l.max.min(i64::MAX as u64) as i64; + let is = |m: tw_api::LimitMeasure| l.measure == m; + tw_config::KeyLimit { + per: l.per.into(), + requests: is(tw_api::LimitMeasure::Requests).then_some(max), + tokens: is(tw_api::LimitMeasure::Tokens).then_some(max), + cost: is(tw_api::LimitMeasure::Cost).then(|| max as f64 / 1_000_000.0), + cache_reads: l.cache_reads, + } +} diff --git a/crates/tw-control/src/lib.rs b/crates/tw-control/src/lib.rs index cb7eb470..eb00f992 100644 --- a/crates/tw-control/src/lib.rs +++ b/crates/tw-control/src/lib.rs @@ -30,6 +30,7 @@ mod contract; pub mod diagnostics; pub mod dryrun; mod gate; +pub mod key_limits; pub mod keys; pub mod listen; pub mod plugins; diff --git a/crates/tw-control/tests/key_limits.rs b/crates/tw-control/tests/key_limits.rs new file mode 100644 index 00000000..d000fe1c --- /dev/null +++ b/crates/tw-control/tests/key_limits.rs @@ -0,0 +1,385 @@ +//! 密钥的用量上限在控制面上:`GET /keys` 里每一条上限用了多少、什么时候重置、哪些模型 +//! 没有价格;保存密钥时写进配置;改名带着用量走;重启之后从请求记录里加回来;存储层 +//! 每记下一行就结算。 + +use std::sync::Arc; + +use axum::body::Body; +use axum::http::{Request, StatusCode}; +use tower::ServiceExt; +use tw_api::Event; +use tw_control::{ConfigManager, ControlState}; + +const CONFIG: &str = "version: 1 +listen: + control: + key: c0ffee00c0ffee00c0ffee00c0ffee00c0ffee00c0ffee00c0ffee00c0ffee00 +clients: + - name: default + key: tw-aaaa + - name: k + key: tw-kkkk + limits: + - { per: day, cost: 5 } + - { per: minute, requests: 30 } + - name: plain + key: tw-pppp + limits: + - { per: week, requests: 100 } + - name: small + key: tw-ssss + limits: + - { per: day, cost: 0.5 } +providers: + - name: 中转 + base_url: http://127.0.0.1:9 + key: sk-x + protocol: anthropic + models: [claude-sonnet-4-5, relay-own-model] +"; + +/// 中午,东八区。天的上限不会碰巧在测试跑到一半时过零点 +const NOON: &str = "2026-10-05T12:00:00+08:00"; + +fn ms(s: &str) -> i64 { + chrono::DateTime::parse_from_rfc3339(s) + .unwrap() + .timestamp_millis() +} + +struct Bed { + _dir: tempfile::TempDir, + app: axum::Router, + gw: tw_gateway::AppState, + rec: Arc>, + path: std::path::PathBuf, +} + +/// `db`:上一次运行留下的请求记录(重启);没有就是一个空库 +fn bed(db: Option) -> Bed { + let d = tempfile::tempdir().unwrap(); + let p = d.path().join("config.yaml"); + std::fs::write(&p, CONFIG).unwrap(); + let cfg = tw_config::try_parse(CONFIG).unwrap(); + let mut gw = tw_gateway::AppState::new(cfg).unwrap(); + gw.set_key_limits_clock(Arc::new(tw_gateway::key_limits::TestClock::new( + ms(NOON), + 8 * 3600, + ))); + let db = db.unwrap_or_else(|| tw_store::Db::in_memory().unwrap()); + // 重启:第一个请求之前把这一期加回来,和 twcore 起来时一样 + tw_control::key_limits::rebuild(&gw, &db); + let rec = tw_store::Recorder::new( + db, + tw_store::Blobs::new(d.path().join("blobs")), + gw.pricing.clone(), + ) + .settling_to(tw_control::key_limits::settle_hook(&gw)); + let rec = Arc::new(tokio::sync::Mutex::new(rec)); + let state = ControlState { + shutdown: Default::default(), + remote: Default::default(), + cfg: Arc::new(ConfigManager::new(p.clone(), gw.clone(), gw.bus.clone())), + gateway: gw.clone(), + store: Some(rec.clone()), + started: std::time::Instant::now(), + price_updater: Default::default(), + chatgpt: Default::default(), + zai: Default::default(), + }; + Bed { + app: tw_control::router(state), + _dir: d, + gw, + rec, + path: p, + } +} + +/// 一个跑完的请求的几条事件:密钥 `key`,开始于 `at`,按 Sonnet 4.5 的价($3/M 输入; +/// 超过 20 万的整个请求换长上下文的价,测试里不碰它) +fn request(id: u64, key: &str, at: i64, input: u64) -> Vec { + vec![ + Event::RequestStarted { + id, + client: key.into(), + client_hint: None, + session: None, + peer: None, + key_masked: None, + route: "default".into(), + rule: "r".into(), + group: None, + rewritten_by: vec![], + provider: "中转".into(), + billing: tw_api::Billing::PerToken, + model: "claude-sonnet-4-5".into(), + method: "POST".into(), + path: "/v1/messages".into(), + input_estimate: None, + session_log_bytes: None, + at_ms: at as u64, + }, + Event::RequestRouted { + id, + route: "default".into(), + rule: "r".into(), + group: None, + rewritten_by: vec![], + denied_by: None, + affinity: None, + attempts: vec![tw_api::AttemptView { + provider: "中转".into(), + model: None, + outcome: tw_api::AttemptOutcome::Served, + status: Some(200), + error: None, + ms: 5, + usage: None, + queued_ms: None, + skipped: None, + }], + billing: tw_api::Billing::PerToken, + }, + Event::RequestFinished { + id, + model: String::new(), + status: 200, + bytes: 1, + duration_ms: 1, + usage: Some(tw_api::UsageView { + input, + output: 0, + cache_read: 0, + cache_write: 0, + cache_1h: false, + }), + tokens_per_sec: None, + answered_model: None, + }, + ] +} + +async fn call( + app: &axum::Router, + method: &str, + path: &str, + body: serde_json::Value, +) -> (StatusCode, serde_json::Value) { + let r = app + .clone() + .oneshot( + Request::builder() + .method(method) + .uri(path) + .header("content-type", "application/json") + .body(Body::from(body.to_string())) + .unwrap(), + ) + .await + .unwrap(); + let st = r.status(); + let b = axum::body::to_bytes(r.into_body(), 1 << 20).await.unwrap(); + let text = String::from_utf8_lossy(&b).to_string(); + ( + st, + serde_json::from_str(&text).unwrap_or(serde_json::Value::String(text)), + ) +} + +async fn keys(b: &Bed) -> serde_json::Value { + let (st, v) = call(&b.app, "GET", "/keys", serde_json::Value::Null).await; + assert_eq!(st, StatusCode::OK, "{v}"); + v +} + +fn key<'a>(list: &'a serde_json::Value, name: &str) -> &'a serde_json::Value { + list.as_array() + .unwrap() + .iter() + .find(|k| k["name"] == name) + .unwrap_or_else(|| panic!("没有密钥「{name}」:{list}")) +} + +/// 存储层每记下一行就结算:`GET /keys` 里看得到这一天花了多少、什么时候重置;设了费用 +/// 上限的密钥带着它用得到、却没有价格的模型。 +#[tokio::test] +async fn the_key_list_shows_what_each_limit_has_used() { + let b = bed(None); + for e in request(1, "k", ms(NOON), 100_000) { + b.rec.lock().await.on_event(&e); + } + // 没有价格的模型、不计费的上游:记下的费用是空的、是 0,费用上限都算 0 + let mut unpriced = request(2, "k", ms(NOON), 100_000); + if let Event::RequestStarted { model, .. } = &mut unpriced[0] { + *model = "relay-own-model".into(); + } + let mut free = request(3, "k", ms(NOON), 100_000); + if let Event::RequestRouted { billing, .. } = &mut free[1] { + *billing = tw_api::Billing::Free; + } + for e in unpriced.iter().chain(&free) { + b.rec.lock().await.on_event(e); + } + let rows = b.rec.lock().await.db().key_usage_since(0).unwrap(); + assert_eq!(rows.iter().map(|r| r.requests).sum::(), 3); + let list = keys(&b).await; + let k = key(&list, "k"); + assert_eq!(k["limits"][0]["per"], "day"); + assert_eq!(k["limits"][0]["measure"], "cost"); + assert_eq!(k["limits"][0]["max"], 5_000_000); + assert_eq!(k["limits"][0]["used"], 300_000, "十万输入 × $3/M"); + assert_eq!(k["limits"][0]["reached"], false); + assert_eq!( + k["limits"][0]["resets_at_ms"], + ms("2026-10-06T00:00:00+08:00") + ); + assert_eq!(k["limits"][1]["per"], "minute"); + assert!( + k["limits"][1].get("resets_at_ms").is_none(), + "滚动的没有重置时刻" + ); + assert_eq!(k["unpriced_models"], serde_json::json!(["relay-own-model"])); + // 没有费用上限的不提醒;没有上限的,列表是空的 + assert!(key(&list, "plain").get("unpriced_models").is_none()); + assert_eq!(key(&list, "default")["limits"], serde_json::json!([])); +} + +/// 重启:这一天、这一周、这个月的数从请求记录里加回来。别的密钥的、前一天的不算进来。 +#[tokio::test] +async fn a_restart_adds_this_periods_usage_back_from_the_records() { + let d = tempfile::tempdir().unwrap(); + let file = d.path().join("data.db"); + let mut earlier = tw_store::Recorder::new( + tw_store::Db::open(&file).unwrap(), + tw_store::Blobs::new(d.path().join("blobs")), + tw_pricing::shared(tw_pricing::PriceBook::builtin().unwrap()), + ); + let events = [ + request(1, "k", ms("2026-10-05T09:00:00+08:00"), 100_000), + request(2, "k", ms("2026-10-04T23:00:00+08:00"), 100_000), + request(3, "plain", ms("2026-10-05T08:00:00+08:00"), 10), + // 上周日的:这一周不算 + request(4, "plain", ms("2026-10-04T08:00:00+08:00"), 10), + ]; + for e in events.iter().flatten() { + earlier.on_event(e); + } + drop(earlier); + // 上一次运行的库原样交给这一次 + let b = bed(Some(tw_store::Db::open(&file).unwrap())); + let list = keys(&b).await; + assert_eq!( + key(&list, "k")["limits"][0]["used"], + 300_000, + "前一天的不算" + ); + assert_eq!(key(&list, "plain")["limits"][0]["used"], 1); +} + +/// 保存密钥时写进配置,费用在界面上是微分、配置里写美元;写错了是配置校验的那句话。 +/// 改名之后用过的跟着走。 +#[tokio::test] +async fn saving_a_key_writes_its_limits_and_a_rename_keeps_what_it_used() { + let b = bed(None); + for e in request(1, "k", ms(NOON), 100_000) { + b.rec.lock().await.on_event(&e); + } + let (st, v) = call( + &b.app, + "PUT", + "/keys/k", + serde_json::json!({ "key": { "name": "k2", "limits": [ + { "per": "day", "measure": "cost", "max": 250_000 }, + { "per": "hour", "measure": "tokens", "max": 1000, "cache_reads": true }, + ] } }), + ) + .await; + assert_eq!(st, StatusCode::OK, "{v}"); + let cfg = tw_config::try_parse(&std::fs::read_to_string(&b.path).unwrap()).unwrap(); + let k2 = cfg.clients.iter().find(|c| c.name == "k2").unwrap(); + assert_eq!(k2.limits[0].cost, Some(0.25)); + assert_eq!(k2.limits[1].tokens, Some(1000)); + assert!(k2.limits[1].cache_reads); + // 改了名:今天花掉的还在,而且已经超过了新的上限 + let list = keys(&b).await; + let k2 = key(&list, "k2"); + assert_eq!(k2["limits"][0]["used"], 300_000); + assert_eq!(k2["limits"][0]["reached"], true); + + let (st, v) = call( + &b.app, + "PUT", + "/keys/k2", + serde_json::json!({ "key": { "name": "k2", "limits": [ + { "per": "day", "measure": "requests", "max": 0 }, + ] } }), + ) + .await; + assert_eq!(st, StatusCode::BAD_REQUEST, "{v}"); + assert!( + v.to_string().contains("config.key_limit_not_positive"), + "{v}" + ); + let (st, v) = call( + &b.app, + "PUT", + "/keys/k2", + serde_json::json!({ "key": { "name": "k2", "limits": [ + { "per": "day", "measure": "cost", "max": 1, "cache_reads": true }, + ] } }), + ) + .await; + assert_eq!(st, StatusCode::BAD_REQUEST, "{v}"); + assert!( + v.to_string().contains("config.key_limit_cache_reads"), + "{v}" + ); +} + +/// 一把密钥这一期到了八成、到了顶:总线上各报一次,界面拿它发系统通知。 +#[tokio::test] +async fn reaching_a_limit_is_told_on_the_bus() { + let b = bed(None); + let mut rx = b.gw.bus.subscribe(); + // 一天 $0.5:$0.45、$0.51,先过八成,再过顶 + for e in request(1, "small", ms(NOON), 150_000) { + b.rec.lock().await.on_event(&e); + } + for e in request(2, "small", ms(NOON), 20_000) { + b.rec.lock().await.on_event(&e); + } + let mut told = Vec::new(); + while let Ok(e) = rx.try_recv() { + if let Event::KeyLimitAlert { + key, + per, + measure, + used, + reached, + .. + } = e + { + told.push((key, per, measure, used, reached)); + } + } + assert_eq!( + told, + [ + ( + "small".to_string(), + tw_api::LimitPer::Day, + tw_api::LimitMeasure::Cost, + 450_000, + false + ), + ( + "small".to_string(), + tw_api::LimitPer::Day, + tw_api::LimitMeasure::Cost, + 510_000, + true + ), + ] + ); +} diff --git a/crates/tw-gateway/src/error.rs b/crates/tw-gateway/src/error.rs index bde5a0c3..0cafb853 100644 --- a/crates/tw-gateway/src/error.rs +++ b/crates/tw-gateway/src/error.rs @@ -88,6 +88,24 @@ pub struct GatewayError { /// 用哪种格式的形状回。**认证失败时还不知道格式**(key 就是没认出 /// 来),所以先按 Anthropic —— 桌面版的主用例是 Claude Code。 pub dialect: Dialect, + /// 什么时候能再来:密钥的用量上限拒绝的(见 [`crate::key_limits`]),和上游都满着的 + /// ([`Source::Busy`])。别的没有 + pub retry: Option, +} + +/// 密钥的用量上限拒绝了一个请求,或者上游都满着([`Source::Busy`]):客户端什么时候能再来。 +/// +/// 写成响应头:`Retry-After`(秒)。滚动窗口另带 `retry-after-ms` 和 `x-should-retry: true` +/// —— Anthropic、OpenAI 的官方 SDK 按毫秒那个等,等完自己重试。**自然周期用完了带 +/// `x-should-retry: false`**:到下一期之前重试多少次都一样,SDK 照默认的退避连试几次只会 +/// 让用户多等;OpenAI 格式的错误体再带 `code: insufficient_quota`,和 OpenAI 自己额度用完 +/// 时一样 —— 认这个码的客户端(Codex)会直接停下来告诉用户,不再重试。 +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub struct Retry { + /// 多久之后,毫秒 + pub after_ms: u64, + /// 这一期用完了,到重置之前重试也没用 + pub until_reset: bool, } impl GatewayError { @@ -96,9 +114,15 @@ impl GatewayError { source, detail, dialect: Dialect::Anthropic, + retry: None, } } + pub fn with_retry(mut self, r: Retry) -> Self { + self.retry = Some(r); + self + } + /// 给人看的那句话。 pub fn message(&self) -> &str { &self.detail.text @@ -127,8 +151,12 @@ impl GatewayError { pub fn rate_limited(detail: Msg) -> Self { Self::new(Source::RateLimited, detail) } + /// 上游都满着:429,[`BUSY_RETRY_AFTER_SECS`] 之后再来,可以重试 pub fn busy(detail: Msg) -> Self { - Self::new(Source::Busy, detail) + Self::new(Source::Busy, detail).with_retry(Retry { + after_ms: BUSY_RETRY_AFTER_SECS * 1000, + until_reset: false, + }) } pub fn denied(detail: Msg) -> Self { Self::new(Source::Denied, detail) @@ -174,11 +202,21 @@ impl GatewayError { /// 发出去了,body 是唯一还能说话的地方 —— 和 `sse_frame` 在流上扮 /// 演的是同一个角色。 pub fn body_bytes(&self) -> Vec { - tw_dialect::convert::error_body( + let body = tw_dialect::convert::error_body( self.dialect, self.source.status().as_u16(), &self.client_message(), - ) + ); + // 自然周期用完了:OpenAI 格式里写成额度用完(见 [`Retry`]) + if self.retry.is_some_and(|r| r.until_reset) + && matches!(self.dialect, Dialect::Chat | Dialect::Responses) + && let Ok(mut v) = serde_json::from_slice::(&body) + && let Some(e) = v.get_mut("error").and_then(|e| e.as_object_mut()) + { + e.insert("code".into(), "insufficient_quota".into()); + return v.to_string().into_bytes(); + } + body } } @@ -194,11 +232,15 @@ impl IntoResponse for GatewayError { "x-thinkwatch-error", HeaderValue::from_static(self.source.slug()), ); - if self.source == Source::Busy { - h.insert( - header::RETRY_AFTER, - HeaderValue::from(BUSY_RETRY_AFTER_SECS), - ); + if let Some(r) = self.retry { + let secs = r.after_ms.div_ceil(1000).max(1); + h.insert(header::RETRY_AFTER, HeaderValue::from(secs)); + if r.until_reset { + h.insert("x-should-retry", HeaderValue::from_static("false")); + } else { + h.insert("retry-after-ms", HeaderValue::from(r.after_ms.max(1))); + h.insert("x-should-retry", HeaderValue::from_static("true")); + } } resp } @@ -293,6 +335,9 @@ mod tests { .into_response(); assert_eq!(r.status(), StatusCode::TOO_MANY_REQUESTS); assert_eq!(r.headers()["retry-after"], "5"); + // 和密钥用量上限的滚动窗口同一套头:可以重试,毫秒的那个给 SDK 用 + assert_eq!(r.headers()["retry-after-ms"], "5000"); + assert_eq!(r.headers()["x-should-retry"], "true"); assert_eq!(r.headers()["x-thinkwatch-error"], "rate_limited"); let b = to_bytes(r.into_body(), 64 * 1024).await.unwrap(); let json: serde_json::Value = serde_json::from_slice(&b).unwrap(); @@ -302,6 +347,61 @@ mod tests { assert!(!r.headers().contains_key("retry-after")); } + /// 密钥的用量上限拒绝时:多久之后能再来,到重置之前该不该重试。 + #[tokio::test] + async fn a_used_up_key_says_when_to_come_back_and_whether_to_retry() { + let e = || GatewayError::rate_limited(msg!("t.x" => "x")); + let head = + |r: &Response, h: &str| r.headers().get(h).map(|v| v.to_str().unwrap().to_string()); + let body = |r: Response| async { + let b = to_bytes(r.into_body(), 64 * 1024).await.unwrap(); + serde_json::from_slice::(&b).unwrap() + }; + // 这一期用完了:到重置那一刻,别重试;OpenAI 的两种格式写成额度用完 + let until_reset = Retry { + after_ms: 3_600_500, + until_reset: true, + }; + for d in [Dialect::Chat, Dialect::Responses] { + let r = e().with_retry(until_reset).in_dialect(d).into_response(); + assert_eq!(r.status(), StatusCode::TOO_MANY_REQUESTS); + assert_eq!(head(&r, "retry-after").as_deref(), Some("3601")); + assert_eq!(head(&r, "x-should-retry").as_deref(), Some("false")); + assert_eq!(head(&r, "retry-after-ms"), None); + let v = body(r).await; + assert_eq!(v["error"]["code"], "insufficient_quota", "{d:?}"); + assert_eq!(v["error"]["type"], "rate_limit_error"); + } + let v = body(e().with_retry(until_reset).into_response()).await; + assert_eq!( + v["error"]["type"], "rate_limit_error", + "Anthropic 的形状不变" + ); + let v = body( + e().with_retry(until_reset) + .in_dialect(Dialect::Gemini) + .into_response(), + ) + .await; + assert_eq!(v["error"]["status"], "RESOURCE_EXHAUSTED"); + // 滚动窗口:准确到毫秒,可以重试 + let r = e() + .with_retry(Retry { + after_ms: 1_500, + until_reset: false, + }) + .in_dialect(Dialect::Chat) + .into_response(); + assert_eq!(head(&r, "retry-after").as_deref(), Some("2")); + assert_eq!(head(&r, "retry-after-ms").as_deref(), Some("1500")); + assert_eq!(head(&r, "x-should-retry").as_deref(), Some("true")); + assert!(body(r).await["error"]["code"].is_null()); + // 别的错误什么都不加 + let r = e().into_response(); + assert_eq!(head(&r, "retry-after"), None); + assert_eq!(head(&r, "x-should-retry"), None); + } + #[test] fn a_responses_stream_is_told_with_response_failed() { // Chat 形状的 `{"error":…}` 在 Responses 的流里是一帧没人认的数据,客户端 diff --git a/crates/tw-gateway/src/key_limits/clock.rs b/crates/tw-gateway/src/key_limits/clock.rs new file mode 100644 index 00000000..85bab1e8 --- /dev/null +++ b/crates/tw-gateway/src/key_limits/clock.rs @@ -0,0 +1,189 @@ +//! 用量上限看的时钟:此刻几点,和某一刻所在的那一天、那一周、那个月从哪儿起、到哪儿止。 +//! +//! **天、周、月按 core 所在机器的本地时区算**:用户说「每天 $5」,说的是他自己那一天, +//! 不是 UTC 那一天 —— 东八区的人早上八点之前花的钱,按 UTC 算会记到前一天去。 +//! +//! 时钟可以换:测试要把时间拨到零点前一秒、拨过周一、拨过一号,还要让滚动窗口的等待 +//! 跟着 tokio 的假时间走。 + +use chrono::{Datelike, Days, NaiveDate, NaiveTime, TimeZone}; +use tw_config::LimitPer; + +pub trait Clock: Send + Sync + 'static { + /// 此刻,Unix 毫秒 + fn now_ms(&self) -> i64; + /// `at_ms` 所在的那一期(天、周、月)的开头和结尾,Unix 毫秒,结尾不含。分钟、小时 + /// 是滚动的,没有「那一期」:给的是以 `at_ms` 结尾的那一段 + fn period(&self, per: LimitPer, at_ms: i64) -> (i64, i64); + /// 某一刻写给人看:`2026-10-06 00:00 +08:00`。带着时区,远程连上来的人也看得懂 + fn show(&self, at_ms: i64) -> String; +} + +/// 系统时钟,本机时区。 +pub struct SystemClock; + +impl Clock for SystemClock { + fn now_ms(&self) -> i64 { + std::time::SystemTime::now() + .duration_since(std::time::UNIX_EPOCH) + .map(|d| d.as_millis() as i64) + .unwrap_or(0) + } + fn period(&self, per: LimitPer, at_ms: i64) -> (i64, i64) { + period_in(&chrono::Local, per, at_ms) + } + fn show(&self, at_ms: i64) -> String { + show_in(&chrono::Local, at_ms) + } +} + +/// 测试用的时钟:固定的时区,时间从 `base_ms` 起跟着 tokio 的时钟走 —— 测试里 +/// `tokio::time::pause` 之后 `advance` 多少,它就走多少,滚动窗口的等待也一起快进。 +pub struct TestClock { + offset: chrono::FixedOffset, + base_ms: i64, + start: tokio::time::Instant, +} + +impl TestClock { + /// `offset_secs`:时区比 UTC 快多少秒(东八区是 `8 * 3600`) + pub fn new(base_ms: i64, offset_secs: i32) -> Self { + Self { + offset: chrono::FixedOffset::east_opt(offset_secs).expect("a valid offset"), + base_ms, + start: tokio::time::Instant::now(), + } + } +} + +impl Clock for TestClock { + fn now_ms(&self) -> i64 { + self.base_ms + self.start.elapsed().as_millis() as i64 + } + fn period(&self, per: LimitPer, at_ms: i64) -> (i64, i64) { + period_in(&self.offset, per, at_ms) + } + fn show(&self, at_ms: i64) -> String { + show_in(&self.offset, at_ms) + } +} + +fn show_in(tz: &Tz, at_ms: i64) -> String +where + Tz::Offset: std::fmt::Display, +{ + match tz.timestamp_millis_opt(at_ms).single() { + Some(t) => t.format("%Y-%m-%d %H:%M %:z").to_string(), + None => at_ms.to_string(), + } +} + +/// 某一刻所在的那一期,按时区 `tz`。周从周一开始。 +pub(crate) fn period_in(tz: &Tz, per: LimitPer, at_ms: i64) -> (i64, i64) { + if let Some(w) = per.rolling_ms() { + return (at_ms - w, at_ms); + } + let Some(t) = tz.timestamp_millis_opt(at_ms).single() else { + return (at_ms, at_ms + 86_400_000); + }; + let d = t.date_naive(); + let (start, end) = match per { + LimitPer::Week => { + let s = d - Days::new(u64::from(d.weekday().num_days_from_monday())); + (s, s + Days::new(7)) + } + LimitPer::Month => { + let s = d.with_day(1).unwrap_or(d); + let e = if s.month() == 12 { + NaiveDate::from_ymd_opt(s.year() + 1, 1, 1) + } else { + NaiveDate::from_ymd_opt(s.year(), s.month() + 1, 1) + }; + (s, e.unwrap_or(s + Days::new(31))) + } + _ => (d, d + Days::new(1)), + }; + (midnight(tz, start), midnight(tz, end)) +} + +/// 这一天零点,Unix 毫秒。 +/// +/// **夏令时在零点切换的地方,零点可能不存在**(时钟从 23:59 直接跳到 01:00):那一天 +/// 从跳过去之后的第一刻算起。零点出现两次的(往回拨),从第一次算起。 +fn midnight(tz: &Tz, d: NaiveDate) -> i64 { + let at = d.and_time(NaiveTime::MIN); + for hours in 0..=3 { + let local = at + chrono::Duration::hours(hours); + if let Some(t) = tz.from_local_datetime(&local).earliest() { + return t.timestamp_millis(); + } + } + // 找不到就按 UTC 算:只差几个小时,总比不算好 + at.and_utc().timestamp_millis() +} + +#[cfg(test)] +mod tests { + use super::*; + + /// 2026-10-05 是周一。东八区的那一天从前一天 16:00 UTC 起 + #[test] + fn a_day_a_week_and_a_month_start_at_local_midnight() { + let tz = chrono::FixedOffset::east_opt(8 * 3600).unwrap(); + let at = |s: &str| { + chrono::DateTime::parse_from_rfc3339(s) + .unwrap() + .timestamp_millis() + }; + // 周三中午 + let now = at("2026-10-07T12:00:00+08:00"); + assert_eq!( + period_in(&tz, LimitPer::Day, now), + ( + at("2026-10-07T00:00:00+08:00"), + at("2026-10-08T00:00:00+08:00") + ) + ); + assert_eq!( + period_in(&tz, LimitPer::Week, now), + ( + at("2026-10-05T00:00:00+08:00"), + at("2026-10-12T00:00:00+08:00") + ), + "周从周一开始" + ); + assert_eq!( + period_in(&tz, LimitPer::Month, now), + ( + at("2026-10-01T00:00:00+08:00"), + at("2026-11-01T00:00:00+08:00") + ) + ); + // 十二月的下一期在明年 + assert_eq!( + period_in(&tz, LimitPer::Month, at("2026-12-31T23:59:59+08:00")).1, + at("2027-01-01T00:00:00+08:00") + ); + // UTC 已经是第二天了,本地还是这一天:按本地算 + assert_eq!( + period_in(&tz, LimitPer::Day, at("2026-10-07T23:30:00+08:00")).0, + at("2026-10-07T00:00:00+08:00") + ); + // 周日还在这一周里 + assert_eq!( + period_in(&tz, LimitPer::Week, at("2026-10-11T23:59:59+08:00")).0, + at("2026-10-05T00:00:00+08:00") + ); + // 滚动的:以此刻结尾的那一段 + assert_eq!(period_in(&tz, LimitPer::Minute, now), (now - 60_000, now)); + } + + #[test] + fn a_moment_is_shown_with_its_offset() { + let tz = chrono::FixedOffset::east_opt(8 * 3600).unwrap(); + let ms = chrono::DateTime::parse_from_rfc3339("2026-10-06T00:00:00+08:00") + .unwrap() + .timestamp_millis(); + assert_eq!(show_in(&tz, ms), "2026-10-06 00:00 +08:00"); + } +} diff --git a/crates/tw-gateway/src/key_limits/mod.rs b/crates/tw-gateway/src/key_limits/mod.rs new file mode 100644 index 00000000..3aee8ef6 --- /dev/null +++ b/crates/tw-gateway/src/key_limits/mod.rs @@ -0,0 +1,988 @@ +//! 每把网关密钥的用量上限:数、等、拒、结算。 +//! +//! 写法在 `tw_config::limits`。这里回答四件事。 +//! +//! **数什么。**请求数:一个准入的请求算一个,数 token 的请求、网关自己答的不算。token: +//! 没走缓存的输入 + 写进缓存的 + 输出,那一条开了 `cache_reads` 再加上从缓存读的。费用: +//! 记下的费用(实测的、估算的都算),没有价格的模型、不计费的上游算 0。**都按存储层记下的 +//! 那一行算**([`KeyLimits::settle`]):重启之后从库里加回来的([`KeyLimits::rebuild`])和 +//! 此刻内存里的是同一份数。 +//! +//! **怎么拒、怎么等。**天、周、月是自然周期:用满了就拒,到下一期之前重试也没用。分钟、 +//! 小时是滚动窗口:用满了先看下一个空位多久之后空出来,等得到(不超过 `slot_wait_secs`) +//! 就等,等不到就拒,并说清楚多久之后再来。顺序见 [`crate::server`] 的管线第 3 步:自然 +//! 周期 → 并发上限 → 滚动窗口。 +//! +//! **在跑的怎么算。**准入时按输入估一个数占着(预留:输入 token 的估算,和按头一个候选算 +//! 的输入费用),存储层记下那一行时换成实数。几个请求同时进来时,超出上限的最多是在跑的 +//! 那些。 +//! +//! **什么时候算到哪一期。**天、周、月按请求开始的时刻(那一行的 `at_ms`)归期,和从库里 +//! 加回来时同一个口径。滚动窗口里请求数记在准入的那一刻,token 和费用记在结算的那一刻: +//! 一个跑了三分钟的请求,它的输出要等它跑完才知道,按开始的时刻记的话,「每分钟多少 +//! token」永远数不到它。 + +use std::collections::{HashMap, HashSet, VecDeque}; +use std::sync::{Arc, Mutex, PoisonError}; +use std::time::Duration; + +use tw_config::{KeyLimit, LimitMeasure, LimitPer}; +use tw_types::{Msg, msg}; + +mod clock; +#[cfg(test)] +mod tests; + +pub use clock::{Clock, SystemClock, TestClock}; + +/// 自然周期。下标就是 [`Books::periods`] 的下标 +const CALENDAR: [LimitPer; 3] = [LimitPer::Day, LimitPer::Week, LimitPer::Month]; + +/// 请求已经结束、存储层却还没交来它那一行时,它的预留再留多久。 +/// +/// 存储层落后、丢了事件时那一行永远不会来:不放掉的话,那份预留一直占着这把密钥的 +/// 额度。存储层正常时几毫秒就到 +const GRACE_MS: i64 = 60_000; + +/// 准入之前就被拒的请求那一行的失败码:**路由**拒绝了它(规则拒绝、选中的上游都服务不了), +/// 一跳都没有。这样的请求没有经过准入,不算进请求数([`Recorded::counts`])。 +/// +/// 只看码不够:同样的码在尝试链的某一跳上也会出现(阶段二的规则拒绝),那时请求已经准入 +/// 过了 —— 所以还要看有没有一跳。 +const NOT_ADMITTED: &[&str] = &[ + "gw.route.denied", + "gw.route.all_selected_disabled", + "gw.model.no_upstream_available", +]; + +/// 上限本身拒绝的请求那一行的失败码前缀。它们也不算进请求数:不然一个被拒的客户端每重试 +/// 一次,窗口就往后推一次,永远等不到空位 +const REFUSED: &str = "gw.key_limit."; + +/// 等滚动窗口空位最多等多久:`failover.slot_wait_secs`,和等上游空位(`crate::slots`) +/// 同一个设置,0 是不等。 +/// +/// **两段各算各的**:这一段在准入时、发给哪一家之前,等上游空位在之后的那几跳里。一个请求 +/// 两样都碰上的话,最多等两倍 +pub fn slot_wait(cfg: &tw_config::Config) -> Duration { + Duration::from_secs(cfg.failover.slot_wait_secs) +} + +/// 这个路径的请求算不算:数 token 的不算(Anthropic 的 `count_tokens`、Gemini 的 +/// `:countTokens`、Responses 的 `input_tokens`)—— 它们不跑模型、不收钱。 +pub fn uncounted(path: &str) -> bool { + if crate::client_api::ClientApi::counts_tokens(path) { + return true; + } + let p = path.trim_end_matches('/'); + p.strip_prefix("/v1").unwrap_or(p) == "/responses/input_tokens" +} + +/// 一份用量。四样分开记,各条上限各取各的。 +#[derive(Debug, Clone, Copy, Default, PartialEq, Eq)] +struct Amount { + requests: i64, + /// 没走缓存的输入 + 写进缓存的 + 输出 + tokens: i64, + cache_read: i64, + /// 微分 + cost: i64, +} + +impl Amount { + fn of(&self, m: LimitMeasure, cache_reads: bool) -> i64 { + match m { + LimitMeasure::Requests => self.requests, + LimitMeasure::Tokens if cache_reads => self.tokens.saturating_add(self.cache_read), + LimitMeasure::Tokens => self.tokens, + LimitMeasure::Cost => self.cost, + } + } + fn add(&mut self, o: &Amount) { + self.requests = self.requests.saturating_add(o.requests); + self.tokens = self.tokens.saturating_add(o.tokens); + self.cache_read = self.cache_read.saturating_add(o.cache_read); + self.cost = self.cost.saturating_add(o.cost); + } +} + +/// 一期(天、周、月)的合计。 +#[derive(Debug, Clone, Copy)] +struct Period { + start: i64, + end: i64, + sum: Amount, +} + +/// 一把密钥的账。**跨重载存活**,按密钥的名字记 —— 存储层那一行记的也是名字。 +#[derive(Debug, Default)] +struct Books { + /// 天、周、月,和 [`CALENDAR`] 同一个顺序。到了下一期就换一本新的 + periods: [Option; 3], + /// 滚动窗口用:最近的每一笔(请求数在准入时,token 和费用在结算时),按记下的先后。 + /// **只有设了分钟、小时上限的密钥才记**,留到最长的那个窗口为止 + recent: VecDeque<(i64, Amount)>, +} + +impl Books { + /// `now` 所在的那一期。到了下一期就从 0 起 + fn period(&mut self, clock: &dyn Clock, per: LimitPer, now: i64) -> &mut Period { + let (start, end) = clock.period(per, now); + let slot = &mut self.periods[calendar_index(per)]; + match slot { + Some(p) if p.start == start => {} + _ => { + *slot = Some(Period { + start, + end, + sum: Amount::default(), + }) + } + } + slot.as_mut().expect("set just above") + } +} + +fn calendar_index(per: LimitPer) -> usize { + CALENDAR.iter().position(|p| *p == per).unwrap_or(0) +} + +/// 一个在跑的请求占着的那一份。 +#[derive(Debug)] +struct Reservation { + key: String, + /// 准入的那一刻。算在哪一期、在不在滚动窗口里看它 + at_ms: i64, + /// 输入 token 的估算 + tokens: i64, + /// 头一个候选的输入费用估算,微分 + cost: i64, + /// 它是哪个请求。准入之后、开始事件之前还没有号 + request: Option, + /// 发现这个请求已经结束的那一刻(见 [`GRACE_MS`]) + closed_since: Option, +} + +/// 报过的那一档:哪把密钥、哪一条、哪一期、八成还是到顶。**上限改了就是另一条**, +/// 重新报 +#[derive(Debug, Clone, PartialEq, Eq, Hash)] +struct Told { + key: String, + per: LimitPer, + measure: LimitMeasure, + cache_reads: bool, + max: i64, + period_start: i64, + reached: bool, +} + +#[derive(Debug, Default)] +struct Inner { + books: HashMap, + /// 在跑的请求的预留,按一个自己的号 + held: HashMap, + /// 请求号 → 预留的号 + by_request: HashMap, + seq: u64, + /// 此刻配置里每把密钥的上限。结算时用:要不要记滚动窗口、报不报到了八成 + limits: HashMap>, + told: HashSet, +} + +/// 准入时一个请求要占的:输入 token 的估算,和按头一个候选算的输入费用。 +#[derive(Debug, Clone, Copy, Default, PartialEq, Eq)] +pub struct Ask { + pub tokens: u64, + pub cost_micros: i64, +} + +/// 存储层记下的一行(结算),或者库里加起来的一组(重建):算不算、算多少,看的是这几样。 +#[derive(Debug, Clone, Default, PartialEq)] +pub struct Recorded { + /// 密钥的名字 + pub client: String, + pub path: String, + /// 网关自己答的 + pub local: bool, + pub error_code: Option, + /// 尝试链上有没有至少一跳 + pub attempted: bool, + /// 几个请求:结算时是 1,重建时是这一组的行数 + pub requests: u64, + pub input: u64, + pub output: u64, + pub cache_read: u64, + pub cache_write: u64, + /// 记下的费用,微分。算不出来的(没有价格、没有用量)是 0 + pub cost_micros: i64, +} + +impl Recorded { + /// 算不算进密钥的用量:**准入过的才算**。 + /// + /// 网关自己答的、数 token 的不经过准入;上限本身拒绝的、路由就拒绝了的没有准入。 + /// 结算一行和从库里加回来用的是这同一个判断,两边的数才对得上。 + pub fn counts(&self) -> bool { + if self.local || uncounted(&self.path) { + return false; + } + match self.error_code.as_deref() { + Some(code) if code.starts_with(REFUSED) => false, + Some(code) if !self.attempted && NOT_ADMITTED.contains(&code) => false, + _ => true, + } + } + + fn amount(&self) -> Amount { + let n = |v: u64| v.min(i64::MAX as u64) as i64; + Amount { + requests: n(self.requests), + tokens: n(self + .input + .saturating_add(self.cache_write) + .saturating_add(self.output)), + cache_read: n(self.cache_read), + // 退款之类的负数不会有;有也不让它把用量往回拨 + cost: self.cost_micros.max(0), + } + } +} + +/// 一次拒绝:哪把密钥、哪一条上限、用了多少、什么时候能再来。 +#[derive(Debug, Clone, PartialEq)] +pub struct Refusal { + pub key: String, + pub limit: KeyLimit, + pub used: i64, + /// 多久之后能再来,毫秒。自然周期是到下一期的时间 + pub retry_after_ms: u64, + /// 自然周期:这一期结束的时刻,和写给人看的样子 + pub resets: Option<(i64, String)>, +} + +impl Refusal { + /// 客户端收到的那个错误:429,带 `Retry-After`。自然周期用完了说到重置之前别再试 + /// (见 [`crate::error::Retry`]) + pub fn error(&self) -> crate::GatewayError { + let retry = crate::error::Retry { + after_ms: self.retry_after_ms, + until_reset: self.resets.is_some(), + }; + crate::GatewayError::rate_limited(self.msg()).with_retry(retry) + } + + /// 那句话。**一种量、一种周期一句**:界面按码翻,参数里的 `per` 是一个英文词 + /// (`day`),界面按它查自己的词表 + pub fn msg(&self) -> Msg { + let key = self.key.clone(); + let per = self.limit.per.word(); + let measure = self.limit.measure(); + let (max, used) = match measure { + LimitMeasure::Cost => (dollars(self.limit.max()), dollars(self.used)), + _ => (self.limit.max().to_string(), self.used.to_string()), + }; + match (&self.resets, measure) { + (Some((at, resets)), LimitMeasure::Requests) => msg!( + "gw.key_limit.requests_per_period", key = key, max = max, per = per, + used = used, resets = resets, resets_at_ms = at => + "Gateway key `{key}` has reached its limit of {max} requests per {per}: {used} so \ + far. It resets at {resets}." + ), + (Some((at, resets)), LimitMeasure::Tokens) => msg!( + "gw.key_limit.tokens_per_period", key = key, max = max, per = per, + used = used, resets = resets, resets_at_ms = at => + "Gateway key `{key}` has reached its limit of {max} tokens per {per}: {used} so \ + far. It resets at {resets}." + ), + (Some((at, resets)), LimitMeasure::Cost) => msg!( + "gw.key_limit.cost_per_period", key = key, max = max, per = per, + used = used, resets = resets, resets_at_ms = at => + "Gateway key `{key}` has reached its limit of {max} per {per}: {used} spent so \ + far. It resets at {resets}." + ), + (None, measure) => { + let retry = self.retry_after_ms.div_ceil(1000).max(1); + match measure { + LimitMeasure::Requests => msg!( + "gw.key_limit.requests_rolling", key = key, max = max, per = per, + used = used, retry = retry => + "Gateway key `{key}` has reached its limit of {max} requests per {per}: \ + {used} in the last {per}. Try again in {retry} s." + ), + LimitMeasure::Tokens => msg!( + "gw.key_limit.tokens_rolling", key = key, max = max, per = per, + used = used, retry = retry => + "Gateway key `{key}` has reached its limit of {max} tokens per {per}: \ + {used} in the last {per}. Try again in {retry} s." + ), + LimitMeasure::Cost => msg!( + "gw.key_limit.cost_rolling", key = key, max = max, per = per, + used = used, retry = retry => + "Gateway key `{key}` has reached its limit of {max} per {per}: {used} \ + spent in the last {per}. Try again in {retry} s." + ), + } + } + } + } +} + +/// 微分写成美元:`$5.00`、`$0.0042`。至少两位小数,多的照实写 +pub(crate) fn dollars(micros: i64) -> String { + let sign = if micros < 0 { "-" } else { "" }; + let m = micros.unsigned_abs(); + let mut frac = format!("{:06}", m % 1_000_000); + while frac.len() > 2 && frac.ends_with('0') { + frac.pop(); + } + format!("{sign}${}.{frac}", m / 1_000_000) +} + +/// 准入过的请求占着的那一份。**丢掉就放掉**:准入之后、开始事件之前出了岔子的请求, +/// 不会一直占着额度。开始之后由 [`Hold::bind`] 交给请求号,等存储层结算。 +#[must_use = "dropping it gives the reservation back"] +pub struct Hold { + owner: Option>, + seq: u64, +} + +impl Hold { + /// 没有预留:这把密钥没设上限,或者这个请求不算(数 token 的) + pub fn none() -> Self { + Self { + owner: None, + seq: 0, + } + } + + /// 请求有号了:预留跟着它,等存储层记下那一行时换成实数 + pub fn bind(mut self, request: u64) { + if let Some(owner) = self.owner.take() { + owner.bind(self.seq, request); + } + } +} + +impl std::fmt::Debug for Hold { + fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + f.debug_struct("Hold") + .field("held", &self.owner.is_some()) + .finish() + } +} + +impl Drop for Hold { + fn drop(&mut self) { + if let Some(owner) = self.owner.take() { + owner.release(self.seq); + } + } +} + +/// 全部密钥的用量账。**跨重载存活**(在 `AppState` 里):改一条路由规则不该让今天花的 +/// 钱归零。 +pub struct KeyLimits { + inner: Mutex, + clock: Arc, + /// 报「到了八成、到了上限」,和认请求有没有结束(见 [`GRACE_MS`]) + bus: tw_observe::EventBus, +} + +impl KeyLimits { + pub fn new(bus: tw_observe::EventBus) -> Self { + Self::with_clock(bus, Arc::new(SystemClock)) + } + + pub fn with_clock(bus: tw_observe::EventBus, clock: Arc) -> Self { + Self { + inner: Mutex::default(), + clock, + bus, + } + } + + fn lock(&self) -> std::sync::MutexGuard<'_, Inner> { + self.inner.lock().unwrap_or_else(PoisonError::into_inner) + } + + /// 配置换了:记下每把密钥此刻的上限。**账不动**:上限改了,已经用掉的照样算 + pub fn configure(&self, cfg: &tw_config::Config) { + let limits = cfg + .clients + .iter() + .filter(|c| !c.limits.is_empty()) + .map(|c| (c.name.clone(), c.limits.clone())) + .collect(); + self.lock().limits = limits; + } + + /// 密钥改了名(控制面改的):账跟着走。手改配置文件改的名认不出来 —— 那和删掉一把、 + /// 新建一把看起来一样,新的那把从 0 起 + pub fn rename(&self, from: &str, to: &str) { + if from == to { + return; + } + let mut g = self.lock(); + if let Some(b) = g.books.remove(from) { + g.books.insert(to.to_string(), b); + } + for r in g.held.values_mut() { + if r.key == from { + r.key = to.to_string(); + } + } + let told: HashSet = g + .told + .drain() + .map(|mut t| { + if t.key == from { + t.key = to.to_string(); + } + t + }) + .collect(); + g.told = told; + if let Some(l) = g.limits.remove(from) { + g.limits.insert(to.to_string(), l); + } + } + + /// 准入第一步:天、周、月的上限,**用满了就拒**。在并发上限之前:一个这一期已经 + /// 用满的请求不用先排一轮队 + pub fn calendar(&self, key: &str, limits: &[KeyLimit]) -> Result<(), Box> { + if !limits.iter().any(|l| l.per.rolling_ms().is_none()) { + return Ok(()); + } + let now = self.clock.now_ms(); + let (out, events) = { + let mut g = self.lock(); + self.sweep(&mut g, now); + let out = self.calendar_locked(&mut g, key, limits, now); + let events = match &out { + Err(r) => self.refused_alert(&mut g, r, now), + Ok(()) => Vec::new(), + }; + (out, events) + }; + self.emit(events); + out + } + + /// 准入第三步:分钟、小时的上限。用满了先等下一个空位,最多等 `wait`;等不到就拒, + /// 带着多久之后能再来。**过了就记上**:这个请求算进滚动窗口,输入的估算占上预留。 + /// + /// 天、周、月在这里再看一遍:等并发名额、等空位的那一阵,别的请求可能把它用满了。 + pub async fn admit( + self: &Arc, + key: &str, + limits: &[KeyLimit], + ask: Ask, + wait: Duration, + ) -> Result> { + if limits.is_empty() { + return Ok(Hold::none()); + } + let deadline = self.clock.now_ms() + wait.as_millis() as i64; + loop { + let now = self.clock.now_ms(); + let step = { + let mut g = self.lock(); + self.sweep(&mut g, now); + match self.calendar_locked(&mut g, key, limits, now) { + Err(r) => { + let events = self.refused_alert(&mut g, &r, now); + Err((r, events)) + } + Ok(()) => match self.rolling_wait(&mut g, key, limits, now) { + None => Ok(self.reserve(&mut g, key, limits, ask, now)), + Some(w) => Err((Box::new(w), Vec::new())), + }, + } + }; + let refusal = match step { + Ok(seq) => { + return Ok(Hold { + owner: Some(self.clone()), + seq, + }); + } + Err((r, events)) => { + self.emit(events); + r + } + }; + // 滚动窗口:空位在等得到的时候空出来就等,等不到就拒 + if refusal.resets.is_some() || now + refusal.retry_after_ms as i64 > deadline { + return Err(refusal); + } + tokio::time::sleep(Duration::from_millis(refusal.retry_after_ms)).await; + } + } + + /// 存储层记下了一行:预留换成实数。`at_ms` 是请求开始的时刻,算在哪一期看它 + pub fn settle(&self, id: u64, at_ms: i64, rec: &Recorded) { + let now = self.clock.now_ms(); + let events = { + let mut g = self.lock(); + let key = match g.by_request.remove(&id).and_then(|s| g.held.remove(&s)) { + Some(r) => r.key, + None => rec.client.clone(), + }; + if !rec.counts() { + return; + } + let amount = rec.amount(); + let rolling = g + .limits + .get(&key) + .and_then(|ls| longest_window(ls)) + .is_some(); + let clock = self.clock.as_ref(); + let books = g.books.entry(key.clone()).or_default(); + for per in CALENDAR { + let p = books.period(clock, per, now); + if at_ms >= p.start && at_ms < p.end { + p.sum.add(&amount); + } + } + // 请求数在准入时已经进了滚动窗口 + if rolling { + books.recent.push_back(( + now, + Amount { + requests: 0, + ..amount + }, + )); + } + self.alerts(&mut g, &key, now) + }; + self.emit(events); + } + + /// 重启之后把天、周、月的数从请求记录里加回来。`since(t)` 给出从 `t` 起每把密钥 + /// 记下的那些(存储层的 `key_usage_since`)。分钟、小时的窗口从空的开始。 + /// + /// **已经过了的档不再报**:重启之前多半报过了,再报一遍就是一条重复的通知。 + pub fn rebuild(&self, mut since: impl FnMut(i64) -> Vec) { + let now = self.clock.now_ms(); + let mut sums: Vec<(LimitPer, i64, i64, HashMap)> = Vec::new(); + for per in CALENDAR { + let (start, end) = self.clock.period(per, now); + let mut by_key: HashMap = HashMap::new(); + for r in since(start).iter().filter(|r| r.counts()) { + by_key.entry(r.client.clone()).or_default().add(&r.amount()); + } + sums.push((per, start, end, by_key)); + } + let mut g = self.lock(); + for (per, start, end, by_key) in sums { + for b in g.books.values_mut() { + b.periods[calendar_index(per)] = Some(Period { + start, + end, + sum: Amount::default(), + }); + } + for (key, sum) in by_key { + g.books.entry(key).or_default().periods[calendar_index(per)] = + Some(Period { start, end, sum }); + } + } + let keys: Vec = g.limits.keys().cloned().collect(); + for key in keys { + // 只记下,不报 + let _ = self.alerts(&mut g, &key, now); + } + } + + /// 一把密钥的每条上限此刻用了多少(`GET /keys`)。在跑的请求按预留算,和准入看的 + /// 是同一个数 + pub fn view(&self, key: &str, limits: &[KeyLimit]) -> Vec { + if limits.is_empty() { + return Vec::new(); + } + let now = self.clock.now_ms(); + let mut g = self.lock(); + self.sweep(&mut g, now); + limits + .iter() + .map(|l| { + let used = self.used(&mut g, key, l, now); + let max = l.max(); + tw_api::KeyLimitView { + per: l.per.into(), + measure: l.measure().into(), + max: max.max(0) as u64, + cache_reads: l.cache_reads, + used: used.max(0) as u64, + resets_at_ms: l + .per + .rolling_ms() + .is_none() + .then(|| self.clock.period(l.per, now).1.max(0) as u64), + reached: used >= max, + } + }) + .collect() + } + + // ------------------------------------------------------------ 锁里面的 + + /// 一条上限此刻用了多少:记下的,加上在跑的请求的预留 + fn used(&self, g: &mut Inner, key: &str, l: &KeyLimit, now: i64) -> i64 { + let (m, cr) = (l.measure(), l.cache_reads); + let reserved = |r: &Reservation| match m { + LimitMeasure::Requests => 1, + LimitMeasure::Tokens => r.tokens, + LimitMeasure::Cost => r.cost, + }; + match l.per.rolling_ms() { + None => { + let p = *g.books.entry(key.to_string()).or_default().period( + self.clock.as_ref(), + l.per, + now, + ); + let held: i64 = g + .held + .values() + .filter(|r| r.key == key && r.at_ms >= p.start && r.at_ms < p.end) + .map(reserved) + .sum(); + p.sum.of(m, cr).saturating_add(held) + } + Some(w) => { + let cutoff = now - w; + let books = g.books.get(key); + let done: i64 = books + .map(|b| { + b.recent + .iter() + .filter(|(t, _)| *t > cutoff) + .map(|(_, a)| a.of(m, cr)) + .sum() + }) + .unwrap_or(0); + // 请求数在准入时已经进了窗口,预留里只算 token 和费用 + let held: i64 = g + .held + .values() + .filter(|r| r.key == key && r.at_ms > cutoff) + .map(|r| match m { + LimitMeasure::Requests => 0, + _ => reserved(r), + }) + .sum(); + done.saturating_add(held) + } + } + } + + fn calendar_locked( + &self, + g: &mut Inner, + key: &str, + limits: &[KeyLimit], + now: i64, + ) -> Result<(), Box> { + for l in limits.iter().filter(|l| l.per.rolling_ms().is_none()) { + let used = self.used(g, key, l, now); + if used >= l.max() { + let end = self.clock.period(l.per, now).1; + return Err(Box::new(Refusal { + key: key.to_string(), + limit: l.clone(), + used, + retry_after_ms: (end - now).max(1) as u64, + resets: Some((end, self.clock.show(end))), + })); + } + } + Ok(()) + } + + /// 滚动窗口要等多久才有空位。不用等是 None;要等的那条里等得最久的那一条的拒绝 + fn rolling_wait( + &self, + g: &mut Inner, + key: &str, + limits: &[KeyLimit], + now: i64, + ) -> Option { + let mut worst: Option = None; + for l in limits { + let Some(w) = l.per.rolling_ms() else { + continue; + }; + let used = self.used(g, key, l, now); + let max = l.max(); + if used < max { + continue; + } + // 窗口里的每一笔按时间先后滑出去,滑到用量低于上限的那一刻就是空位 + let (m, cr) = (l.measure(), l.cache_reads); + let cutoff = now - w; + let mut items: Vec<(i64, i64)> = g + .books + .get(key) + .map(|b| { + b.recent + .iter() + .filter(|(t, _)| *t > cutoff) + .map(|(t, a)| (*t, a.of(m, cr))) + .collect() + }) + .unwrap_or_default(); + if m != LimitMeasure::Requests { + items.extend( + g.held + .values() + .filter(|r| r.key == key && r.at_ms > cutoff) + .map(|r| { + let v = if m == LimitMeasure::Tokens { + r.tokens + } else { + r.cost + }; + (r.at_ms, v) + }), + ); + } + items.retain(|(_, v)| *v > 0); + items.sort_by_key(|(t, _)| *t); + let mut left = used; + let mut free_at = now + w; + for (t, v) in items { + left -= v; + if left < max { + free_at = t + w; + break; + } + } + let wait = (free_at - now).max(1) as u64; + if worst.as_ref().is_none_or(|r| wait > r.retry_after_ms) { + worst = Some(Refusal { + key: key.to_string(), + limit: l.clone(), + used, + retry_after_ms: wait, + resets: None, + }); + } + } + worst + } + + /// 记上这个请求:滚动窗口里的一个请求,和它的预留。交回预留的号 + fn reserve(&self, g: &mut Inner, key: &str, limits: &[KeyLimit], ask: Ask, now: i64) -> u64 { + if let Some(w) = longest_window(limits) { + let books = g.books.entry(key.to_string()).or_default(); + books.recent.push_back(( + now, + Amount { + requests: 1, + ..Default::default() + }, + )); + while books.recent.front().is_some_and(|(t, _)| *t <= now - w) { + books.recent.pop_front(); + } + } + g.seq += 1; + let seq = g.seq; + g.held.insert( + seq, + Reservation { + key: key.to_string(), + at_ms: now, + tokens: ask.tokens.min(i64::MAX as u64) as i64, + cost: ask.cost_micros.max(0), + request: None, + closed_since: None, + }, + ); + seq + } + + fn bind(&self, seq: u64, request: u64) { + let mut g = self.lock(); + if let Some(r) = g.held.get_mut(&seq) { + r.request = Some(request); + g.by_request.insert(request, seq); + } + } + + fn release(&self, seq: u64) { + let mut g = self.lock(); + if let Some(r) = g.held.remove(&seq) + && let Some(id) = r.request + { + g.by_request.remove(&id); + } + } + + /// 放掉结束了太久、存储层却一直没交来那一行的预留(见 [`GRACE_MS`])。顺手把滚动 + /// 窗口里滑出去的扔掉 + fn sweep(&self, g: &mut Inner, now: i64) { + let mut gone = Vec::new(); + for (seq, r) in g.held.iter_mut() { + let Some(id) = r.request else { continue }; + if self.bus.is_open(id) { + r.closed_since = None; + continue; + } + let since = *r.closed_since.get_or_insert(now); + if now - since >= GRACE_MS { + gone.push((*seq, id)); + } + } + for (seq, id) in gone { + g.held.remove(&seq); + g.by_request.remove(&id); + } + let windows: HashMap = g + .limits + .iter() + .filter_map(|(k, ls)| longest_window(ls).map(|w| (k.clone(), w))) + .collect(); + for (key, b) in g.books.iter_mut() { + match windows.get(key) { + Some(w) => { + while b.recent.front().is_some_and(|(t, _)| *t <= now - w) { + b.recent.pop_front(); + } + } + None => b.recent.clear(), + } + } + } + + /// 天、周、月的上限到了八成、到了顶:**每一期、每一档报一次**。只看记下的数 —— + /// 预留是估的,按它报的话结算之后可能又退回去,那是一条假警报 + fn alerts(&self, g: &mut Inner, key: &str, now: i64) -> Vec { + let Some(limits) = g.limits.get(key).cloned() else { + return Vec::new(); + }; + let mut out = Vec::new(); + for l in limits.iter().filter(|l| l.per.rolling_ms().is_none()) { + let p = *g.books.entry(key.to_string()).or_default().period( + self.clock.as_ref(), + l.per, + now, + ); + let used = p.sum.of(l.measure(), l.cache_reads); + let max = l.max(); + let reached = used >= max; + let near = used.saturating_mul(5) >= max.saturating_mul(4); + if !near { + continue; + } + let told = |reached: bool| Told { + key: key.to_string(), + per: l.per, + measure: l.measure(), + cache_reads: l.cache_reads, + max, + period_start: p.start, + reached, + }; + // 一下子从八成以下跳过了顶:只报到顶,八成那一档一并算报过 + let new_near = g.told.insert(told(false)); + let new_reached = reached && g.told.insert(told(true)); + if new_reached || (new_near && !reached) { + out.push(self.alert(key, l, used, reached, p.end, now)); + } + } + out + } + + /// 被拒了:用满了的那一条一定到顶了,没报过就报 + fn refused_alert(&self, g: &mut Inner, r: &Refusal, now: i64) -> Vec { + let Some((end, _)) = r.resets else { + return Vec::new(); + }; + let start = self.clock.period(r.limit.per, now).0; + let told = |reached: bool| Told { + key: r.key.clone(), + per: r.limit.per, + measure: r.limit.measure(), + cache_reads: r.limit.cache_reads, + max: r.limit.max(), + period_start: start, + reached, + }; + g.told.insert(told(false)); + if g.told.insert(told(true)) { + vec![self.alert(&r.key, &r.limit, r.used, true, end, now)] + } else { + Vec::new() + } + } + + fn alert( + &self, + key: &str, + l: &KeyLimit, + used: i64, + reached: bool, + resets: i64, + now: i64, + ) -> tw_api::Event { + tw_api::Event::KeyLimitAlert { + id: self.bus.next_id(), + key: key.to_string(), + per: l.per.into(), + measure: l.measure().into(), + max: l.max().max(0) as u64, + used: used.max(0) as u64, + cache_reads: l.cache_reads, + reached, + resets_at_ms: resets.max(0) as u64, + at_ms: now.max(0) as u64, + } + } + + /// 在锁外面报 + fn emit(&self, events: Vec) { + for e in events { + self.bus.emit(e); + } + } +} + +/// 一把密钥的分钟、小时上限里最长的那个窗口。没有就是 None +fn longest_window(limits: &[KeyLimit]) -> Option { + limits.iter().filter_map(|l| l.per.rolling_ms()).max() +} + +/// 这把密钥用得到、却没有价格的模型(`ClientView.unpriced_models`)。 +/// +/// **和 `GET /v1/models` 同一份清单**:目录里这把密钥的 `allow` 放行的名称(不挑格式 —— +/// 一把密钥给哪种客户端用都行),每一个看提供它的每一家:按量计费、价目表里查不到发给 +/// 那一家的名字的,就是它。别名按对到那一家的名称查。只是查表,不发请求。 +pub fn unpriced_models(state: &crate::AppState, cfg: &tw_config::Config, key: &str) -> Vec { + let catalog = state.catalog.load(); + let book = state.pricing.load(); + let allow = crate::models::key_allow(cfg, key); + catalog + .resolve_allowed(None, allow) + .into_iter() + .filter(|name| { + catalog.providers_for(name).iter().any(|p| { + let Some(provider) = cfg.providers.iter().find(|x| &x.name == p) else { + return false; + }; + if provider.billing == tw_config::Billing::Free { + return false; + } + let sent = catalog + .served(name) + .iter() + .find(|(sp, _)| sp == p) + .map_or(name.as_str(), |(_, m)| m.as_str()); + book.resolve_for(p, sent).is_none() + }) + }) + .collect() +} diff --git a/crates/tw-gateway/src/key_limits/tests.rs b/crates/tw-gateway/src/key_limits/tests.rs new file mode 100644 index 00000000..21ef5f0e --- /dev/null +++ b/crates/tw-gateway/src/key_limits/tests.rs @@ -0,0 +1,539 @@ +use super::*; + +/// 东八区。2026-10-05 是周一 +const CST: i32 = 8 * 3600; + +fn at(s: &str) -> i64 { + chrono::DateTime::parse_from_rfc3339(s) + .unwrap() + .timestamp_millis() +} + +fn parse(yaml: &str) -> Vec { + serde_yaml_ng::from_str(yaml).unwrap() +} + +struct Bed { + limits: Arc, + bus: tw_observe::EventBus, + rx: tokio::sync::broadcast::Receiver, + set: Vec, +} + +/// 密钥 `k` 带着这几条上限,时钟从 `now` 起 +fn bed(now: &str, yaml: &str) -> Bed { + let bus = tw_observe::EventBus::new(); + let rx = bus.subscribe(); + let limits = Arc::new(KeyLimits::with_clock( + bus.clone(), + Arc::new(TestClock::new(at(now), CST)), + )); + let set = parse(yaml); + let cfg = tw_config::Config { + clients: vec![tw_config::Client { + name: "k".into(), + key: "tw-k".into(), + limits: set.clone(), + ..Default::default() + }], + ..Default::default() + }; + limits.configure(&cfg); + Bed { + limits, + bus, + rx, + set, + } +} + +impl Bed { + fn now(&self) -> i64 { + self.limits.clock.now_ms() + } + + /// 不等:准入,过了就交回预留 + async fn try_admit(&self, ask: Ask) -> Result> { + self.limits.admit("k", &self.set, ask, Duration::ZERO).await + } + + /// 准入一个请求、给它一个号、让它「在跑」(总线上开始了) + async fn run(&self, id: u64, ask: Ask) -> Result<(), Box> { + self.limits.calendar("k", &self.set)?; + let hold = self.try_admit(ask).await?; + self.bus.emit(started(id)); + hold.bind(id); + Ok(()) + } + + /// 存储层记下了这一行:请求在总线上结束,然后结算 + fn done(&self, id: u64, rec: Recorded) { + self.bus.emit(tw_api::Event::RequestFinished { + id, + model: "m".into(), + status: 200, + bytes: 0, + duration_ms: 0, + usage: None, + tokens_per_sec: None, + answered_model: None, + }); + self.limits.settle(id, self.now(), &rec); + } + + fn used(&self) -> Vec { + self.limits + .view("k", &self.set) + .iter() + .map(|v| v.used) + .collect() + } + + fn alerts(&mut self) -> Vec<(u64, bool)> { + let mut out = Vec::new(); + while let Ok(e) = self.rx.try_recv() { + if let tw_api::Event::KeyLimitAlert { used, reached, .. } = e { + out.push((used, reached)); + } + } + out + } +} + +fn started(id: u64) -> tw_api::Event { + tw_api::Event::RequestStarted { + id, + client: "k".into(), + client_hint: None, + session: None, + peer: None, + key_masked: None, + route: "default".into(), + rule: "r".into(), + group: None, + rewritten_by: vec![], + provider: "p".into(), + billing: tw_api::Billing::PerToken, + model: "m".into(), + method: "POST".into(), + path: "/v1/messages".into(), + input_estimate: None, + session_log_bytes: None, + at_ms: 0, + } +} + +/// 记下的一行:发往了上游、成功了 +fn row(input: u64, output: u64, cache_read: u64, cost: i64) -> Recorded { + Recorded { + client: "k".into(), + path: "/v1/messages".into(), + local: false, + error_code: None, + attempted: true, + requests: 1, + input, + output, + cache_read, + cache_write: 0, + cost_micros: cost, + } +} + +fn ask(tokens: u64, cost: i64) -> Ask { + Ask { + tokens, + cost_micros: cost, + } +} + +// ───────────────────────────────────────────────── 自然周期 + +#[tokio::test(start_paused = true)] +async fn a_day_is_used_up_then_refused_until_local_midnight() { + let b = bed("2026-10-05T23:59:00+08:00", "[{per: day, requests: 2}]"); + b.run(1, Ask::default()).await.unwrap(); + b.run(2, Ask::default()).await.unwrap(); + // 两个都还在跑:预留已经占满 + let r = b.run(3, Ask::default()).await.unwrap_err(); + assert_eq!(r.used, 2); + assert_eq!(r.retry_after_ms, 60_000, "到本地零点还有一分钟"); + let (end, text) = r.resets.clone().unwrap(); + assert_eq!(end, at("2026-10-06T00:00:00+08:00")); + assert_eq!(text, "2026-10-06 00:00 +08:00"); + let m = r.msg(); + assert_eq!(m.code, "gw.key_limit.requests_per_period"); + assert_eq!( + m.text, + "Gateway key `k` has reached its limit of 2 requests per day: 2 so far. It resets at \ + 2026-10-06 00:00 +08:00." + ); + assert_eq!(m.arg("resets_at_ms"), end.to_string()); + // 结算之后还是两个:一个请求就是一个 + b.done(1, row(10, 10, 0, 0)); + b.done(2, row(10, 10, 0, 0)); + assert_eq!(b.used(), [2]); + assert!(b.run(3, Ask::default()).await.is_err()); + // 过了零点,新的一天从 0 起 + tokio::time::advance(Duration::from_secs(61)).await; + assert_eq!(b.used(), [0]); + b.run(4, Ask::default()).await.unwrap(); +} + +#[tokio::test(start_paused = true)] +async fn a_week_starts_on_monday_and_a_month_on_the_first() { + // 周日晚上,也是这个月的最后一天 + let b = bed( + "2026-05-31T23:00:00+08:00", + "[{per: week, requests: 1}, {per: month, requests: 5}]", + ); + b.run(1, Ask::default()).await.unwrap(); + b.done(1, row(1, 1, 0, 0)); + let r = b.run(2, Ask::default()).await.unwrap_err(); + assert_eq!(r.limit.per, LimitPer::Week); + assert_eq!(r.resets.unwrap().0, at("2026-06-01T00:00:00+08:00")); + assert_eq!(b.used(), [1, 1]); + // 周一零点:周和月都重新算 + tokio::time::advance(Duration::from_secs(3600)).await; + assert_eq!(b.used(), [0, 0]); + b.run(2, Ask::default()).await.unwrap(); + b.done(2, row(1, 1, 0, 0)); + // 周二:这一周用过了,这个月还有 + tokio::time::advance(Duration::from_secs(86_400)).await; + assert_eq!(b.used(), [1, 1]); + assert_eq!( + b.run(3, Ask::default()).await.unwrap_err().limit.per, + LimitPer::Week + ); +} + +/// 天、周、月按请求开始的时刻归期:零点前开始、零点后才记下的那一行算前一天的 —— +/// 从库里加回来时也是这样算的 +#[tokio::test(start_paused = true)] +async fn a_request_that_started_yesterday_counts_for_yesterday() { + let b = bed("2026-10-05T23:59:59+08:00", "[{per: day, cost: 1}]"); + let begun = b.now(); + b.run(1, Ask::default()).await.unwrap(); + tokio::time::advance(Duration::from_secs(5)).await; + b.limits.settle(1, begun, &row(10, 10, 0, 900_000)); + assert_eq!(b.used(), [0]); +} + +// ───────────────────────────────────────────────── 滚动窗口 + +#[tokio::test(start_paused = true)] +async fn a_rolling_minute_waits_for_its_next_slot_or_refuses() { + let b = bed("2026-10-05T10:00:00+08:00", "[{per: minute, requests: 2}]"); + let t0 = b.now(); + b.run(1, Ask::default()).await.unwrap(); + tokio::time::advance(Duration::from_secs(10)).await; + b.run(2, Ask::default()).await.unwrap(); + // 第一个要到 60 秒时才滑出去:50 秒,等不了 30 秒的就拒 + let r = b + .limits + .admit("k", &b.set, Ask::default(), Duration::from_secs(30)) + .await + .unwrap_err(); + assert_eq!(r.retry_after_ms, 50_000); + assert!(r.resets.is_none()); + let m = r.msg(); + assert_eq!(m.code, "gw.key_limit.requests_rolling"); + assert_eq!( + m.text, + "Gateway key `k` has reached its limit of 2 requests per minute: 2 in the last minute. \ + Try again in 50 s." + ); + assert_eq!(b.limits.clock.now_ms(), t0 + 10_000, "拒的时候不等"); + // 再过 25 秒,空位 25 秒后出来:等得到就等 + tokio::time::advance(Duration::from_secs(25)).await; + let hold = b + .limits + .admit("k", &b.set, Ask::default(), Duration::from_secs(30)) + .await + .unwrap(); + drop(hold); + assert_eq!(b.limits.clock.now_ms(), t0 + 60_000, "等到第一个滑出窗口"); +} + +#[tokio::test(start_paused = true)] +async fn tokens_in_a_rolling_hour_count_when_they_are_settled() { + let b = bed("2026-10-05T10:00:00+08:00", "[{per: hour, tokens: 1000}]"); + // 估了 600 占着 + b.run(1, ask(600, 0)).await.unwrap(); + assert_eq!(b.used(), [600]); + // 跑了 50 分钟才结束,用了 900:记在结束的那一刻 + tokio::time::advance(Duration::from_secs(50 * 60)).await; + b.done(1, row(500, 400, 0, 0)); + assert_eq!(b.used(), [900]); + b.run(2, ask(200, 0)).await.unwrap(); + // 这时候已经超了:要等 900 那一笔滑出去,一小时 + let r = b.try_admit(ask(1, 0)).await.unwrap_err(); + assert_eq!(r.used, 1100); + assert_eq!(r.retry_after_ms, 3_600_000); + assert_eq!(r.msg().code, "gw.key_limit.tokens_rolling"); +} + +// ───────────────────────────────────────────────── 预留与结算 + +#[tokio::test(start_paused = true)] +async fn a_reservation_is_replaced_by_what_was_recorded() { + let b = bed( + "2026-10-05T10:00:00+08:00", + "[{per: day, tokens: 1000}, {per: day, cost: 1}]", + ); + b.run(1, ask(800, 300_000)).await.unwrap(); + b.run(2, ask(100, 50_000)).await.unwrap(); + assert_eq!(b.used(), [900, 350_000], "在跑的按估算占着"); + // 流式的跑完了:实数比估的多 + b.done(1, row(700, 300, 0, 420_000)); + assert_eq!(b.used(), [1100, 470_000]); + // 用满了:新的请求被拒,哪怕只要一点点 + assert!(b.run(3, ask(1, 1)).await.is_err()); + // 客户端走掉的、断在半路的:记下多少算多少 + b.done(2, row(10, 0, 0, 3_000)); + assert_eq!(b.used(), [1010, 423_000]); +} + +#[tokio::test(start_paused = true)] +async fn a_hold_that_never_got_a_request_number_gives_its_share_back() { + let b = bed("2026-10-05T10:00:00+08:00", "[{per: day, tokens: 1000}]"); + let hold = b.try_admit(ask(900, 0)).await.unwrap(); + assert_eq!(b.used(), [900]); + drop(hold); + assert_eq!(b.used(), [0]); +} + +/// 存储层丢了那一行(总线落后):请求结束一分钟之后,预留放掉 +#[tokio::test(start_paused = true)] +async fn a_reservation_whose_row_never_comes_is_let_go_after_its_request_ended() { + let b = bed("2026-10-05T10:00:00+08:00", "[{per: day, tokens: 1000}]"); + b.run(1, ask(900, 0)).await.unwrap(); + tokio::time::advance(Duration::from_secs(600)).await; + assert_eq!(b.used(), [900], "还在跑的一直占着"); + b.bus.emit(tw_api::Event::RequestCancelled { + id: 1, + model: "m".into(), + status: None, + bytes: 0, + duration_ms: 0, + usage: None, + answered_model: None, + }); + assert_eq!(b.used(), [900], "结束之后先等存储层"); + tokio::time::advance(Duration::from_secs(61)).await; + assert_eq!(b.used(), [0]); +} + +#[tokio::test(start_paused = true)] +async fn what_does_not_count_does_not_count() { + let b = bed("2026-10-05T10:00:00+08:00", "[{per: day, requests: 100}]"); + let not = |f: &dyn Fn(&mut Recorded)| { + let mut r = row(1, 1, 0, 0); + f(&mut r); + r + }; + let cases = [ + not(&|r| r.local = true), + not(&|r| r.path = "/v1/messages/count_tokens".into()), + not(&|r| r.path = "/v1beta/models/gemini-2.5-pro:countTokens".into()), + not(&|r| r.path = "/v1/responses/input_tokens".into()), + // 上限自己拒的 + not(&|r| { + r.attempted = false; + r.error_code = Some("gw.key_limit.requests_per_period".into()); + }), + // 路由就拒了,一跳都没有 + not(&|r| { + r.attempted = false; + r.error_code = Some("gw.route.denied".into()); + }), + not(&|r| { + r.attempted = false; + r.error_code = Some("gw.model.no_upstream_available".into()); + }), + ]; + for (i, r) in cases.iter().enumerate() { + assert!(!r.counts(), "{r:?}"); + b.limits.settle(100 + i as u64, b.now(), r); + } + assert_eq!(b.used(), [0]); + // 准入过了、在某一跳上被阶段二的规则拒了:算 + let mut hop_denied = row(0, 0, 0, 0); + hop_denied.error_code = Some("gw.route.denied".into()); + // 准入过了、第一跳还没回话客户端就走了:一跳都没记下,也算 + let mut gone = row(0, 0, 0, 0); + gone.attempted = false; + for r in [hop_denied, gone] { + assert!(r.counts(), "{r:?}"); + b.limits.settle(200, b.now(), &r); + } + assert_eq!(b.used(), [2]); +} + +#[tokio::test(start_paused = true)] +async fn unpriced_and_free_cost_nothing_and_cache_reads_count_only_when_asked() { + let b = bed( + "2026-10-05T10:00:00+08:00", + "[{per: day, cost: 1}, {per: day, tokens: 1000}, \ + {per: day, tokens: 1000, cache_reads: true}]", + ); + // 没有价格、不计费:存储层记的是 None / 0,交到这里都是 0 + b.done(1, row(100, 10, 500, 0)); + b.done(2, row(100, 10, 500, 0)); + assert_eq!(b.used(), [0, 220, 1220]); + let m = b.run(3, Ask::default()).await.unwrap_err().msg(); + assert_eq!(m.code, "gw.key_limit.tokens_per_period"); + assert_eq!(m.arg("used"), "1220"); +} + +#[test] +fn cost_reads_as_dollars() { + assert_eq!(dollars(5_000_000), "$5.00"); + assert_eq!(dollars(4_200), "$0.0042"); + assert_eq!(dollars(1_234_567), "$1.234567"); + assert_eq!(dollars(0), "$0.00"); +} + +#[tokio::test(start_paused = true)] +async fn a_cost_refusal_names_the_amounts_in_dollars() { + let b = bed("2026-10-05T10:00:00+08:00", "[{per: day, cost: 5}]"); + b.done(1, row(1, 1, 0, 5_030_000)); + let r = b.run(2, Ask::default()).await.unwrap_err(); + let m = r.msg(); + assert_eq!(m.code, "gw.key_limit.cost_per_period"); + assert_eq!( + m.text, + "Gateway key `k` has reached its limit of $5.00 per day: $5.03 spent so far. It resets \ + at 2026-10-06 00:00 +08:00." + ); + let e = r.error(); + assert_eq!(e.source, crate::error::Source::RateLimited); + assert_eq!( + e.retry, + Some(crate::error::Retry { + after_ms: (at("2026-10-06T00:00:00+08:00") - at("2026-10-05T10:00:00+08:00")) as u64, + until_reset: true + }) + ); +} + +// ───────────────────────────────────────────────── 重启、改名 + +#[tokio::test(start_paused = true)] +async fn a_restart_adds_the_periods_back_from_the_store() { + let mut b = bed( + "2026-10-07T10:00:00+08:00", + "[{per: day, cost: 0.12}, {per: week, cost: 10}, {per: month, cost: 10}, \ + {per: minute, requests: 1}]", + ); + let day = at("2026-10-07T00:00:00+08:00"); + let week = at("2026-10-05T00:00:00+08:00"); + let month = at("2026-10-01T00:00:00+08:00"); + let mut asked = Vec::new(); + b.limits.rebuild(|since| { + asked.push(since); + // 按开始时刻:今天的、这周早些时候的、这个月早些时候的 + let mut rows = vec![row(0, 0, 0, 100_000)]; + if since <= week { + rows.push(row(0, 0, 0, 200_000)); + } + if since <= month { + rows.push(row(0, 0, 0, 400_000)); + } + // 不算的那几种,加回来时一样不算 + let mut refused = row(0, 0, 0, 999_000); + refused.attempted = false; + refused.error_code = Some("gw.route.denied".into()); + rows.push(refused); + rows + }); + assert_eq!(asked, [day, week, month]); + assert_eq!( + b.used(), + [100_000, 300_000, 700_000, 0], + "分钟窗口从空的开始" + ); + // 天的那一条加回来就过了八成:重启前多半报过了,不再报 + assert!(b.alerts().is_empty()); + // 再花一点就到顶了:这一档没报过,报;八成那一档不补 + b.done(1, row(0, 0, 0, 50_000)); + assert_eq!(b.alerts(), [(150_000, true)]); +} + +#[tokio::test(start_paused = true)] +async fn a_renamed_key_keeps_what_it_used() { + let b = bed("2026-10-05T10:00:00+08:00", "[{per: day, requests: 3}]"); + b.run(1, Ask::default()).await.unwrap(); + b.done(1, row(1, 1, 0, 0)); + b.run(2, Ask::default()).await.unwrap(); + b.limits.rename("k", "k2"); + assert_eq!(b.limits.view("k2", &b.set)[0].used, 2); + // 在跑的那一个结算到新名字上(存储层那一行记的还是旧名字) + b.done(2, row(1, 1, 0, 0)); + assert_eq!(b.limits.view("k2", &b.set)[0].used, 2); + assert_eq!(b.limits.view("k", &b.set)[0].used, 0); +} + +// ───────────────────────────────────────────────── 提醒 + +#[tokio::test(start_paused = true)] +async fn eighty_percent_and_the_limit_are_each_told_once_a_period() { + let mut b = bed("2026-10-05T23:00:00+08:00", "[{per: day, requests: 5}]"); + for id in 1..=3 { + b.run(id, Ask::default()).await.unwrap(); + b.done(id, row(1, 1, 0, 0)); + } + assert!(b.alerts().is_empty(), "六成不报"); + b.run(4, Ask::default()).await.unwrap(); + b.done(4, row(1, 1, 0, 0)); + assert_eq!(b.alerts(), [(4, false)]); + b.run(5, Ask::default()).await.unwrap(); + b.done(5, row(1, 1, 0, 0)); + assert_eq!(b.alerts(), [(5, true)]); + // 之后被拒:不再报 + assert!(b.run(6, Ask::default()).await.is_err()); + assert!(b.alerts().is_empty()); + // 下一天重新算 + tokio::time::advance(Duration::from_secs(3600)).await; + for id in 7..=11 { + b.run(id, Ask::default()).await.unwrap(); + b.done(id, row(1, 1, 0, 0)); + } + assert_eq!(b.alerts(), [(4, false), (5, true)]); +} + +#[tokio::test(start_paused = true)] +async fn a_refusal_from_requests_still_in_flight_is_told_too() { + let mut b = bed("2026-10-05T10:00:00+08:00", "[{per: day, requests: 1}]"); + b.run(1, Ask::default()).await.unwrap(); + // 还没结算,靠预留拒的:到顶这一档照样报,带着拒绝时看到的数 + assert!(b.run(2, Ask::default()).await.is_err()); + assert_eq!(b.alerts(), [(1, true)]); + b.done(1, row(1, 1, 0, 0)); + assert!(b.alerts().is_empty(), "已经报过了"); +} + +#[tokio::test(start_paused = true)] +async fn the_view_says_how_much_is_left_and_when_it_resets() { + let b = bed( + "2026-10-05T10:00:00+08:00", + "[{per: day, cost: 2.5}, {per: minute, requests: 1}]", + ); + b.run(1, ask(0, 0)).await.unwrap(); + b.done(1, row(1, 1, 0, 2_500_000)); + let v = b.limits.view("k", &b.set); + assert_eq!(v[0].measure, tw_api::LimitMeasure::Cost); + assert_eq!( + (v[0].max, v[0].used, v[0].reached), + (2_500_000, 2_500_000, true) + ); + assert_eq!( + v[0].resets_at_ms, + Some(at("2026-10-06T00:00:00+08:00") as u64) + ); + assert_eq!(v[1].per, tw_api::LimitPer::Minute); + assert_eq!( + (v[1].used, v[1].reached, v[1].resets_at_ms), + (1, true, None) + ); +} diff --git a/crates/tw-gateway/src/lib.rs b/crates/tw-gateway/src/lib.rs index 315fcb7c..3a111a64 100644 --- a/crates/tw-gateway/src/lib.rs +++ b/crates/tw-gateway/src/lib.rs @@ -26,6 +26,7 @@ pub mod glm; pub mod guard; pub mod health; pub mod hint; +pub mod key_limits; pub mod l1; pub mod l3; pub mod latency; diff --git a/crates/tw-gateway/src/server/pipeline.rs b/crates/tw-gateway/src/server/pipeline.rs index 2424c32e..e528dc6b 100644 --- a/crates/tw-gateway/src/server/pipeline.rs +++ b/crates/tw-gateway/src/server/pipeline.rs @@ -20,6 +20,7 @@ use crate::forward; use crate::state::{AppState, Runtime}; use tw_types::msg; +mod admission; mod hop; mod opening; mod plug; @@ -116,18 +117,22 @@ pub(super) async fn pipeline( } }; - // 管线第 3 步:这把密钥自己的并发上限。**等,不拒绝** —— 理由在 - // `crate::limits`。放在路由之后:被规则挡下的请求不用先等一轮。 + // 管线第 3 步:这把密钥的用量上限和并发上限(见 `admission`)。放在路由之后:被规则 + // 挡下的请求不用先等一轮。被用量上限拒绝的照样留一行 // - // **通行证交给回程,跟着响应体走**(见 `relay`):放在这里的话它在响应头交出去的那一刻 - // 就还了,一条还在流的回答不再算数,上限管的只是等响应头的那一段 - let limit = rt - .config - .clients - .iter() - .find(|c| c.name == req.client_name) - .and_then(|c| c.max_concurrent); - let pass = state.gate.acquire(&req.client_name, limit).await; + // **并发的通行证交给回程,跟着响应体走**(见 `relay`):放在这里的话它在响应头交出去的 + // 那一刻就还了,一条还在流的回答不再算数,上限管的只是等响应头的那一段 + let admission::Admitted { pass, hold } = admission::admit( + &state, + &rt, + &req, + &reading, + &choice, + &decision, + fp.as_deref(), + ending, + ) + .await?; // 管线第 4 步:内容过滤先下结论,不发事件。删过的话,后面一律用删过的那一份 let screening = screen(&rt, &mut req, &mut reading); @@ -141,6 +146,8 @@ pub(super) async fn pipeline( fp.as_deref(), ending, ); + // 用量上限的预留跟着请求号走,等存储层记下这一行时换成实数 + hold.bind(started.id); // 结论挂在请求号上报。**拒绝的也在开始之后**:被拒是一次来源为 `denied` 的失败, // 流量里照样留一行;一个字节都不发 let provider = started.alive.first().map(String::as_str).unwrap_or(""); diff --git a/crates/tw-gateway/src/server/pipeline/admission.rs b/crates/tw-gateway/src/server/pipeline/admission.rs new file mode 100644 index 00000000..8ccdea62 --- /dev/null +++ b/crates/tw-gateway/src/server/pipeline/admission.rs @@ -0,0 +1,146 @@ +//! 管线第 3 步:这把密钥的用量上限和并发上限。 +//! +//! **顺序是设计**: +//! +//! 1. 天、周、月的上限([`crate::key_limits`])—— 这一期用满了直接拒:到下一期之前等多久 +//! 都一样,不用先排一轮并发的队; +//! 2. 并发上限([`crate::limits`])—— 等,不拒,理由在那儿; +//! 3. 分钟、小时的上限 —— 下一个空位在 `slot_wait_secs` 之内空出来就等,等不到就拒,并说清 +//! 多久之后再来。过了就把这个请求记上,按输入估一个数占着,等存储层记下它那一行时换成 +//! 实数。 +//! +//! 被上限拒绝的请求**照样开始、照样留一行**(和路由拒绝的一样,见 +//! [`crate::server::routed_nowhere`]):流量里看得见它被哪一条上限拒了。数 token 的请求 +//! 不算用量,只过并发上限。 + +use super::{Inbound, look, open, redaction}; +use crate::error::GatewayError; +use crate::key_limits::{Ask, Hold, Refusal}; +use crate::server::Choice; +use crate::state::{AppState, Runtime}; + +/// 过了这一步的请求手里拿着的。 +pub(super) struct Admitted { + /// 并发闸门的通行证。丢掉就归还 + pub(super) pass: crate::limits::Pass, + /// 用量上限的预留。开始事件之后交给请求号([`Hold::bind`]),丢掉就放掉 + pub(super) hold: Hold, +} + +#[allow(clippy::too_many_arguments)] +pub(super) async fn admit( + state: &AppState, + rt: &Runtime, + req: &Inbound, + reading: &crate::client_api::Reading, + choice: &Choice, + decision: &tw_engine::Decision, + fp: Option<&str>, + ending: &mut Option, +) -> Result { + let key = rt.config.clients.iter().find(|c| c.name == req.client_name); + let max_concurrent = key.and_then(|c| c.max_concurrent); + let limits: &[tw_config::KeyLimit] = match key { + Some(c) if !crate::key_limits::uncounted(req.uri.path()) => &c.limits, + _ => &[], + }; + if let Err(r) = state.key_limits.calendar(&req.client_name, limits) { + return Err(refused(state, rt, req, reading, choice, fp, ending, &r)); + } + let pass = state.gate.acquire(&req.client_name, max_concurrent).await; + let ask = if limits.is_empty() { + Ask::default() + } else { + ask(state, rt, reading, decision, limits) + }; + let wait = crate::key_limits::slot_wait(&rt.config); + match state + .key_limits + .admit(&req.client_name, limits, ask, wait) + .await + { + Ok(hold) => Ok(Admitted { pass, hold }), + Err(r) => Err(refused(state, rt, req, reading, choice, fp, ending, &r)), + } +} + +/// 这个请求要占多少:输入 token 的估算(解不开的请求没有估算,占 0),和设了费用上限 +/// 时按头一个候选、发给它的名字算的输入费用。没有价格、不计费的是 0 —— 和结算时一样。 +fn ask( + state: &AppState, + rt: &Runtime, + reading: &crate::client_api::Reading, + decision: &tw_engine::Decision, + limits: &[tw_config::KeyLimit], +) -> Ask { + let facts = &reading.facts; + let tokens = if matches!(reading.decoded, Some(Ok(_))) { + facts.input_tokens + } else { + 0 + }; + let priced = limits + .iter() + .any(|l| l.measure() == tw_config::LimitMeasure::Cost); + let cost_micros = if priced && tokens > 0 { + input_cost(state, rt, facts, decision, tokens).unwrap_or(0) + } else { + 0 + }; + Ask { + tokens, + cost_micros, + } +} + +fn input_cost( + state: &AppState, + rt: &Runtime, + facts: &tw_engine::RequestFacts, + decision: &tw_engine::Decision, + tokens: u64, +) -> Option { + let first = decision.candidates.first()?; + let p = rt.config.providers.iter().find(|p| &p.name == first)?; + let asked = rt.engine.asked_of(facts, decision, &p.name, None); + let catalog = state.catalog.load(); + let allow = crate::models::key_allow(&rt.config, &facts.client); + let model = crate::sent::name(&rt.config, &catalog, decision, p, &asked, allow).ok()?; + let usage = tw_pricing::Usage { + input: tokens, + ..Default::default() + }; + crate::quote::quote(&state.pricing.load(), &p.name, &model, &usage, p.billing).cost_micros +} + +/// 被上限拒了:照样开始、留一行(尝试链是空的),交回给客户端的那个 429。 +#[allow(clippy::too_many_arguments)] +fn refused( + state: &AppState, + rt: &Runtime, + req: &Inbound, + reading: &crate::client_api::Reading, + choice: &Choice, + fp: Option<&str>, + ending: &mut Option, + r: &Refusal, +) -> GatewayError { + tracing::info!(key = %req.client_name, per = r.limit.per.word(), "a usage limit of the key refused the request"); + let why = r.error(); + // 一个字节都没发出去;存下来的请求照样按这一档换、打码 + let (_, ledger) = look(rt, req); + let (id, _) = open( + state, + req, + reading, + choice, + ("", tw_api::Billing::PerToken), + fp, + ending, + redaction(rt, ledger), + ); + state + .bus + .emit(crate::server::routed_nowhere(id, choice.clone())); + why +} diff --git a/crates/tw-gateway/src/server/upgrade.rs b/crates/tw-gateway/src/server/upgrade.rs index 53b742e8..796fa773 100644 --- a/crates/tw-gateway/src/server/upgrade.rs +++ b/crates/tw-gateway/src/server/upgrade.rs @@ -228,7 +228,38 @@ pub(super) async fn ws_upgrade( .headers_for(provider, http) .await .map_err(|e| GatewayError::config(crate::state::credential_failed(e, &name)))?; + // 这把密钥的用量上限:**一条连接算一个请求**,连上之前看一遍,和 HTTP 那条路的准入 + // 同一套(见 `crate::key_limits`)。连接上的每个 `response.create` 不再分开数:存储层给 + // 整条连接记一行、不带用量,分开数的话,重启之后从记录里加回来的数就对不上了 + let limits = rt + .config + .clients + .iter() + .find(|c| c.name == client_name) + .map(|c| c.limits.as_slice()) + .unwrap_or_default(); + let hold = match state + .key_limits + .admit( + &client_name, + limits, + Default::default(), + crate::key_limits::slot_wait(&rt.config), + ) + .await + { + Ok(hold) => hold, + // 被拒的照样留一行,和规则拒绝的一样 + Err(r) => { + let why = r.error(); + let (id, ending) = open(&choice, "", tw_api::Billing::PerToken); + state.bus.emit(super::routed_nowhere(id, choice)); + ending.failed(why.source.into(), why.detail.clone()); + return Err(why); + } + }; let (id, ending) = open(&choice, &name, provider.billing.into()); + hold.bind(id); // 插件:升级那一刻的那一份表,一条连接用到底。**插件只管 Responses 的 WebSocket**(每个 // `response.create` 是一次对话请求);别的路径上的连接(比如 Realtime 的 `/v1/realtime`) // 不属于插件处理的任何一种请求,所有插件都不管:原样接上,什么都不记 diff --git a/crates/tw-gateway/src/state.rs b/crates/tw-gateway/src/state.rs index 25fd0c3d..d8013aa1 100644 --- a/crates/tw-gateway/src/state.rs +++ b/crates/tw-gateway/src/state.rs @@ -139,6 +139,12 @@ pub struct AppState { /// 每家上游自己的并发上限(见 [`crate::slots`])。**不在 Runtime 里**,理由和 `gate` /// 一样:它握着在跑的请求占着的位置。配置换了由 [`Self::reload`] 在原地改上限 pub slots: Arc, + /// 每把密钥的用量上限(见 [`crate::key_limits`]):今天、这周、这个月用了多少,最近 + /// 一分钟、一小时用了多少,在跑的请求占着多少。 + /// + /// **跨重载存活**,理由和并发闸门一样:改一条路由规则不该让今天花的钱归零,在跑的 + /// 请求的预留也不该变成孤儿。上限本身跟着配置换([`crate::key_limits::KeyLimits::configure`]) + pub key_limits: Arc, /// 一家上游不在当前运行时里时顶上的 Client(直连,不读系统代理) pub http: reqwest::Client, /// 观测事件往这里丢。没有订阅者时是零成本的 —— 数据面不该知道有 @@ -261,12 +267,16 @@ impl AppState { health.configure(&rt.config.failover); let slots = Arc::new(crate::slots::Slots::default()); slots.configure(&rt.config.providers); + let bus = tw_observe::EventBus::new(); + let key_limits = Arc::new(crate::key_limits::KeyLimits::new(bus.clone())); + key_limits.configure(&rt.config); let state = Self { rt: Arc::new(arc_swap::ArcSwap::from_pointee(rt)), gate: Default::default(), slots, + key_limits, http, - bus: tw_observe::EventBus::new(), + bus, health, catalog: Arc::new(arc_swap::ArcSwap::from_pointee(Default::default())), models, @@ -337,6 +347,14 @@ impl AppState { self.rt.load().config.clone() } + /// 换一个看用量上限的时钟(测试把时间拨到零点前后)。**账从空的开始**:只在测试里、 + /// 第一个请求之前调 + pub fn set_key_limits_clock(&mut self, clock: Arc) { + let limits = crate::key_limits::KeyLimits::with_clock(self.bus.clone(), clock); + limits.configure(&self.config()); + self.key_limits = Arc::new(limits); + } + /// 直接换一份插件进去,配置照旧:正在跑的请求用完它们手上那一份,新请求看到的是 /// 新的。**测试装插件替身走这里**;生产上装哪些插件由配置和插件文件决定(见 /// [`Self::reload_plugins`]),下一次重载就照那个重建 @@ -483,6 +501,7 @@ impl AppState { .rcu(|book| book.with_config(sheets.clone(), assign.clone())); self.health.configure(&next.config.failover); self.slots.configure(&next.config.providers); + self.key_limits.configure(&next.config); self.rt.store(Arc::new(next)); self.announce_broken(broken); // 模型汇总马上按新配置重算:删掉、停用的上游的模型必须立刻消失(列表 diff --git a/crates/tw-gateway/tests/key_limits.rs b/crates/tw-gateway/tests/key_limits.rs new file mode 100644 index 00000000..9456dd20 --- /dev/null +++ b/crates/tw-gateway/tests/key_limits.rs @@ -0,0 +1,260 @@ +//! 密钥的用量上限,走真的网关:被拒的请求在每一种客户端格式里长什么样、带什么响应头, +//! 流量里有没有它那一行,数 token 的请求和 WebSocket 连接怎么算。 +//! +//! 数和等的细节在 `tw_gateway::key_limits` 的单元测试里;这里看的是它接在管线上的样子。 + +use std::net::SocketAddr; +use std::sync::Arc; +use std::sync::atomic::{AtomicUsize, Ordering}; +use std::time::Duration; + +use axum::Router; +use axum::routing::any; +use tw_config::{Client, Config, Provider}; + +/// 什么都回 200 的上游,数着收到了几个请求 +async fn upstream() -> (SocketAddr, Arc) { + let hits = Arc::new(AtomicUsize::new(0)); + let h = hits.clone(); + let app = Router::new().fallback(any(move || { + let h = h.clone(); + async move { + h.fetch_add(1, Ordering::SeqCst); + axum::response::Response::builder() + .header("content-type", "application/json") + .body(axum::body::Body::from( + r#"{"id":"m","type":"message","role":"assistant","model":"claude-sonnet-4-5", + "content":[{"type":"text","text":"hi"}],"stop_reason":"end_turn", + "usage":{"input_tokens":10,"output_tokens":2},"input_tokens":10}"#, + )) + .unwrap() + } + })); + let l = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap(); + let a = l.local_addr().unwrap(); + tokio::spawn(async move { axum::serve(l, app).await.unwrap() }); + (a, hits) +} + +/// 网关,密钥 `k` 带着这几条上限。**时钟定在中午**:天的上限不会碰巧在测试跑到一半时 +/// 过零点。先订阅事件再起服务 +async fn gateway( + up: SocketAddr, + limits: &str, +) -> (SocketAddr, tokio::sync::broadcast::Receiver) { + let cfg = Config { + clients: vec![Client { + name: "k".into(), + key: "tw-k".into(), + limits: serde_yaml_ng::from_str(limits).unwrap(), + ..Default::default() + }], + providers: vec![Provider { + name: "官方".into(), + base_url: format!("http://{up}"), + key: Some("sk-x".into()), + protocol: Some(tw_config::Protocol::Anthropic), + ..Default::default() + }], + ..Default::default() + }; + let mut state = tw_gateway::AppState::new(cfg).unwrap(); + let noon = chrono::DateTime::parse_from_rfc3339("2026-10-05T12:00:00+08:00") + .unwrap() + .timestamp_millis(); + state.set_key_limits_clock(Arc::new(tw_gateway::key_limits::TestClock::new( + noon, + 8 * 3600, + ))); + let rx = state.bus.subscribe(); + let addr = tw_gateway::serve(state, ([127, 0, 0, 1], 0).into()) + .await + .unwrap(); + tokio::time::sleep(Duration::from_millis(50)).await; + (addr, rx) +} + +const MESSAGE: &str = + r#"{"model":"claude-sonnet-4-5","max_tokens":5,"messages":[{"role":"user","content":"hi"}]}"#; + +async fn anthropic(gw: SocketAddr, path: &str) -> reqwest::Response { + reqwest::Client::new() + .post(format!("http://{gw}{path}")) + .header("x-api-key", "tw-k") + .header("anthropic-version", "2023-06-01") + .body(MESSAGE) + .send() + .await + .unwrap() +} + +fn head(r: &reqwest::Response, h: &str) -> Option { + r.headers().get(h).map(|v| v.to_str().unwrap().to_string()) +} + +/// 一天一个请求:第二个在四种格式里都是 429,说清是哪把密钥、哪一条、用了多少、什么 +/// 时候重置;到重置之前别重试,OpenAI 的两种格式写成额度用完。**一个字节都没发给上游** +#[tokio::test] +async fn a_used_up_day_is_refused_in_each_clients_own_shape() { + let (up, hits) = upstream().await; + let (gw, _rx) = gateway(up, "[{per: day, requests: 1}]").await; + assert_eq!(anthropic(gw, "/v1/messages").await.status(), 200); + assert_eq!(hits.load(Ordering::SeqCst), 1); + + let r = anthropic(gw, "/v1/messages").await; + assert_eq!(r.status(), 429); + assert_eq!( + head(&r, "x-thinkwatch-error").as_deref(), + Some("rate_limited") + ); + assert_eq!(head(&r, "x-should-retry").as_deref(), Some("false")); + // 中午到零点:十二个小时 + assert_eq!(head(&r, "retry-after").as_deref(), Some("43200")); + let v: serde_json::Value = r.json().await.unwrap(); + assert_eq!(v["error"]["type"], "rate_limit_error"); + assert_eq!( + v["error"]["message"], + "[ThinkWatch] Gateway key `k` has reached its limit of 1 requests per day: 1 so far. It \ + resets at 2026-10-06 00:00 +08:00." + ); + + let client = reqwest::Client::new(); + for path in ["/v1/chat/completions", "/v1/responses"] { + let r = client + .post(format!("http://{gw}{path}")) + .header("authorization", "Bearer tw-k") + .body(r#"{"model":"claude-sonnet-4-5","messages":[{"role":"user","content":"hi"}],"input":"hi"}"#) + .send() + .await + .unwrap(); + assert_eq!(r.status(), 429, "{path}"); + assert_eq!(head(&r, "x-should-retry").as_deref(), Some("false")); + let v: serde_json::Value = r.json().await.unwrap(); + assert_eq!(v["error"]["code"], "insufficient_quota", "{path}: {v}"); + assert_eq!(v["error"]["type"], "rate_limit_error"); + } + let r = client + .post(format!( + "http://{gw}/v1beta/models/claude-sonnet-4-5:generateContent" + )) + .header("x-goog-api-key", "tw-k") + .body(r#"{"contents":[{"role":"user","parts":[{"text":"hi"}]}]}"#) + .send() + .await + .unwrap(); + assert_eq!(r.status(), 429); + let v: serde_json::Value = r.json().await.unwrap(); + assert_eq!(v["error"]["status"], "RESOURCE_EXHAUSTED"); + assert_eq!(hits.load(Ordering::SeqCst), 1, "被拒的发给了上游"); +} + +/// 被拒的请求流量里照样有一行:开始、一条空的尝试链、一个带着上限那句话的失败。 +#[tokio::test] +async fn a_refused_request_is_recorded_like_other_refusals() { + let (up, _hits) = upstream().await; + let (gw, mut rx) = gateway(up, "[{per: day, requests: 1}]").await; + anthropic(gw, "/v1/messages").await; + let r = anthropic(gw, "/v1/messages").await; + assert_eq!(r.status(), 429); + let mut seen = Vec::new(); + let deadline = tokio::time::Instant::now() + Duration::from_secs(2); + while let Ok(Ok(e)) = tokio::time::timeout_at(deadline, rx.recv()).await { + seen.push(e); + } + let failed = seen + .iter() + .find_map(|e| match e { + tw_api::Event::RequestFailed { + id, + message, + source, + .. + } => Some((*id, message.clone(), *source)), + _ => None, + }) + .expect("被拒的请求没有结局"); + assert_eq!(failed.1.code, "gw.key_limit.requests_per_period"); + assert_eq!(failed.1.arg("key"), "k"); + assert_eq!(failed.1.arg("per"), "day"); + assert_eq!(failed.2, tw_api::FailureSource::RateLimited); + assert!(seen.iter().any(|e| matches!( + e, + tw_api::Event::RequestStarted { id, .. } if *id == failed.0 + ))); + assert!(seen.iter().any(|e| matches!( + e, + tw_api::Event::RequestRouted { id, attempts, .. } if *id == failed.0 && attempts.is_empty() + ))); + // 被拒的那一刻到了顶:报一条,给界面发系统通知 + assert!(seen.iter().any(|e| matches!( + e, + tw_api::Event::KeyLimitAlert { key, reached: true, .. } if key == "k" + ))); +} + +/// 一分钟一个:第二个等不到空位(要等快一分钟,超过 30 秒),拒,并说准多久之后再来, +/// 可以重试。 +#[tokio::test] +async fn a_rolling_limit_says_exactly_when_to_retry() { + let (up, _hits) = upstream().await; + let (gw, _rx) = gateway(up, "[{per: minute, requests: 1}]").await; + assert_eq!(anthropic(gw, "/v1/messages").await.status(), 200); + let r = anthropic(gw, "/v1/messages").await; + assert_eq!(r.status(), 429); + assert_eq!(head(&r, "x-should-retry").as_deref(), Some("true")); + let ms: u64 = head(&r, "retry-after-ms").unwrap().parse().unwrap(); + assert!((50_000..=60_000).contains(&ms), "{ms}"); + let secs: u64 = head(&r, "retry-after").unwrap().parse().unwrap(); + assert_eq!(secs, ms.div_ceil(1000)); + let v: serde_json::Value = r.json().await.unwrap(); + assert!( + v["error"]["message"] + .as_str() + .unwrap() + .contains("1 requests per minute"), + "{v}" + ); +} + +/// 数 token 的请求不跑模型、不收钱:不算用量,用满了也照样放行。 +#[tokio::test] +async fn counting_tokens_is_not_counted_and_not_refused() { + let (up, _hits) = upstream().await; + let (gw, _rx) = gateway(up, "[{per: day, requests: 1}]").await; + for _ in 0..3 { + assert_eq!( + anthropic(gw, "/v1/messages/count_tokens").await.status(), + 200 + ); + } + assert_eq!(anthropic(gw, "/v1/messages").await.status(), 200); + assert_eq!(anthropic(gw, "/v1/messages").await.status(), 429); + assert_eq!( + anthropic(gw, "/v1/messages/count_tokens").await.status(), + 200 + ); +} + +/// WebSocket 一条连接算一个请求,**连上之前**看:用满了,升级就是一个 429。 +#[tokio::test] +async fn a_websocket_connection_is_admitted_when_it_opens() { + let (up, _hits) = upstream().await; + let (gw, _rx) = gateway(up, "[{per: day, requests: 1}]").await; + assert_eq!(anthropic(gw, "/v1/messages").await.status(), 200); + use tokio_tungstenite::tungstenite::client::IntoClientRequest; + let mut req = format!("ws://{gw}/v1/responses") + .into_client_request() + .unwrap(); + req.headers_mut() + .insert("authorization", "Bearer tw-k".parse().unwrap()); + match tokio_tungstenite::connect_async(req).await { + Err(tokio_tungstenite::tungstenite::Error::Http(r)) => { + assert_eq!(r.status(), 429); + assert_eq!( + r.headers().get("x-should-retry").unwrap().to_str().unwrap(), + "false" + ); + } + other => panic!("升级没有被拒:{:?}", other.map(|_| ())), + } +} diff --git a/crates/tw-observe/src/bus.rs b/crates/tw-observe/src/bus.rs index 60fcbbfc..d0dd2f61 100644 --- a/crates/tw-observe/src/bus.rs +++ b/crates/tw-observe/src/bus.rs @@ -72,6 +72,8 @@ fn about_the_request(ev: &tw_api::Event) -> bool { | E::RequestPriced { .. } | E::QuotaSeen { .. } | E::QuotaExhausted { .. } + // 一把密钥的用量到了上限:说的是那把密钥,不是哪一个请求 + | E::KeyLimitAlert { .. } | E::LocallyAnswered { .. } | E::CredentialRotated { .. } | E::CredentialExpired { .. } @@ -244,6 +246,16 @@ impl EventBus { } } + /// 这个请求还在跑吗:发过开始事件、还没有结局。**只看一个号,不拷事件** —— + /// [`EventBus::in_flight`] 要把整张表连事件一起拷出来。 + /// + /// 给密钥用量上限的预留用(见 `tw_gateway::key_limits`):一个请求的预留要等存储层 + /// 记下它那一行才换成实数,存储层落后丢了那一行时,靠它认出这个请求早已结束 + pub fn is_open(&self, id: u64) -> bool { + let t = self.tally.lock().unwrap_or_else(|p| p.into_inner()); + t.open.contains_key(&id) + } + /// 此刻的实时读数:在跑的请求(和 [`EventBus::in_flight`] 同一批),和最近 /// 一分钟跑完的请求平均每秒生成多少 token。 /// @@ -400,6 +412,9 @@ mod tests { let open = b.in_flight().requests; assert_eq!(open.iter().map(|r| r.id).collect::>(), [4]); + // 只问一个号的那一问,和快照说的是同一件事 + assert!(b.is_open(4)); + assert!(!b.is_open(1) && !b.is_open(2) && !b.is_open(3) && !b.is_open(99)); // **原样**:听的人拿它当补发的事件,字段一个都不能少 assert!( matches!(&open[0].events[0], tw_api::Event::RequestStarted { model, at_ms: 1004, .. } if model == "m"), diff --git a/crates/tw-store/src/db.rs b/crates/tw-store/src/db.rs index ae06ae00..80641ccb 100644 --- a/crates/tw-store/src/db.rs +++ b/crates/tw-store/src/db.rs @@ -146,6 +146,26 @@ pub struct RequestRow { pub session_log_bytes: Option, } +/// 一把密钥从某一刻起用了多少,按几样分开数(见 [`Db::key_usage_since`])。 +/// +/// **分开的那几样就是「算不算」要看的**:数 token 的请求、准入之前就被拒的不算进密钥的 +/// 用量上限,判断在网关那边(`tw_gateway::key_limits`),和它结算一行时是同一个判断。 +#[derive(Debug, Clone, PartialEq)] +pub struct KeyUsage { + pub client: String, + pub path: String, + pub error_code: Option, + /// 尝试链上有没有至少一跳 + pub attempted: bool, + pub requests: i64, + pub input_tokens: i64, + pub output_tokens: i64, + pub cache_read_tokens: i64, + pub cache_write_tokens: i64, + /// 记下的费用合计,微分。算不出来的那几行算 0 + pub cost_micros: i64, +} + #[derive(Debug)] pub struct Db { /// 上游体检的查询在 `crate::health`,和这里共用一个连接 @@ -1204,6 +1224,42 @@ impl Db { rows.collect::, _>>().map_err(Into::into) } + /// 每把密钥从 `since_ms` 起用了多少:重启之后,密钥的用量上限按它把这一天、这一周、 + /// 这个月的数加回来。 + /// + /// 网关自己答的不在里面。**按「算不算」要看的几样分组**(路径、失败的码、有没有发往 + /// 上游),组数和密钥、路径的个数相当,一把密钥一个月的记录也只有几十组。 + pub fn key_usage_since(&self, since_ms: i64) -> Result, DbError> { + let mut st = self.conn.prepare( + "SELECT client, path, error_code, + (CASE WHEN json_valid(routing) + THEN COALESCE(json_array_length(routing, '$.attempts'), 0) + ELSE 0 END) > 0 AS attempted, + COUNT(*), + COALESCE(SUM(input_tokens), 0), COALESCE(SUM(output_tokens), 0), + COALESCE(SUM(cache_read_tokens), 0), COALESCE(SUM(cache_write_tokens), 0), + COALESCE(SUM(cost_micros), 0) + FROM requests + WHERE at_ms >= ?1 AND local = 0 + GROUP BY client, path, error_code, attempted", + )?; + let rows = st.query_map(params![since_ms], |r| { + Ok(KeyUsage { + client: r.get(0)?, + path: r.get(1)?, + error_code: r.get(2)?, + attempted: r.get(3)?, + requests: r.get(4)?, + input_tokens: r.get(5)?, + output_tokens: r.get(6)?, + cache_read_tokens: r.get(7)?, + cache_write_tokens: r.get(8)?, + cost_micros: r.get(9)?, + }) + })?; + rows.collect::, _>>().map_err(Into::into) + } + /// 各条路由走了多少请求、各条规则命中了多少,以及记录从哪一刻起是全的(见 /// [`tw_api::RouteStats`])。 pub fn route_stats(&self, since_ms: i64, until_ms: i64) -> Result { @@ -1820,6 +1876,76 @@ pub(crate) mod tests { ); } + /// 密钥用量上限重启之后加回来的数:按密钥、路径、失败的码、有没有发往上游分组, + /// 从那一刻起,网关自己答的不算。 + #[test] + fn key_usage_adds_up_each_key_from_a_moment_on() { + let db = Db::in_memory().unwrap(); + let t0 = 1_000_000_000i64; + let attempt = r#"{"route":"default","rule":"r","rewritten_by":[], + "attempts":[{"provider":"官方","outcome":"served","status":200,"ms":5}]}"#; + let nowhere = r#"{"route":"default","rule":"r","rewritten_by":[],"attempts":[]}"#; + let mut early = row(1, t0 - 1); + early.routing = Some(attempt.into()); + let mut a = row(2, t0); + a.routing = Some(attempt.into()); + let mut b = row(3, t0 + 5); + b.routing = Some(attempt.into()); + b.cost_micros = None; + b.cache_write_tokens = Some(7); + let mut refused = row(4, t0 + 6); + refused.routing = Some(nowhere.into()); + refused.error = Some(Msg { + code: "gw.route.denied".into(), + args: Default::default(), + text: "x".into(), + }); + refused.input_tokens = None; + refused.output_tokens = None; + refused.cache_read_tokens = None; + refused.cost_micros = None; + let mut local = row(5, t0 + 7); + local.local = true; + let mut other = row(6, t0 + 8); + other.client = "codex".into(); + for r in [&early, &a, &b, &refused, &local, &other] { + db.insert(r).unwrap(); + } + let mut got = db.key_usage_since(t0).unwrap(); + got.sort_by(|x, y| (&x.client, x.attempted).cmp(&(&y.client, y.attempted))); + assert_eq!(got.len(), 3, "{got:?}"); + let (refused_g, served, codex) = (&got[0], &got[1], &got[2]); + assert_eq!( + (served.client.as_str(), served.attempted, served.requests), + ("claude-code", true, 2), + "早于那一刻的、本地答的都不算" + ); + assert_eq!( + ( + served.input_tokens, + served.output_tokens, + served.cache_read_tokens, + served.cache_write_tokens, + served.cost_micros + ), + (2000, 1000, 400, 7, 12_000), + "算不出钱的那一行算 0" + ); + assert_eq!( + ( + refused_g.attempted, + refused_g.error_code.as_deref(), + refused_g.requests + ), + (false, Some("gw.route.denied"), 1) + ); + assert_eq!( + (codex.client.as_str(), codex.attempted), + ("codex", false), + "没有路由那一列的当作没发往上游" + ); + } + #[test] fn a_window_narrows_the_history_and_no_window_means_the_most_recent() { let db = Db::in_memory().unwrap(); diff --git a/crates/tw-store/src/lib.rs b/crates/tw-store/src/lib.rs index eb73756f..a5c995cf 100644 --- a/crates/tw-store/src/lib.rs +++ b/crates/tw-store/src/lib.rs @@ -20,8 +20,10 @@ pub mod task; pub mod transcript; pub use blobs::{Blobs, Which}; -pub use db::{Db, DbError, Latency, PluginRunRow, RequestRow, SecurityEvent, Summary, TokenRate}; -pub use recorder::{Recorder, price_source}; +pub use db::{ + Db, DbError, KeyUsage, Latency, PluginRunRow, RequestRow, SecurityEvent, Summary, TokenRate, +}; +pub use recorder::{Recorder, SettleHook, Settled, price_source}; pub use task::StoredBody; use std::path::Path; diff --git a/crates/tw-store/src/recorder.rs b/crates/tw-store/src/recorder.rs index fa30a15d..057819b6 100644 --- a/crates/tw-store/src/recorder.rs +++ b/crates/tw-store/src/recorder.rs @@ -93,8 +93,42 @@ pub struct Recorder { /// /// `None` 表示没人要听(测试、以及不带总线的调用方)。 bus: Option, + /// 每记下一行请求就交一份 [`Settled`] 出去:密钥的用量上限拿它把预留换成实数。 + /// + /// **直接调,不走总线。**总线上丢了事件的话,这一行也就不在库里,重启之后从库里 + /// 加回来的数和此刻内存里的数对得上;走总线再绕一圈,两边丢的不是同一批。 + settled: Option, +} + +/// 记下的一行里,密钥的用量上限要的那几样(见 `tw_gateway::key_limits`)。 +/// +/// **和库里那一行是同一份数**:用量、费用就是写进去的那些,重启之后从库里加回来的和 +/// 此刻结算的一样多。 +#[derive(Debug, Clone, PartialEq)] +pub struct Settled { + pub id: u64, + /// 请求开始的时刻,那一行的 `at_ms`。算在哪一天、哪一周看它 + pub at_ms: i64, + /// 密钥的名字 + pub client: String, + pub path: String, + /// 网关自己答的(本地估的 token 数) + pub local: bool, + /// 失败的原因的码。没失败的是 None + pub error_code: Option, + /// 尝试链上有没有至少一跳。路由就拒绝了的、准入没过的没有 + pub attempted: bool, + pub input: u64, + pub output: u64, + pub cache_read: u64, + pub cache_write: u64, + /// 记下的费用,微分。算不出来的(没有价格、没有用量)是 None + pub cost_micros: Option, } +/// 结算往哪儿交。由 twcore 接到网关的密钥用量上限上(`tw_control::key_limits`) +pub type SettleHook = std::sync::Arc; + /// 一个价格的来源,给界面看的样子。`date`:当时默认价目表的数据日期。 pub fn price_source(source: &tw_pricing::Source, date: &str) -> tw_api::PriceSourceView { match source { @@ -120,9 +154,16 @@ impl Recorder { pricing, inflight: HashMap::new(), bus: None, + settled: None, } } + /// 每记下一行请求,交一份 [`Settled`] 给 `hook`。 + pub fn settling_to(mut self, hook: SettleHook) -> Self { + self.settled = Some(hook); + self + } + /// 把算出来的价钱报回总线上。 pub fn reporting_to(mut self, bus: tw_observe::EventBus) -> Self { self.bus = Some(bus); @@ -580,6 +621,8 @@ impl Recorder { | Event::ConfigRejected { .. } | Event::QuotaSeen { .. } | Event::QuotaExhausted { .. } + // 密钥用量到了上限:说的是那把密钥现在的样子,它的用量就是这张表里的那些行 + | Event::KeyLimitAlert { .. } // 凭据轮换说的是配置文件该改了,跟哪一次请求无关 | Event::CredentialRotated { .. } | Event::CredentialExpired { .. } @@ -693,6 +736,25 @@ impl Recorder { at_ms: p.at_ms as u64, }); } + if let Some(hook) = &self.settled { + hook(&Settled { + id, + at_ms: p.at_ms, + client: p.client.clone(), + path: p.path.clone(), + local, + error_code: match &how { + Ending::Failed(message) => Some(message.code.clone()), + _ => None, + }, + attempted: !p.routing.attempts.is_empty(), + input: u.map_or(0, |u| u.input), + output: u.map_or(0, |u| u.output), + cache_read: u.map_or(0, |u| u.cache_read), + cache_write: u.map_or(0, |u| u.cache_write), + cost_micros, + }); + } self.write(RequestRow { id: id as i64, at_ms: p.at_ms, @@ -2585,3 +2647,134 @@ mod failure_tests { assert_eq!(priced, Some((1, Some(300_015), true))); } } + +#[cfg(test)] +mod settle_hook_tests { + use super::tests::{finished, rec, started}; + use super::*; + use std::sync::{Arc, Mutex}; + use tw_api::UsageView; + + fn hooked() -> (tempfile::TempDir, Recorder, Arc>>) { + let (d, r) = rec(); + let seen: Arc>> = Arc::default(); + let into = seen.clone(); + let r = r.settling_to(Arc::new(move |s: &Settled| { + into.lock().unwrap().push(s.clone()) + })); + (d, r, seen) + } + + fn usage() -> Option { + Some(UsageView { + input: 100_000, + output: 1, + cache_read: 40, + cache_write: 7, + cache_1h: false, + }) + } + + fn served(id: u64) -> Event { + Event::RequestRouted { + id, + route: "default".into(), + rule: "catch-all".into(), + group: None, + rewritten_by: vec![], + denied_by: None, + affinity: None, + attempts: vec![tw_api::AttemptView { + provider: "官方".into(), + model: None, + outcome: tw_api::AttemptOutcome::Served, + status: Some(200), + error: None, + ms: 5, + usage: None, + queued_ms: None, + skipped: None, + }], + billing: tw_api::Billing::PerToken, + } + } + + /// 跑完的、客户端走掉的、断在半路的,**三种结局都交一份**,用量和费用就是写进库的 + /// 那些:密钥的用量上限拿它把预留换成实数,重启之后从库里加回来的也是这些数。 + #[test] + fn every_ending_hands_over_what_was_written() { + let (_d, mut r, seen) = hooked(); + for id in 1..=3 { + r.on_event(&started(id, "claude-sonnet-4-5")); + r.on_event(&served(id)); + } + r.on_event(&finished(1, usage())); + r.on_event(&Event::RequestCancelled { + id: 2, + model: String::new(), + status: Some(200), + bytes: 1, + duration_ms: 1, + usage: usage(), + answered_model: None, + }); + r.on_event(&Event::RequestFailed { + id: 3, + model: String::new(), + source: tw_api::FailureSource::Upstream, + message: tw_api::Msg { + code: "t.broke".into(), + args: Default::default(), + text: "broke".into(), + }, + bytes: None, + duration_ms: None, + usage: usage(), + answered_model: None, + }); + let seen = seen.lock().unwrap(); + assert_eq!(seen.iter().map(|s| s.id).collect::>(), [1, 2, 3]); + for s in seen.iter() { + let row = r.db().get(s.id as i64).unwrap().unwrap(); + assert_eq!(s.cost_micros, row.cost_micros, "和库里那一行是同一个数"); + assert!(s.cost_micros.is_some()); + assert_eq!( + (s.input, s.output, s.cache_read, s.cache_write), + (100_000, 1, 40, 7) + ); + assert_eq!( + (s.client.as_str(), s.path.as_str()), + ("claude-code", "/v1/messages") + ); + assert_eq!(s.at_ms, 1_000_000); + assert!(s.attempted && !s.local); + } + assert_eq!(seen[2].error_code.as_deref(), Some("t.broke")); + assert_eq!(seen[0].error_code, None); + } + + /// 路由就拒绝了的:没有一跳。没有用量的:费用是 None,不是 0 + #[test] + fn a_request_refused_before_any_hop_says_so() { + let (_d, mut r, seen) = hooked(); + r.on_event(&started(1, "claude-sonnet-4-5")); + r.on_event(&Event::RequestFailed { + id: 1, + model: String::new(), + source: tw_api::FailureSource::Denied, + message: tw_api::Msg { + code: "gw.route.denied".into(), + args: Default::default(), + text: "denied".into(), + }, + bytes: None, + duration_ms: None, + usage: None, + answered_model: None, + }); + let s = seen.lock().unwrap()[0].clone(); + assert!(!s.attempted); + assert_eq!(s.cost_micros, None); + assert_eq!(s.error_code.as_deref(), Some("gw.route.denied")); + } +} diff --git a/docs/config.md b/docs/config.md index 1114e853..93c891c0 100644 --- a/docs/config.md +++ b/docs/config.md @@ -308,6 +308,7 @@ gateway. A key is an identity. Limits, model scope and route are per key. | `route` | string | — | Name of the route requests with this key take. Unset: `default_route`. | | `client` | string | — | The client this key was made for (`claude-code`, `codex`, …), recorded when the desktop app points a client at the gateway. A client has at most one. | | `disabled` | bool | `false` | Refuse every request made with this key, and keep the key. | +| `limits` | list of [`clients[].limits[]`](#cfg-clients-limits) | — | Usage limits: requests, tokens or cost per minute, hour, day, week or month. A request has to pass every one. Unset: no limit. | ```yaml @@ -321,6 +322,50 @@ clients: route: cheap ``` +#### `clients[].limits` + +Usage limits for a key: at most so many requests, tokens or dollars per +minute, hour, day, week or month. Each entry counts one of the three; a key +can have several, and a request has to pass every one. + +`minute` and `hour` are rolling: the last 60 seconds, the last 60 minutes. +When one is used up, a request waits for the next free slot if it frees +within `failover.slot_wait_secs` (30 seconds by default), and is refused +otherwise. `day`, `week` and `month` follow the calendar in the time zone of +the machine twcore runs on and start again at midnight, on Monday and on the +1st. When one is used up, requests are refused until it starts again. + +A refused request gets HTTP 429 in the client's own error format, naming the +key, the limit, the amount used and when it resets, and it shows in the +traffic list. Cost is what is recorded for each request, so a model without a +price and an upstream with `billing: free` count as $0. A request still +running counts with an estimate of its input until it is recorded. After a +restart, the day, week and month are added up again from the request records; +a monthly limit therefore needs `retention.row_days` of at least 31. Minute +and hour limits start empty. + + + + +| Field | Type | Default | Description | +|---|---|---|---| +| `per` | `minute` \| `hour` \| `day` \| `week` \| `month` | **required** | The period. `minute` and `hour` are rolling (the last 60 seconds, the last 60 minutes); `day`, `week` and `month` start again at local midnight, on Monday and on the 1st. | +| `requests` | integer | — | At most this many requests. Token counts and answers the gateway gives itself do not count. | +| `tokens` | integer | — | At most this many tokens: uncached input, cache writes and output. | +| `cost` | number | — | At most this much, in US dollars, as recorded for each request. Models without a price and upstreams with `billing: free` count as 0. | +| `cache_reads` | bool | `false` | Count cache reads too. Only for a `tokens` limit. | + + +```yaml +clients: + - name: build-server + key: tw-q8r2s4t6u8v2w4x6y8z2a4b6 + limits: + - { per: minute, requests: 30 } + - { per: day, cost: 5 } + - { per: month, tokens: 20000000, cache_reads: true } +``` + ### `providers` Upstreams: the APIs requests are forwarded to. diff --git a/docs/config.zh-CN.md b/docs/config.zh-CN.md index 6cfeb41d..9b6ddd79 100644 --- a/docs/config.zh-CN.md +++ b/docs/config.zh-CN.md @@ -217,6 +217,7 @@ listen: | `route` | 字符串 | — | 这把密钥的请求走哪条路由。不写:`default_route`。 | | `client` | 字符串 | — | 这把密钥是为哪个客户端生成的(`claude-code`、`codex` 等),由桌面应用接管客户端时写入。一个客户端最多一把。 | | `disabled` | 布尔 | `false` | 拒绝使用这把密钥的所有请求,密钥本身保留。 | +| `limits` | 对象列表,见 [`clients[].limits[]`](#cfg-clients-limits) | — | 用量上限:每分钟、每小时、每天、每周或每月的请求数、token 数或费用。请求要通过每一条。不写:不限。 | ```yaml @@ -230,6 +231,44 @@ clients: route: cheap ``` +#### `clients[].limits` + +一把密钥的用量上限:每分钟、每小时、每天、每周或每月最多多少个请求、多少 token、多少美元。 +每一条只数其中一种;一把密钥可以有好几条,请求要通过每一条。 + +`minute`、`hour` 是滚动的:最近 60 秒、最近 60 分钟。用满之后,下一个空位在 +`failover.slot_wait_secs`(默认 30 秒)之内空出来,请求就等它;等不到就拒绝。`day`、 +`week`、`month` 按 twcore 所在机器的本地时区算,在零点、周一零点、每月一号零点重新开始; +用满之后,到重新开始之前的请求一律拒绝。 + +被拒的请求收到 HTTP 429,错误格式和客户端自己的一致,写明是哪把密钥、哪一条上限、用了多少、 +什么时候重置;流量列表里也有这一条。费用按每个请求记下的费用算,所以没有价格的模型、 +`billing: free` 的上游算 0。还在进行的请求先按输入的估算计入,记下之后换成实际用量。 +重启之后,这一天、这一周、这个月的用量从请求记录里重新加起来,所以设了按月上限时 +`retention.row_days` 至少要 31;按分钟、按小时的上限从零开始。 + + + + +| 字段 | 类型 | 默认值 | 说明 | +|---|---|---|---| +| `per` | `minute` \| `hour` \| `day` \| `week` \| `month` | **必填** | 按多长一段时间算。`minute`、`hour` 是滚动的(最近 60 秒、最近 60 分钟);`day`、`week`、`month` 在本地时间的零点、周一零点、每月一号零点重新算。 | +| `requests` | 整数 | — | 最多这么多个请求。数 token 的请求和网关自己答的不算。 | +| `tokens` | 整数 | — | 最多这么多 token:未命中缓存的输入、写入缓存的和输出。 | +| `cost` | 数字 | — | 最多花这么多美元,按每个请求记下的费用算。没有价格的模型、`billing: free` 的上游算 0。 | +| `cache_reads` | 布尔 | `false` | 把读取缓存的 token 也算进去。只有 `tokens` 上限能写。 | + + +```yaml +clients: + - name: build-server + key: tw-q8r2s4t6u8v2w4x6y8z2a4b6 + limits: + - { per: minute, requests: 30 } + - { per: day, cost: 5 } + - { per: month, tokens: 20000000, cache_reads: true } +``` + ### `providers` 上游,即请求被转发到的接口。 From d6e7146d19fc40890554d4a96cab22cd4d8ead8f Mon Sep 17 00:00:00 2001 From: fylorn <249551762+fylorn@users.noreply.github.com> Date: Mon, 5 Oct 2026 17:40:23 +0800 Subject: [PATCH 08/22] Measure upstream speed from sending to first content; share one wait budget Latency samples (url-test ordering and the load-balance latency factor) were taken when relaying began and counted from the request's arrival, so they included key-limit and slot waits, plugins and every earlier failed or abandoned hop, and ended at a point that depended on the hop's position (first content behind the opening watch, response headers for the last candidate, nearly the whole answer when not streamed). An upstream dropped by the slow-start switch recorded nothing and kept its old fast samples, while the next upstream was charged the wait. A sample is now the time from sending that hop to the first content of a streamed answer, detected by the same first-token reader that reports RequestFirstToken, wherever the hop sits in the candidate list. Answers that are not streamed give no sample. An upstream given up on for a slow start records the time it was given, a lower bound that ranks it slow. L1 handshake seeds stand in only until an upstream has enough real samples and are then dropped, never mixed into the median. Load-balance: a member at its max_concurrent at decision time sits the round out like a paused one. Charging it for a request the slot cap then skipped made a busy member drift below its share. failover.slot_wait_secs is now one deadline per request, set when the key gate lets it in and shared by the rolling key-limit wait and the upstream slot wait; before, each could use the full time, so one request could wait twice as long. Co-Authored-By: Claude Opus 5.5 --- crates/tw-api/src/lib.rs | 7 +- crates/tw-config/src/failover.rs | 7 +- crates/tw-config/tests/manual/schema.rs | 8 +- crates/tw-control/src/dryrun.rs | 6 +- crates/tw-engine/src/engine.rs | 19 +- crates/tw-engine/src/weighted.rs | 40 ++- crates/tw-gateway/src/ending.rs | 80 ++++++ crates/tw-gateway/src/key_limits/mod.rs | 42 ++- crates/tw-gateway/src/key_limits/tests.rs | 26 ++ crates/tw-gateway/src/latency.rs | 162 ++++++++---- crates/tw-gateway/src/server/pipeline.rs | 14 +- .../src/server/pipeline/admission.rs | 24 +- crates/tw-gateway/src/server/pipeline/hop.rs | 20 +- .../tw-gateway/src/server/pipeline/relay.rs | 16 +- crates/tw-gateway/src/server/pipeline/slow.rs | 11 + crates/tw-gateway/src/slots.rs | 13 +- crates/tw-gateway/src/state.rs | 13 +- crates/tw-gateway/tests/latency_samples.rs | 244 ++++++++++++++++++ crates/tw-gateway/tests/upstream_slots.rs | 100 ++++++- docs/config.md | 32 ++- docs/config.zh-CN.md | 19 +- 21 files changed, 771 insertions(+), 132 deletions(-) create mode 100644 crates/tw-gateway/tests/latency_samples.rs diff --git a/crates/tw-api/src/lib.rs b/crates/tw-api/src/lib.rs index 5b0d3caa..fee07b85 100644 --- a/crates/tw-api/src/lib.rs +++ b/crates/tw-api/src/lib.rs @@ -1905,7 +1905,8 @@ pub struct FailoverView { pub stream_start_wait_secs: u64, /// 等过 `stream_start_wait_secs` 还没有内容就换下一家(最后一家照常等) pub next_on_slow_start: bool, - /// 上游满着(`max_concurrent`)时,一个请求合计最多等多少秒空位。0 是不等 + /// 一个请求合计最多等多少秒:等密钥的分钟、小时上限空出名额,和等满着(`max_concurrent`) + /// 的上游空出位置,共用这一段。0 是不等 pub slot_wait_secs: u64, } @@ -4961,8 +4962,8 @@ pub struct DryRunCandidate { /// 经过的是 `load-balance` 组时,它在组里的权重(没写权重的是 1)。别的时候没有 #[serde(default, skip_serializing_if = "Option::is_none")] pub weight: Option, - /// 它的典型首字节时间(最近样本的中位数),毫秒。只在顺序看它时有:`url-test`, - /// 按快慢分的 `load-balance`。样本不够时没有 + /// 它典型的快慢:从发出去到回答的第一段内容(最近样本的中位数),毫秒。只在顺序看它时 + /// 有:`url-test`,按快慢分的 `load-balance`。样本不够时没有 #[serde(default, skip_serializing_if = "Option::is_none")] pub ttfb_ms: Option, /// 它最近的成功率,0 到 1(最近 50 次、30 分钟以内)。只在按成败分的 `load-balance` diff --git a/crates/tw-config/src/failover.rs b/crates/tw-config/src/failover.rs index 48666c08..9b29b7da 100644 --- a/crates/tw-config/src/failover.rs +++ b/crates/tw-config/src/failover.rs @@ -43,9 +43,10 @@ pub struct Failover { /// 输出的模型开头本来就慢,开着时要把等待调长 #[serde(default)] pub next_on_slow_start: bool, - /// 上游的并发数满了(`providers[].max_concurrent`)时,一个请求最多等多少秒空位, - /// **整个请求合起来算**。留在那一家的对话等它空出来,候选都满了时等先空出来的那一家; - /// 等不到的换下一家,或者回 429。0 是不等 + /// 一个请求最多等多少秒,**整个请求合起来算**:准入时等密钥的分钟、小时上限空出名额, + /// 之后上游的并发数满了(`providers[].max_concurrent`)时等空位,共用这一段。留在那一家 + /// 的对话等它空出来,候选都满了时等先空出来的那一家;等不到的换下一家,或者回 429。 + /// 0 是不等 #[serde(default = "d_slot_wait_secs")] pub slot_wait_secs: u64, } diff --git a/crates/tw-config/tests/manual/schema.rs b/crates/tw-config/tests/manual/schema.rs index 8a0ef785..b175a74b 100644 --- a/crates/tw-config/tests/manual/schema.rs +++ b/crates/tw-config/tests/manual/schema.rs @@ -1330,8 +1330,8 @@ pub fn sections() -> Vec
{ Kind::Int, Def::Is("30"), t( - "Seconds a request waits in total for a free slot on upstreams that are at their `max_concurrent`. After that it goes to the next upstream, or, when every candidate is full, is answered with 429. `0`: never wait. From 0 to 300.", - "上游的并发数满了(`max_concurrent`)时,一个请求等空位合计最多等的秒数。等不到就换下一家;候选全满时回 429。`0`:不等。取值 0 到 300。", + "Seconds a request waits in all, counted once the key's own `max_concurrent` lets it in: for a key's `minute` or `hour` limit to free up, and for a free slot on upstreams at their `max_concurrent`. A key limit that does not free up in time refuses the request; without an upstream slot in time it goes to the next upstream, or, when every candidate is full, is answered with 429. `0`: never wait. From 0 to 300.", + "一个请求合计最多等的秒数,从过了密钥自己的 `max_concurrent` 时算起:等密钥的 `minute`、`hour` 上限空出名额,和等并发数满了(`max_concurrent`)的上游空出位置,都算在里面。密钥的上限到时空不出来就拒绝;等不到上游的空位就换下一家,候选全满时回 429。`0`:不等。取值 0 到 300。", ), ), ], @@ -1355,8 +1355,8 @@ pub fn sections() -> Vec
{ Kind::Enum(group_types), Def::Is("fallback"), t( - "`fallback`: the first healthy member, in order. `select`: the member named in `selected`. `load-balance`: new conversations take turns, in proportion to the members' weights. `url-test`: the fastest by measured time to first byte. `cheapest`: the lowest input price.", - "`fallback`:按顺序取第一个健康的。`select`:取 `selected` 指定的那个。`load-balance`:新对话按成员的权重轮流。`url-test`:按实测首字节时间取最快的。`cheapest`:取输入单价最低的。", + "`fallback`: the first healthy member, in order. `select`: the member named in `selected`. `load-balance`: new conversations take turns, in proportion to the members' weights. `url-test`: the fastest by measured time from sending a request to the first content of the answer. `cheapest`: the lowest input price.", + "`fallback`:按顺序取第一个健康的。`select`:取 `selected` 指定的那个。`load-balance`:新对话按成员的权重轮流。`url-test`:按实测从发出请求到回答第一段内容的时间取最快的。`cheapest`:取输入单价最低的。", ), ), row( diff --git a/crates/tw-control/src/dryrun.rs b/crates/tw-control/src/dryrun.rs index 3c0d1f75..98f2563d 100644 --- a/crates/tw-control/src/dryrun.rs +++ b/crates/tw-control/src/dryrun.rs @@ -50,8 +50,8 @@ fn order_like_the_data_plane( } } -/// 一家候选的排序依据:典型首字节时间、最近的成功率、按它们算出的系数。**只给顺序 -/// 真用到的那几样**([`runtime_facts`] 只取用得上的)—— `url-test` 看首字节时间; +/// 一家候选的排序依据:典型的快慢(从发出去到回答的第一段内容)、最近的成功率、按它们算出 +/// 的系数。**只给顺序真用到的那几样**([`runtime_facts`] 只取用得上的)—— `url-test` 看快慢; /// `load-balance` 按 `balance_by` 看快慢、成败,系数乘在权重上。用不着的、没有样本的 /// 是 None。 struct Basis { @@ -325,7 +325,7 @@ pub async fn dry_run( // `load-balance` 的候选带上各自的权重:排头的为什么是它,一半在这个数里 let balanced = group.filter(|g| g.kind == tw_engine::GroupType::LoadBalance); // 每一家收到的模型名,和为什么不是请求里写的那个;顺序看运行时数字的,再加上 - // 每一家的那几个数字(权重、首字节时间、成功率、系数)—— 「为什么轮到它」 + // 每一家的那几个数字(权重、快慢、成功率、系数)—— 「为什么轮到它」 // 要能从这里看出来 let mut basis_of = bases(&d, runtime.as_ref()); out.candidate_models = out diff --git a/crates/tw-engine/src/engine.rs b/crates/tw-engine/src/engine.rs index 91c04561..1f4eae7e 100644 --- a/crates/tw-engine/src/engine.rs +++ b/crates/tw-engine/src/engine.rs @@ -30,8 +30,8 @@ pub enum GroupType { /// `balance_by` 还可以按快慢、成败给权重乘一个系数([`BalanceBy`])。 /// **轮的是新对话**:已经有人回答过、缓存还热着的对话留在那一家 LoadBalance, - /// 选最快的。判据是**真实流量测出来的 TTFB**,样本不够时用启动时 - /// 那次零成本的 L1 握手计时补。 + /// 选最快的。判据是**真实流量测出来的快慢**:从发出去到回答的第一段内容(流式回答才有), + /// 样本不够时用启动时那次零成本的 L1 握手计时补。 /// /// **测不到的那些排最后,而不是排最前。**「没测到」不等于「慢」, /// 但把它排前面就等于放弃了「选最快的」这个承诺;排最后它仍然是 @@ -78,7 +78,7 @@ pub enum BalanceBy { /// 只按成员的权重 #[default] Weights, - /// 越快的分得越多:看首字节时间,和 `url-test` 同一份样本 + /// 越快的分得越多:看从发出去到回答的第一段内容用了多久,和 `url-test` 同一份样本 Latency, /// 越少失败的分得越多:看最近的成功率 Health, @@ -102,7 +102,7 @@ impl BalanceBy { *self == BalanceBy::Weights } - /// 要不要每家的首字节时间([`Facts::ttfb_ms`]) + /// 要不要每家的快慢([`Facts::ttfb_ms`]) pub fn uses_latency(&self) -> bool { matches!(self, BalanceBy::Latency | BalanceBy::LatencyHealth) } @@ -125,7 +125,7 @@ const HEALTH_FACTOR_FLOOR: f64 = 0.05; /// `load-balance` 每一家的系数,和 `members` 一一对应:**成员的权重乘上它**,就是这一家 /// 这一轮的有效权重([`crate::weighted`] 按它轮)。`weights` 全是 1。 /// -/// - 快慢:(测到的那几家首字节时间的中位数 ÷ 这一家的)²,夹在 0.1 到 10 之间。平方让差别 +/// - 快慢:(测到的那几家典型快慢的中位数 ÷ 这一家的)²,夹在 0.1 到 10 之间。平方让差别 /// 看得出来:快一倍的分到四倍。 /// - 成败:成功率²,最低 0.05。 /// - **没有样本的那一项算 1,就是「中等」**:新加的上游会被试到,但不会一上来就被灌满。 @@ -156,7 +156,7 @@ pub fn balance_factors(by: BalanceBy, members: &[String], f: &Facts) -> Vec .collect() } -/// 测到了的那几家首字节时间的中位数,毫秒。一家都没测到是 `None`。 +/// 测到了的那几家典型快慢([`Facts::ttfb_ms`])的中位数,毫秒。一家都没测到是 `None`。 /// /// 偶数家时取中间两家的平均:系数围着它往两头夹,取其中一家的话,夹的那一刀会偏向一边 fn median_ttfb(members: &[String], f: &Facts) -> Option { @@ -202,7 +202,10 @@ pub struct Facts { /// 此刻停着的上游:熔断着、失败之后冷却着。`load-balance` 这一轮不算它们 —— /// 网关反正会跳过它们,轮到它们的那一次会落到组里排在后面的那一家头上 pub paused: std::collections::HashSet, - /// 每家的典型 TTFB(毫秒)。**缺席 = 样本不够**,不是「很快」 + /// 此刻并发数满着的上游(`providers[].max_concurrent`)。`load-balance` 这一轮也不算它们, + /// 和停着的一样:网关会当场跳过它们,轮到它们的那一次记了账却没答 + pub busy: std::collections::HashSet, + /// 每家典型的快慢:从发出去到回答的第一段内容,毫秒。**缺席 = 样本不够**,不是「很快」 pub ttfb_ms: std::collections::HashMap, /// 每家跑这个模型的单价,(输入, 输出),微分/百万 token。 /// **缺席 = 算不出价钱**,不是「免费」 @@ -263,7 +266,7 @@ pub fn order_by(g: &Group, members: &[String], f: &Facts) -> Vec { GroupType::Fallback | GroupType::Select => members.to_vec(), GroupType::LoadBalance => balance(g, members, f), GroupType::UrlTest => { - // 有样本的按 TTFB 升序;没样本的保持原有相对次序排在后面。 + // 有样本的按快慢升序;没样本的保持原有相对次序排在后面。 // **`sort_by_key` 是稳定排序**,所以同速的两家不会每次换位 // —— 那会让 prompt cache 白白多断一次 let mut out = members.to_vec(); diff --git a/crates/tw-engine/src/weighted.rs b/crates/tw-engine/src/weighted.rs index 6c031753..7ee14091 100644 --- a/crates/tw-engine/src/weighted.rs +++ b/crates/tw-engine/src/weighted.rs @@ -172,8 +172,10 @@ pub(crate) mod members { /// 和试算列出来的那一份一样;算完再去掉停着的。 /// /// **停着的不算**([`Facts::paused`]):熔断、冷却着的那一家这一次本来就会被跳过, -/// 轮到它的那一次落到组里排在它后面的那一家,那一家就平白多拿一份。全都停着时都算 -/// —— 那时网关照样一家家试(fail-open),排头的还是按权重来。 +/// 轮到它的那一次落到组里排在它后面的那一家,那一家就平白多拿一份。**并发数满着的也 +/// 不算**([`Facts::busy`]):网关当场跳过它,这一份记在它头上的话,它答得越多越满、 +/// 越满越被记空账,拿到的比它的权重少。停着的、满着的加起来是全部时都算 —— 那时网关 +/// 照样一家家试(fail-open)、等先空出来的那一家,排头的还是按权重来。 fn round<'a>(g: &Group, members: &'a [String], f: &Facts) -> Vec<(&'a str, i64)> { let factors = balance_factors(g.balance_by, members, f); let all: Vec<(&'a str, i64)> = members @@ -184,7 +186,7 @@ fn round<'a>(g: &Group, members: &'a [String], f: &Facts) -> Vec<(&'a str, i64)> let up: Vec<(&'a str, i64)> = all .iter() .copied() - .filter(|(m, _)| !f.paused.contains(*m)) + .filter(|(m, _)| !f.paused.contains(*m) && !f.busy.contains(*m)) .collect(); if up.is_empty() { all } else { up } } @@ -319,6 +321,34 @@ mod tests { assert_eq!(f.current_weight.get("乙").copied(), before); } + /// 并发数满着的同样不参加,当前权重也不动:空出来之后接着按权重轮,不欠也不多。 + /// 满着的和停着的加起来是全部时都参加 + #[test] + fn busy_members_sit_out_unless_every_member_is_busy_or_paused() { + let g = group(&[("甲", 1), ("乙", 1), ("丙", 1)]); + let mut f = Facts { + busy: ["甲".to_string()].into_iter().collect(), + ..Default::default() + }; + assert_eq!( + run(&g, &names(&g), &mut f, 4), + ["乙", "丙", "乙", "丙"], + "满着的那一家轮不到,没答的不记在它头上" + ); + assert_eq!(f.current_weight.get("甲"), None, "满着的不记账"); + // 空出来了:从它没欠账的样子接着轮 + f.busy.clear(); + let got = run(&g, &names(&g), &mut f, 30); + assert_eq!(got.iter().filter(|m| *m == "甲").count(), 10, "{got:?}"); + + let mut f = Facts { + busy: ["甲".to_string(), "乙".to_string()].into_iter().collect(), + paused: ["丙".to_string()].into_iter().collect(), + ..Default::default() + }; + assert_eq!(run(&g, &names(&g), &mut f, 3), ["甲", "乙", "丙"]); + } + /// 停着的(熔断、冷却)同样不参加;全都停着时都参加 #[test] fn paused_members_sit_out_unless_every_member_is_paused() { @@ -431,7 +461,7 @@ mod tests { assert_eq!(effective(1, 1e-9), 1); } - /// 按快慢分:首字节快的那一家分到的新对话多。一整圈(有效权重之和那么多次)下来, + /// 按快慢分:第一段内容来得快的那一家分到的新对话多。一整圈(有效权重之和那么多次)下来, /// 各家排头的次数正好是各自的有效权重,当前权重回到零 #[test] fn latency_gives_the_faster_member_more_new_conversations() { @@ -498,7 +528,7 @@ mod tests { fn equal_factors_keep_the_exact_interleaving() { let want = ["甲", "乙", "甲", "甲", "甲", "乙", "甲", "甲", "乙", "甲"]; let cases = [ - // 首字节一样:系数都是 1 + // 一样快:系数都是 1 Facts { ttfb_ms: ttfb(&[("甲", 200), ("乙", 200)]), ..Default::default() diff --git a/crates/tw-gateway/src/ending.rs b/crates/tw-gateway/src/ending.rs index e0ef4eb2..041fffe7 100644 --- a/crates/tw-gateway/src/ending.rs +++ b/crates/tw-gateway/src/ending.rs @@ -63,6 +63,8 @@ pub struct Ending { first: Option, /// 第一个 token 是什么时候、以什么开的头 opened: Option, + /// 第一个 token 到的时候给回答的这一家记一个快慢样本(见 [`Ending::timed`]) + lap: Option, /// 看上游有没有在流里报错。**只在上游回的是成功的流时才有**(见 [`Ending::streaming`]), /// 认出来就扔掉 watch: Option, @@ -168,6 +170,16 @@ struct FirstToken { hidden_thought: bool, } +/// 回答的这一家的快慢样本记到哪儿、从什么时候算起(见 [`crate::latency`])。 +pub struct Lap { + pub latency: std::sync::Arc, + /// 回答的那一家 + pub provider: String, + /// 这一跳发出去的那一刻。**不是请求进来的那一刻**:之前的等待、插件、失败了的几跳 + /// 都不是这一家慢 + pub sent: Instant, +} + /// 第一个 token 到的那一刻。 #[derive(Debug, Clone, Copy)] struct Opened { @@ -202,6 +214,7 @@ impl Ending { tap: ResponseTap::new(), first: None, opened: None, + lap: None, watch: None, upstream_error: None, refusal: None, @@ -231,6 +244,14 @@ impl Ending { }); } + /// 第一个 token 到的时候,给回答的这一家记一个快慢样本:从这一跳发出去到这一刻(见 + /// [`crate::latency`])。**认第一个 token 的是同一个**([`Ending::streaming`] 之后才有, + /// 没调过它的什么都不记):样本、请求列表里的首 token、开头慢不慢,说的是同一件事。 + /// 一个 token 都没等到的(流断了、上游只报了错),不记。 + pub fn timed(&mut self, lap: Lap) { + self.lap = Some(lap); + } + /// 上游 `provider` 回的不是 2xx,原样交给了客户端(4xx 是请求本身的问题,或者没有 /// 下一家可换了;3xx 交还客户端,由它决定跟不跟)。 /// @@ -316,6 +337,10 @@ impl Ending { return; }; self.first = None; + if let Some(lap) = self.lap.take() { + lap.latency + .record(&lap.provider, crate::latency::ms(lap.sent.elapsed())); + } let ms = self.duration_ms(); self.opened = Some(Opened { ms, @@ -1062,6 +1087,61 @@ mod tests { assert!(first_tokens(&drain(&mut rx)).is_empty()); } + fn lap(latency: &std::sync::Arc, sent: Instant) -> Lap { + Lap { + latency: latency.clone(), + provider: "up".into(), + sent, + } + } + + /// 快慢样本记在第一个 token 到的那一刻,从这一跳发出去算起 —— 不是从请求进来、也不是 + /// 从响应头到的那一刻。一个请求只记一个 + #[test] + fn the_sample_runs_from_sending_the_hop_to_the_first_token() { + let bus = tw_observe::EventBus::new(); + let latency = std::sync::Arc::new(crate::latency::Latency::new()); + // 请求进来之后过了好一阵才发出去(等空位、前面的几跳失败) + let arrived = Instant::now() - std::time::Duration::from_secs(5); + for _ in 0..3 { + let mut e = Ending::new(bus.clone(), 7, MODEL.into(), arrived, 1_000, None); + e.responded(200); + e.streaming(ir::Dialect::Anthropic, "up"); + let sent = Instant::now() - std::time::Duration::from_millis(300); + e.timed(lap(&latency, sent)); + e.feed(MESSAGE_START); + e.feed(TEXT_BLOCK); + assert_eq!(latency.typical("up"), None, "开场帧被当成了内容"); + e.feed(TEXT_DELTA); + e.feed(TEXT_DELTA); + e.finished(200); + } + let t = latency.typical("up").expect("三次都该记"); + assert!((300..2_000).contains(&t), "{t} 毫秒"); + } + + /// 没有内容的(流里只报了错、流断了)、不是流的,都不记 + #[test] + fn no_first_token_no_sample() { + let bus = tw_observe::EventBus::new(); + let latency = std::sync::Arc::new(crate::latency::Latency::new()); + for _ in 0..3 { + let mut e = responding(&bus); + e.streaming(ir::Dialect::Anthropic, "up"); + e.timed(lap(&latency, Instant::now())); + e.feed(MESSAGE_START); + e.feed(b"event: error\ndata: {\"type\":\"error\",\"error\":{\"type\":\"overloaded_error\",\"message\":\"Overloaded\"}}\n\n"); + e.finished(200); + // 整包的回答:没说是流,内容到了也不算 + let mut e = responding(&bus); + e.timed(lap(&latency, Instant::now())); + e.feed(TEXT_DELTA); + e.finished(200); + } + assert_eq!(latency.snapshot(&["up".to_string()]).len(), 0); + assert_eq!(latency.typical("up"), None); + } + /// 工具调用一开头就算:块开头就带着工具名,参数的第一段常常是空的 #[test] fn a_tool_call_opens_the_answer_before_its_arguments() { diff --git a/crates/tw-gateway/src/key_limits/mod.rs b/crates/tw-gateway/src/key_limits/mod.rs index 3aee8ef6..cf7eaf52 100644 --- a/crates/tw-gateway/src/key_limits/mod.rs +++ b/crates/tw-gateway/src/key_limits/mod.rs @@ -9,9 +9,9 @@ //! 此刻内存里的是同一份数。 //! //! **怎么拒、怎么等。**天、周、月是自然周期:用满了就拒,到下一期之前重试也没用。分钟、 -//! 小时是滚动窗口:用满了先看下一个空位多久之后空出来,等得到(不超过 `slot_wait_secs`) -//! 就等,等不到就拒,并说清楚多久之后再来。顺序见 [`crate::server`] 的管线第 3 步:自然 -//! 周期 → 并发上限 → 滚动窗口。 +//! 小时是滚动窗口:用满了先看下一个空位多久之后空出来,等得到(在这个请求的等待期限之前, +//! 期限见 [`slot_wait`])就等,等不到就拒,并说清楚多久之后再来。顺序见 [`crate::server`] +//! 的管线第 3 步:自然周期 → 并发上限 → 滚动窗口。 //! //! **在跑的怎么算。**准入时按输入估一个数占着(预留:输入 token 的估算,和按头一个候选算 //! 的输入费用),存储层记下那一行时换成实数。几个请求同时进来时,超出上限的最多是在跑的 @@ -59,11 +59,11 @@ const NOT_ADMITTED: &[&str] = &[ /// 一次,窗口就往后推一次,永远等不到空位 const REFUSED: &str = "gw.key_limit."; -/// 等滚动窗口空位最多等多久:`failover.slot_wait_secs`,和等上游空位(`crate::slots`) -/// 同一个设置,0 是不等。 +/// 一个请求最多等多久:`failover.slot_wait_secs`,0 是不等。 /// -/// **两段各算各的**:这一段在准入时、发给哪一家之前,等上游空位在之后的那几跳里。一个请求 -/// 两样都碰上的话,最多等两倍 +/// **一个请求合起来算**:等滚动窗口的空位在准入时、发给哪一家之前,等上游空位 +/// (`crate::slots`)在之后的那几跳里,两段共用准入时定下的同一个期限(见 +/// `server::pipeline::admission`)。各给一份的话,两样都碰上的请求能等两倍那么久 pub fn slot_wait(cfg: &tw_config::Config) -> Duration { Duration::from_secs(cfg.failover.slot_wait_secs) } @@ -468,21 +468,34 @@ impl KeyLimits { out } - /// 准入第三步:分钟、小时的上限。用满了先等下一个空位,最多等 `wait`;等不到就拒, - /// 带着多久之后能再来。**过了就记上**:这个请求算进滚动窗口,输入的估算占上预留。 - /// - /// 天、周、月在这里再看一遍:等并发名额、等空位的那一阵,别的请求可能把它用满了。 + /// 准入第三步,从此刻起最多等 `wait`(见 [`Self::admit_by`])。WebSocket 的连接用它: + /// 那条路之后不再等上游的空位,这一段就是全部 pub async fn admit( self: &Arc, key: &str, limits: &[KeyLimit], ask: Ask, wait: Duration, + ) -> Result> { + self.admit_by(key, limits, ask, tokio::time::Instant::now() + wait) + .await + } + + /// 准入第三步:分钟、小时的上限。用满了先等下一个空位,最多等到 `until`(这个请求的 + /// 等待期限,见 [`slot_wait`]);等不到就拒,带着多久之后能再来。**过了就记上**:这个 + /// 请求算进滚动窗口,输入的估算占上预留。 + /// + /// 天、周、月在这里再看一遍:等并发名额、等空位的那一阵,别的请求可能把它用满了。 + pub async fn admit_by( + self: &Arc, + key: &str, + limits: &[KeyLimit], + ask: Ask, + until: tokio::time::Instant, ) -> Result> { if limits.is_empty() { return Ok(Hold::none()); } - let deadline = self.clock.now_ms() + wait.as_millis() as i64; loop { let now = self.clock.now_ms(); let step = { @@ -511,8 +524,9 @@ impl KeyLimits { r } }; - // 滚动窗口:空位在等得到的时候空出来就等,等不到就拒 - if refusal.resets.is_some() || now + refusal.retry_after_ms as i64 > deadline { + // 滚动窗口:空位在期限之前空出来就等,等不到就拒 + let left = until.saturating_duration_since(tokio::time::Instant::now()); + if refusal.resets.is_some() || u128::from(refusal.retry_after_ms) > left.as_millis() { return Err(refusal); } tokio::time::sleep(Duration::from_millis(refusal.retry_after_ms)).await; diff --git a/crates/tw-gateway/src/key_limits/tests.rs b/crates/tw-gateway/src/key_limits/tests.rs index 21ef5f0e..b11821db 100644 --- a/crates/tw-gateway/src/key_limits/tests.rs +++ b/crates/tw-gateway/src/key_limits/tests.rs @@ -255,6 +255,32 @@ async fn a_rolling_minute_waits_for_its_next_slot_or_refuses() { assert_eq!(b.limits.clock.now_ms(), t0 + 60_000, "等到第一个滑出窗口"); } +/// 等待期限是整个请求的那一个(见 [`slot_wait`]),调用方给的是那一刻,不是从此刻起再给 +/// 一整段:期限之前空不出来的就拒 +#[tokio::test(start_paused = true)] +async fn a_rolling_wait_ends_at_the_requests_own_deadline() { + let b = bed("2026-10-05T10:00:00+08:00", "[{per: minute, requests: 1}]"); + b.run(1, Ask::default()).await.unwrap(); + let until = tokio::time::Instant::now() + Duration::from_secs(30); + // 35 秒之后,空位还要 25 秒才出来:期限已经过了,拒 + tokio::time::advance(Duration::from_secs(35)).await; + let r = b + .limits + .admit_by("k", &b.set, Ask::default(), until) + .await + .unwrap_err(); + assert_eq!(r.retry_after_ms, 25_000); + // 从此刻算起的一整段(`admit`)就等得到 + let t = b.now(); + let hold = b + .limits + .admit("k", &b.set, Ask::default(), Duration::from_secs(30)) + .await + .unwrap(); + drop(hold); + assert_eq!(b.now(), t + 25_000); +} + #[tokio::test(start_paused = true)] async fn tokens_in_a_rolling_hour_count_when_they_are_settled() { let b = bed("2026-10-05T10:00:00+08:00", "[{per: hour, tokens: 1000}]"); diff --git a/crates/tw-gateway/src/latency.rs b/crates/tw-gateway/src/latency.rs index 35145702..deb42165 100644 --- a/crates/tw-gateway/src/latency.rs +++ b/crates/tw-gateway/src/latency.rs @@ -1,17 +1,28 @@ -//! 每家上游的典型首字节时间(`url-test`)。 +//! 每家上游典型的快慢:**从这一跳发出去,到回答的第一段内容**。`url-test` 按它挑最快的, +//! `load-balance` 按快慢分时(`balance_by: latency`)按它算系数。 //! -//! # 为什么是 TTFB 而不是总时长 +//! # 量的是哪一段 //! -//! 总时长里最大的一项是「模型说了多少字」,那和上游快不快无关 —— -//! 一个回答长的请求会让最快的上游看起来最慢。**TTFB 才是这一层能控制、 -//! 用户也真的在等的那段**(L3 也是这个判据)。 +//! - **起点是这一跳发出去的那一刻**,不是请求进网关的那一刻:等密钥的上限、等上游的空位、 +//! 跑插件、前面失败了的几跳,都不是这一家慢。记到它头上的话,排在第二位的那家永远替 +//! 第一家背着那段时间。 +//! - **终点是第一段内容**(正文、推理的字、工具调用的开头 —— 和首 token 同一个认法,见 +//! [`crate::ending`]),不是响应头:不少中转收下请求马上回 200,然后在流里排队,响应头 +//! 来得快什么都说明不了。也不是总时长:总时长里最大的一项是「模型说了多少字」,那和 +//! 上游快不快无关 —— 一个回答长的请求会让最快的上游看起来最慢。 +//! - **和这一跳排在第几家无关**:排在前面的那几家,开头要先看一眼有没有报错(见 +//! `server::pipeline::opening`),最后一家直接转发 —— 量的都是同一段。 +//! - **只有流式回答有样本**:整包的回答只有「全到了」那一个时刻,分不出哪段在排队、哪段 +//! 在说话。 +//! - **开头慢被放弃的那一家**(`failover.next_on_slow_start`)记它被给的那段时间:它至少 +//! 这么慢,记成这个数,它就排到慢的那一头去。什么都不记的话,它留着的还是从前快的 +//! 样本,下一个请求还先发给它。 //! //! # 为什么是中位数而不是 EWMA //! //! EWMA 的问题是**说不清**:用户问「为什么切到乙了」,我们只能给一个 //! 被历史加权过的数,而它对应不到任何一次真实的请求。中位数返回的 -//! 永远是一个**真实发生过的值**,用户能在请求列表里找到它 —— 和分位数 -//! 用最近秩法是同一条理由。 +//! 永远是一个**真实发生过的值** —— 和分位数用最近秩法是同一条理由。 //! //! 这也是为什么当初说「EWMA 有意不做」:那时它没有消费者。 //! 现在 `url-test` 就是它的消费者,而消费者想要的是能解释的数。 @@ -19,6 +30,7 @@ use std::collections::HashMap; use std::collections::VecDeque; use std::sync::Mutex; +use std::time::Duration; /// 每家留多少个样本。 /// @@ -32,9 +44,36 @@ const WINDOW: usize = 32; /// DNS 没缓存 —— 那个数字比没有更误导。 const MIN_SAMPLES: usize = 3; +/// 一家的样本。 +#[derive(Default)] +struct Window { + /// 真实请求量到的,最近的在后面 + real: VecDeque, + /// 启动时 L1 握手垫的底(见 [`seed_url_test`])。**真实样本够数之前顶着用,够数就扔掉**, + /// 从不和真实样本放在一起取中位数 + seed: Option, +} + +impl Window { + /// 这家的典型值:真实样本够数就是它们的中位数,不够时有垫底的用垫底的,都没有是 `None` + fn typical(&self) -> Option { + if self.real.len() >= MIN_SAMPLES { + let mut v: Vec = self.real.iter().copied().collect(); + v.sort_unstable(); + return Some(v[v.len() / 2]); + } + self.seed + } +} + +/// 一段时间写成样本的毫秒数 +pub fn ms(d: Duration) -> u32 { + d.as_millis().min(u128::from(u32::MAX)) as u32 +} + #[derive(Default)] pub struct Latency { - inner: Mutex>>, + inner: Mutex>, } impl Latency { @@ -42,57 +81,49 @@ impl Latency { Self::default() } - /// 记一次。**转发路径上调,所以必须便宜**:一次哈希加一次 push。 - pub fn record(&self, provider: &str, ttfb_ms: u32) { + /// 记一个真实样本,毫秒。**转发路径上调,所以必须便宜**:一次哈希加一次 push。 + pub fn record(&self, provider: &str, ms: u32) { let mut g = self.inner.lock().expect("lock not poisoned"); let w = g.entry(provider.to_string()).or_default(); - if w.len() == WINDOW { - w.pop_front(); + if w.real.len() == WINDOW { + w.real.pop_front(); + } + w.real.push_back(ms); + // 真实样本够数了:垫的底用不着了。**扔掉,不是混进去** —— L1 量的是建连,不含上游 + // 排队和推理,混进中位数会把一家慢的上游拉成快的 + if w.real.len() >= MIN_SAMPLES { + w.seed = None; } - w.push_back(ttfb_ms); } - /// 启动时用 L1 握手的结果垫一个底(样本不够时用零成本的 L1 补)。 + /// 用 L1 握手的结果垫一个底:真实样本还不够数时,[`Self::typical`] 先用它。 /// - /// **只在完全没有真实样本时垫**。真实流量一到就该由它说了算 —— - /// L1 量的是建连,不含上游排队和推理,天生偏乐观。 - pub fn seed(&self, provider: &str, ttfb_ms: u32) { + /// **真实样本够数之后不再垫**:真实流量一到就该由它说了算。够数之前再垫一次(每天跟着 + /// 模型清单刷新一次)换成新测的那个数。 + /// + /// 不够数时用垫底的而不是当作没测过:`url-test` 把没测过的排在最后,一家刚收到第一个 + /// 请求的上游要是因此沉到最后,就再也轮不到它攒够样本了。 + pub fn seed(&self, provider: &str, ms: u32) { let mut g = self.inner.lock().expect("lock not poisoned"); let w = g.entry(provider.to_string()).or_default(); - if w.is_empty() { - // 垫满 `MIN_SAMPLES` 才算数 —— 否则它自己也是「样本不够」 - for _ in 0..MIN_SAMPLES { - w.push_back(ttfb_ms); - } + if w.real.len() < MIN_SAMPLES { + w.seed = Some(ms); } } - /// 这家的典型 TTFB。`None` = 样本不够,**不是「很快」**。 + /// 这家的典型值,毫秒。`None` = 样本不够、也没垫过底,**不是「很快」**。 pub fn typical(&self, provider: &str) -> Option { let g = self.inner.lock().expect("lock not poisoned"); - let w = g.get(provider)?; - if w.len() < MIN_SAMPLES { - return None; - } - let mut v: Vec = w.iter().copied().collect(); - v.sort_unstable(); - Some(v[v.len() / 2]) + g.get(provider)?.typical() } /// 一次取一批 —— 排序时要用到,逐个取会连着锁好几次。 pub fn snapshot(&self, names: &[String]) -> HashMap { let g = self.inner.lock().expect("lock not poisoned"); - let mut out = HashMap::new(); - for n in names { - if let Some(w) = g.get(n) - && w.len() >= MIN_SAMPLES - { - let mut v: Vec = w.iter().copied().collect(); - v.sort_unstable(); - out.insert(n.clone(), v[v.len() / 2]); - } - } - out + names + .iter() + .filter_map(|n| Some((n.clone(), g.get(n)?.typical()?))) + .collect() } } @@ -103,8 +134,8 @@ impl Latency { /// 没有这一步的话,`url-test` 在攒够真实样本之前完全等同于 `fallback` /// —— 用户配了「选最快的」,而头几十个请求全落在配置里排第一那家。 /// -/// L1 是握手计时,不发一个 API 请求、不花一分钱;也**不是定期 -/// 跑的** —— 真实流量一到就该由它说了算。 +/// L1 是握手计时,不发一个 API 请求、不花一分钱。**真实样本够数之后它就不算了** +/// (见 [`Latency::seed`])—— 真实流量一到就该由它说了算。 pub async fn seed_url_test(state: &crate::state::AppState) { let rt = state.runtime(); let mut want: Vec = Vec::new(); @@ -169,7 +200,7 @@ mod tests { #[test] fn the_median_is_always_a_number_that_really_happened() { - // 用户问「为什么切到乙了」时,这个数要能在请求列表里找得到 + // 用户问「为什么切到乙了」时,这个数要是某一次请求真的量到的 let l = Latency::new(); for ms in [100, 105, 110, 5000, 108] { l.record("甲", ms); @@ -210,6 +241,38 @@ mod tests { assert_eq!(l.typical("甲"), Some(300)); } + #[test] + fn a_seed_stands_in_until_real_samples_are_enough_and_is_never_mixed_in() { + let l = Latency::new(); + l.seed("甲", 20); + // 真实样本还不够数:顶着用垫的底。当作没测过的话,`url-test` 把它排到最后,它就再也 + // 攒不够样本了 + l.record("甲", 900); + l.record("甲", 50); + assert_eq!(l.typical("甲"), Some(20)); + // 够数了:只看真实的。混在一起取中位数是 [20, 20, 20, 50, 900, 900] 里的 50 + l.record("甲", 900); + assert_eq!(l.typical("甲"), Some(900)); + let names = vec!["甲".to_string()]; + assert_eq!(l.snapshot(&names).get("甲"), Some(&900)); + } + + #[test] + fn a_new_seed_replaces_the_old_one_while_real_samples_are_too_few() { + // 每天跟着模型清单再测一次:还没攒够真实样本的,换成新测的数 + let l = Latency::new(); + l.seed("甲", 20); + l.record("甲", 500); + l.seed("甲", 80); + assert_eq!(l.typical("甲"), Some(80)); + } + + #[test] + fn a_duration_is_written_in_milliseconds() { + assert_eq!(ms(Duration::from_millis(1_234)), 1_234); + assert_eq!(ms(Duration::from_secs(u64::MAX)), u32::MAX); + } + #[test] fn a_snapshot_leaves_out_the_ones_with_too_few_samples() { let l = Latency::new(); @@ -217,9 +280,16 @@ mod tests { for _ in 0..3 { l.record("乙", 200); } - let names = vec!["甲".to_string(), "乙".to_string(), "丙".to_string()]; + l.seed("丁", 30); + let names = vec![ + "甲".to_string(), + "乙".to_string(), + "丙".to_string(), + "丁".to_string(), + ]; let s = l.snapshot(&names); - assert_eq!(s.len(), 1, "{s:?}"); + assert_eq!(s.len(), 2, "{s:?}"); assert_eq!(s.get("乙"), Some(&200)); + assert_eq!(s.get("丁"), Some(&30), "垫过底的算测过"); } } diff --git a/crates/tw-gateway/src/server/pipeline.rs b/crates/tw-gateway/src/server/pipeline.rs index e528dc6b..bb7417e3 100644 --- a/crates/tw-gateway/src/server/pipeline.rs +++ b/crates/tw-gateway/src/server/pipeline.rs @@ -62,6 +62,9 @@ struct Started { ledger: tw_guard::redact::replace::Ledger, /// 出站脱敏在客户端原文里找到的。插件改过的那一跳只再报插件写进来的(见 [`plug`]) found: Vec, + /// 这个请求最多等到什么时候:准入时定下(见 [`admission`]),等密钥的分钟、小时上限和 + /// 等上游的空位共用这一段(`failover.slot_wait_secs`) + wait_until: tokio::time::Instant, } pub(super) async fn pipeline( @@ -122,7 +125,11 @@ pub(super) async fn pipeline( // // **并发的通行证交给回程,跟着响应体走**(见 `relay`):放在这里的话它在响应头交出去的 // 那一刻就还了,一条还在流的回答不再算数,上限管的只是等响应头的那一段 - let admission::Admitted { pass, hold } = admission::admit( + let admission::Admitted { + pass, + hold, + wait_until, + } = admission::admit( &state, &rt, &req, @@ -142,6 +149,7 @@ pub(super) async fn pipeline( &req, &reading, choice, + wait_until, &decision, fp.as_deref(), ending, @@ -235,7 +243,7 @@ pub(super) async fn pipeline( /// 数 token 由网关估了数(见 [`crate::count`]):把它交给客户端,照常报响应头和结局。 /// -/// **不经过 `relay`**:那里按上游的回答记首字节时间、额度、凭据和代理的状态,而这个 +/// **不经过 `relay`**:那里按上游的回答记快慢样本、额度、凭据和代理的状态,而这个 /// 回答不是上游给的 —— 记上去的话,一家从没被问过的上游会显示成「刚刚答得飞快」。 /// 响应头上带 `x-thinkwatch-local`,和本地应答的一样。 fn estimated( @@ -736,6 +744,7 @@ fn start( req: &Inbound, reading: &crate::client_api::Reading, choice: Choice, + wait_until: tokio::time::Instant, decision: &tw_engine::Decision, fp: Option<&str>, ending: &mut Option, @@ -795,6 +804,7 @@ fn start( conversation: crate::affinity::identity(&req.headers, fp), ledger, found, + wait_until, } } diff --git a/crates/tw-gateway/src/server/pipeline/admission.rs b/crates/tw-gateway/src/server/pipeline/admission.rs index 8ccdea62..41f28cfd 100644 --- a/crates/tw-gateway/src/server/pipeline/admission.rs +++ b/crates/tw-gateway/src/server/pipeline/admission.rs @@ -5,9 +5,14 @@ //! 1. 天、周、月的上限([`crate::key_limits`])—— 这一期用满了直接拒:到下一期之前等多久 //! 都一样,不用先排一轮并发的队; //! 2. 并发上限([`crate::limits`])—— 等,不拒,理由在那儿; -//! 3. 分钟、小时的上限 —— 下一个空位在 `slot_wait_secs` 之内空出来就等,等不到就拒,并说清 -//! 多久之后再来。过了就把这个请求记上,按输入估一个数占着,等存储层记下它那一行时换成 -//! 实数。 +//! 3. 分钟、小时的上限 —— 下一个空位在这个请求的等待期限之前空出来就等,等不到就拒,并 +//! 说清多久之后再来。过了就把这个请求记上,按输入估一个数占着,等存储层记下它那一行时 +//! 换成实数。 +//! +//! **等待期限一个请求只有一个**:过了并发上限那一刻起算 `failover.slot_wait_secs`,这里等 +//! 滚动窗口的空位、之后等上游的空位(见 `hop`)都算在里面 —— 两段各给一份的话,一个请求能 +//! 等两倍那么久,而等的时候客户端一个字节都收不到。并发上限那一段不算:它等前面的请求结束, +//! 不拒绝,等多久由客户端决定(见 [`crate::limits`])。 //! //! 被上限拒绝的请求**照样开始、照样留一行**(和路由拒绝的一样,见 //! [`crate::server::routed_nowhere`]):流量里看得见它被哪一条上限拒了。数 token 的请求 @@ -25,6 +30,8 @@ pub(super) struct Admitted { pub(super) pass: crate::limits::Pass, /// 用量上限的预留。开始事件之后交给请求号([`Hold::bind`]),丢掉就放掉 pub(super) hold: Hold, + /// 这个请求最多等到什么时候。这一步等滚动窗口用掉的,之后等上游空位就少等那么久 + pub(super) wait_until: tokio::time::Instant, } #[allow(clippy::too_many_arguments)] @@ -48,18 +55,23 @@ pub(super) async fn admit( return Err(refused(state, rt, req, reading, choice, fp, ending, &r)); } let pass = state.gate.acquire(&req.client_name, max_concurrent).await; + // 等待期限从这里起算:之后的两段等待共用它 + let wait_until = tokio::time::Instant::now() + crate::key_limits::slot_wait(&rt.config); let ask = if limits.is_empty() { Ask::default() } else { ask(state, rt, reading, decision, limits) }; - let wait = crate::key_limits::slot_wait(&rt.config); match state .key_limits - .admit(&req.client_name, limits, ask, wait) + .admit_by(&req.client_name, limits, ask, wait_until) .await { - Ok(hold) => Ok(Admitted { pass, hold }), + Ok(hold) => Ok(Admitted { + pass, + hold, + wait_until, + }), Err(r) => Err(refused(state, rt, req, reading, choice, fp, ending, &r)), } } diff --git a/crates/tw-gateway/src/server/pipeline/hop.rs b/crates/tw-gateway/src/server/pipeline/hop.rs index 688d96b3..4633a75f 100644 --- a/crates/tw-gateway/src/server/pipeline/hop.rs +++ b/crates/tw-gateway/src/server/pipeline/hop.rs @@ -46,6 +46,8 @@ pub(super) struct Served<'a> { /// 这一跳在这家占着的位置(见 [`crate::slots`])。**跟着回答走**:交完、或者客户端 /// 走掉,响应体被丢掉时才还回去 pub(super) slot: crate::slots::Slot, + /// 这一跳发出去的那一刻。这一家的快慢样本从这里算起(见 [`crate::latency`]) + pub(super) sent_at: std::time::Instant, } /// 这个请求的着落。 @@ -146,10 +148,10 @@ pub(super) async fn try_upstreams<'a>( // 满着没发的那几家(见 `crate::slots`),按候选的顺序。候选都看过一遍还没有着落时, // 在它们里面等先空出来的那一家 let mut busy: Vec<&String> = Vec::new(); - // 等空位等到什么时候:头一次要等时定下,**整个请求共用这一段** —— 等的时候客户端 - // 一个字节都收不到 - let wait = std::time::Duration::from_secs(rt.config.failover.slot_wait_secs); - let mut deadline: Option = None; + // 等空位等到什么时候:准入时定下的那一刻(见 `super::admission`),**整个请求共用** + // —— 等密钥的分钟、小时上限已经用掉的,这里就少等那么久。等的时候客户端一个字节都 + // 收不到 + let until = started.wait_until; // 等过空位的那一跳在尝试链上的位置和等了多久。那一跳怎么收场都只进一行,进了之后补上 let mut queued: Option<(usize, u64)> = None; // 等到最后,剩下的候选还都满着 @@ -162,7 +164,6 @@ pub(super) async fn try_upstreams<'a>( Some(name) => (name, None), None if busy.is_empty() => break, None => { - let until = *deadline.get_or_insert_with(|| tokio::time::Instant::now() + wait); let t = std::time::Instant::now(); match state.slots.first_free(&busy, until).await { Some((k, slot)) => (busy.remove(k), Some((slot, t.elapsed()))), @@ -325,7 +326,6 @@ pub(super) async fn try_upstreams<'a>( // 这段对话留在这家是为了它的缓存:等它空出来,等不到再换下一家(缓存就丢在 // 这家了)。别的候选满着当场跳过 if slot.is_none() && started.choice.stayed_on.as_ref() == Some(name) { - let until = *deadline.get_or_insert_with(|| tokio::time::Instant::now() + wait); // 这个请求能等的已经等完了(或者配置的是不等):不再等 if tokio::time::Instant::now() < until { let t = std::time::Instant::now(); @@ -511,11 +511,13 @@ pub(super) async fn try_upstreams<'a>( let rest = queue.iter().chain(busy.iter()).copied(); successor(state, rt, req, reading, decision, &catalog, allow, rest) }; + // 这一跳发出去的那一刻:这一家的快慢样本从这里算起(见 `crate::latency`) + let sent_at = std::time::Instant::now(); // 开头慢就换下一家:等到什么时候,从这一刻(请求发出去)算起。最后一家不换。到点时 // 问过、后面没有接得下的,清掉它:这一跳从此和不开时一样 let mut slow_deadline = slow_wait .filter(|_| !last) - .map(|w| tokio::time::Instant::now() + w); + .map(|w| tokio::time::Instant::from_std(sent_at) + w); // 发出去、等响应头。**等着的这个 future 只活在这一块里**:放弃这一家时它跟着丢掉, // 连接随之断开 @@ -556,6 +558,7 @@ pub(super) async fn try_upstreams<'a>( // `super::slow`) Err(_) if others() => { let waited = slow_wait.unwrap_or_default(); + super::slow::timed_out(state, &provider.name, waited); chain.push(super::slow::abandoned( &provider.name, model.clone(), @@ -687,6 +690,7 @@ pub(super) async fn try_upstreams<'a>( session: out.session, refusal, slot, + sent_at, }); break; } @@ -718,6 +722,7 @@ pub(super) async fn try_upstreams<'a>( // 上游回了话,说明代理是通的 state.note_proxy_ok(&provider.proxy); let waited = slow_wait.unwrap_or_default(); + super::slow::timed_out(state, &provider.name, waited); chain.push(super::slow::abandoned( &provider.name, model.clone(), @@ -819,6 +824,7 @@ pub(super) async fn try_upstreams<'a>( session: out.session, refusal: None, slot, + sent_at, }); break; } diff --git a/crates/tw-gateway/src/server/pipeline/relay.rs b/crates/tw-gateway/src/server/pipeline/relay.rs index b36b6a45..db683e98 100644 --- a/crates/tw-gateway/src/server/pipeline/relay.rs +++ b/crates/tw-gateway/src/server/pipeline/relay.rs @@ -46,6 +46,7 @@ pub(super) fn respond( session, refusal, slot, + sent_at, .. } = served; let status = @@ -54,13 +55,6 @@ pub(super) fn respond( // 到结束可能还有好几分钟,UI 要能在这个点就把行画出来并标「进行中」, // 而不是等它结束才出现。 let ttfb_ms = req.started.elapsed().as_millis() as u64; - // `url-test` 的判据。**只记成功的那些** —— 一个 500 在 - // 十毫秒内返回,会让最坏的上游看起来最快 - if status.is_success() { - state - .latency - .record(&provider.name, ttfb_ms.min(u32::MAX as u64) as u32); - } state.bus.emit(tw_api::Event::RequestHeaders { id, status: status.as_u16(), @@ -130,6 +124,14 @@ pub(super) fn respond( // 成功的流才有「第一个 token」:整包的一起到,错误不是回答 if generates && plan.is_sse && status.is_success() { ending.streaming(upstream_dialect, &provider.name); + // `url-test` 和按快慢分的 `load-balance` 的判据:从这一跳发出去到第一个 token(见 + // `crate::latency`)。**只有成功的流有**:一个 500 在十毫秒内返回,会让最坏的上游 + // 看起来最快;整包的回答分不出排队和说话 + ending.timed(crate::ending::Lap { + latency: state.latency.clone(), + provider: provider.name.clone(), + sent: sent_at, + }); } // 回的不是 2xx:原样交给客户端,**这个请求照样是失败的** —— 客户端拿到的是上游的 // 错误,不是回答。原因在错误正文里,交完时读(见 `Ending::refused`) diff --git a/crates/tw-gateway/src/server/pipeline/slow.rs b/crates/tw-gateway/src/server/pipeline/slow.rs index f9bb0c13..ce7e1830 100644 --- a/crates/tw-gateway/src/server/pipeline/slow.rs +++ b/crates/tw-gateway/src/server/pipeline/slow.rs @@ -11,6 +11,8 @@ //! `hop::successor`):停用着的、这一跳发不出去的不算 —— 否则放弃了一个慢的,换来的是 //! 一个注定失败的。 //! - **这一家不停用、不算失败**:慢不是坏,下一个请求它可能就快了。 +//! - **它的快慢样本记它被给的那段时间**(见 [`timed_out`]):`url-test` 和按快慢分的 +//! `load-balance` 照这个把它往后排。 //! - **尝试链上记一跳 `slow_start`**,带着上游可能已经收了钱的输入(见 //! [`tw_api::AttemptUsage`])。 //! - 只管客户端要流式的请求:整包的请求本来就要等全部生成完,开头慢说明不了什么。 @@ -33,6 +35,15 @@ pub(super) fn wait(rt: &Runtime, reading: &crate::client_api::Reading) -> Option (f.next_on_slow_start && streams).then(|| Duration::from_secs(f.stream_start_wait_secs)) } +/// 放弃了这一家:给它记一个快慢样本,就是它被给的那段时间(见 [`crate::latency`])。 +/// +/// **它至少这么慢**,这是个下限:记成这个数,`url-test` 和按快慢分的 `load-balance` 就把它 +/// 排到慢的那一头。什么都不记的话,它留着的还是从前快的样本,下一个请求照样先发给它, +/// 而等它的这段时间算到了接下来那一家头上。 +pub(super) fn timed_out(state: &crate::state::AppState, provider: &str, waited: Duration) { + state.latency.record(provider, crate::latency::ms(waited)); +} + /// 放弃了的那一跳:尝试链上的一行。`status` 是上游回的(响应头没到的没有),`seen` 是流 /// 开头里上游报的用量。 pub(super) fn abandoned( diff --git a/crates/tw-gateway/src/slots.rs b/crates/tw-gateway/src/slots.rs index 229c1473..61f51cc0 100644 --- a/crates/tw-gateway/src/slots.rs +++ b/crates/tw-gateway/src/slots.rs @@ -9,8 +9,9 @@ //! - 别的候选满着:当场跳过,试下一家; //! - 候选都满着:在它们里面等先空出来的那一家,等不到就回 429。 //! -//! 等多久是 `failover.slot_wait_secs`,**一个请求合起来算**:等的时候客户端一个字节都 -//! 收不到。**等不是失败**:满着的上游不停用、不进熔断的账。 +//! 等多久是 `failover.slot_wait_secs`,**一个请求合起来算**,准入时等密钥的分钟、小时上限 +//! 用掉的也算在里面(见 `crate::key_limits::slot_wait`):等的时候客户端一个字节都收不到。 +//! **等不是失败**:满着的上游不停用、不进熔断的账。 //! //! 一个位置从发出请求占到回答交完、或者客户端走掉,由 [`Slot`] 的 Drop 还回去。和每把 //! 密钥的闸([`crate::limits`])一样,**跨重载存活**,上限改了在原来那个信号量上加减: @@ -131,6 +132,14 @@ impl Slots { Some(books.limit) } + /// 这家此刻满着:设了上限、一个空位都没有。`load-balance` 排这一轮时看它(满着的不 + /// 参加,见 `tw_engine::weighted`)。**只是此刻的样子**:真要发的时候还是 [`Self::try_take`] + /// 说了算 + pub fn is_full(&self, name: &str) -> bool { + self.pool(name) + .is_some_and(|pool| !pool.sem.is_closed() && pool.sem.available_permits() == 0) + } + /// 不等:有空位就占一个,这家满着是 `None`。不限并发的上游一律给一个空的。 pub fn try_take(&self, name: &str) -> Option { let Some(pool) = self.pool(name) else { diff --git a/crates/tw-gateway/src/state.rs b/crates/tw-gateway/src/state.rs index d8013aa1..75af7f45 100644 --- a/crates/tw-gateway/src/state.rs +++ b/crates/tw-gateway/src/state.rs @@ -189,7 +189,8 @@ pub struct AppState { /// 上游各自多打一次往返。用户真的改了 refresh token 时,缓存自己认 /// 得出来(指纹对不上就重换)。 pub oauth: Arc, - /// 每家的典型首字节时间。`url-test` 策略靠它排序。 + /// 每家典型的快慢:从发出去到回答的第一段内容(见 [`crate::latency`])。`url-test` 靠它 + /// 排序,按快慢分的 `load-balance` 靠它算系数。 /// /// **跨重载存活**:改一条规则不该让所有上游回到「没测过」。 pub latency: Arc, @@ -415,6 +416,16 @@ impl AppState { } else { Default::default() }, + // 并发数满着的也不参加这一轮:它们会被当场跳过(见 `crate::slots`) + busy: if balanced { + candidates + .iter() + .filter(|p| self.slots.is_full(p)) + .cloned() + .collect() + } else { + Default::default() + }, // `url-test` 选最快的;`load-balance` 按快慢分的也看它 ttfb_ms: if g.kind == GroupType::UrlTest || (balanced && g.balance_by.uses_latency()) { self.latency.snapshot(candidates) diff --git a/crates/tw-gateway/tests/latency_samples.rs b/crates/tw-gateway/tests/latency_samples.rs new file mode 100644 index 00000000..3f451ef8 --- /dev/null +++ b/crates/tw-gateway/tests/latency_samples.rs @@ -0,0 +1,244 @@ +//! 快慢样本(`url-test` 和按快慢分的 `load-balance` 看的那个数)量的是哪一段,端到端:从 +//! 这一跳发出去到回答的第一段内容 —— 不含之前失败了的几跳,不是响应头,和这一跳排在第几家 +//! 无关;整包的回答不记;开头慢被放弃的那一家记它被给的那段时间。 + +use std::net::SocketAddr; +use std::sync::Arc; +use std::time::Duration; + +use axum::Router; +use bytes::Bytes; +use serde_json::{Value, json}; +use tw_config::{Client, Config, Failover, Protocol, Provider}; +use tw_gateway::latency::Latency; + +/// 一个假上游的样子 +#[derive(Clone, Default)] +struct Script { + /// 响应头之前等多久 + header_delay_ms: u64, + /// 回一个错误(这个状态码),不回答 + status: Option, + /// 流式:开头的例行帧之后隔多久出第一段内容 + content_after_ms: u64, + /// 流式:发完开头就只发心跳,一直不出内容 + stall: bool, +} + +const MESSAGE_START: &str = "event: message_start\ndata: {\"type\":\"message_start\",\"message\":{\"usage\":{\"input_tokens\":12,\"output_tokens\":1}}}\n\n"; +const PING: &str = "event: ping\ndata: {\"type\":\"ping\"}\n\n"; +const ANSWER: &str = "event: content_block_start\ndata: {\"type\":\"content_block_start\",\"index\":0,\"content_block\":{\"type\":\"text\",\"text\":\"\"}}\n\n\ + event: content_block_delta\ndata: {\"type\":\"content_block_delta\",\"index\":0,\"delta\":{\"type\":\"text_delta\",\"text\":\"hello\"}}\n\n\ + event: content_block_stop\ndata: {\"type\":\"content_block_stop\",\"index\":0}\n\n\ + event: message_stop\ndata: {\"type\":\"message_stop\"}\n\n"; +const WHOLE: &str = r#"{"type":"message","role":"assistant","content":[{"type":"text","text":"hello"}],"usage":{"input_tokens":12,"output_tokens":1}}"#; + +async fn upstream(script: Script) -> SocketAddr { + let app = Router::new().fallback(axum::routing::any(move |body: Bytes| { + let script = script.clone(); + async move { + tokio::time::sleep(Duration::from_millis(script.header_delay_ms)).await; + if let Some(status) = script.status { + return axum::response::Response::builder() + .status(status) + .header("content-type", "application/json") + .body(axum::body::Body::from( + r#"{"type":"error","error":{"type":"api_error","message":"boom"}}"#, + )) + .unwrap(); + } + let streams = serde_json::from_slice::(&body) + .ok() + .and_then(|v| v["stream"].as_bool()) + .unwrap_or(false); + if !streams { + tokio::time::sleep(Duration::from_millis(script.content_after_ms)).await; + return axum::response::Response::builder() + .header("content-type", "application/json") + .body(axum::body::Body::from(WHOLE)) + .unwrap(); + } + let (after, stall) = (script.content_after_ms, script.stall); + let body = async_stream::stream! { + yield Ok::<_, std::convert::Infallible>(Bytes::from_static(MESSAGE_START.as_bytes())); + if stall { + loop { + tokio::time::sleep(Duration::from_millis(100)).await; + yield Ok(Bytes::from_static(PING.as_bytes())); + } + } else { + tokio::time::sleep(Duration::from_millis(after)).await; + yield Ok(Bytes::from_static(ANSWER.as_bytes())); + } + }; + axum::response::Response::builder() + .header("content-type", "text/event-stream") + .body(axum::body::Body::from_stream(body)) + .unwrap() + } + })); + let l = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap(); + let addr = l.local_addr().unwrap(); + tokio::spawn(async move { axum::serve(l, app).await.unwrap() }); + addr +} + +fn provider(name: &str, at: SocketAddr) -> Provider { + Provider { + name: name.into(), + base_url: format!("http://{at}"), + key: Some("sk-upstream".into()), + protocol: Some(Protocol::Anthropic), + ..Default::default() + } +} + +/// 起网关,按声明的顺序故障转移。开头最多等 1 秒(测试里图快;配置校验要求换家时至少 +/// 5 秒,网关自己不查) +async fn gateway(providers: Vec, switch: bool) -> (SocketAddr, Arc) { + let cfg = Config { + version: 1, + clients: vec![Client { + name: "c".into(), + key: "tw-k".into(), + ..Default::default() + }], + providers, + failover: Failover { + stream_start_wait_secs: 1, + next_on_slow_start: switch, + ..Default::default() + }, + ..Default::default() + }; + let state = tw_gateway::AppState::new(cfg).unwrap(); + let latency = state.latency.clone(); + let addr = tw_gateway::serve(state, ([127, 0, 0, 1], 0).into()) + .await + .unwrap(); + tokio::time::sleep(Duration::from_millis(50)).await; + (addr, latency) +} + +/// 发 `n` 个请求,每个都读完 +async fn ask(gw: SocketAddr, stream: bool, n: usize) { + let body = json!({"model": "claude-sonnet-5", "max_tokens": 64, "stream": stream, + "messages": [{"role": "user", "content": "Say hello."}]}); + let client = reqwest::Client::builder().no_proxy().build().unwrap(); + for _ in 0..n { + let resp = client + .post(format!("http://{gw}/v1/messages")) + .header("content-type", "application/json") + .header("x-api-key", "tw-k") + .body(body.to_string()) + .send() + .await + .unwrap(); + let status = resp.status().as_u16(); + let text = resp.text().await.unwrap(); + assert_eq!(status, 200, "{text}"); + assert!(text.contains("hello"), "{text}"); + } +} + +#[tokio::test] +async fn the_sample_runs_to_the_first_content_not_to_the_response_headers() { + // 响应头和开头的例行帧马上就到,内容 400 毫秒之后才来:上游在排队。它是唯一的一家(最后 + // 一家,开头不看直接转发),以前记的是响应头到的那一刻 + let up = upstream(Script { + content_after_ms: 400, + ..Default::default() + }) + .await; + let (gw, latency) = gateway(vec![provider("排队", up)], false).await; + ask(gw, true, 3).await; + let t = latency.typical("排队").expect("三个流式回答该有三个样本"); + assert!((400..1_500).contains(&t), "{t} 毫秒"); +} + +#[tokio::test] +async fn earlier_hops_are_not_charged_to_the_upstream_that_answers() { + // 第一家 700 毫秒后回 500,第二家 100 毫秒出内容:第二家的样本里没有第一家那段 + let bad = upstream(Script { + header_delay_ms: 700, + status: Some(500), + ..Default::default() + }) + .await; + let good = upstream(Script { + content_after_ms: 100, + ..Default::default() + }) + .await; + let (gw, latency) = gateway(vec![provider("坏", bad), provider("好", good)], false).await; + ask(gw, true, 3).await; + let t = latency.typical("好").expect("三个样本"); + assert!( + (100..600).contains(&t), + "{t} 毫秒:第一家的时间算到了它头上" + ); + assert_eq!(latency.typical("坏"), None, "失败的不记"); +} + +#[tokio::test] +async fn the_position_of_the_hop_does_not_change_what_is_measured() { + // 排在前面的那一家开头要先看一眼有没有报错,再转发;量到的和它排在最后时是同一段 + let first = upstream(Script { + content_after_ms: 300, + ..Default::default() + }) + .await; + let spare = upstream(Script::default()).await; + let (gw, latency) = gateway(vec![provider("先", first), provider("备", spare)], false).await; + ask(gw, true, 3).await; + let t = latency.typical("先").expect("三个样本"); + assert!((300..1_200).contains(&t), "{t} 毫秒"); + assert_eq!(latency.typical("备"), None, "没用上的不记"); +} + +#[tokio::test] +async fn an_answer_that_is_not_streamed_leaves_no_sample() { + // 整包的回答只有「全到了」那一个时刻:分不出排队和说话 + let up = upstream(Script { + content_after_ms: 50, + ..Default::default() + }) + .await; + let (gw, latency) = gateway(vec![provider("整包", up)], false).await; + ask(gw, false, 3).await; + assert_eq!(latency.typical("整包"), None); +} + +#[tokio::test] +async fn an_upstream_given_up_on_is_charged_the_time_it_was_given() { + // 慢的那一家以前很快(留着三个快的样本),现在开了流就不出内容:每次等满 1 秒被放弃。 + // 它记 1 秒,中位数就落到慢的那一头;接下来答的那一家不替它背这 1 秒 + let slow = upstream(Script { + stall: true, + ..Default::default() + }) + .await; + let quick = upstream(Script::default()).await; + let (gw, latency) = gateway(vec![provider("慢", slow), provider("快", quick)], true).await; + for _ in 0..3 { + latency.record("慢", 50); + } + ask(gw, true, 3).await; + assert_eq!(latency.typical("慢"), Some(1_000)); + let t = latency.typical("快").expect("三个样本"); + assert!(t < 800, "{t} 毫秒:放弃的那一家的等待算到了它头上"); +} + +#[tokio::test] +async fn an_upstream_whose_headers_never_come_is_charged_the_time_it_was_given() { + let mute = upstream(Script { + header_delay_ms: 10_000, + ..Default::default() + }) + .await; + let quick = upstream(Script::default()).await; + let (gw, latency) = gateway(vec![provider("无声", mute), provider("快", quick)], true).await; + ask(gw, true, 3).await; + assert_eq!(latency.typical("无声"), Some(1_000)); + assert!(latency.typical("快").is_some_and(|t| t < 800)); +} diff --git a/crates/tw-gateway/tests/upstream_slots.rs b/crates/tw-gateway/tests/upstream_slots.rs index 02d73c8c..fc8ccfa2 100644 --- a/crates/tw-gateway/tests/upstream_slots.rs +++ b/crates/tw-gateway/tests/upstream_slots.rs @@ -190,7 +190,10 @@ impl Log { } async fn serve(cfg: tw_config::Config) -> (SocketAddr, tw_gateway::AppState, Log) { - let state = tw_gateway::AppState::new(cfg).unwrap(); + serve_state(tw_gateway::AppState::new(cfg).unwrap()).await +} + +async fn serve_state(state: tw_gateway::AppState) -> (SocketAddr, tw_gateway::AppState, Log) { let mut rx = state.bus.subscribe(); let log = Log(Arc::default()); let into = log.0.clone(); @@ -616,3 +619,98 @@ async fn a_keys_limit_holds_until_the_streamed_answer_ends() { .unwrap(); assert_eq!(st, 200); } + +/// `load-balance` 组里满着的那一家不参加这一轮:轮到它的话它被当场跳过,这一份却记在它头上, +/// 它答得越多越满、越满越被记空账,拿到的比它的权重少 +#[tokio::test] +async fn a_full_member_of_a_load_balance_group_sits_its_turns_out() { + let (a, b) = (upstream().await, upstream().await); + let mut cfg = config(&a, &b, (Some(1), None), 5); + cfg.groups[0].kind = tw_engine::GroupType::LoadBalance; + cfg.default_route = Some("默认".into()); + let (gw, state, log) = serve(cfg).await; + let group = state.runtime().engine.groups()[0].clone(); + // 头一个轮到甲,占着它的那个位置 + let held = hold(gw, &log, "占着").await; + assert_eq!(a.hits(), 1); + let before = state.balance.peek(&group); + + // 甲满着的时候,每个新对话都排给乙:没有一个先轮到甲、再被跳过 + for i in 0..4 { + let model = format!("新的{i}"); + let st = ask(gw, &model, None, &format!("[{}]", user(&model)), false) + .send() + .await + .unwrap() + .status(); + assert_eq!(st, 200); + let (chain, _) = log.routed(&model).await; + assert_eq!(chain.len(), 1, "轮到了满着的甲:{chain:#?}"); + assert_eq!(chain[0].provider, "乙"); + } + assert_eq!( + state.balance.peek(&group).get("甲"), + before.get("甲"), + "满着的甲被记了账" + ); + a.release(); + held.await.unwrap(); +} + +/// 用量上限看的时钟:跟着真实时间走,还能往后拨 +struct Ahead(std::sync::atomic::AtomicI64); + +impl tw_gateway::key_limits::Clock for Ahead { + fn now_ms(&self) -> i64 { + tw_gateway::key_limits::SystemClock.now_ms() + self.0.load(Ordering::SeqCst) + } + fn period(&self, per: tw_config::LimitPer, at_ms: i64) -> (i64, i64) { + tw_gateway::key_limits::SystemClock.period(per, at_ms) + } + fn show(&self, at_ms: i64) -> String { + tw_gateway::key_limits::SystemClock.show(at_ms) + } +} + +/// 一个请求只有一段可等的时间:等密钥的分钟上限用掉的,等上游空位时就少等那么久。两段各给 +/// 一份的话,这个请求要等两倍那么久才收到 429 +#[tokio::test] +async fn the_key_limit_wait_and_the_slot_wait_share_one_budget() { + use tw_gateway::key_limits::Clock as _; + let (a, b) = (upstream().await, upstream().await); + let mut cfg = config(&a, &b, (Some(1), Some(1)), 2); + cfg.clients[0].limits = serde_yaml_ng::from_str("[{per: minute, requests: 2}]").unwrap(); + let mut state = tw_gateway::AppState::new(cfg).unwrap(); + let clock = Arc::new(Ahead(Default::default())); + state.set_key_limits_clock(clock.clone()); + let (gw, _state, log) = serve_state(state).await; + let _on_a = hold(gw, &log, "占着甲").await; + let _on_b = hold(gw, &log, "占着乙").await; + // 拨到第一个请求之后 59 秒:这一分钟的两个用满了,第一个再过不到一秒滑出去 + clock.0.store(59_000, Ordering::SeqCst); + assert!(clock.now_ms() > 0); + + let t = std::time::Instant::now(); + let r = ask(gw, "挤不进", None, &format!("[{}]", user("你好")), false) + .send() + .await + .unwrap(); + let took = t.elapsed(); + assert_eq!(r.status(), 429); + let id = log.id("挤不进").await; + let code = log + .until("失败", |evs| { + evs.iter().find_map(|e| match e { + Event::RequestFailed { id: i, message, .. } if *i == id => { + Some(message.code.clone()) + } + _ => None, + }) + }) + .await; + assert_eq!(code, "gw.busy_all", "该是等过了分钟上限、再等上游的空位"); + assert!( + took >= Duration::from_millis(1_800) && took < Duration::from_millis(2_500), + "等了 {took:?}:两段该共用 2 秒" + ); +} diff --git a/docs/config.md b/docs/config.md index 93c891c0..f7abf437 100644 --- a/docs/config.md +++ b/docs/config.md @@ -331,9 +331,9 @@ can have several, and a request has to pass every one. `minute` and `hour` are rolling: the last 60 seconds, the last 60 minutes. When one is used up, a request waits for the next free slot if it frees within `failover.slot_wait_secs` (30 seconds by default), and is refused -otherwise. `day`, `week` and `month` follow the calendar in the time zone of -the machine twcore runs on and start again at midnight, on Monday and on the -1st. When one is used up, requests are refused until it starts again. +otherwise. Any later wait for a busy upstream comes out of the same time. +`day`, `week` and `month` follow the calendar in the time zone of the machine +twcore runs on and start again at midnight, on Monday and on the 1st. When one is used up, requests are refused until it starts again. A refused request gets HTTP 429 in the client's own error format, naming the key, the limit, the amount used and when it resets, and it shows in the @@ -998,7 +998,9 @@ failover: ``` When upstreams are at their `max_concurrent`, a request waits for a free slot -for at most `slot_wait_secs` in all. An ongoing conversation waits for the +for at most `slot_wait_secs` in all. The same time also covers waiting for a +key's `minute` or `hour` limit, so a request never waits longer than +`slot_wait_secs` for the two together. An ongoing conversation waits for the upstream it stays on and, if no slot frees in time, moves on to the next one, where its cache starts over. A new conversation skips a full upstream at once. When every candidate is full, the request waits for whichever frees first; if @@ -1018,7 +1020,7 @@ busy. | `rate_limit_max_pause_secs` | integer | `3600` | A rate-limited upstream is set aside for the time its `Retry-After` gives, at most this many seconds. Without `Retry-After` it counts as a failure without a stated reason. | | `stream_start_wait_secs` | integer | `15` | Seconds to hold a streamed answer until its first content arrives. An error before then moves the request to the next upstream; after this long, what has arrived is passed on. From 1 to 120. | | `next_on_slow_start` | bool | `false` | When a streamed answer still has no content `stream_start_wait_secs` after the request was sent, give up on that upstream and send the request to the next one. The last upstream always waits. The upstream given up on is not set aside. Needs `stream_start_wait_secs` of at least 5. | -| `slot_wait_secs` | integer | `30` | Seconds a request waits in total for a free slot on upstreams that are at their `max_concurrent`. After that it goes to the next upstream, or, when every candidate is full, is answered with 429. `0`: never wait. From 0 to 300. | +| `slot_wait_secs` | integer | `30` | Seconds a request waits in all, counted once the key's own `max_concurrent` lets it in: for a key's `minute` or `hour` limit to free up, and for a free slot on upstreams at their `max_concurrent`. A key limit that does not free up in time refuses the request; without an upstream slot in time it goes to the next upstream, or, when every candidate is full, is answered with 429. `0`: never wait. From 0 to 300. | ### `aliases` @@ -1066,7 +1068,7 @@ group with `to`. | Field | Type | Default | Description | |---|---|---|---| | `name` | string | **required** | Name of the group; unique, and not the name of an upstream. | -| `type` | `fallback` \| `select` \| `load-balance` \| `url-test` \| `cheapest` | `fallback` | `fallback`: the first healthy member, in order. `select`: the member named in `selected`. `load-balance`: new conversations take turns, in proportion to the members' weights. `url-test`: the fastest by measured time to first byte. `cheapest`: the lowest input price. | +| `type` | `fallback` \| `select` \| `load-balance` \| `url-test` \| `cheapest` | `fallback` | `fallback`: the first healthy member, in order. `select`: the member named in `selected`. `load-balance`: new conversations take turns, in proportion to the members' weights. `url-test`: the fastest by measured time from sending a request to the first content of the answer. `cheapest`: the lowest input price. | | `providers` | list of strings or [`groups[].providers[]`](#cfg-groups-providers) | **required** | Member upstreams, by name; not groups. Each upstream appears once in a group. In a `load-balance` group, a member can be written as `{name, weight}`. | | `selected` | string | — | For `select`: the chosen member. | | `balance_by` | `weights` \| `latency` \| `health` \| `latency-health` | `weights` | For `load-balance`: what the members' weights are multiplied by. `weights`: nothing; the weights alone. `latency`: faster upstreams get more. `health`: upstreams that fail less get more. `latency-health`: both. Other group types take only `weights`. | @@ -1081,9 +1083,9 @@ requests are shared out: with `{ name: anthropic, weight: 7 }` and `relay`, the official API serves seven requests in ten. Conversations in progress stay on the upstream that answers them (see below) and count toward its share, so the balance is kept by where new conversations start. An upstream -that is cooling down after failures, or cannot serve a request, sits that -request out, and the others share it by their weights. Other group types take -no weights. +that is cooling down after failures, is at its `max_concurrent`, or cannot +serve a request, sits that request out, and the others share it by their +weights. Other group types take no weights. @@ -1122,9 +1124,10 @@ group shares out requests by the result in the same way as above. - `weights` (the default): the weights alone. - `latency`: faster upstreams get a larger share. Speed is the typical time - to first byte, the same measurement `url-test` uses. An upstream twice as - fast as the middle of the group has its weight multiplied by four, by at - most ten and at least a tenth. + from sending a request to the first content of the answer, the same + measurement `url-test` uses. An upstream twice as fast as the middle of the + group has its weight multiplied by four, by at most ten and at least a + tenth. - `health`: upstreams that fail less get a larger share. It looks at the last 50 requests within the past 30 minutes. Server errors, rate limits, used-up quota or balance, rejected credentials, timeouts and connection @@ -1137,6 +1140,11 @@ group shares out requests by the result in the same way as above. [`failover`](#cfg-failover) as before. - `latency-health`: both factors, multiplied. +Speed is measured on streamed answers only, from the moment the request is +sent to that upstream, so waiting and upstreams that failed before it do not +count. An upstream given up on because its stream was slow to start +(`failover.next_on_slow_start`) counts as taking the whole wait. + An upstream without enough measurements yet counts as average. As with weights alone, conversations in progress stay where they are, and new conversations make up the difference. diff --git a/docs/config.zh-CN.md b/docs/config.zh-CN.md index 9b6ddd79..b5e17c79 100644 --- a/docs/config.zh-CN.md +++ b/docs/config.zh-CN.md @@ -237,9 +237,9 @@ clients: 每一条只数其中一种;一把密钥可以有好几条,请求要通过每一条。 `minute`、`hour` 是滚动的:最近 60 秒、最近 60 分钟。用满之后,下一个空位在 -`failover.slot_wait_secs`(默认 30 秒)之内空出来,请求就等它;等不到就拒绝。`day`、 -`week`、`month` 按 twcore 所在机器的本地时区算,在零点、周一零点、每月一号零点重新开始; -用满之后,到重新开始之前的请求一律拒绝。 +`failover.slot_wait_secs`(默认 30 秒)之内空出来,请求就等它;等不到就拒绝。之后若还要等 +上游空出位置,用的是同一段时间里剩下的部分。`day`、`week`、`month` 按 twcore 所在机器的 +本地时区算,在零点、周一零点、每月一号零点重新开始;用满之后,到重新开始之前的请求一律拒绝。 被拒的请求收到 HTTP 429,错误格式和客户端自己的一致,写明是哪把密钥、哪一条上限、用了多少、 什么时候重置;流量列表里也有这一条。费用按每个请求记下的费用算,所以没有价格的模型、 @@ -783,7 +783,8 @@ failover: ``` 上游的并发数满了(`max_concurrent`)时,一个请求等空位合计最多 `slot_wait_secs` -秒。进行中的对话等它留在的那一家,到时还没有空位就换下一家,缓存在那边从头建; +秒。等密钥的 `minute`、`hour` 上限也算在这段时间里,两样加起来不超过 `slot_wait_secs`。 +进行中的对话等它留在的那一家,到时还没有空位就换下一家,缓存在那边从头建; 新的对话遇到满着的上游直接跳过。候选全满时,请求等先空出来的那一家;都没有空出来, 客户端收到 429 和 `Retry-After`,说明上游都忙。 @@ -800,7 +801,7 @@ failover: | `rate_limit_max_pause_secs` | 整数 | `3600` | 被限流的上游按它给的 `Retry-After` 停用,最多这么多秒。没有 `Retry-After` 的按没有说明原因的失败计。 | | `stream_start_wait_secs` | 整数 | `15` | 流式回答在第一段内容到达前最多暂存的秒数。在此之前上游报错,请求换到下一家;超过这个时间,已收到的部分照常交给客户端。取值 1 到 120。 | | `next_on_slow_start` | 布尔 | `false` | 流式回答在请求发出 `stream_start_wait_secs` 秒后仍没有内容时,放弃这家上游,把请求交给下一家。最后一家总是等下去。被放弃的上游不会停用。开启时 `stream_start_wait_secs` 至少为 5。 | -| `slot_wait_secs` | 整数 | `30` | 上游的并发数满了(`max_concurrent`)时,一个请求等空位合计最多等的秒数。等不到就换下一家;候选全满时回 429。`0`:不等。取值 0 到 300。 | +| `slot_wait_secs` | 整数 | `30` | 一个请求合计最多等的秒数,从过了密钥自己的 `max_concurrent` 时算起:等密钥的 `minute`、`hour` 上限空出名额,和等并发数满了(`max_concurrent`)的上游空出位置,都算在里面。密钥的上限到时空不出来就拒绝;等不到上游的空位就换下一家,候选全满时回 429。`0`:不等。取值 0 到 300。 | ### `aliases` @@ -834,7 +835,7 @@ aliases: | 字段 | 类型 | 默认值 | 说明 | |---|---|---|---| | `name` | 字符串 | **必填** | 策略组的名字,不能重复,也不能和上游同名。 | -| `type` | `fallback` \| `select` \| `load-balance` \| `url-test` \| `cheapest` | `fallback` | `fallback`:按顺序取第一个健康的。`select`:取 `selected` 指定的那个。`load-balance`:新对话按成员的权重轮流。`url-test`:按实测首字节时间取最快的。`cheapest`:取输入单价最低的。 | +| `type` | `fallback` \| `select` \| `load-balance` \| `url-test` \| `cheapest` | `fallback` | `fallback`:按顺序取第一个健康的。`select`:取 `selected` 指定的那个。`load-balance`:新对话按成员的权重轮流。`url-test`:按实测从发出请求到回答第一段内容的时间取最快的。`cheapest`:取输入单价最低的。 | | `providers` | 列表,每项是字符串或对象,对象见 [`groups[].providers[]`](#cfg-groups-providers) | **必填** | 成员上游的名字,不能是策略组。同一个上游在一个策略组中只出现一次。`load-balance` 组的成员可以写成 `{name, weight}`。 | | `selected` | 字符串 | — | `select` 类型选中的成员。 | | `balance_by` | `weights` \| `latency` \| `health` \| `latency-health` | `weights` | `load-balance` 类型用:成员的权重再乘上什么。`weights`:不乘,只按权重。`latency`:越快的上游分得越多。`health`:越少失败的上游分得越多。`latency-health`:两者都看。其他类型只能是 `weights`。 | @@ -842,7 +843,7 @@ aliases: 默认类型为 `fallback`:单个使用者的机器上没有需要分散的负载。 -`load-balance` 组的成员可以带权重,取值 1 到 100;只写名字的成员权重为 1。权重决定组内请求怎么分:写成 `{ name: anthropic, weight: 7 }` 和 `relay` 时,每十个请求有七个由官方 API 服务。进行中的对话留在回答它的那一家(见下文),也算进那一家的份额,因此份额靠新对话从哪一家开始来补齐。因失败处于冷却、或服务不了某个请求的上游不参与这一次分配,其余成员按各自的权重分。其他类型的策略组不用权重。 +`load-balance` 组的成员可以带权重,取值 1 到 100;只写名字的成员权重为 1。权重决定组内请求怎么分:写成 `{ name: anthropic, weight: 7 }` 和 `relay` 时,每十个请求有七个由官方 API 服务。进行中的对话留在回答它的那一家(见下文),也算进那一家的份额,因此份额靠新对话从哪一家开始来补齐。因失败处于冷却、并发数已满(`max_concurrent`)、或服务不了某个请求的上游不参与这一次分配,其余成员按各自的权重分。其他类型的策略组不用权重。 @@ -867,10 +868,12 @@ groups: `balance_by` 让 `load-balance` 组再看各上游最近的表现:每个成员的权重乘上一个系数,组内请求按乘出来的结果照上文的方式分。 - `weights`(默认):只按权重。 -- `latency`:越快的上游分得越多。快慢看典型的首字节时间,与 `url-test` 使用同一份测量。比组内居中者快一倍的上游,权重乘以四;最多乘以十,最少乘以十分之一。 +- `latency`:越快的上游分得越多。快慢看典型的从发出请求到回答第一段内容的时间,与 `url-test` 使用同一份测量。比组内居中者快一倍的上游,权重乘以四;最多乘以十,最少乘以十分之一。 - `health`:越少失败的上游分得越多。依据是最近 30 分钟内的最近 50 次请求:服务器错误、限流、额度或余额用尽、凭据被拒、超时和连接失败算作失败;请求本身导致的错误不算,客户端取消、因开头太慢而换走、因并发数满了(`max_concurrent`)而跳过也不算。经常失败的上游至少保留权重的二十分之一,仍会偶尔分到新对话,以便发现它已经恢复;完全失败的上游照旧由 [`failover`](#cfg-failover) 暂停。 - `latency-health`:两个系数相乘。 +快慢只在流式回答上测,从请求发给这家上游的那一刻算起:之前的等待、之前失败的上游都不算在内。因开头太慢而被放弃的上游(`failover.next_on_slow_start`),按等满的那段时间计。 + 测量还不够的上游按中等对待。与只按权重时一样,进行中的对话留在原来的上游,差额由新对话补齐。 ```yaml From 174f419eafbda9521fe2032ed60ad1f52687cd26 Mon Sep 17 00:00:00 2001 From: fylorn <249551762+fylorn@users.noreply.github.com> Date: Mon, 5 Oct 2026 17:58:37 +0800 Subject: [PATCH 09/22] Count each turn on a Responses WebSocket as its own request A Responses WebSocket connection used to be one request row for its whole life, with no usage and no cost: ws.rs never read `usage`, the key's usage limits were checked once when the connection opened, and neither the key's `max_concurrent` nor the upstream's slot applied at all. A client configured for WebSocket could therefore spend without limit and without a trace of what it cost. Now each `response.create` is a request, from the frame to the answer's `response.completed` / `response.failed` / `response.incomplete` (or an `error`): RequestStarted, RequestHeaders, RequestRouted (one hop, the connection's upstream, with queued_ms), first token and an ending that carries the turn's usage, so the recorder prices it like any HTTP request and key usage, checkup and traffic all see it. The session comes from the frame (Codex sends `prompt_cache_key`), so a conversation's turns group like HTTP ones. A connection closed mid-turn cancels the turn; an upstream that breaks fails it. Each turn goes through the same admission as HTTP, in the same order: calendar limits refuse, the key's `max_concurrent` waits, rolling limits wait within `slot_wait_secs` or refuse, then the upstream's slot waits within `slot_wait_secs` (the connection is bound to one upstream, so there is no next candidate) and fails as busy. A refused turn leaves a row and is answered with `response.failed` (code `insufficient_quota` when a calendar period is used up); the connection stays open. The key's permit and the upstream slot are held from sending a turn until its answer ends, so an idle connection holds nothing. While a turn waits, the upstream side keeps relaying and later client frames queue behind it, so a turn waiting for a slot held by the previous turn on the same connection cannot deadlock. The connection itself leaves no row, except when there is no turn to carry the outcome: an upgrade a rule denies and an upstream that cannot be connected keep their single row. Realtime and other WebSocket paths keep one row per connection and the admission at open. Co-Authored-By: Claude Opus 5.5 --- Cargo.lock | 1 + crates/tw-api/src/lib.rs | 28 +- crates/tw-control/Cargo.toml | 2 + crates/tw-control/tests/ws_turns.rs | 172 ++++++ crates/tw-gateway/src/ending.rs | 60 +- crates/tw-gateway/src/server/upgrade.rs | 143 ++--- crates/tw-gateway/src/slots.rs | 7 +- crates/tw-gateway/src/ws.rs | 779 ++++++++++++++++++------ crates/tw-gateway/src/ws/turn.rs | 545 +++++++++++++++++ crates/tw-gateway/tests/alias_ws.rs | 53 +- crates/tw-gateway/tests/endings.rs | 56 +- crates/tw-gateway/tests/key_limits.rs | 9 +- crates/tw-gateway/tests/ws.rs | 647 +++++++++++++++++++- docs/config.md | 10 +- docs/config.zh-CN.md | 5 +- 15 files changed, 2169 insertions(+), 348 deletions(-) create mode 100644 crates/tw-control/tests/ws_turns.rs create mode 100644 crates/tw-gateway/src/ws/turn.rs diff --git a/Cargo.lock b/Cargo.lock index 61f48374..bcda9c87 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -3229,6 +3229,7 @@ dependencies = [ "thiserror", "tokio", "tokio-stream", + "tokio-tungstenite 0.28.0", "tower", "tracing", "tw-api", diff --git a/crates/tw-api/src/lib.rs b/crates/tw-api/src/lib.rs index fee07b85..bcbacaea 100644 --- a/crates/tw-api/src/lib.rs +++ b/crates/tw-api/src/lib.rs @@ -867,7 +867,8 @@ pub enum Event { /// 的指纹)、离这段对话的上一个请求不超过半小时,就还是那一次;隔久了算 /// 新的一次。**开始时就给出来**,界面才能把一个还在跑的请求放进它的会话、 /// 把那次会话标成进行中。认不出会话的没有:正文里没有任何能认人的东西, - /// 或者是 WebSocket 升级(升级请求没有正文) + /// 或者是整条连接一行的 WebSocket(升级请求没有正文)。Responses 连接上的每一轮 + /// 按那一帧认,和 HTTP 的请求一样 #[serde(default, skip_serializing_if = "Option::is_none")] session: Option, /// 请求从哪台机器来:**这条连接对面的地址**,不可伪造。本机(回环) @@ -938,9 +939,11 @@ pub enum Event { /// `response.created` 这类开场帧不算:那是上游收到请求就发的。 /// /// 只有成功的流式响应有。非流式的整段一起到,没有「第一个」;不带 `alt=sse` 的 - /// Gemini 流和 WebSocket 那条路不在这里解析,也没有。 + /// Gemini 流和整条连接一行的 WebSocket 不在这里解析,也没有。Responses 连接上的每一轮 + /// 有。 RequestFirstToken { id: u64, ttft_ms: u64 }, - /// 结束了:上游回的是成功的状态码(2xx;WebSocket 那条路是升级成功的 101),回答 + /// 结束了:上游回的是成功的状态码(2xx;整条连接一行的 WebSocket 是升级成功的 101, + /// Responses 连接上的一轮是 200),回答 /// 交完了。**上游回了别的、原样交给了客户端的不是这一条**,是 `RequestFailed` —— /// 客户端拿到的是上游的错误,不是回答 RequestFinished { @@ -952,7 +955,8 @@ pub enum Event { /// 在结局里才到的。模型名只在开始事件里的话,一个开始时没人在听、 /// 结束时有人在听的请求,它的用量就不知道该记在哪个模型上。 /// - /// WebSocket 那条路是空串:升级请求里没有模型名(和开始事件一样)。 + /// 整条连接一行的 WebSocket 和开始事件一样:Realtime 是查询串里的那个,别的连接 + /// 升级时还不知道,是空串。 model: String, status: u16, bytes: u64, @@ -972,7 +976,7 @@ pub enum Event { /// 上游在回答里写的模型名:Anthropic 和 Chat 的 `model`、Responses 的 /// `response.model`、Gemini 的 `modelVersion`。**原样,不归一。** /// - /// 回答里没写的没有:Bedrock 的 Converse 不写,WebSocket 那条路不看。和 + /// 回答里没写的没有:Bedrock 的 Converse 不写,整条连接一行的 WebSocket 不看。和 /// `model` 不是一回事 —— 那是客户端要的,这是上游说它用的 #[serde(default, skip_serializing_if = "Option::is_none")] answered_model: Option, @@ -1045,8 +1049,9 @@ pub enum Event { /// `RequestFinished` 上的话,失败的那条路就没有尝试链 —— 而那恰恰 /// 是最需要看它的时候。 /// - /// WebSocket 那条路也发:和上游的握手有了结果就发。那条路不做故障转移, - /// 尝试链只有一跳。 + /// WebSocket 那条路也发,不做故障转移,尝试链只有一跳:整条连接一行的,和上游的 + /// 握手有了结果就发;Responses 连接上的一轮,上游这一轮的第一帧到了就发(没等到就在 + /// 结局之前),那一跳是这条连接连着的那一家。 /// /// **规则做了决定、请求却一家上游都没到的也发**:规则拒绝了它(第一阶段), /// 或者规则选中的上游都服务不了这个模型 —— 那时尝试链是空的,紧跟着一条 @@ -1496,8 +1501,9 @@ pub struct AttemptView { /// **费用按它算**:请求改写成另一个模型发出去,上游按那个模型收钱。 #[serde(default, skip_serializing_if = "Option::is_none")] pub model: Option, - /// WebSocket 的那一跳是一次握手:上游同意升级(101)是 `served`,回了别的 - /// 状态码是 `status`,连不上是 `error`。 + /// WebSocket 连接的那一跳是一次握手:上游同意升级(101)是 `served`,回了别的 + /// 状态码是 `status`,连不上是 `error`。Responses 的连接上每一轮是一个请求,那一跳是这条 + /// 已经接下的连接:发出去了是 `served`,状态码记 200。 pub outcome: AttemptOutcome, /// 上游返回的状态码。`error` 时没有 #[serde(default, skip_serializing_if = "Option::is_none")] @@ -3824,7 +3830,7 @@ pub struct Summary { /// 失败的、没有用量的、不计费的都不在这里 —— 配价格对它们没用。 pub unpriced_requests: i64, /// 有多少条请求**没有拿到用量**,所以同样算不出钱:上游没报,或者连接 - /// 在它报之前就结束了(客户端取消、WebSocket 会话)。 + /// 在它报之前就结束了(客户端取消、整条连接一行的 WebSocket 会话)。 /// /// 和 `unpriced_requests` 一样让金额合计偏低,但配价格解决不了它 —— /// 界面上是两句不同的话。上游确实接下了的才算:成功的响应和客户端 @@ -4100,7 +4106,7 @@ pub struct HistoryRow { /// 任务;看着一次很贵的任务,也回不到具体是哪一条。库里这一列一直 /// 都在(`requests.session`,还建了索引),只是没有交出来。 /// - /// 认不出会话的请求(拼不出指纹的,比如 WebSocket、本地应答)是 `None`。 + /// 认不出会话的请求(拼不出指纹的,比如整条连接一行的 WebSocket、本地应答)是 `None`。 #[serde(default, skip_serializing_if = "Option::is_none")] pub session: Option, /// 按请求头推测是哪个应用发的(`claude-code`、`codex`…)。**可以伪造**, diff --git a/crates/tw-control/Cargo.toml b/crates/tw-control/Cargo.toml index a94ced22..151a148c 100644 --- a/crates/tw-control/Cargo.toml +++ b/crates/tw-control/Cargo.toml @@ -65,3 +65,5 @@ tw-config = { workspace = true } # 测试挑端口:从远程控制端口那一段里随机挑(见 tests/common) rand = { workspace = true } tokio = { workspace = true, features = ["rt", "macros"] } +# Responses 的 WebSocket 上每一轮怎么记账,要走真的连接:客户端那一侧需要 connect +tokio-tungstenite = { workspace = true, features = ["connect"] } diff --git a/crates/tw-control/tests/ws_turns.rs b/crates/tw-control/tests/ws_turns.rs new file mode 100644 index 00000000..77483c73 --- /dev/null +++ b/crates/tw-control/tests/ws_turns.rs @@ -0,0 +1,172 @@ +//! Responses 的 WebSocket 上**每一轮是一个请求**:存储层给每一轮记一行,带着这一轮回答的 +//! 用量,**和 HTTP 的请求同一套查价**;密钥的用量上限按这一行结算。 +//! +//! 网关、存储层、密钥用量的结算和 `twcore` 一样接(见 `tw_control::key_limits`)。上游是一个 +//! 像 Responses 的 WebSocket 那样回答的假服务。 + +use std::net::SocketAddr; +use std::time::Duration; + +use axum::extract::ws::{Message, WebSocketUpgrade}; +use futures::{SinkExt, StreamExt}; + +/// 每一轮回答的用量:输入 1200(其中 1000 走了缓存)、输出 30。Responses 的输入数包含缓存读 +const USAGE: &str = r#"{"input_tokens":1200,"input_tokens_details":{"cached_tokens":1000},"output_tokens":30,"output_tokens_details":{"reasoning_tokens":0},"total_tokens":1230}"#; + +/// 每个 `response.create` 回 created、一段文字、completed(带用量) +async fn upstream() -> SocketAddr { + let app = axum::Router::new().route( + "/v1/responses", + axum::routing::any(|ws: WebSocketUpgrade| async move { + ws.on_upgrade(|mut sock| async move { + while let Some(Ok(m)) = sock.recv().await { + let Message::Text(_) = m else { continue }; + let usage: serde_json::Value = serde_json::from_str(USAGE).unwrap(); + let frames = [ + serde_json::json!({"type":"response.created","response":{"id":"resp_1","status":"in_progress","model":"gpt-5","output":[]}}), + serde_json::json!({"type":"response.output_item.added","output_index":0,"item":{"type":"message","id":"msg","role":"assistant","content":[]}}), + serde_json::json!({"type":"response.output_text.delta","output_index":0,"content_index":0,"item_id":"msg","delta":"hello"}), + serde_json::json!({"type":"response.completed","response":{"id":"resp_1","status":"completed","model":"gpt-5","output":[],"usage":usage}}), + ]; + for f in frames { + if sock.send(Message::Text(f.to_string().into())).await.is_err() { + return; + } + } + } + }) + }), + ); + let l = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap(); + let addr = l.local_addr().unwrap(); + tokio::spawn(async move { axum::serve(l, app).await.unwrap() }); + addr +} + +/// 两轮,每一轮一行:用量是那一轮回答的,费用按 gpt-5 的价($1.25/M 输入、$0.125/M 缓存读、 +/// $10/M 输出)算:200 × 1.25 + 1000 × 0.125 + 30 × 10 = 675 微美元,**是实数不是估算**。密钥 +/// 这一天的 token 和费用就是两行加起来 +#[tokio::test] +async fn each_websocket_turn_is_recorded_priced_and_counted_against_the_key() { + let up = upstream().await; + let d = tempfile::tempdir().unwrap(); + let yaml = format!( + "version: 1 +listen: + control: + key: c0ffee00c0ffee00c0ffee00c0ffee00c0ffee00c0ffee00c0ffee00c0ffee00 +clients: + - name: codex + key: tw-k + limits: + - {{ per: day, tokens: 100000 }} + - {{ per: day, cost: 5 }} +providers: + - name: openai + base_url: http://{up} + key: sk-x + protocol: openai-responses +" + ); + let cfg = tw_config::try_parse(&yaml).unwrap(); + let limits = cfg.clients[0].limits.clone(); + let gw = tw_gateway::AppState::new(cfg).unwrap(); + let (_bodies, rx) = tokio::sync::mpsc::channel(1); + let store = tw_store::task::spawn( + tw_store::Recorder::new( + tw_store::Db::open(&d.path().join("data.db")).unwrap(), + tw_store::Blobs::new(d.path().join("blobs")), + gw.pricing.clone(), + ) + .settling_to(tw_control::key_limits::settle_hook(&gw)), + gw.bus.subscribe(), + rx, + ); + let addr = tw_gateway::serve(gw.clone(), ([127, 0, 0, 1], 0).into()) + .await + .unwrap(); + tokio::time::sleep(Duration::from_millis(50)).await; + + use tokio_tungstenite::tungstenite::client::IntoClientRequest; + let mut req = format!("ws://{addr}/v1/responses") + .into_client_request() + .unwrap(); + req.headers_mut() + .insert("authorization", "Bearer tw-k".parse().unwrap()); + let (mut c, _) = tokio_tungstenite::connect_async(req).await.unwrap(); + for text in ["one", "two"] { + let frame = serde_json::json!({ + "type": "response.create", + "model": "gpt-5", + "prompt_cache_key": "conversation-1", + "input": [{"role": "user", "content": [{"type": "input_text", "text": text}]}] + }); + c.send(tokio_tungstenite::tungstenite::Message::Text( + frame.to_string().into(), + )) + .await + .unwrap(); + loop { + let m = tokio::time::timeout(Duration::from_secs(5), c.next()) + .await + .expect("the answer did not end") + .unwrap() + .unwrap(); + if m.into_text().unwrap().contains("response.completed") { + break; + } + } + } + + // 两行都落库了(落库之前已经结算过) + let deadline = std::time::Instant::now() + Duration::from_secs(10); + let rows = loop { + let rows = store.lock().await.db().recent(None, 10).unwrap(); + if rows.len() >= 2 { + break rows; + } + assert!(std::time::Instant::now() < deadline, "rows: {rows:?}"); + tokio::time::sleep(Duration::from_millis(20)).await; + }; + // 连接还开着:连接本身不留行,两行都是那两轮 + assert_eq!(rows.len(), 2, "{rows:?}"); + assert_ne!(rows[0].id, rows[1].id); + for r in &rows { + assert_eq!( + (r.client.as_str(), r.provider.as_str()), + ("codex", "openai") + ); + assert_eq!( + (r.model.as_str(), r.sent_model.as_str()), + ("gpt-5", "gpt-5") + ); + assert_eq!(r.path, "/v1/responses"); + assert_eq!(r.status, Some(200)); + assert_eq!( + (r.input_tokens, r.cache_read_tokens, r.output_tokens), + (Some(200), Some(1000), Some(30)), + "{r:?}" + ); + assert_eq!(r.cost_micros, Some(675), "{r:?}"); + assert!(!r.cost_estimated); + assert!(r.ttft_ms.is_some(), "第一个 token 的时刻没记下:{r:?}"); + assert!(r.error.is_none() && !r.cancelled, "{r:?}"); + let routing = r.routing.as_deref().unwrap(); + assert!(routing.contains("\"provider\":\"openai\""), "{routing}"); + } + // 同一段对话的两轮归到同一次会话里 + assert!(rows[0].session.is_some()); + assert_eq!(rows[0].session, rows[1].session); + + // 密钥这一天:两行的 token(没走缓存的输入 + 输出)和费用 + let view = gw.key_limits.view("codex", &limits); + let used: Vec<(tw_api::LimitMeasure, u64)> = view.iter().map(|v| (v.measure, v.used)).collect(); + assert_eq!( + used, + [ + (tw_api::LimitMeasure::Tokens, 2 * (200 + 30)), + (tw_api::LimitMeasure::Cost, 2 * 675), + ] + ); + drop(c); +} diff --git a/crates/tw-gateway/src/ending.rs b/crates/tw-gateway/src/ending.rs index 041fffe7..fdd2112e 100644 --- a/crates/tw-gateway/src/ending.rs +++ b/crates/tw-gateway/src/ending.rs @@ -355,14 +355,27 @@ impl Ending { /// 只数字节,不嗅用量、不留档。 /// - /// WebSocket 那条路用它。一条连接上跑着好几轮回答,每轮各报一次用量, - /// 而嗅探器是「每个字段取最大值」—— 喂给它,得到的是其中某一轮的数, - /// 看起来却像整条连接的;模型名也不知道(升级请求里没有),算不了钱。 - /// **与其报一个错的数,不如说没有。** + /// 整条连接一行的 WebSocket 用它(Realtime 和别的路径,见 [`crate::ws`])。一条连接上 + /// 跑着好几轮回答,每轮各报一次用量,而嗅探器是「每个字段取最大值」—— 喂给它,得到的 + /// 是其中某一轮的数,看起来却像整条连接的。**与其报一个错的数,不如说没有。** + /// Responses 的连接每一轮各是一个请求,用的是 [`Ending::frame`]。 pub fn count(&mut self, bytes: usize) { self.bytes += bytes as u64; } + /// WebSocket 上上游的一帧文本:Responses 连接上的一轮(见 `crate::ws::turn`)。一条消息 + /// 就是一个事件,**按 SSE 的一帧喂**给认第一个 token、嗅用量、看错误的那几样 —— 它们 + /// 读的是 SSE;字节只数消息本身。已经是 SSE 形状的(桥接过来的)原样喂 + pub fn frame(&mut self, text: &str) { + let sse = if text.lines().any(|l| l.starts_with("data: ")) { + std::borrow::Cow::Borrowed(text) + } else { + std::borrow::Cow::Owned(format!("data: {text}\n\n")) + }; + self.feed(sse.as_bytes()); + self.bytes = self.bytes - sse.len() as u64 + text.len() as u64; + } + /// 走完了。上游在流里报过错的、回的不是 2xx 的,报的是失败(见 [`Ending::streaming`]、 /// [`Ending::refused`])。 pub fn finished(mut self, status: u16) { @@ -1269,6 +1282,45 @@ mod tests { ); } + /// Responses 连接上的一轮(见 `crate::ws::turn`):一条消息按 SSE 的一帧喂,第一个 token、 + /// 用量照认,**字节只数消息本身** + #[test] + fn a_websocket_frame_is_read_as_one_event_and_counted_as_itself() { + let bus = tw_observe::EventBus::new(); + let mut rx = bus.subscribe(); + let mut e = responding(&bus); + e.streaming(ir::Dialect::Responses, "up"); + let frames = [ + r#"{"type":"response.created","response":{"id":"r","status":"in_progress","model":"gpt-5","output":[]}}"#, + r#"{"type":"response.output_item.added","output_index":0,"item":{"type":"message","id":"m","role":"assistant","content":[]}}"#, + r#"{"type":"response.output_text.delta","output_index":0,"content_index":0,"item_id":"m","delta":"hi"}"#, + r#"{"type":"response.completed","response":{"id":"r","status":"completed","model":"gpt-5","output":[],"usage":{"input_tokens":50,"input_tokens_details":{"cached_tokens":20},"output_tokens":5,"total_tokens":55}}}"#, + ]; + for f in frames { + e.frame(f); + } + e.finished(200); + + let got = drain(&mut rx); + assert!( + matches!(got.first(), Some(Event::RequestFirstToken { id: 7, .. })), + "{got:?}" + ); + match got.last() { + Some(Event::RequestFinished { + bytes, + usage: Some(u), + answered_model, + .. + }) => { + assert_eq!(*bytes, frames.iter().map(|f| f.len() as u64).sum::()); + assert_eq!((u.input, u.cache_read, u.output), (30, 20, 5)); + assert_eq!(answered_model.as_deref(), Some("gpt-5")); + } + other => panic!("该是一次带着用量的结束,实际 {other:?}"), + } + } + /// 流在网关自己的代码里崩掉了。**Drop 同样会跑**(unwind 会丢掉流里的 /// 局部变量),而这时报「客户端取消」是在冤枉客户端。 /// diff --git a/crates/tw-gateway/src/server/upgrade.rs b/crates/tw-gateway/src/server/upgrade.rs index 796fa773..4481508d 100644 --- a/crates/tw-gateway/src/server/upgrade.rs +++ b/crates/tw-gateway/src/server/upgrade.rs @@ -16,6 +16,10 @@ use tw_types::msg; /// 有效。之后把连接交给 [`crate::ws::proxy`],那里会在 /// 每一帧上重新点一遍管线的保护。路由事件也由那边发:选中的那一家接没 /// 接下,要和它握完手才知道。 +/// +/// **Responses 的连接上每个 `response.create` 是一个请求**(见 `crate::ws::turn`):连接 +/// 本身不留行,密钥的用量上限、并发上限按轮算,升级时不看。Realtime 和别的路径的连接照旧 +/// 整条连接一行,升级时过一遍用量上限。 #[allow(clippy::too_many_arguments)] pub(super) async fn ws_upgrade( state: AppState, @@ -45,46 +49,32 @@ pub(super) async fn ws_upgrade( ..Default::default() }; let route = rt.engine.route_of(&client_name).to_string(); + // 这条连接上开始的请求都一样的那几项:整条连接一行的那一行,Responses 的连接上的每一轮 + let opener = crate::ws::turn::Opener { + bus: state.bus.clone(), + client: client_name.clone(), + client_hint: crate::hint::client_hint(&headers), + peer: from.peer.clone(), + key_masked: from.key.clone(), + path: uri.path().to_string(), + }; // 开始事件。**被规则拒绝的升级也发**(和 HTTP 那条路一样,见 // `super::routed_nowhere`):流量里要有这一行,规则的命中也要数得到它 + // + // 这条连接怎么断的,就是这个请求的结局。**跟着连接走**:升级没完成就被丢掉的 —— + // 客户端没等到 101 就走了 —— 由 Drop 报成取消 let open = |choice: &Choice, provider: &str, billing: tw_api::Billing| { - let id = state.bus.next_id(); - let at_ms = now_ms(); - state.bus.emit(tw_api::Event::RequestStarted { - id, - client: client_name.clone(), - client_hint: crate::hint::client_hint(&headers), - // 升级请求没有正文,认不出是哪段对话 - session: None, - peer: from.peer.clone(), - key_masked: from.key.clone(), - route: choice.route.clone(), - rule: choice.rule.clone(), - group: choice.group.clone(), - rewritten_by: choice.rewritten_by.clone(), - provider: provider.to_string(), - billing, + opener.open(crate::ws::turn::Opening { + choice, + to: (provider, billing), // 客户端写的模型名:Realtime 的连接写在查询串里,别的连接升级时还不知道 model: facts.model.clone(), - method: "WS".to_string(), - path: uri.path().to_string(), - // 升级请求没有正文,没有可估的 + // 升级请求没有正文,认不出是哪段对话,也没有可估的 + session: None, input_estimate: None, - session_log_bytes: None, - at_ms, - }); - // 这条连接怎么断的,就是这个请求的结局。**跟着连接走**:升级没完成 - // 就被丢掉的 —— 客户端没等到 101 就走了 —— 由 Drop 报成取消。WS 帧 - // 不留档,所以没有 body 的去处 - let ending = crate::ending::Ending::new( - state.bus.clone(), - id, - String::new(), started, - at_ms as i64, - None, - ); - (id, ending) + at_ms: now_ms(), + }) }; let decision = match rt.engine.route(&facts).map_err(|e| { GatewayError::config(msg!("gw.route.failed", detail = e => "Routing failed: {detail}")) @@ -228,44 +218,61 @@ pub(super) async fn ws_upgrade( .headers_for(provider, http) .await .map_err(|e| GatewayError::config(crate::state::credential_failed(e, &name)))?; - // 这把密钥的用量上限:**一条连接算一个请求**,连上之前看一遍,和 HTTP 那条路的准入 - // 同一套(见 `crate::key_limits`)。连接上的每个 `response.create` 不再分开数:存储层给 - // 整条连接记一行、不带用量,分开数的话,重启之后从记录里加回来的数就对不上了 - let limits = rt - .config - .clients - .iter() - .find(|c| c.name == client_name) - .map(|c| c.limits.as_slice()) - .unwrap_or_default(); - let hold = match state - .key_limits - .admit( - &client_name, - limits, - Default::default(), - crate::key_limits::slot_wait(&rt.config), - ) - .await - { - Ok(hold) => hold, - // 被拒的照样留一行,和规则拒绝的一样 - Err(r) => { - let why = r.error(); - let (id, ending) = open(&choice, "", tw_api::Billing::PerToken); - state.bus.emit(super::routed_nowhere(id, choice)); - ending.failed(why.source.into(), why.detail.clone()); - return Err(why); + // Responses 的连接:每个 `response.create` 是一个请求(见 `crate::ws::turn`) + let responses = crate::client_api::ClientApi::of_path(uri.path()) + == Some(crate::client_api::ClientApi::OpenaiResponses) + && crate::client_api::ClientApi::generates(uri.path()); + let rows = if responses { + // 连接本身不留行,上限按轮看。连不上上游时按升级的这一刻补上这一行 + crate::ws::Rows::Turns { + line: std::sync::Arc::new(crate::ws::turn::Line { + opener: opener.clone(), + choice: choice.clone(), + provider: name.clone(), + billing: provider.billing, + }), + upgraded: (started, now_ms()), + } + } else { + // 这把密钥的用量上限:**整条连接算一个请求**,连上之前看一遍,和 HTTP 那条路的准入 + // 同一套(见 `crate::key_limits`)。这一行不带用量,用量的上限只数得到它的请求数 + let limits = rt + .config + .clients + .iter() + .find(|c| c.name == client_name) + .map(|c| c.limits.as_slice()) + .unwrap_or_default(); + let hold = match state + .key_limits + .admit( + &client_name, + limits, + Default::default(), + crate::key_limits::slot_wait(&rt.config), + ) + .await + { + Ok(hold) => hold, + // 被拒的照样留一行,和规则拒绝的一样 + Err(r) => { + let why = r.error(); + let (id, ending) = open(&choice, "", tw_api::Billing::PerToken); + state.bus.emit(super::routed_nowhere(id, choice)); + ending.failed(why.source.into(), why.detail.clone()); + return Err(why); + } + }; + let (id, ending) = open(&choice, &name, provider.billing.into()); + hold.bind(id); + crate::ws::Rows::Connection { + id, + ending: Box::new(ending), } }; - let (id, ending) = open(&choice, &name, provider.billing.into()); - hold.bind(id); // 插件:升级那一刻的那一份表,一条连接用到底。**插件只管 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()); let plugins = (responses && !rt.plugins.is_empty()).then(|| crate::ws::Plugins { pool: state.plugin_pool.clone(), set: rt.plugins.clone(), @@ -299,8 +306,6 @@ pub(super) async fn ws_upgrade( Ok(ws.on_upgrade(move |sock| async move { // 一条 WS 连接活多久,这个请求就算在服务中多久 let _live = live; - let mut ending = ending; - ending.responded(101); - crate::ws::proxy(state, sock, upstream, rules, id, ending, plugins, naming).await; + crate::ws::proxy(state, sock, upstream, rules, rows, plugins, naming).await; })) } diff --git a/crates/tw-gateway/src/slots.rs b/crates/tw-gateway/src/slots.rs index 61f51cc0..11a164d4 100644 --- a/crates/tw-gateway/src/slots.rs +++ b/crates/tw-gateway/src/slots.rs @@ -17,9 +17,10 @@ //! 密钥的闸([`crate::limits`])一样,**跨重载存活**,上限改了在原来那个信号量上加减: //! 换一个新的,在跑的请求占着的位置就不算数了。上限去掉了的,等着的请求当场放行。 //! -//! **不占位置的两种**:数 token 的请求,它不跑模型、一眨眼就回来;WebSocket 的连接(见 -//! [`crate::ws`]),一条连接跑好几轮、中间可以闲着很久,占着的话一条没在答话的连接就能 -//! 把这家堵死。 +//! **不占位置的**:数 token 的请求,它不跑模型、一眨眼就回来;WebSocket 的连接本身 —— 一条 +//! 连接跑好几轮、中间可以闲着很久,占着的话一条没在答话的连接就能把这家堵死。Responses 的 +//! 连接上**每一轮各占一个**,从发出去占到这一次回答完(见 `crate::ws::turn`),闲着的连接 +//! 什么都不占;Realtime 和别的路径的连接不占。 use std::collections::HashMap; use std::sync::{Arc, Mutex}; diff --git a/crates/tw-gateway/src/ws.rs b/crates/tw-gateway/src/ws.rs index 8ac6d655..e7edea17 100644 --- a/crates/tw-gateway/src/ws.rs +++ b/crates/tw-gateway/src/ws.rs @@ -27,6 +27,14 @@ //! Realtime 的连接(`/v1/realtime`)模型写在升级请求的查询串里:升级时按它路由、过密钥的 //! 模型范围、对别名([`Naming::connect`]),发给上游的查询串写这一家自己的名称。 //! +//! # 一轮一个请求 +//! +//! **Responses 的连接上每个 `response.create` 是一个请求**(见 [`turn`]):从这一帧到这一次 +//! 回答完,开始、路由、结局三条事件,存储层记一行,带着这一次回答的用量,照 HTTP 那条路查价; +//! 密钥的用量上限、并发上限、这一家的位置(`max_concurrent`)都按轮算,闲着的连接什么都不占。 +//! 连接本身不留行,连不上上游的除外。Realtime 和别的路径的连接照旧**整条连接一行**、不带 +//! 用量:它们的回答不按 Responses 的事件收尾,分不出一轮一轮。 +//! //! 脚本插件也在这条路上跑(见 [`crate::plugin`]):客户端发来的每个 //! `response.create` 是一次请求。**这条路只有一跳**(升级时就连定了那一家,不换), //! 所以每个 `response.create` 过一遍请求钩子:上游是这条连接连的那一家,模型名是发给它的 @@ -59,6 +67,8 @@ use crate::error::GatewayError; use crate::state::AppState; use tw_types::{Msg, msg}; +pub(crate) mod turn; + /// `Option` 的替身。 /// /// axum 0.8 只给 `WebSocketUpgrade` 实现了 `FromRequestParts`,没有 @@ -261,6 +271,31 @@ struct Outgoing { model: String, /// 规则的参数改写:阶段一(按这一帧求)和阶段二累积的。模型名不看这里,看 `model` set: tw_engine::SetAction, + /// 附加了参数改写的规则:阶段一的,加上阶段二的。这一轮的开始和路由事件里报 + rewritten_by: Vec, +} + +/// 一帧为什么没发出去([`Naming::frame`])。 +struct NotSent { + why: GatewayError, + /// 规则拒绝了它:**这一帧照样留一行**,和 HTTP 那条路被规则拒绝的请求一样(见 + /// [`turn::denied`])。别的 —— 密钥不让用要发的模型、这一家服务不了要的别名、规则求不了 + /// 值 —— 不留,和 HTTP 那条路准入没过一样 + denied: Option, +} + +impl NotSent { + fn plain(why: GatewayError) -> Box { + Box::new(Self { why, denied: None }) + } +} + +/// 拒绝了一帧的那条规则,在哪个阶段。 +pub(crate) enum Denied { + /// 阶段一:它就是决定这一帧去向的那条 + PhaseOne(String), + /// 阶段二:记在路由事件的 `denied_by` 上 + PhaseTwo(String), } /// 要发的名字为什么发不出去([`Naming::name`])。 @@ -376,37 +411,60 @@ impl Naming { .ok_or_else(|| Unsent::Unserved(self.unserved(client, &asked.model))) } - /// 一帧 `response.create`(`frame`)发给这一家的样子。这一帧发不出去时是告诉客户端的那个 - /// 错误:规则拒绝了它、规则求不了值、密钥不让用要发的模型,或者这一家服务不了要的别名。 + /// 一帧按这条路径上的一次请求读出的样子:规则的条件按它求值,开始事件里的模型名、输入的 + /// 估算也从它来 + fn read(&self, frame: &serde_json::Value) -> crate::client_api::Reading { + let mut reading = crate::client_api::read(&self.path, None, Some(frame)); + reading.facts.client = self.client.clone(); + reading + } + + /// 一帧 `response.create`(读出的性质是 `facts`,见 [`Self::read`])发给这一家的样子。这一帧 + /// 发不出去时是告诉客户端的那个错误:规则拒绝了它、规则求不了值、密钥不让用要发的模型, + /// 或者这一家服务不了要的别名。 fn frame( &self, catalog: &tw_engine::Catalog, - frame: &serde_json::Value, - ) -> Result { - // 一帧的其余字段就是一个 Responses 请求:规则的条件按它读出的性质求值 - let mut facts = crate::client_api::read(&self.path, None, Some(frame)).facts; - facts.client = self.client.clone(); - let decision = self.this_frame(&facts)?; + facts: &tw_engine::RequestFacts, + ) -> Result> { + let decision = self.this_frame(facts)?; let p = &self.provider; - let (mut set, renamed) = match self.engine.phase_two(&facts, &p.name, &decision.set) { - Ok(tw_engine::Outcome2::Proceed { set, model, .. }) => (set, model), + let (mut set, renamed, two) = match self.engine.phase_two(facts, &p.name, &decision.set) { + Ok(tw_engine::Outcome2::Proceed { + set, + model, + rewritten_by, + }) => (set, model, rewritten_by), Ok(tw_engine::Outcome2::Deny { rule, reason }) => { tracing::info!(%rule, provider = %p.name, "a phase-two rule denied a WebSocket request"); - return Err(denied(rule, reason)); + return Err(Box::new(NotSent { + why: denied(rule.clone(), reason), + denied: Some(Denied::PhaseTwo(rule)), + })); } - Err(e) => return Err(rule_failed(e)), + Err(e) => return Err(NotSent::plain(rule_failed(e))), }; let asked = self .engine - .asked_of(&facts, &decision, &p.name, renamed.as_deref()); + .asked_of(facts, &decision, &p.name, renamed.as_deref()); let model = self .name(catalog, &decision, &facts.model, &asked) - .map_err(Unsent::into_error)?; + .map_err(|u| NotSent::plain(u.into_error()))?; // Codex 后端不认最大输出:HTTP 那条路发给它之前也会删掉(`crate::chatgpt::shape_passthrough`) if p.effective_protocol() == Some(tw_config::Protocol::Chatgpt) { set.max_tokens = None; } - Ok(Outgoing { model, set }) + let mut rewritten_by = decision.rewritten_by; + for r in two { + if !rewritten_by.contains(&r) { + rewritten_by.push(r); + } + } + Ok(Outgoing { + model, + set, + rewritten_by, + }) } /// 这一帧的决定。**去向是升级时的**(候选、经过的组、指定的模型),参数改写按这一帧重新 @@ -417,7 +475,7 @@ impl Naming { fn this_frame( &self, facts: &tw_engine::RequestFacts, - ) -> Result { + ) -> Result> { match self.engine.route(facts) { Ok(tw_engine::Outcome::Route(d)) => Ok(tw_engine::Decision { set: d.set, @@ -426,11 +484,14 @@ impl Naming { }), Ok(tw_engine::Outcome::Deny { rule, reason }) => { tracing::info!(%rule, "a rule denied a WebSocket request"); - Err(denied(rule, reason)) + Err(Box::new(NotSent { + why: denied(rule.clone(), reason), + denied: Some(Denied::PhaseOne(rule)), + })) } - Err(e) => Err(GatewayError::config(msg!( + Err(e) => Err(NotSent::plain(GatewayError::config(msg!( "gw.route.failed", detail = e => "Routing failed: {detail}" - ))), + )))), } } @@ -565,7 +626,11 @@ struct Pipes { dropping: bool, rules: Rules, provider: String, + /// 这条连接的号。整条连接一行的(Realtime 和别的路径)是那一行的;Responses 的连接自己 + /// 不留行,不在哪一轮里的帧(上游的、客户端的)报出去的事件挂在它上面(见 [`Self::event_id`]) id: u64, + /// Responses 的连接上在跑的几轮(见 [`turn`])。别的连接没有 + turns: Option, /// 范围里可能有插件时才有 plugins: Option, /// 每个 `response.create` 发出去的模型名怎么定。Responses 的连接才有 @@ -582,23 +647,46 @@ struct Pipes { rename: Option, } +impl Pipes { + /// 这一帧报出去的事件(脱敏、内容过滤、工具调用审查、插件的运行)挂在哪个请求上:上游 + /// 此刻在回答的那一轮,没有就是这条连接 + fn event_id(&self) -> u64 { + self.turns + .as_ref() + .and_then(turn::Turns::front_id) + .unwrap_or(self.id) + } +} + +/// 这条连接在流量里怎么记(见 [`turn`])。 +pub(crate) enum Rows { + /// 整条连接一行:Realtime 和别的路径的连接。升级时已经开始了,`ending` 是它欠着的结局 + Connection { + id: u64, + ending: Box, + }, + /// 每一轮一行:Responses 的连接。连接本身不留行 —— 连不上上游的除外,那时按升级的那一刻 + /// (`upgraded`:用时从哪一刻算起、那一刻的 Unix 毫秒)补上这一行 + Turns { + line: Arc, + upgraded: (std::time::Instant, u64), + }, +} + /// 接管一次升级。 /// /// 路由、鉴权都在调用方做完了 —— 这里把两条流接起来,并且**在每一帧 /// 上重新点一遍管线的保护**。 /// -/// 路由事件也在这里发:选中的那一家接没接下,要等和它握完手才知道。 -/// -/// **这条连接怎么断的,就是这个请求的结局**(`ending`)。每一条收场的 -/// 路径都先报结局、再去关连接:关连接要等对面,而对面可能已经不在了。 -#[allow(clippy::too_many_arguments)] -pub async fn proxy( +/// 整条连接一行的,路由事件在这里发:选中的那一家接没接下,要等和它握完手才知道。 +/// **这条连接怎么断的,就是这个请求的结局**。每一条收场的路径都先报结局、再去关连接:关连接 +/// 要等对面,而对面可能已经不在了。每一轮一行的(Responses),每一轮各报各的(见 [`turn`])。 +pub(crate) async fn proxy( state: AppState, client: WebSocket, upstream: Upstream, rules: Rules, - id: u64, - ending: crate::ending::Ending, + rows: Rows, plugins: Option, naming: Option, ) { @@ -607,49 +695,76 @@ pub async fn proxy( // **连不上也要报。**和 HTTP 那条路一样,失败的时候恰恰最需要看这一跳; // 报在结局之前,存储层落库时手上才有它 let name = &upstream.provider.name; - // 这一跳记发出去的模型名,和 HTTP 那条路一样。一条连接跑好几轮、每一帧写的模型可能 - // 不一样,升级时定得下来的只有每一帧都发的那个(指定模型、阶段一的改写,见 - // [`Naming::fixed`]);Realtime 的连接是查询串里的模型对过的名字 - let model = upstream - .model - .clone() - .or_else(|| naming.as_ref().and_then(|n| n.fixed(&state.catalog.load()))); - let attempt = match &connected { - Ok(_) => crate::server::hop( - name, - model, - tw_api::AttemptOutcome::Served, - 101, - hop_started, - ), - Err(NotConnected { - status: Some(s), .. - }) => crate::server::hop(name, model, tw_api::AttemptOutcome::Status, *s, hop_started), - Err(e) => crate::server::hop_failed(name, model, e.why.clone(), hop_started), - }; - // 和 HTTP 那条路同一个规矩:没接下的不按那一家记账 - let billing = match &connected { - Ok(_) => upstream.provider.billing, - Err(_) => tw_config::Billing::PerToken, + // 每一轮一行的连接连上了:连接本身不留行,每一轮各有各的号(见 `turn`)。连不上的补上 + // 这一行,和整条连接一行的一样报 + let (id, ending, turns) = match (rows, &connected) { + (Rows::Connection { id, mut ending }, _) => { + ending.responded(101); + (id, Some(*ending), None) + } + (Rows::Turns { line, .. }, Ok(_)) => { + (state.bus.next_id(), None, Some(turn::Turns::new(line))) + } + (Rows::Turns { line, upgraded }, Err(_)) => { + let (id, mut ending) = line.opener.open(turn::Opening { + choice: &line.choice, + to: (&line.provider, line.billing.into()), + model: String::new(), + session: None, + input_estimate: None, + started: upgraded.0, + at_ms: upgraded.1, + }); + ending.responded(101); + (id, Some(ending), None) + } }; - state.bus.emit(tw_api::Event::RequestRouted { - id, - route: upstream.route, - rule: upstream.rule, - group: upstream.group, - // 升级时就作用上的那几条(Realtime 的连接改写了查询串里的模型)。Responses 的连接上 - // 参数改写每一帧按这一帧求(见 [`Naming`]),那时路由事件早就发了 - rewritten_by: upstream.rewritten_by, - denied_by: None, - affinity: None, - attempts: vec![attempt], - billing: billing.into(), - }); + if ending.is_some() { + // 这一跳记发出去的模型名,和 HTTP 那条路一样。一条连接跑好几轮、每一帧写的模型可能 + // 不一样,升级时定得下来的只有每一帧都发的那个(指定模型、阶段一的改写,见 + // [`Naming::fixed`]);Realtime 的连接是查询串里的模型对过的名字 + let model = upstream + .model + .clone() + .or_else(|| naming.as_ref().and_then(|n| n.fixed(&state.catalog.load()))); + let attempt = match &connected { + Ok(_) => crate::server::hop( + name, + model, + tw_api::AttemptOutcome::Served, + 101, + hop_started, + ), + Err(NotConnected { + status: Some(s), .. + }) => crate::server::hop(name, model, tw_api::AttemptOutcome::Status, *s, hop_started), + Err(e) => crate::server::hop_failed(name, model, e.why.clone(), hop_started), + }; + // 和 HTTP 那条路同一个规矩:没接下的不按那一家记账 + let billing = match &connected { + Ok(_) => upstream.provider.billing, + Err(_) => tw_config::Billing::PerToken, + }; + state.bus.emit(tw_api::Event::RequestRouted { + id, + route: upstream.route, + rule: upstream.rule, + group: upstream.group, + // 升级时就作用上的那几条(Realtime 的连接改写了查询串里的模型) + rewritten_by: upstream.rewritten_by, + denied_by: None, + affinity: None, + attempts: vec![attempt], + billing: billing.into(), + }); + } let up = match connected { Ok(up) => up, Err(e) => { let text = e.why.text.clone(); - ending.failed(e.source, e.why); + if let Some(ending) = ending { + ending.failed(e.source, e.why); + } close_with(client, &text).await; return; } @@ -665,6 +780,7 @@ pub async fn proxy( rules, provider: upstream.provider.name, id, + turns, plugins, naming, requested_model: String::new(), @@ -796,98 +912,66 @@ enum End { Cut(Msg), } +/// 在等准入的那一轮(见 [`turn::admit`]):等到了交回这一轮和这一帧接下来要用的。 +type Waiting = std::pin::Pin< + Box, Next)> + Send>, +>; + +/// 一帧 `response.create` 过了准入之后要用的:查过内容过滤的那一帧(删过的话是删过的样子)、 +/// 客户端要的模型名、发给这一家的样子。 +struct Next { + text: String, + requested: String, + out: Outgoing, +} + +/// 客户端的一帧处理完之后怎么办。 +enum Step { + Go, + End(End), +} + async fn pump( state: AppState, client: WebSocket, up: Stream, p: &mut Pipes, - mut ending: crate::ending::Ending, + mut ending: Option, ) { let (mut c_tx, mut c_rx) = client.split(); let (mut u_tx, mut u_rx) = up.split(); - let end = loop { + // 一轮在等准入(上限、并发、这一家的位置)。**等的时候上游那一边照常转发**:前一轮的 + // 回答要接着交给客户端,它答完了,这一轮等的位置才空得出来 + let mut waiting: Option = None; + // 等的时候客户端接着发来的帧:排在那一轮后面,轮到了按顺序处理 + let mut held: std::collections::VecDeque = Default::default(); + let end = 'pump: loop { + while waiting.is_none() + && let Some(m) = held.pop_front() + { + if let Step::End(end) = + client_frame(&state, p, m, &mut c_tx, &mut u_tx, &mut waiting).await + { + break 'pump end; + } + } tokio::select! { - // 客户端 → 上游:**和普通请求同一个脱敏函数** msg = c_rx.next() => { let Some(Ok(m)) = msg else { break End::Closed }; - let out = match m { - Message::Text(t) => { - // 内容过滤在插件和脱敏之前:看的是客户端的原话。删过的话,后面用删过的 - // 那一帧 - let text = match screen_frame(&state, p, t.as_str()) { - Ok(text) => text, - Err(why) => { - let _ = c_tx.send(Message::Text( - format!("[ThinkWatch] {}", why.text).into(), - )).await; - break End::Cut(why); - } - }; - // 发出去的模型名和插件的请求钩子:钩子拿到的是查过(删过)的那一帧。改过 - // 的那一版再查一遍内容过滤(只报插件加进来的),脱敏换的是改过的那一版 - let text = match request(&state, p, &text).await { - Ok(text) => text, - Err(Refusal::Cut(why)) => { - let _ = c_tx.send(Message::Text( - format!("[ThinkWatch] {}", why.text).into(), - )).await; - break End::Cut(why); - } - // 只是这一帧不发:替它回一个 `response.failed`,连接照常 - Err(Refusal::Frame(err)) => { - tracing::info!(provider = %p.provider, why = %err.detail.text, - "a WebSocket request was not sent"); - let failed = failed_frame(err, None); - if c_tx.send(Message::Text(failed.into())).await.is_err() { - break End::Closed; - } - continue; - } - }; - let mode = p.rules.redact_mode; - let found = crate::guard::find(mode, &p.rules.redact, text.as_bytes()); - 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, 0), - at_ms: crate::server::now_ms(), - }); - if mode.acts() { - // 换的和报出去的是同一批:我们自己的占位符、base64 载荷不换 - let hits = crate::guard::hits(&text, &p.rules.redact); - let r = tw_guard::redact::replace::apply( - &text, - &hits, - std::mem::replace( - &mut p.ledger, - tw_guard::redact::replace::Ledger::new( - tw_guard::redact::replace::Scheme::SECRET, - ), - ), - ); - p.ledger = r.ledger; - UpMsg::Text(r.text.into()) - } else { - UpMsg::Text(text.into()) - } - } - } - // 二进制不检查,也不假装检查过 - Message::Binary(b) => UpMsg::Binary(b), - Message::Ping(b) => UpMsg::Ping(b), - Message::Pong(b) => UpMsg::Pong(b), - Message::Close(_) => break End::Closed, - }; - if let Err(e) = u_tx.send(out).await { - break End::Broke(msg!( - "gw.ws.send_failed", detail = e => "Sending to the upstream failed: {detail}" - )); + if waiting.is_some() { + if matches!(m, Message::Close(_)) { break End::Closed } + held.push_back(m); + continue; + } + if let Step::End(end) = client_frame(&state, p, m, &mut c_tx, &mut u_tx, &mut waiting).await { + break end; + } + } + got = async { waiting.as_mut().expect("polled only while one is waiting").await }, + if waiting.is_some() => { + waiting = None; + if let Step::End(end) = admitted(&state, p, got, &mut c_tx, &mut u_tx).await { + break end; } } // 上游 → 客户端:先还原占位符,再过工具墙 @@ -911,7 +995,9 @@ async fn pump( } } UpMsg::Binary(b) => { - ending.count(b.len()); + if let Some(e) = ending.as_mut() { + e.count(b.len()); + } Message::Binary(b) } UpMsg::Ping(b) => Message::Ping(b), @@ -924,16 +1010,274 @@ async fn pump( } } }; - // **先报结局,再关连接。**关连接要等对面回话,而对面可能早就不在了 - match end { - End::Closed => ending.finished(101), - End::Broke(why) => ending.failed(tw_api::FailureSource::Upstream, why), - End::Cut(why) => ending.failed(tw_api::FailureSource::Denied, why), + // **先报结局,再关连接。**关连接要等对面回话,而对面可能早就不在了。在等准入的那一轮 + // 开始了的话记成取消;在跑的几轮,上游断了的是失败,别的是取消 + drop(waiting); + if let Some(t) = p.turns.as_mut() { + match &end { + End::Broke(why) => t.fail_all(tw_api::FailureSource::Upstream, why.clone()), + _ => t.clear(), + } + } + if let Some(ending) = ending { + match end { + End::Closed => ending.finished(101), + End::Broke(why) => ending.failed(tw_api::FailureSource::Upstream, why), + End::Cut(why) => ending.failed(tw_api::FailureSource::Denied, why), + } } let _ = c_tx.close().await; let _ = u_tx.close().await; } +type UpstreamSink = futures::stream::SplitSink; + +/// 客户端 → 上游的一帧:**和普通请求同一个脱敏函数**。Responses 的连接上,一帧 +/// `response.create` 是一轮的开头([`begin_turn`]),要过准入时放进 `waiting`。 +async fn client_frame( + state: &AppState, + p: &mut Pipes, + m: Message, + c_tx: &mut ClientSink, + u_tx: &mut UpstreamSink, + waiting: &mut Option, +) -> Step { + let out = match m { + Message::Text(t) => { + if p.turns.is_some() + && let Some(frame) = create_frame(t.as_str()) + { + return begin_turn(state, p, t.as_str(), frame, c_tx, waiting).await; + } + // 内容过滤在脱敏之前:看的是客户端的原话。删过的话,后面用删过的那一帧 + let text = match screen_frame(state, p, t.as_str()) { + Ok(text) => text, + Err(why) => { + let _ = c_tx + .send(Message::Text(format!("[ThinkWatch] {}", why.text).into())) + .await; + return Step::End(End::Cut(why)); + } + }; + outbound(state, p, p.event_id(), text) + } + // 二进制不检查,也不假装检查过 + Message::Binary(b) => UpMsg::Binary(b), + Message::Ping(b) => UpMsg::Ping(b), + Message::Pong(b) => UpMsg::Pong(b), + Message::Close(_) => return Step::End(End::Closed), + }; + match u_tx.send(out).await { + Ok(()) => Step::Go, + Err(e) => Step::End(End::Broke(send_failed(e))), + } +} + +fn send_failed(e: impl std::fmt::Display) -> Msg { + msg!( + "gw.ws.send_failed", detail = e => "Sending to the upstream failed: {detail}" + ) +} + +/// 是一帧 `response.create` 的话,解出来的样子 +fn create_frame(text: &str) -> Option { + serde_json::from_str::(text) + .ok() + .filter(|v| v.get("type").and_then(|t| t.as_str()) == Some("response.create")) +} + +/// 一轮的开头:一帧 `response.create`(`raw` 是客户端的原话,`frame` 是它解出来的样子)。 +/// +/// 位置和 HTTP 那条路一样:内容过滤先下结论(不报:结论挂在这一轮的号上报),定发给这一家的 +/// 模型名和参数改写([`Naming`]:规则拒绝了的照样留一行),然后交去准入([`turn::admit`])。 +/// 插件的请求钩子、脱敏、发出在过了准入之后([`admitted`])。 +async fn begin_turn( + state: &AppState, + p: &mut Pipes, + raw: &str, + frame: serde_json::Value, + c_tx: &mut ClientSink, + waiting: &mut Option, +) -> Step { + let arrived = std::time::Instant::now(); + let at_ms = crate::server::now_ms(); + let (Some(naming), Some(turns)) = (p.naming.as_ref(), p.turns.as_ref()) else { + return Step::Go; + }; + let screening = { + let s = &p.rules.screen; + if s.mode.detects() { + crate::guard::screen(s, tw_dialect::ir::Dialect::Responses, raw.as_bytes()) + } else { + Default::default() + } + }; + // 删过的话,后面一律用删过的那一帧 + let (text, frame) = match &screening.body { + Some(b) => { + let text = String::from_utf8_lossy(b).into_owned(); + let frame = serde_json::from_str(&text).unwrap_or(frame); + (text, frame) + } + None => (raw.to_string(), frame), + }; + let reading = naming.read(&frame); + let requested = reading.facts.model.clone(); + let input_estimate = + matches!(reading.decoded, Some(Ok(_))).then_some(reading.facts.input_tokens); + let fingerprint = crate::session::fingerprint(&frame); + let line = turns.line.clone(); + // 别名对到这一家、密钥的模型范围继承时看的清单:这一帧(一次请求)用同一份 + let catalog = state.catalog.load(); + let out = match naming.frame(&catalog, &reading.facts) { + Ok(out) => out, + // 只是这一帧不发:替它回一个 `response.failed`,连接照常 + Err(no) => { + tracing::info!(provider = %p.provider, why = %no.why.detail.text, + "a WebSocket request was not sent"); + if let Some(by) = no.denied { + turn::denied( + state, + &line, + by, + &no.why, + requested, + fingerprint.as_deref(), + input_estimate, + arrived, + at_ms, + ); + } + return reply_failed(c_tx, no.why).await; + } + }; + let admit = turn::Admit { + state: state.clone(), + line, + rewritten_by: out.rewritten_by.clone(), + requested: requested.clone(), + sent: out.model.clone(), + fingerprint, + input_estimate, + screening, + arrived, + at_ms, + }; + let next = Next { + text, + requested, + out, + }; + *waiting = Some(Box::pin(async move { (turn::admit(admit).await, next) })); + Step::Go +} + +/// 一轮过了准入(或者没过):过插件的请求钩子、脱敏,发给上游,排进在跑的那几轮里。没过的、 +/// 插件拒绝的替它回一个 `response.failed`(被内容过滤、插件拒绝而切断的除外)。 +async fn admitted( + state: &AppState, + p: &mut Pipes, + (got, next): (Result, Next), + c_tx: &mut ClientSink, + u_tx: &mut UpstreamSink, +) -> Step { + let mut turn = match got { + Ok(turn) => turn, + Err(turn::NotAdmitted::Failed(err)) => { + tracing::info!(provider = %p.provider, why = %err.detail.text, + "a WebSocket request was not admitted"); + return reply_failed(c_tx, err).await; + } + Err(turn::NotAdmitted::Cut(why)) => { + let _ = c_tx + .send(Message::Text(format!("[ThinkWatch] {}", why.text).into())) + .await; + return Step::End(End::Cut(why)); + } + }; + let text = match request(state, p, turn.id, next).await { + Ok(text) => text, + Err(Refusal::Cut(why)) => { + turn.unsent(Vec::new(), tw_api::FailureSource::Denied, why.clone()); + let _ = c_tx + .send(Message::Text(format!("[ThinkWatch] {}", why.text).into())) + .await; + return Step::End(End::Cut(why)); + } + // 只是这一帧不发:替它回一个 `response.failed`,连接照常 + Err(Refusal::Frame(err)) => { + tracing::info!(provider = %p.provider, why = %err.detail.text, + "a WebSocket request was not sent"); + turn.unsent(Vec::new(), err.source.into(), err.detail.clone()); + return reply_failed(c_tx, err).await; + } + }; + let model = Some(p.sent_model.clone()).filter(|m| !m.is_empty() && *m != p.requested_model); + let out = outbound(state, p, turn.id, text); + turn.sent(model.clone()); + match u_tx.send(out).await { + Ok(()) => { + if let Some(t) = p.turns.as_mut() { + t.push(turn); + } + Step::Go + } + Err(e) => { + let why = send_failed(e); + let hop = crate::server::hop_failed( + &p.provider, + model, + why.clone(), + std::time::Instant::now(), + ); + turn.unsent(vec![hop], tw_api::FailureSource::Upstream, why.clone()); + Step::End(End::Broke(why)) + } + } +} + +/// 替没发出去的那一帧回一个 `response.failed`(见 [`failed_frame`]),连接照常 +async fn reply_failed(c_tx: &mut ClientSink, err: GatewayError) -> Step { + let failed = failed_frame(err, None); + if c_tx.send(Message::Text(failed.into())).await.is_err() { + return Step::End(End::Closed); + } + Step::Go +} + +/// 发给上游之前的最后一步:出站脱敏,**和普通请求同一个函数、同一份全局规则**。客户端发来 +/// 的一帧是一次请求,找到的挂在请求 `id` 上各报各的(一次最多报几个见 +/// `crate::guard::REPORTED_MAX`) +fn outbound(state: &AppState, p: &mut Pipes, id: u64, text: String) -> UpMsg { + let mode = p.rules.redact_mode; + let found = crate::guard::find(mode, &p.rules.redact, text.as_bytes()); + if found.is_empty() { + return UpMsg::Text(text.into()); + } + state.bus.emit(tw_api::Event::SecretsFound { + id, + provider: p.provider.clone(), + replaced: mode.acts(), + items: crate::guard::items(&found, 0), + at_ms: crate::server::now_ms(), + }); + if !mode.acts() { + return UpMsg::Text(text.into()); + } + // 换的和报出去的是同一批:我们自己的占位符、base64 载荷不换 + let hits = crate::guard::hits(&text, &p.rules.redact); + let r = tw_guard::redact::replace::apply( + &text, + &hits, + std::mem::replace( + &mut p.ledger, + tw_guard::redact::replace::Ledger::new(tw_guard::redact::replace::Scheme::SECRET), + ), + ); + p.ledger = r.ledger; + UpMsg::Text(r.text.into()) +} + type ClientSink = futures::stream::SplitSink; /// 上游的一帧文本处理完之后怎么办。 @@ -944,12 +1288,51 @@ enum Flow { } /// 上游的一帧文本:还原占位符、回答钩子、工具墙,然后发给客户端。 +/// +/// Responses 的连接上它属于上游此刻在回答的那一轮(见 [`turn`]):这一轮的结局按上游原话认 +/// (用量、第一个 token、上游报的错),回答完了的那一帧交给客户端之后,这一轮收场、放掉它 +/// 占着的。被工具墙切断的,这一轮记成拒绝。 async fn upstream_text( state: &AppState, p: &mut Pipes, t: &str, c_tx: &mut ClientSink, - ending: &mut crate::ending::Ending, + ending: &mut Option, +) -> Flow { + let kind = frame_kind(t); + if let Some(turn) = p.turns.as_mut().and_then(turn::Turns::front) { + turn.upstream(t); + } + let flow = relay(state, p, t, kind.as_deref(), c_tx, ending).await; + if let Some(turns) = p.turns.as_mut() { + match &flow { + Flow::Sent if ends_turn(kind.as_deref()) => turns.finish_front(), + Flow::End(End::Cut(why)) => { + turns.fail_front(tw_api::FailureSource::Denied, why.clone()) + } + _ => {} + } + } + flow +} + +/// 上游的这一帧(`type` 是 `kind`)是不是一次回答的结尾:完成、失败、没答完,或者一个错误 +/// (没开始回答就出错的,上游只回一个 `error`) +fn ends_turn(kind: Option<&str>) -> bool { + matches!( + kind, + Some("response.completed" | "response.failed" | "response.incomplete" | "error") + ) +} + +/// [`upstream_text`] 的转发那一半。`kind` 是这一帧的 `type` +async fn relay( + state: &AppState, + p: &mut Pipes, + t: &str, + kind: Option<&str>, + c_tx: &mut ClientSink, + ending: &mut Option, ) -> Flow { let restored = tw_guard::redact::replace::restore(t, &p.ledger); // 模型名换回客户端用的名称:一条消息是一个完整的 JSON,整条过一遍。排在回答钩子之前, @@ -963,12 +1346,11 @@ async fn upstream_text( None => restored, }; // 回答的边界:一次新的回答起一组回答钩子的实例;被切掉的那次剩下的帧不发 - let kind = frame_kind(&restored); let terminal = matches!( - kind.as_deref(), + kind, Some("response.completed" | "response.failed" | "response.incomplete") ); - if kind.as_deref() == Some("response.created") { + if kind == Some("response.created") { p.response = response_id(&restored); p.dropping = false; // 回答钩子:这一次回答起一组实例 @@ -1029,7 +1411,7 @@ async fn upstream_text( ledger: p.ledger.clone(), }; state.bus.emit(crate::server::flagged( - p.id, + p.event_id(), &p.provider, h, blocked, @@ -1045,7 +1427,9 @@ async fn upstream_text( .await; return Flow::End(End::Cut(why)); } - ending.count(msg.len()); + if let Some(e) = ending.as_mut() { + e.count(msg.len()); + } // 发不给客户端,就是客户端已经走了 if c_tx.send(Message::Text(msg.into())).await.is_err() { return Flow::End(End::Closed); @@ -1058,8 +1442,11 @@ async fn upstream_text( Flow::Sent } -/// 切掉这一次回答:替它发 `response.failed`,它剩下的帧不再发 +/// 切掉这一次回答:替它发 `response.failed`,它剩下的帧不再发。这一轮的结局记成拒绝 async fn fail_response(p: &mut Pipes, c_tx: &mut ClientSink, why: Msg) -> Flow { + if let Some(t) = p.turns.as_mut() { + t.cut_front(why.clone()); + } let failed = failed_frame(GatewayError::denied(why), p.response.as_deref()); p.dropping = true; p.reply = None; @@ -1069,44 +1456,35 @@ async fn fail_response(p: &mut Pipes, c_tx: &mut ClientSink, why: Msg) -> Flow { Flow::Sent } -/// 客户端发来的一帧为什么不发。 +/// 过了准入的一帧为什么还是不发(插件的请求钩子,见 [`plugin_request`])。 enum Refusal { /// 切断这条连接,告诉客户端的是这句话:插件拒绝了这个请求,和内容过滤拒掉一帧一样 Cut(Msg), - /// 只是这一帧不发,替它回一个 `response.failed`(见 [`failed_frame`]),连接照常:要的 - /// 别名这一家服务不了、密钥不让用要发的模型、规则拒绝了它。下一帧要的可能就是能发的 + /// 只是这一帧不发,替它回一个 `response.failed`(见 [`failed_frame`]),连接照常:插件换上的 + /// 别名这一家服务不了、密钥不让用插件换上的模型。下一帧要的可能就是能发的 Frame(GatewayError), } -/// 一次 `response.create`:定发给这一家的模型名和参数改写([`Naming`]),过插件的请求钩子。 -/// `text` 是查过内容过滤的那一帧(删过的话是删过的样子)。返回要发给上游的那一帧:插件改过 -/// 的话是改过的,模型名写成发给这一家的那个,规则的参数改写写进去。别的帧原样。 +/// 过了准入的一帧 `response.create`:过插件的请求钩子,写上发给这一家的模型名和规则的参数 +/// 改写([`Naming`] 在准入之前定好的,见 [`begin_turn`])。返回要发给上游的那一帧:插件改过的 +/// 话是改过的。`id` 是这一轮的号。 /// /// 这条路只有一跳:上游是这条连接连的那一家,插件的运行记在第 0 跳上。插件改过的那一版 -/// **再查一遍内容过滤**,只报插件加进来的(客户端的原话已经在 [`screen_frame`] 查过了,见 +/// **再查一遍内容过滤**,只报插件加进来的(客户端的原话在 [`begin_turn`] 查过了,见 /// [`screen_changed`])。 -async fn request(state: &AppState, p: &mut Pipes, text: &str) -> Result { - let Some(naming) = p.naming.as_ref() else { - return Ok(text.to_string()); - }; - let Some(frame) = serde_json::from_str::(text) - .ok() - .filter(|v| v.get("type").and_then(|t| t.as_str()) == Some("response.create")) - else { - return Ok(text.to_string()); - }; - let requested = frame - .get("model") - .and_then(|m| m.as_str()) - .unwrap_or_default() - .to_string(); - // 别名对到这一家、密钥的模型范围继承时看的清单:这一帧(一次请求)用同一份 +async fn request(state: &AppState, p: &mut Pipes, id: u64, next: Next) -> Result { + let Next { + text, + requested, + out: Outgoing { + model: sent, set, .. + }, + } = next; let catalog = state.catalog.load(); - let Outgoing { model: sent, set } = naming.frame(&catalog, &frame).map_err(Refusal::Frame)?; let (out, sent) = if p.plugins.is_some() { - plugin_request(state, p, &catalog, text, &requested, sent).await? + plugin_request(state, p, id, &catalog, &text, &requested, sent).await? } else { - (text.to_string(), sent) + (text, sent) }; // 和 HTTP 那条路一样,参数改写作用在插件改过的那一版上 let out = rewrite(out, &sent, &set); @@ -1127,6 +1505,7 @@ async fn request(state: &AppState, p: &mut Pipes, text: &str) -> Result plugged, Err(refused) => { - crate::plugin::request::record(state, p.id, &refused.runs); + crate::plugin::request::record(state, id, &refused.runs); return Err(Refusal::Cut(refused.why)); } }; - crate::plugin::request::record(state, p.id, &plugged.runs); + crate::plugin::request::record(state, id, &plugged.runs); p.bridge = plugged.bridge; let Some(c) = plugged.changed else { return Ok((text.to_string(), sent)); @@ -1190,6 +1569,7 @@ async fn plugin_request( let out = screen_changed( state, p, + id, text, String::from_utf8_lossy(&c.body).into_owned(), ) @@ -1229,7 +1609,13 @@ fn rewrite(text: String, model: &str, set: &tw_engine::SetAction) -> String { /// /// 返回要发出去的那一帧:处置档下插件加进来的字命中了删除规则的,是删过的样子。要拒绝时 /// 是告诉客户端的那句话 -fn screen_changed(state: &AppState, p: &Pipes, before: &str, after: String) -> Result { +fn screen_changed( + state: &AppState, + p: &Pipes, + id: u64, + before: &str, + after: String, +) -> Result { let s = &p.rules.screen; if !s.mode.detects() { return Ok(after); @@ -1240,7 +1626,7 @@ fn screen_changed(state: &AppState, p: &Pipes, before: &str, after: String) -> R before.as_bytes(), after.as_bytes(), ); - if let Some(why) = crate::guard::report(&state.bus, p.id, &p.provider, &sc) { + if let Some(why) = crate::guard::report(&state.bus, id, &p.provider, &sc) { return Err(why); } Ok(match sc.body { @@ -1265,7 +1651,7 @@ async fn start_reply(state: &AppState, p: &mut Pipes) -> Result<(), Msg> { model: &p.sent_model, requested_model: &p.requested_model, upstream: &p.provider, - request_id: p.id, + request_id: p.event_id(), attempt: 0, }; match crate::plugin::reply::Chain::start(state, &pc.set, bridge, &ctx).await { @@ -1304,6 +1690,7 @@ fn response_id(frame: &str) -> Option { /// 替被切掉的那次回答(或者没发出去的那一帧)发的 `response.failed`:和 SSE 那条路同一个 /// 形状(`tw_dialect` 的错误帧),id 换成这次回答的。没发出去的那一帧没有回答,id 是新的 fn failed_frame(err: GatewayError, response: Option<&str>) -> String { + let until_reset = err.retry.is_some_and(|r| r.until_reset); let sse = err .in_dialect(tw_dialect::ir::Dialect::Responses) .sse_frame(); @@ -1315,16 +1702,22 @@ fn failed_frame(err: GatewayError, response: Option<&str>) -> String { if let Some(id) = response { v["response"]["id"] = serde_json::Value::String(id.to_string()); } + // 密钥这一期的上限用完了:和 HTTP 那条路的 OpenAI 格式一样写成额度用完(见 + // `crate::error::Retry`)。Codex 按 `code` 决定退不退避,认这个码的直接停下来告诉用户 + if until_reset { + v["response"]["error"]["code"] = serde_json::Value::String("insufficient_quota".into()); + } v.to_string() } /// 客户端发来的一帧过一遍内容过滤:处置档下该拒的话是告诉客户端的那句话,否则是要发 -/// 出去的那一帧(删过的话是删过的样子)。 +/// 出去的那一帧(删过的话是删过的样子)。命中的挂在 [`Pipes::event_id`] 上报。 /// /// Codex 在 WS 上发的是 `{"type":"response.create", …}`,其余字段就是一个 Responses /// 请求:**按消息结构看**,和 HTTP 那条路一样只看调用方的消息、删也只删那里(见 /// [`crate::guard::screen`])。别的帧只用码位规则查整段原文(见 -/// [`crate::guard::screen_raw`])。 +/// [`crate::guard::screen_raw`])。Responses 的连接上一帧 `response.create` 是一轮的开头, +/// 在 [`begin_turn`] 里查,结论挂在那一轮上报。 fn screen_frame(state: &AppState, p: &Pipes, text: &str) -> Result { let s = &p.rules.screen; if !s.mode.detects() { @@ -1338,7 +1731,7 @@ fn screen_frame(state: &AppState, p: &Pipes, text: &str) -> Result } else { crate::guard::screen_raw(s, text) }; - if let Some(why) = crate::guard::report(&state.bus, p.id, &p.provider, &sc) { + if let Some(why) = crate::guard::report(&state.bus, p.event_id(), &p.provider, &sc) { return Err(why); } Ok(match sc.body { diff --git a/crates/tw-gateway/src/ws/turn.rs b/crates/tw-gateway/src/ws/turn.rs new file mode 100644 index 00000000..c4977247 --- /dev/null +++ b/crates/tw-gateway/src/ws/turn.rs @@ -0,0 +1,545 @@ +//! Responses 的 WebSocket 连接上的一轮:**每个 `response.create` 是一个请求**。 +//! +//! 从客户端发来这一帧,到上游回完这一次回答(`response.completed`、`response.failed`、 +//! `response.incomplete`,或者一个 `error`),和 HTTP 那条路的一个请求一样:开始、路由、 +//! 结局三条事件,存储层记一行。用量是这一次回答里的 `usage`(输入含缓存读、输出),结局事件 +//! 交给同一个记录器,按发给这一家的模型名、照 HTTP 那条路同一套查价:费用、密钥的用量、 +//! 体检、流量看到的都是它。第一个 token 什么时候到、回答里写的是哪个模型,也和 HTTP 那条路 +//! 一样认(见 [`Ending::frame`])。连接半路断了,这一轮记成取消;上游断了,记成失败。尝试链 +//! 只有一跳:这条连接连着的那一家。 +//! +//! **连接本身不留行**:一条连接跑好几轮、中间可以闲着很久,流量里该看的是每一轮。会话照 +//! HTTP 那条路按每一帧认(Codex 每段对话带着 `prompt_cache_key`),同一段对话的几轮归到同一 +//! 次会话里。没有哪一轮可挂的 —— 升级时就被规则拒绝的、连不上上游的 —— 那条连接自己留一行 +//! (见 `server::upgrade`)。 +//! +//! **上限和并发按轮算,和 HTTP 那条路的一个请求同一套**([`admit`]):天、周、月用满了就拒, +//! 密钥的并发上限等前面的结束,分钟、小时等得到就等;然后占这一家的一个位置 +//! (`max_concurrent`)—— 这条连接只连着这一家,没有下一家可换:满着就等,等不到回 +//! `response.failed`。被拒的这一轮替它回一个 `response.failed`,连接照常。占着的(密钥的 +//! 并发通行证、这一家的位置)从发出去占到这一次回答完、或者连接断了;**闲着的连接什么都 +//! 不占**。 +//! +//! 上限看的是这一轮开始那一刻的配置,和 HTTP 那条路每个请求一样:去向和防护按升级时的走 +//! 到底,上限改了,下一轮就照新的数。 + +use std::collections::VecDeque; +use std::sync::Arc; +use std::time::Instant; + +use crate::ending::Ending; +use crate::error::GatewayError; +use crate::server::Choice; +use crate::state::AppState; +use tw_types::{Msg, msg}; + +/// 一条连接上开始一个请求时都一样的那几项:谁发的、从哪儿来、哪条路径。整条连接一行的 +/// (Realtime 和别的路径)和每一轮一行的(Responses)都由它开始。 +#[derive(Clone)] +pub(crate) struct Opener { + pub(crate) bus: tw_observe::EventBus, + /// 网关密钥的名字 + pub(crate) client: String, + pub(crate) client_hint: Option, + pub(crate) peer: Option, + pub(crate) key_masked: Option, + /// 升级的路径 + pub(crate) path: String, +} + +/// 一个请求开始时各不相同的那几项。 +pub(crate) struct Opening<'a> { + pub(crate) choice: &'a Choice, + /// 要发往的那一家和它怎么收钱。一家都不会去的(被拒了)是空的名字 + pub(crate) to: (&'a str, tw_api::Billing), + /// 客户端要的模型名 + pub(crate) model: String, + pub(crate) session: Option, + pub(crate) input_estimate: Option, + /// 用时从哪一刻算起,和那一刻的 Unix 毫秒(那一行的 `at_ms`) + pub(crate) started: Instant, + pub(crate) at_ms: u64, +} + +impl Opener { + /// 发 `RequestStarted`,交回这个请求的号和它欠着的结局。**WS 的帧不留档**,结局没有正文 + /// 的去处 + pub(crate) fn open(&self, o: Opening<'_>) -> (u64, Ending) { + let id = self.bus.next_id(); + self.bus.emit(tw_api::Event::RequestStarted { + id, + client: self.client.clone(), + client_hint: self.client_hint.clone(), + session: o.session, + peer: self.peer.clone(), + key_masked: self.key_masked.clone(), + route: o.choice.route.clone(), + rule: o.choice.rule.clone(), + group: o.choice.group.clone(), + rewritten_by: o.choice.rewritten_by.clone(), + provider: o.to.0.to_string(), + billing: o.to.1, + model: o.model.clone(), + method: "WS".to_string(), + path: self.path.clone(), + input_estimate: o.input_estimate, + session_log_bytes: None, + at_ms: o.at_ms, + }); + let ending = Ending::new( + self.bus.clone(), + id, + o.model, + o.started, + o.at_ms as i64, + None, + ); + (id, ending) + } +} + +/// 一条 Responses 连接上每一轮都一样的:开始事件里的那几项、升级时路由的结论、连着的那一家。 +pub(crate) struct Line { + pub(crate) opener: Opener, + /// 升级时路由的结论:走的路由、决定去向的规则、经过的组。参数改写每一轮按那一帧求,换掉 + /// 这里的 `rewritten_by` + pub(crate) choice: Choice, + /// 这条连接连着的那一家 + pub(crate) provider: String, + pub(crate) billing: tw_config::Billing, +} + +/// 这条连接上在跑的几轮,按发出去的先后。**上游按顺序回答**:上游来的帧都算头一轮的,头一轮 +/// 回答完了,下一轮接上。 +pub(crate) struct Turns { + pub(crate) line: Arc, + queue: VecDeque, +} + +impl Turns { + pub(crate) fn new(line: Arc) -> Self { + Self { + line, + queue: VecDeque::new(), + } + } + + /// 上游此刻在回答的那一轮 + pub(crate) fn front(&mut self) -> Option<&mut Turn> { + self.queue.front_mut() + } + + pub(crate) fn front_id(&self) -> Option { + self.queue.front().map(|t| t.id) + } + + /// 发出去了:排在后面等上游回答 + pub(crate) fn push(&mut self, t: Turn) { + self.queue.push_back(t); + } + + /// 头一轮回答完了:报结局,放掉它占着的 + pub(crate) fn finish_front(&mut self) { + if let Some(t) = self.queue.pop_front() { + t.finish(); + } + } + + /// 头一轮失败了:被防护切断了 + pub(crate) fn fail_front(&mut self, source: tw_api::FailureSource, why: Msg) { + if let Some(t) = self.queue.pop_front() { + t.fail(source, why); + } + } + + /// 上游断了:在跑的几轮都没答完,一样失败 + pub(crate) fn fail_all(&mut self, source: tw_api::FailureSource, why: Msg) { + while let Some(t) = self.queue.pop_front() { + t.fail(source, why.clone()); + } + } + + /// 回答钩子切掉了头一轮的回答(替它发过 `response.failed`):上游收尾时记成拒绝 + pub(crate) fn cut_front(&mut self, why: Msg) { + if let Some(t) = self.queue.front_mut() { + t.cut = Some(why); + } + } + + /// 连接断了:没答完的几轮记成取消(结局的 Drop) + pub(crate) fn clear(&mut self) { + self.queue.clear(); + } +} + +/// 还没报的路由事件:尝试链上那一跳。**上游这一轮的第一帧到的时候报**,那一跳的用时就是等 +/// 它的那一段,和 HTTP 那条路等响应头一样;等不到的在结局之前报。 +struct Route { + choice: Choice, + attempt: tw_api::AttemptView, + billing: tw_api::Billing, + /// 这一跳从什么时候算:发出去的那一刻 + since: Instant, +} + +/// 在跑的一轮。**丢掉就是取消**(结局的 Drop),占着的跟着还回去。 +pub(crate) struct Turn { + pub(crate) id: u64, + /// 报了就没有了 + ending: Option, + bus: tw_observe::EventBus, + /// 这一帧到的那一刻:首字节时间从它算,和 HTTP 那条路从请求进来算一样 + started: Instant, + /// 上游这一轮的第一帧到了没有。到了报响应头 + responded: bool, + route: Option, + /// 回答钩子切掉了这一次回答:结局记成拒绝,原因是这一句 + cut: Option, + /// 密钥的并发通行证和这一家的位置:**跟着这一轮走**,回答完了、连接断了就还 + _pass: crate::limits::Pass, + _slot: crate::slots::Slot, +} + +impl Turn { + /// 这一帧发出去了,发给这一家的模型名是 `model`(和客户端要的不一样时才有:别名对过的、 + /// 规则或插件换过的)。尝试链上那一跳从这一刻算 + pub(crate) fn sent(&mut self, model: Option) { + if let Some(r) = self.route.as_mut() { + r.attempt.model = model; + r.since = Instant::now(); + } + } + + /// 上游这一轮来了一帧,**上游原话**(带占位符的那一版):头一帧报响应头,每一帧喂给结局 + /// 认用量、第一个 token、上游报的错 + pub(crate) fn upstream(&mut self, text: &str) { + if !self.responded { + self.responded = true; + self.routed(); + self.bus.emit(tw_api::Event::RequestHeaders { + id: self.id, + status: 200, + ttfb_ms: self.started.elapsed().as_millis() as u64, + }); + if let Some(e) = self.ending.as_mut() { + e.responded(200); + } + } + if let Some(e) = self.ending.as_mut() { + e.frame(text); + } + } + + fn routed(&mut self) { + let Some(r) = self.route.take() else { return }; + let mut attempt = r.attempt; + attempt.ms = r.since.elapsed().as_millis() as u64; + self.bus + .emit(routed(self.id, r.choice, vec![attempt], r.billing)); + } + + /// 这一次回答完了。上游在回答里报了错的(`response.failed`、`error`)是失败,原因是它说的 + /// 那句(见 [`Ending::streaming`]) + fn finish(mut self) { + self.routed(); + let Some(e) = self.ending.take() else { return }; + match self.cut.take() { + Some(why) => e.failed(tw_api::FailureSource::Denied, why), + None => e.finished(200), + } + } + + /// 失败了:上游断了、被防护切断了。用量照样带着(上游已经计了费) + pub(crate) fn fail(mut self, source: tw_api::FailureSource, why: Msg) { + self.routed(); + if let Some(e) = self.ending.take() { + e.failed(source, why); + } + } + + /// 准入过了、这一帧却没发出去:插件拒绝了它,或者写不过去。尝试链上是 `attempts`(写不 + /// 过去的那一跳;插件拒绝的没有),没接下的不按那一家记账 + pub(crate) fn unsent( + mut self, + attempts: Vec, + source: tw_api::FailureSource, + why: Msg, + ) { + if let Some(r) = self.route.take() { + self.bus.emit(routed( + self.id, + r.choice, + attempts, + tw_api::Billing::PerToken, + )); + } + if let Some(e) = self.ending.take() { + e.failed(source, why); + } + } +} + +impl Drop for Turn { + /// 没答完就被丢掉了(连接断了):路由事件要在结局之前到 —— 存储层落库时手上没有它的话, + /// 这一行没有尝试链。结局随后由它自己的 Drop 报成取消 + fn drop(&mut self) { + self.routed(); + } +} + +fn routed( + id: u64, + choice: Choice, + attempts: Vec, + billing: tw_api::Billing, +) -> tw_api::Event { + tw_api::Event::RequestRouted { + id, + route: choice.route, + rule: choice.rule, + group: choice.group, + rewritten_by: choice.rewritten_by, + denied_by: None, + affinity: None, + attempts, + billing, + } +} + +/// 一帧 `response.create` 要过准入时手上的。 +pub(crate) struct Admit { + pub(crate) state: AppState, + pub(crate) line: Arc, + /// 这一帧的参数改写附加的规则:阶段一按这一帧求的,加上阶段二的 + pub(crate) rewritten_by: Vec, + /// 客户端要的模型名 + pub(crate) requested: String, + /// 发给这一家的([`super::Naming`] 定的)。插件之后可能还会换 + pub(crate) sent: String, + /// 这一帧的会话指纹(见 [`crate::session::fingerprint`]) + pub(crate) fingerprint: Option, + /// 输入 token 的估算。解不开的帧没有 + pub(crate) input_estimate: Option, + /// 内容过滤的结论:开始之后挂在这一轮的号上报 + pub(crate) screening: tw_guard::content::Screening, + /// 这一帧到的那一刻 + pub(crate) arrived: Instant, + pub(crate) at_ms: u64, +} + +/// 这一轮没过准入。**行已经留下了**(开始、路由、结局),这里只说怎么告诉客户端。 +pub(crate) enum NotAdmitted { + /// 这一轮不发,替它回一个 `response.failed`,连接照常:上限拒了,这一家满着等不到 + Failed(GatewayError), + /// 内容过滤拒了这一帧:切断连接,告诉客户端的是这句话 + Cut(Msg), +} + +/// 一帧 `response.create` 的准入,**和 HTTP 那条路的一个请求同一套、同一个顺序**(见 +/// `server::pipeline::admission` 和 `hop`): +/// +/// 1. 天、周、月的上限 —— 用满了就拒; +/// 2. 密钥的并发上限 —— 等前面的结束; +/// 3. 分钟、小时的上限 —— 下一个空位在 `slot_wait_secs` 之内空出来就等,等不到就拒;过了 +/// 就按输入的估算占着,存储层记下这一行时换成实数; +/// 4. 开始:发开始事件,内容过滤的结论挂在这一轮上报(拒绝的切断连接); +/// 5. 这一家的位置(`max_concurrent`)—— 满着就等,最多 `slot_wait_secs`;这条连接只连着 +/// 这一家,等不到就是这一家忙,回 429 那句话。 +/// +/// 被上限拒的、等不到位置的**照样留一行**,流量里看得见它为什么没发出去。拿到的通行证和 +/// 位置交给这一轮([`Turn`]),跟着它走。 +pub(crate) async fn admit(a: Admit) -> Result { + let state = &a.state; + let line = a.line.clone(); + let key_name = line.opener.client.as_str(); + let rt = state.runtime(); + let key = rt.config.clients.iter().find(|c| c.name == key_name); + let max_concurrent = key.and_then(|c| c.max_concurrent); + let limits: &[tw_config::KeyLimit] = key.map(|c| c.limits.as_slice()).unwrap_or_default(); + // 等滚动窗口的空位、等这一家的位置,各等最多这么久(见 `crate::key_limits::slot_wait`) + let wait = crate::key_limits::slot_wait(&rt.config); + let choice = Choice { + rewritten_by: a.rewritten_by.clone(), + ..line.choice.clone() + }; + if let Err(r) = state.key_limits.calendar(key_name, limits) { + return Err(refused(&a, &choice, r.error())); + } + let pass = state.gate.acquire(key_name, max_concurrent).await; + let ask = ask(state, &a, limits); + let hold = match state.key_limits.admit(key_name, limits, ask, wait).await { + Ok(hold) => hold, + Err(r) => return Err(refused(&a, &choice, r.error())), + }; + let (id, mut ending) = line.opener.open(Opening { + choice: &choice, + to: (&line.provider, line.billing.into()), + model: a.requested.clone(), + session: session(&a), + input_estimate: a.input_estimate, + started: a.arrived, + at_ms: a.at_ms, + }); + hold.bind(id); + ending.streaming(tw_dialect::ir::Dialect::Responses, &line.provider); + // 拒绝的也在开始之后:被拒是一次来源为 `denied` 的失败,一个字节都不发 + if let Some(why) = crate::guard::report(&state.bus, id, &line.provider, &a.screening) { + ending.failed(tw_api::FailureSource::Denied, why.clone()); + return Err(NotAdmitted::Cut(why)); + } + let model = Some(a.sent.clone()).filter(|m| !m.is_empty() && *m != a.requested); + let hop_started = Instant::now(); + let (slot, queued_ms) = match state.slots.try_take(&line.provider) { + Some(slot) => (slot, None), + None => { + let mut waited = None; + let mut slot = None; + if !wait.is_zero() { + slot = state + .slots + .take_by(&line.provider, tokio::time::Instant::now() + wait) + .await; + waited = Some(hop_started.elapsed().as_millis() as u64); + } + match slot { + Some(slot) => (slot, waited), + None => { + let limit = state.slots.limit(&line.provider).unwrap_or_default(); + let busy = + crate::server::hop_busy(&line.provider, model, limit, waited, hop_started); + state + .bus + .emit(routed(id, choice, vec![busy], tw_api::Billing::PerToken)); + // 和 HTTP 那条路候选全满时同一句:能服务它的就这一家 + let err = GatewayError::busy(msg!( + "gw.busy_all", upstreams = format!("`{}`", line.provider) => + "Every upstream that can serve this request is at its concurrency limit \ + (max_concurrent): {upstreams}. None had a free slot in time; try again shortly." + )); + ending.failed(err.source.into(), err.detail.clone()); + return Err(NotAdmitted::Failed(err)); + } + } + } + }; + let mut attempt = crate::server::hop( + &line.provider, + model, + tw_api::AttemptOutcome::Served, + 200, + hop_started, + ); + attempt.queued_ms = queued_ms; + Ok(Turn { + id, + ending: Some(ending), + bus: state.bus.clone(), + started: a.arrived, + responded: false, + route: Some(Route { + choice, + attempt, + billing: line.billing.into(), + since: Instant::now(), + }), + cut: None, + _pass: pass, + _slot: slot, + }) +} + +/// 这一帧归到哪一次会话(见 [`crate::session::Sessions`])。认不出来是 None +fn session(a: &Admit) -> Option { + a.fingerprint + .as_deref() + .map(|fp| a.state.sessions.assign(fp, a.at_ms)) +} + +/// 这一轮要占多少:输入 token 的估算,和设了费用上限时按这一家、发给它的名字算的输入费用。 +/// 没有价格、不计费的是 0 —— 和结算时一样(见 `server::pipeline::admission`) +fn ask(state: &AppState, a: &Admit, limits: &[tw_config::KeyLimit]) -> crate::key_limits::Ask { + let tokens = a.input_estimate.unwrap_or(0); + let priced = limits + .iter() + .any(|l| l.measure() == tw_config::LimitMeasure::Cost); + let cost_micros = if priced && tokens > 0 { + let usage = tw_pricing::Usage { + input: tokens, + ..Default::default() + }; + crate::quote::quote( + &state.pricing.load(), + &a.line.provider, + &a.sent, + &usage, + a.line.billing, + ) + .cost_micros + .unwrap_or(0) + } else { + 0 + }; + crate::key_limits::Ask { + tokens, + cost_micros, + } +} + +/// 被上限拒了:照样开始、留一行(尝试链是空的),交回替它回 `response.failed` 的那个错误。 +fn refused(a: &Admit, choice: &Choice, why: GatewayError) -> NotAdmitted { + tracing::info!(key = %a.line.opener.client, "a usage limit of the key refused a WebSocket request"); + let (id, ending) = a.line.opener.open(Opening { + choice, + to: ("", tw_api::Billing::PerToken), + model: a.requested.clone(), + session: session(a), + input_estimate: a.input_estimate, + started: a.arrived, + at_ms: a.at_ms, + }); + a.state + .bus + .emit(crate::server::routed_nowhere(id, choice.clone())); + ending.failed(why.source.into(), why.detail.clone()); + NotAdmitted::Failed(why) +} + +/// 规则拒绝了这一帧:**和 HTTP 那条路被规则拒绝的请求一样留一行**(开始、空尝试链的路由、 +/// 一条 `denied` 的失败)。阶段二拒绝的,拒绝它的那条规则记在路由事件的 `denied_by` 上;阶段 +/// 一拒绝的,它就是决定去向的那条 +#[allow(clippy::too_many_arguments)] +pub(crate) fn denied( + state: &AppState, + line: &Line, + by: super::Denied, + why: &GatewayError, + requested: String, + fingerprint: Option<&str>, + input_estimate: Option, + started: Instant, + at_ms: u64, +) { + let mut choice = line.choice.clone(); + let denied_by = match by { + super::Denied::PhaseOne(rule) => { + choice.rule = rule; + None + } + super::Denied::PhaseTwo(rule) => Some(rule), + }; + let (id, ending) = line.opener.open(Opening { + choice: &choice, + to: ("", tw_api::Billing::PerToken), + model: requested, + session: fingerprint.map(|fp| state.sessions.assign(fp, at_ms)), + input_estimate, + started, + at_ms, + }); + let mut ev = crate::server::routed_nowhere(id, choice); + if let tw_api::Event::RequestRouted { denied_by: d, .. } = &mut ev { + *d = denied_by; + } + state.bus.emit(ev); + ending.failed(why.source.into(), why.detail.clone()); +} diff --git a/crates/tw-gateway/tests/alias_ws.rs b/crates/tw-gateway/tests/alias_ws.rs index f437a6f3..8b9d5241 100644 --- a/crates/tw-gateway/tests/alias_ws.rs +++ b/crates/tw-gateway/tests/alias_ws.rs @@ -188,7 +188,8 @@ fn refusal(frames: &[Value]) -> String { .to_string() } -/// 握手后发的路由事件里那一跳记下的模型名 +/// 下一轮的路由事件里那一跳记下的模型名:发给这一家的那个,和这一帧写的不一样时才有。 +/// **一轮是一个请求**(见 `tw_gateway::ws::turn`),尝试链按轮记 async fn attempt_model(rx: &mut Receiver) -> Option { loop { let ev = tokio::time::timeout(Duration::from_secs(5), rx.recv()) @@ -215,11 +216,12 @@ async fn an_alias_goes_out_as_the_upstreams_own_name_and_the_answer_shows_the_al let (up, seen) = upstream().await; let (gw, mut events) = gateway(config(provider(up, SCOPE), FAST, ""), vec![]).await; let mut c = connect(gw).await; - // 别名要看每一帧写的是什么:升级时说不上来 - assert_eq!(attempt_model(&mut events).await, None); - c.send(create("codex-fast")).await.unwrap(); let frames = one_answer(&mut c).await; + assert_eq!( + attempt_model(&mut events).await.as_deref(), + Some("gpt-5.1-codex-mini") + ); let sent = seen.lock().unwrap()[0].clone(); assert_eq!(sent["model"], "gpt-5.1-codex-mini"); // 别的字段原样 @@ -233,6 +235,7 @@ async fn an_alias_goes_out_as_the_upstreams_own_name_and_the_answer_shows_the_al c.send(create("gpt-5.1-codex")).await.unwrap(); let frames = one_answer(&mut c).await; + assert_eq!(attempt_model(&mut events).await, None, "原样发的不记"); assert_eq!(seen.lock().unwrap()[1]["model"], "gpt-5.1-codex"); assert_eq!( answered(&frames), @@ -251,14 +254,14 @@ async fn a_rule_that_pins_a_model_for_the_client_sends_the_pinned_name() { "; let (gw, mut events) = gateway(config(provider(up, &[]), FAST, rules), vec![]).await; let mut c = connect(gw).await; - assert_eq!( - attempt_model(&mut events).await.as_deref(), - Some("gpt-5.1-codex-max") - ); // 指定的原样发,写的是别名也不对 for asked in ["gpt-5.1-codex", "codex-fast"] { c.send(create(asked)).await.unwrap(); let frames = one_answer(&mut c).await; + assert_eq!( + attempt_model(&mut events).await.as_deref(), + Some("gpt-5.1-codex-max") + ); assert_eq!( seen.lock().unwrap().last().unwrap()["model"], "gpt-5.1-codex-max" @@ -280,13 +283,12 @@ async fn a_rule_rewrite_to_an_alias_is_resolved_and_a_phase_two_name_goes_out_as "; let (gw, mut events) = gateway(config(provider(up, SCOPE), FAST, rules), vec![]).await; let mut c = connect(gw).await; + c.send(create("gpt-5.1-codex")).await.unwrap(); + let frames = one_answer(&mut c).await; assert_eq!( attempt_model(&mut events).await.as_deref(), Some("gpt-5.1-codex-mini") ); - - c.send(create("gpt-5.1-codex")).await.unwrap(); - let frames = one_answer(&mut c).await; assert_eq!(seen.lock().unwrap()[0]["model"], "gpt-5.1-codex-mini"); assert_eq!(answered(&frames), ["gpt-5.1-codex", "gpt-5.1-codex"]); @@ -302,6 +304,24 @@ async fn a_rule_rewrite_to_an_alias_is_resolved_and_a_phase_two_name_goes_out_as "[ThinkWatch] Rule `不给` denied this request: 这个模型不走这里" ); assert_eq!(seen.lock().unwrap().len(), 2); + // 被规则拒绝的这一轮**照样留一行**,和 HTTP 那条路一样:空的尝试链,阶段二拒绝它的是 + // 哪条规则 + let (denied_by, attempts) = loop { + let ev = tokio::time::timeout(Duration::from_secs(5), events.recv()) + .await + .expect("no routing event for the denied request") + .unwrap(); + if let Event::RequestRouted { + denied_by: Some(rule), + attempts, + .. + } = ev + { + break (rule, attempts); + } + }; + assert_eq!(denied_by, "不给"); + assert!(attempts.is_empty(), "{attempts:?}"); c.send(create("gpt-5.1-codex")).await.unwrap(); let frames = one_answer(&mut c).await; @@ -344,7 +364,6 @@ async fn an_alias_the_upstream_cannot_serve_fails_that_request_and_the_connectio "; let (gw, mut events) = gateway(config(provider(up, SCOPE), &aliases, rules), vec![]).await; let mut c = connect(gw).await; - assert_eq!(attempt_model(&mut events).await, None); c.send(create("gpt-5.1-codex")).await.unwrap(); let frames = one_answer(&mut c).await; assert_eq!( @@ -354,6 +373,16 @@ async fn an_alias_the_upstream_cannot_serve_fails_that_request_and_the_connectio (claude-sonnet-5, us.anthropic.claude-sonnet-5-v1:0), so the request was not sent." ); assert!(seen.lock().unwrap().is_empty()); + // 这一家服务不了要的别名:**不留这一行**,和 HTTP 那条路准入没过一样 + while let Ok(Ok(ev)) = tokio::time::timeout(Duration::from_millis(200), events.recv()).await { + assert!( + !matches!( + ev, + Event::RequestStarted { .. } | Event::RequestRouted { .. } + ), + "{ev:?}" + ); + } } /// 插件把模型名换成别名:发给这一家的是它自己的名称,回答里写回客户端要的。插件的 diff --git a/crates/tw-gateway/tests/endings.rs b/crates/tw-gateway/tests/endings.rs index 42d60fa1..85322ca2 100644 --- a/crates/tw-gateway/tests/endings.rs +++ b/crates/tw-gateway/tests/endings.rs @@ -666,40 +666,44 @@ async fn an_error_answer_passed_on_to_the_client_is_failed_once_in_the_upstreams // ---------------------------------------------------------------- WebSocket -/// 一次 Codex 会话结束了。**以前 WS 这条路只有开始、没有结局**,每一条连接 -/// 在界面上都永远是「进行中」。 +/// 一帧 `response.create`:Responses 的连接上它是一轮的开头,**一轮是一个请求** +fn create() -> tokio_tungstenite::tungstenite::Message { + tokio_tungstenite::tungstenite::Message::Text( + r#"{"type":"response.create","model":"gpt-5","input":"hi"}"#.into(), + ) +} + +/// 一轮没答完客户端就关了连接:**这一轮记成取消,恰好一条**。连接本身不留行(一轮一个 +/// 请求,见 `tw_gateway::ws::turn`):关连接不再多出一条结局。 /// -/// 用量是 None:一条连接上跑着好几轮回答,而且升级请求里没有模型名 —— -/// 报一个数就是在编。 +/// 回显的上游把这一帧原样回过来,那不是一次回答的结尾:这一轮一直没答完。用量是 None, +/// 上游什么都没报 #[tokio::test] -async fn a_websocket_session_the_client_closes_is_finished() { +async fn a_websocket_turn_the_client_walks_away_from_is_cancelled() { let (gw, mut events) = serve(cfg(provider(ws_upstream("echo").await))).await; let mut c = ws_connect(gw).await; - c.send(tokio_tungstenite::tungstenite::Message::Text("hi".into())) - .await - .unwrap(); - tokio::time::timeout(Duration::from_secs(3), c.next()) + c.send(create()).await.unwrap(); + let echoed = tokio::time::timeout(Duration::from_secs(3), c.next()) .await .expect("等回帧超时") .unwrap() + .unwrap() + .into_text() .unwrap(); c.close(None).await.unwrap(); let got = endings(&mut events).await; assert_eq!(got.len(), 1, "该恰好有一个结局:{got:?}"); - assert!( - matches!( - &got[0], - Event::RequestFinished { - status: 101, - bytes: 2, - usage: None, - .. - } - ), - "该是一次带着回帧字节数、没有用量的结束:{got:?}" - ); - assert_eq!(model_of(&got[0]), "", "升级请求里没有模型名,不该编一个"); + match &got[0] { + Event::RequestCancelled { + status: Some(200), + bytes, + usage: None, + .. + } => assert_eq!(*bytes, echoed.len() as u64), + other => panic!("该是一次带着回帧字节数、没有用量的取消:{other:?}"), + } + assert_eq!(model_of(&got[0]), "gpt-5", "这一轮要的模型名"); } #[tokio::test] @@ -739,12 +743,8 @@ async fn a_websocket_cut_for_a_dangerous_tool_call_is_failed_as_denied() { }; let (gw, mut events) = serve(c).await; let mut client = ws_connect(gw).await; - client - .send(tokio_tungstenite::tungstenite::Message::Text( - "随便问一句".into(), - )) - .await - .unwrap(); + // 被切断的是这一轮:一轮是一个请求 + client.send(create()).await.unwrap(); let said = tokio::time::timeout(Duration::from_secs(3), client.next()) .await diff --git a/crates/tw-gateway/tests/key_limits.rs b/crates/tw-gateway/tests/key_limits.rs index 9456dd20..d4608330 100644 --- a/crates/tw-gateway/tests/key_limits.rs +++ b/crates/tw-gateway/tests/key_limits.rs @@ -1,5 +1,5 @@ //! 密钥的用量上限,走真的网关:被拒的请求在每一种客户端格式里长什么样、带什么响应头, -//! 流量里有没有它那一行,数 token 的请求和 WebSocket 连接怎么算。 +//! 流量里有没有它那一行,数 token 的请求怎么算。WebSocket 上每一轮怎么算在 `ws.rs`。 //! //! 数和等的细节在 `tw_gateway::key_limits` 的单元测试里;这里看的是它接在管线上的样子。 @@ -235,14 +235,15 @@ async fn counting_tokens_is_not_counted_and_not_refused() { ); } -/// WebSocket 一条连接算一个请求,**连上之前**看:用满了,升级就是一个 429。 +/// Realtime 的连接**整条连接算一个请求**,连上之前看:用满了,升级就是一个 429。 +/// (Responses 的连接每一轮各算一个,见 `ws.rs`) #[tokio::test] -async fn a_websocket_connection_is_admitted_when_it_opens() { +async fn a_realtime_connection_is_admitted_when_it_opens() { let (up, _hits) = upstream().await; let (gw, _rx) = gateway(up, "[{per: day, requests: 1}]").await; assert_eq!(anthropic(gw, "/v1/messages").await.status(), 200); use tokio_tungstenite::tungstenite::client::IntoClientRequest; - let mut req = format!("ws://{gw}/v1/responses") + let mut req = format!("ws://{gw}/v1/realtime?model=gpt-realtime") .into_client_request() .unwrap(); req.headers_mut() diff --git a/crates/tw-gateway/tests/ws.rs b/crates/tw-gateway/tests/ws.rs index 8b348afc..c649d2fe 100644 --- a/crates/tw-gateway/tests/ws.rs +++ b/crates/tw-gateway/tests/ws.rs @@ -364,38 +364,27 @@ fn the_route(evs: &[Event]) -> (String, Option, tw_api::AttemptView, Str ) } -/// **一次升级也报路由**:命中了哪条规则、经过哪个组、那一家接没接下、按什么 -/// 记账 —— 和 HTTP 那条路同一个形状。以前 WS 这条路只有开始和结局,详情里说 -/// 「没有路由信息」,上游的计费方式也没跟着报。 +/// **每一轮都报路由**:命中了哪条规则、经过哪个组、发给了哪一家、按什么记账 —— 和 HTTP +/// 那条路同一个形状。一轮是一个请求(见 `tw_gateway::ws::turn`),连接本身不留行 #[tokio::test] -async fn a_websocket_session_reports_its_route_and_its_upstreams_billing() { - let (up, _seen) = start_upstream("echo").await; +async fn a_websocket_turn_reports_its_route_and_its_upstreams_billing() { + let (up, _seen) = responder(None).await; let (gw, mut events) = serve(routed_to_an_account(up)).await; let mut c = connect(gw).await; - c.send(tokio_tungstenite::tungstenite::Message::Text("hi".into())) - .await - .unwrap(); - tokio::time::timeout(Duration::from_secs(3), c.next()) - .await - .expect("等回帧超时") - .unwrap() - .unwrap(); + c.send(create("hi")).await.unwrap(); + let frames = one_answer(&mut c).await; + assert_eq!(frames.last().unwrap()["type"], "response.completed"); c.close(None).await.unwrap(); let evs = until_the_ending(&mut events).await; - // 开始事件就说清按什么记账:升级没完成客户端就走了的,只有这一个 assert!( - matches!(&evs[0], Event::RequestStarted { billing, .. } if billing == "free"), + matches!(&evs[0], Event::RequestStarted { billing, route, method, model, session: Some(_), .. } + if billing == "free" && route == "default" && method == "WS" && model == "gpt-5"), "{evs:?}" ); let (rule, group, hop, billing) = the_route(&evs); assert_eq!(rule, "Codex 走账号"); assert_eq!(group.as_deref(), Some("账号池")); - // 走的哪条路由,开始和路由两条事件都说;升级请求没有正文,认不出会话 - assert!( - matches!(&evs[0], Event::RequestStarted { route, session: None, .. } if route == "default"), - "{evs:?}" - ); assert!( evs.iter() .any(|e| matches!(e, Event::RequestRouted { route, .. } if route == "default")), @@ -403,13 +392,30 @@ async fn a_websocket_session_reports_its_route_and_its_upstreams_billing() { ); assert_eq!( (hop.provider.as_str(), hop.outcome.slug(), hop.status), - ("订阅账号", "served", Some(101)) + ("订阅账号", "served", Some(200)) ); assert_eq!(billing, "free"); assert!( - matches!(evs.last(), Some(Event::RequestFinished { status: 101, .. })), + matches!( + evs.last(), + Some(Event::RequestFinished { + status: 200, + usage: Some(_), + .. + }) + ), "{evs:?}" ); + // 关掉连接不再多出一行 + while let Ok(Ok(ev)) = tokio::time::timeout(Duration::from_millis(300), events.recv()).await { + assert!( + !matches!( + ev, + Event::RequestStarted { .. } | Event::RequestFinished { .. } + ), + "{ev:?}" + ); + } } /// 规则拒绝了这次升级:**和 HTTP 那条路一样留一行** —— 开始、空尝试链的路由、 @@ -699,3 +705,600 @@ async fn a_secret_restored_into_a_flagged_call_is_masked_in_the_event() { assert!(!e.contains("USERSOWNKEY"), "**事件里是明文的密钥**:{e}"); } } + +// ---------------------------------------------------------------- 一轮一个请求 + +/// 这一轮回答的用量:输入 1200(其中 1000 走了缓存)、输出 30。Responses 的输入数包含缓存读 +const USAGE: &str = r#"{"input_tokens":1200,"input_tokens_details":{"cached_tokens":1000},"output_tokens":30,"output_tokens_details":{"reasoning_tokens":0},"total_tokens":1230}"#; + +/// 像 Responses 的 WebSocket 那样回答的上游:每个 `response.create` 回 created、一段文字, +/// 然后 completed(带用量)。`hold` 给了的话,回完那段文字之后等它变成 true 再收尾。记下收到的 +/// 每一帧 +async fn responder( + hold: Option>, +) -> (SocketAddr, Arc>>) { + let seen: Arc>> = Arc::default(); + let s = seen.clone(); + let app = Router::new().route( + "/backend-api/codex/responses", + axum::routing::any(move |ws: WebSocketUpgrade| { + let (seen, hold) = (s.clone(), hold.clone()); + async move { ws.on_upgrade(move |sock| answer(sock, seen, hold)) } + }), + ); + let l = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap(); + let addr = l.local_addr().unwrap(); + tokio::spawn(async move { axum::serve(l, app).await.unwrap() }); + (addr, seen) +} + +async fn answer( + mut sock: WebSocket, + seen: Arc>>, + mut hold: Option>, +) { + while let Some(Ok(m)) = sock.recv().await { + let Message::Text(t) = m else { continue }; + let n = { + let mut s = seen.lock().unwrap(); + s.push(t.to_string()); + s.len() + }; + let id = format!("resp_{n}"); + let head = [ + serde_json::json!({"type":"response.created","response":{"id":id,"status":"in_progress","model":"gpt-5","output":[]}}), + serde_json::json!({"type":"response.output_item.added","output_index":0,"item":{"type":"message","id":"msg","role":"assistant","content":[]}}), + serde_json::json!({"type":"response.output_text.delta","output_index":0,"content_index":0,"item_id":"msg","delta":"hello"}), + ]; + for f in head { + if sock + .send(Message::Text(f.to_string().into())) + .await + .is_err() + { + return; + } + } + if let Some(rx) = hold.as_mut() { + while !*rx.borrow_and_update() { + if rx.changed().await.is_err() { + return; + } + } + } + let usage: serde_json::Value = serde_json::from_str(USAGE).unwrap(); + let done = serde_json::json!({"type":"response.completed","response":{"id":id,"status":"completed","model":"gpt-5","output":[],"usage":usage}}); + if sock + .send(Message::Text(done.to_string().into())) + .await + .is_err() + { + return; + } + } +} + +/// 用量上限看的时钟:跟着真的时间走(东八区,从中午起),测试可以把它往后拨 +struct JumpClock { + base: tw_gateway::key_limits::TestClock, + extra: std::sync::atomic::AtomicI64, +} + +impl JumpClock { + fn new() -> Arc { + let noon = chrono::DateTime::parse_from_rfc3339("2026-10-05T12:00:00+08:00") + .unwrap() + .timestamp_millis(); + Arc::new(Self { + base: tw_gateway::key_limits::TestClock::new(noon, 8 * 3600), + extra: Default::default(), + }) + } + fn jump(&self, ms: i64) { + self.extra + .fetch_add(ms, std::sync::atomic::Ordering::SeqCst); + } +} + +impl tw_gateway::key_limits::Clock for JumpClock { + fn now_ms(&self) -> i64 { + self.base.now_ms() + self.extra.load(std::sync::atomic::Ordering::SeqCst) + } + fn period(&self, per: tw_config::LimitPer, at_ms: i64) -> (i64, i64) { + self.base.period(per, at_ms) + } + fn show(&self, at_ms: i64) -> String { + self.base.show(at_ms) + } +} + +/// 一家 Responses 上游 `up`、一把密钥 `codex`;`tweak` 改配置。交回网关的状态:测试要看这一家 +/// 此刻占着几个位置。**先订阅事件再起服务** +async fn turns_gateway( + up: SocketAddr, + tweak: impl FnOnce(&mut Config), + clock: Option>, +) -> (SocketAddr, Receiver, tw_gateway::AppState) { + let mut cfg = Config { + version: 1, + clients: vec![Client { + name: "codex".into(), + key: "tw-wskey".into(), + ..Default::default() + }], + providers: vec![Provider { + name: "up".into(), + base_url: format!("http://{up}"), + key: Some("sk-upstream".into()), + protocol: Some(tw_config::Protocol::OpenaiResponses), + ..Default::default() + }], + ..Default::default() + }; + tweak(&mut cfg); + let mut state = tw_gateway::AppState::new(cfg).unwrap(); + if let Some(c) = clock { + state.set_key_limits_clock(c); + } + let events = state.bus.subscribe(); + let addr = tw_gateway::serve(state.clone(), ([127, 0, 0, 1], 0).into()) + .await + .unwrap(); + tokio::time::sleep(Duration::from_millis(60)).await; + (addr, events, state) +} + +type Socket = + tokio_tungstenite::WebSocketStream>; + +/// 收到这一次回答的结尾为止的每一帧(解成 JSON) +async fn one_answer(c: &mut Socket) -> Vec { + let mut out = Vec::new(); + loop { + let m = tokio::time::timeout(Duration::from_secs(5), c.next()) + .await + .unwrap_or_else(|_| panic!("the answer did not end: {out:?}")) + .expect("the connection closed") + .unwrap(); + let tokio_tungstenite::tungstenite::Message::Text(t) = m else { + continue; + }; + let v: serde_json::Value = + serde_json::from_str(&t).unwrap_or_else(|_| serde_json::json!({ "raw": t.as_str() })); + let end = matches!( + v["type"].as_str(), + Some("response.completed" | "response.failed") + ); + out.push(v); + if end { + return out; + } + } +} + +/// 读到这一轮的第一段文字为止:上游已经在回答这一轮了 +async fn until_text(c: &mut Socket) { + loop { + let m = tokio::time::timeout(Duration::from_secs(5), c.next()) + .await + .expect("no text") + .unwrap() + .unwrap(); + if m.into_text().unwrap().contains("output_text.delta") { + return; + } + } +} + +fn id_of(e: &Event) -> Option { + match e { + Event::RequestStarted { id, .. } + | Event::RequestHeaders { id, .. } + | Event::RequestRouted { id, .. } + | Event::RequestFirstToken { id, .. } + | Event::RequestFinished { id, .. } + | Event::RequestFailed { id, .. } + | Event::RequestCancelled { id, .. } => Some(*id), + _ => None, + } +} + +fn is_ending(e: &Event) -> bool { + matches!( + e, + Event::RequestFinished { .. } + | Event::RequestFailed { .. } + | Event::RequestCancelled { .. } + ) +} + +/// 请求的事件,按到达的顺序,到第 `n` 个结局为止 +async fn requests(rx: &mut Receiver, n: usize) -> Vec { + let mut got = Vec::new(); + let mut ended = 0; + while ended < n { + let ev = tokio::time::timeout(Duration::from_secs(10), rx.recv()) + .await + .unwrap_or_else(|_| panic!("only {ended} of {n} endings: {got:?}")) + .unwrap(); + if id_of(&ev).is_none() { + continue; + } + ended += usize::from(is_ending(&ev)); + got.push(ev); + } + got +} + +/// 一个请求的那几条事件,按到达的顺序 +fn of(evs: &[Event], id: u64) -> Vec<&Event> { + evs.iter().filter(|e| id_of(e) == Some(id)).collect() +} + +fn started_ids(evs: &[Event]) -> Vec { + evs.iter() + .filter_map(|e| match e { + Event::RequestStarted { id, .. } => Some(*id), + _ => None, + }) + .collect() +} + +/// **每个 `response.create` 是一个请求**:开始、响应头、路由、第一个 token、结局,结局带着 +/// 这一轮回答的用量 —— 存储层按它查价、密钥的用量按它结算。两轮是两行,各有各的号;连接 +/// 本身不留行 +#[tokio::test] +async fn each_response_create_is_its_own_request_with_its_usage() { + let (up, _seen) = responder(None).await; + let (gw, mut rx, _state) = turns_gateway(up, |_| {}, None).await; + let mut c = connect(gw).await; + for text in ["one", "two"] { + c.send(create(text)).await.unwrap(); + let frames = one_answer(&mut c).await; + assert_eq!(frames.last().unwrap()["type"], "response.completed"); + } + c.close(None).await.unwrap(); + + let evs = requests(&mut rx, 2).await; + let ids = started_ids(&evs); + assert_eq!(ids.len(), 2, "{evs:?}"); + assert_ne!(ids[0], ids[1]); + for id in ids { + let mine = of(&evs, id); + assert!( + matches!(mine[0], Event::RequestStarted { method, model, provider, session: Some(_), input_estimate: Some(n), .. } + if method == "WS" && model == "gpt-5" && provider == "up" && *n > 0), + "{mine:?}" + ); + assert!( + mine.iter() + .any(|e| matches!(e, Event::RequestHeaders { status: 200, .. })), + "{mine:?}" + ); + assert!( + mine.iter() + .any(|e| matches!(e, Event::RequestFirstToken { .. })), + "第一个 token 没认出来:{mine:?}" + ); + let routed = mine + .iter() + .position(|e| matches!(e, Event::RequestRouted { .. })) + .unwrap_or_else(|| panic!("{mine:?}")); + assert!(routed < mine.len() - 1, "路由事件到在了结局之后:{mine:?}"); + let Event::RequestRouted { attempts, .. } = mine[routed] else { + unreachable!() + }; + assert_eq!(attempts.len(), 1, "{attempts:?}"); + assert_eq!( + ( + attempts[0].provider.as_str(), + attempts[0].outcome.slug(), + attempts[0].status + ), + ("up", "served", Some(200)) + ); + assert_eq!(attempts[0].model, None, "发出去的就是客户端要的那个"); + match mine.last().unwrap() { + Event::RequestFinished { + status: 200, + model, + usage: Some(u), + answered_model, + .. + } => { + assert_eq!(model, "gpt-5"); + assert_eq!((u.input, u.cache_read, u.output), (200, 1000, 30)); + assert_eq!(answered_model.as_deref(), Some("gpt-5")); + } + other => panic!("该是一次带着用量的结束:{other:?}"), + } + } + // 关掉连接不再多出一行 + while let Ok(Ok(ev)) = tokio::time::timeout(Duration::from_millis(300), rx.recv()).await { + assert!(id_of(&ev).is_none(), "{ev:?}"); + } +} + +/// 一轮没答完客户端就关了连接:这一轮记成取消,路由事件照样在结局之前到 +#[tokio::test] +async fn a_turn_the_connection_closes_on_is_cancelled() { + let (_release, hold) = tokio::sync::watch::channel(false); + let (up, _seen) = responder(Some(hold)).await; + let (gw, mut rx, _state) = turns_gateway(up, |_| {}, None).await; + let mut c = connect(gw).await; + c.send(create("hi")).await.unwrap(); + until_text(&mut c).await; + drop(c); + + let evs = requests(&mut rx, 1).await; + let kinds: Vec<&str> = evs + .iter() + .filter_map(|e| match e { + Event::RequestStarted { .. } => Some("started"), + Event::RequestRouted { .. } => Some("routed"), + Event::RequestCancelled { .. } => Some("cancelled"), + Event::RequestFinished { .. } | Event::RequestFailed { .. } => Some("other"), + _ => None, + }) + .collect(); + assert_eq!(kinds, ["started", "routed", "cancelled"], "{evs:?}"); +} + +/// 密钥这一天的上限用完了:**这一轮替它回一个 `response.failed`**,带着上限的那句话、写成 +/// 额度用完(Codex 认这个码,不再重试),连接照常;流量里照样有这一行。到了第二天,同一条 +/// 连接上的下一轮照常发出 +#[tokio::test] +async fn a_used_up_key_limit_fails_the_turn_and_the_connection_stays_usable() { + let (up, seen) = responder(None).await; + let clock = JumpClock::new(); + let (gw, mut rx, _state) = turns_gateway( + up, + |c| c.clients[0].limits = serde_yaml_ng::from_str("[{per: day, requests: 1}]").unwrap(), + Some(clock.clone()), + ) + .await; + let mut c = connect(gw).await; + c.send(create("one")).await.unwrap(); + one_answer(&mut c).await; + requests(&mut rx, 1).await; + + c.send(create("two")).await.unwrap(); + let frames = one_answer(&mut c).await; + let failed = frames.last().unwrap(); + assert_eq!(failed["type"], "response.failed", "{frames:?}"); + assert_eq!( + failed["response"]["error"]["message"], + "[ThinkWatch] Gateway key `codex` has reached its limit of 1 requests per day: 1 so far. \ + It resets at 2026-10-06 00:00 +08:00." + ); + assert_eq!(failed["response"]["error"]["code"], "insufficient_quota"); + assert_eq!(seen.lock().unwrap().len(), 1, "被拒的这一轮到了上游"); + let evs = requests(&mut rx, 1).await; + assert!( + matches!(&evs[0], Event::RequestStarted { provider, .. } if provider.is_empty()), + "{evs:?}" + ); + assert!( + evs.iter() + .any(|e| matches!(e, Event::RequestRouted { attempts, .. } if attempts.is_empty())), + "{evs:?}" + ); + match evs.last().unwrap() { + Event::RequestFailed { + source, message, .. + } => { + assert_eq!(source.slug(), "rate_limited"); + assert_eq!(message.code, "gw.key_limit.requests_per_period"); + } + other => panic!("{other:?}"), + } + + // 第二天:同一条连接,照常发出 + clock.jump(13 * 3600 * 1000); + c.send(create("three")).await.unwrap(); + let frames = one_answer(&mut c).await; + assert_eq!(frames.last().unwrap()["type"], "response.completed"); + assert_eq!(seen.lock().unwrap().len(), 2); +} + +/// 每分钟的上限:下一个空位在等得到的时候空出来,这一轮**等它**,然后照常发出;等不到的 +/// (`slot_wait_secs: 0`)当场替它回 `response.failed`,说清多久之后再来 +#[tokio::test] +async fn a_rolling_key_limit_waits_within_the_turn_or_refuses_it() { + let (up, seen) = responder(None).await; + let clock = JumpClock::new(); + let (gw, _rx, _state) = turns_gateway( + up, + |c| c.clients[0].limits = serde_yaml_ng::from_str("[{per: minute, requests: 1}]").unwrap(), + Some(clock.clone()), + ) + .await; + let mut c = connect(gw).await; + c.send(create("one")).await.unwrap(); + one_answer(&mut c).await; + // 一分钟差 600 毫秒:下一个空位 600 毫秒后空出来 + clock.jump(59_400); + let t = std::time::Instant::now(); + c.send(create("two")).await.unwrap(); + let frames = one_answer(&mut c).await; + assert_eq!(frames.last().unwrap()["type"], "response.completed"); + assert!( + t.elapsed() >= Duration::from_millis(500), + "没等空位就发了:{:?}", + t.elapsed() + ); + assert_eq!(seen.lock().unwrap().len(), 2); + + // 不等的配置:当场拒,连接照常 + let (up, seen) = responder(None).await; + let (gw, _rx, _state) = turns_gateway( + up, + |c| { + c.clients[0].limits = serde_yaml_ng::from_str("[{per: minute, requests: 1}]").unwrap(); + c.failover.slot_wait_secs = 0; + }, + Some(JumpClock::new()), + ) + .await; + let mut c = connect(gw).await; + c.send(create("one")).await.unwrap(); + one_answer(&mut c).await; + c.send(create("two")).await.unwrap(); + let frames = one_answer(&mut c).await; + let failed = frames.last().unwrap(); + assert_eq!(failed["type"], "response.failed", "{frames:?}"); + let said = failed["response"]["error"]["message"].as_str().unwrap(); + assert!( + said.starts_with( + "[ThinkWatch] Gateway key `codex` has reached its limit of 1 requests per minute: 1 in \ + the last minute. Try again in " + ), + "{said}" + ); + // 滚动窗口过一会儿就空出来:可以重试 + assert_eq!(failed["response"]["error"]["code"], "rate_limit_exceeded"); + assert_eq!(seen.lock().unwrap().len(), 1); +} + +/// 这一家的并发上限(`max_concurrent`)**按轮占**:闲着的连接什么都不占,一轮从发出去占到 +/// 这一次回答完 +#[tokio::test] +async fn an_upstream_slot_is_held_for_a_turn_and_not_by_an_idle_connection() { + let (release, hold) = tokio::sync::watch::channel(false); + let (up, _seen) = responder(Some(hold)).await; + let (gw, mut rx, state) = + turns_gateway(up, |c| c.providers[0].max_concurrent = Some(1), None).await; + let mut c = connect(gw).await; + // 连着、闲着:位置是空的 + drop( + state + .slots + .try_take("up") + .expect("an idle connection holds a slot"), + ); + + c.send(create("hi")).await.unwrap(); + until_text(&mut c).await; + assert!( + state.slots.try_take("up").is_none(), + "回答着的这一轮没占位置" + ); + + release.send_replace(true); + one_answer(&mut c).await; + requests(&mut rx, 1).await; + // 答完了:还回来了,连接还开着 + assert!(state.slots.try_take("up").is_some(), "答完的这一轮没还位置"); + drop(c); +} + +/// 密钥的并发上限(`max_concurrent`)也**按轮占**:一条连接上的一轮在答,同一把密钥另一条 +/// 连接上的一轮等它答完再发;答完之后闲着的那条连接不挡别人 +#[tokio::test] +async fn a_keys_max_concurrent_is_held_per_turn() { + let (release, hold) = tokio::sync::watch::channel(false); + let (up, seen) = responder(Some(hold)).await; + let (gw, _rx, _state) = + turns_gateway(up, |c| c.clients[0].max_concurrent = Some(1), None).await; + let mut a = connect(gw).await; + let mut b = connect(gw).await; + a.send(create("a")).await.unwrap(); + until_text(&mut a).await; + + b.send(create("b")).await.unwrap(); + tokio::time::sleep(Duration::from_millis(300)).await; + assert_eq!(seen.lock().unwrap().len(), 1, "密钥的并发上限没挡住第二轮"); + + release.send_replace(true); + one_answer(&mut a).await; + let frames = one_answer(&mut b).await; + assert_eq!(frames.last().unwrap()["type"], "response.completed"); + assert_eq!(seen.lock().unwrap().len(), 2); + + // a 连着、闲着:b 的下一轮当场就发 + b.send(create("b again")).await.unwrap(); + let frames = one_answer(&mut b).await; + assert_eq!(frames.last().unwrap()["type"], "response.completed"); +} + +/// 这一家满着:这一轮**等它空出位置**(最多 `slot_wait_secs`),空出来了照常发,尝试链上记着 +/// 等了多久;等不到替它回 `response.failed`(忙,可以重试),流量里留一行,连接照常 +#[tokio::test] +async fn a_turn_waits_for_a_full_upstream_and_fails_as_busy_when_none_frees() { + let (up, seen) = responder(None).await; + let (gw, mut rx, state) = turns_gateway( + up, + |c| { + c.providers[0].max_concurrent = Some(1); + c.failover.slot_wait_secs = 1; + }, + None, + ) + .await; + let mut c = connect(gw).await; + + // 别的请求占着,300 毫秒后还回来:这一轮等到了 + let other = state.slots.try_take("up").unwrap(); + c.send(create("one")).await.unwrap(); + tokio::time::sleep(Duration::from_millis(300)).await; + assert!(seen.lock().unwrap().is_empty(), "满着就发了"); + drop(other); + let frames = one_answer(&mut c).await; + assert_eq!(frames.last().unwrap()["type"], "response.completed"); + let evs = requests(&mut rx, 1).await; + let queued = evs.iter().find_map(|e| match e { + Event::RequestRouted { attempts, .. } => attempts[0].queued_ms, + _ => None, + }); + assert!( + queued.is_some_and(|ms| ms >= 250), + "尝试链上没记等了多久:{evs:?}" + ); + + // 一直占着:等满一秒,这一轮是忙 + let other = state.slots.try_take("up").unwrap(); + let t = std::time::Instant::now(); + c.send(create("two")).await.unwrap(); + let frames = one_answer(&mut c).await; + assert!( + t.elapsed() >= Duration::from_millis(900), + "{:?}", + t.elapsed() + ); + let failed = frames.last().unwrap(); + assert_eq!(failed["type"], "response.failed", "{frames:?}"); + assert_eq!(failed["response"]["error"]["code"], "rate_limit_exceeded"); + assert!( + failed["response"]["error"]["message"] + .as_str() + .unwrap() + .contains("(max_concurrent): `up`"), + "{failed}" + ); + assert_eq!(seen.lock().unwrap().len(), 1); + let evs = requests(&mut rx, 1).await; + let hop = evs + .iter() + .find_map(|e| match e { + Event::RequestRouted { attempts, .. } => Some(attempts[0].clone()), + _ => None, + }) + .unwrap_or_else(|| panic!("{evs:?}")); + assert_eq!(hop.skipped, Some(tw_api::ServeSkip::Busy)); + assert!(hop.queued_ms.is_some_and(|ms| ms >= 900), "{hop:?}"); + match evs.last().unwrap() { + Event::RequestFailed { + source, message, .. + } => { + assert_eq!(source.slug(), "rate_limited"); + assert_eq!(message.code, "gw.busy_all"); + } + other => panic!("{other:?}"), + } + + // 空出来了:同一条连接,照常 + drop(other); + c.send(create("three")).await.unwrap(); + let frames = one_answer(&mut c).await; + assert_eq!(frames.last().unwrap()["type"], "response.completed"); +} diff --git a/docs/config.md b/docs/config.md index f7abf437..a0d049fb 100644 --- a/docs/config.md +++ b/docs/config.md @@ -344,6 +344,11 @@ restart, the day, week and month are added up again from the request records; a monthly limit therefore needs `retention.row_days` of at least 31. Minute and hour limits start empty. +On a Responses WebSocket connection, each `response.create` is a request of +its own: it is recorded with its usage and cost and checked against these +limits, and a refused one is answered with `response.failed` while the +connection stays open. + @@ -442,7 +447,10 @@ been passed on in full or the client has gone. When the upstream is full, a conversation that stays on it to reuse its prompt cache waits for a slot; any other request goes straight to the next upstream. How long a request waits is `failover.slot_wait_secs`. Waiting is not a failure: the upstream is -not set aside. Requests that only count tokens do not take a slot. +not set aside. Requests that only count tokens do not take a slot. On a +Responses WebSocket connection, each `response.create` takes a slot from when +it is sent until its answer ends, and so does a key's `max_concurrent`; an +idle connection takes none. #### `providers[].oauth` diff --git a/docs/config.zh-CN.md b/docs/config.zh-CN.md index b5e17c79..7e0e5ced 100644 --- a/docs/config.zh-CN.md +++ b/docs/config.zh-CN.md @@ -247,6 +247,9 @@ clients: 重启之后,这一天、这一周、这个月的用量从请求记录里重新加起来,所以设了按月上限时 `retention.row_days` 至少要 31;按分钟、按小时的上限从零开始。 +Responses 的 WebSocket 连接上,每个 `response.create` 各算一个请求:带着各自的用量和费用记下, +按这些上限检查;被拒的那一个收到 `response.failed`,连接保持不断。 + @@ -325,7 +328,7 @@ providers: ChatGPT 账号上游(`protocol: chatgpt`)只接受桌面应用登录得到的凭据,不能手写。不支持 Claude 和 Google 的订阅登录,请使用 API 密钥。 -有的中转站和账号同时只接受几个请求,多出来的直接拒绝。`max_concurrent` 让网关守住这个数:请求发出时占用这家的一个位置,回答完整交给客户端、或者客户端断开时归还。这家满了的时候,为复用提示缓存而留在这家的对话等空位,别的请求直接换下一家。最多等多久由 `failover.slot_wait_secs` 决定。等待不算失败,这家不会因此停用。只计算 token 数的请求不占位置。 +有的中转站和账号同时只接受几个请求,多出来的直接拒绝。`max_concurrent` 让网关守住这个数:请求发出时占用这家的一个位置,回答完整交给客户端、或者客户端断开时归还。这家满了的时候,为复用提示缓存而留在这家的对话等空位,别的请求直接换下一家。最多等多久由 `failover.slot_wait_secs` 决定。等待不算失败,这家不会因此停用。只计算 token 数的请求不占位置。Responses 的 WebSocket 连接上,每个 `response.create` 从发出起占一个位置,直到它的回答结束,密钥的 `max_concurrent` 也一样;空闲的连接不占位置。 #### `providers[].oauth` From 9c81886cb876d2823c3002e12de1115877a239dc Mon Sep 17 00:00:00 2001 From: fylorn <249551762+fylorn@users.noreply.github.com> Date: Mon, 5 Oct 2026 18:40:31 +0800 Subject: [PATCH 10/22] Count WebSocket turns like HTTP hops: one wait deadline, speed and success samples MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Responses WebSocket turns became requests of their own in the previous commit, but they still lived outside the routing signals the HTTP path now uses: - turn::admit gave the rolling key-limit wait and the upstream-slot wait the full slot_wait_secs each, so a turn could wait twice as long as an HTTP request. It now sets one deadline after the key's max_concurrent gate and shares it between the two waits, as HTTP admission does. - No turn ever fed the latency tracker, the breaker or the success rate: WebSocket did not touch Health at all, so an upstream failing every turn was never paused and kept its full share in a health-balanced group. A turn now records a speed sample from the moment the upstream starts on it (sent, or the previous turn on the connection ended) to its first content, read by the same first-token reader as HTTP. It records success or failure once, like an HTTP hop: first content is success; a turn that ends before any content is judged by its last frame with the same reader and classifier as an HTTP stream that errors before content (upstream fault = failure, the request's own fault = success); a broken upstream connection is a failure; a busy turn, a refused one, a client that leaves and a guard cut record nothing. The connection itself counts like a hop too: a handshake the upstream refuses or that cannot be made, and credentials that cannot be read at upgrade, record a failure; a whole-connection row (Realtime) that connects records a success. - A turn that ends on an upstream `error` frame may be followed by a `response.failed` for the same response. It was attributed to the next queued turn, which then ended as that failure while its own answer fell to the turn after it. Such a frame is now recognised by its response id and passed to the client without being counted for any turn. Only turns that ended on `error` are remembered, each for one closing frame, and a turn's own `response.created` or id always wins, so an upstream that reuses one id for every response cannot leave a turn waiting forever. The upstream slot tests named their route "默认" without default_route, so every request went through the built-in all-upstreams group instead of the group the tests are about. The config now sets it, and the helper that reads each request's route asserts the request went through that group. Co-Authored-By: Claude Opus 5.5 --- crates/tw-gateway/src/ending.rs | 6 + crates/tw-gateway/src/key_limits/mod.rs | 5 +- crates/tw-gateway/src/server.rs | 3 +- crates/tw-gateway/src/server/pipeline.rs | 2 + crates/tw-gateway/src/server/pipeline/hop.rs | 18 ++ .../tw-gateway/src/server/pipeline/opening.rs | 10 + crates/tw-gateway/src/server/upgrade.rs | 15 +- crates/tw-gateway/src/ws.rs | 78 ++++- crates/tw-gateway/src/ws/turn.rs | 184 +++++++++-- crates/tw-gateway/tests/upstream_slots.rs | 47 ++- crates/tw-gateway/tests/ws.rs | 292 ++++++++++++++++++ docs/config.md | 9 +- docs/config.zh-CN.md | 6 +- 13 files changed, 619 insertions(+), 56 deletions(-) diff --git a/crates/tw-gateway/src/ending.rs b/crates/tw-gateway/src/ending.rs index fdd2112e..06fdf16a 100644 --- a/crates/tw-gateway/src/ending.rs +++ b/crates/tw-gateway/src/ending.rs @@ -252,6 +252,12 @@ impl Ending { self.lap = Some(lap); } + /// 第一段内容到了没有:认第一个 token 的那一个认出来了([`Ending::streaming`] 之后才 + /// 认)。WebSocket 上 Responses 的一轮靠它判断这一家答上了没有(见 `crate::ws::turn`) + pub fn has_content(&self) -> bool { + self.opened.is_some() + } + /// 上游 `provider` 回的不是 2xx,原样交给了客户端(4xx 是请求本身的问题,或者没有 /// 下一家可换了;3xx 交还客户端,由它决定跟不跟)。 /// diff --git a/crates/tw-gateway/src/key_limits/mod.rs b/crates/tw-gateway/src/key_limits/mod.rs index cf7eaf52..34960a64 100644 --- a/crates/tw-gateway/src/key_limits/mod.rs +++ b/crates/tw-gateway/src/key_limits/mod.rs @@ -468,8 +468,9 @@ impl KeyLimits { out } - /// 准入第三步,从此刻起最多等 `wait`(见 [`Self::admit_by`])。WebSocket 的连接用它: - /// 那条路之后不再等上游的空位,这一段就是全部 + /// 准入第三步,从此刻起最多等 `wait`(见 [`Self::admit_by`])。整条连接一行的 WebSocket + /// 连接(Realtime 和别的路径)用它:那条路不等上游的空位,这一段就是全部。Responses 连接 + /// 上的每一轮和 HTTP 的请求一样用 [`Self::admit_by`] pub async fn admit( self: &Arc, key: &str, diff --git a/crates/tw-gateway/src/server.rs b/crates/tw-gateway/src/server.rs index 259cc6a7..e161fb79 100644 --- a/crates/tw-gateway/src/server.rs +++ b/crates/tw-gateway/src/server.rs @@ -13,6 +13,7 @@ use crate::error::GatewayError; use crate::health::Health; use crate::state::AppState; use listing::{get_model, list_models}; +pub(crate) use pipeline::stream_fault; use tw_types::msg; mod listing; @@ -228,7 +229,7 @@ async fn passthrough( /// (全都熔断时我们是放行的,所以那些失败照样打得到这家)。只睡一次 /// 就下结论的话,它醒来时看到的是一段还没走完的冷却 —— 然后闭嘴,而 /// 真正到点的那一刻再也没有人报。 -fn note_health( +pub(crate) fn note_health( bus: &tw_observe::EventBus, health: &Arc, provider: &str, diff --git a/crates/tw-gateway/src/server/pipeline.rs b/crates/tw-gateway/src/server/pipeline.rs index bb7417e3..bee7fb95 100644 --- a/crates/tw-gateway/src/server/pipeline.rs +++ b/crates/tw-gateway/src/server/pipeline.rs @@ -27,6 +27,8 @@ mod plug; mod relay; mod slow; +pub(crate) use hop::stream_fault; + /// 256 MiB。大到能装下几张 4K 图的 base64(膨胀 33%),小到失控的 /// 客户端打不爆内存。 const MAX_BODY: usize = 256 * 1024 * 1024; diff --git a/crates/tw-gateway/src/server/pipeline/hop.rs b/crates/tw-gateway/src/server/pipeline/hop.rs index 4633a75f..db76d7fa 100644 --- a/crates/tw-gateway/src/server/pipeline/hop.rs +++ b/crates/tw-gateway/src/server/pipeline/hop.rs @@ -1706,6 +1706,24 @@ fn known_reset(state: &AppState, provider: &str, cause: Cause) -> Cause { Cause::QuotaUsedUp { resets_at_ms } } +/// 一个流式回答在第一段内容之前就收了尾,收尾的那个事件是 `data`:这算不算 `provider` 的错、 +/// 按什么原因停用。**和这里开头报错的一跳同一个判据**([`super::opening`] 读出状态码, +/// [`crate::failure::classify`] 判断):5xx、限流、额度用完、凭据被拒是 `Some`;请求本身的 +/// 问题、不是错误的是 `None` —— 那是这一家答上了。WebSocket 上 Responses 的一轮用它(见 +/// `crate::ws::turn`) +pub(crate) fn stream_fault( + state: &AppState, + provider: &str, + dialect: tw_dialect::ir::Dialect, + data: &str, +) -> Option { + let (status, body) = super::opening::stream_error(dialect, data)?; + match crate::failure::classify(status, &http::HeaderMap::new(), &body, now_ms()) { + Verdict::Failed(cause) => Some(known_reset(state, provider, cause)), + Verdict::ClientError => None, + } +} + /// 这个回答要不要等开头:生成回答的流式响应才等。等的话,上游说的是哪种格式、 /// 是不是 Bedrock 的二进制帧 fn opening_of( diff --git a/crates/tw-gateway/src/server/pipeline/opening.rs b/crates/tw-gateway/src/server/pipeline/opening.rs index 7f74e6ab..6fdcacba 100644 --- a/crates/tw-gateway/src/server/pipeline/opening.rs +++ b/crates/tw-gateway/src/server/pipeline/opening.rs @@ -251,6 +251,16 @@ fn judge(dialect: Dialect, event: Option<&str>, data: &str) -> Judge { } } +/// 流里的一个事件(`data`)是不是上游报的错:是的话,那个错误对应的状态码和上游的原话,交给 +/// [`crate::failure::classify`]。**和开头报错是同一个判据**([`judge`]):WebSocket 上 Responses +/// 的一轮在第一段内容之前收尾时按它给这一家记成败(见 `hop::stream_fault`) +pub(super) fn stream_error(dialect: Dialect, data: &str) -> Option<(u16, Bytes)> { + match judge(dialect, None, data) { + Judge::Error { status, body, .. } => Some((status, body)), + Judge::Preamble | Judge::Content => None, + } +} + fn error(status: u16, data: &str, kind: String, message: String) -> Judge { Judge::Error { status, diff --git a/crates/tw-gateway/src/server/upgrade.rs b/crates/tw-gateway/src/server/upgrade.rs index 4481508d..dc615aa3 100644 --- a/crates/tw-gateway/src/server/upgrade.rs +++ b/crates/tw-gateway/src/server/upgrade.rs @@ -214,10 +214,17 @@ pub(super) async fn ws_upgrade( ))); } let http = rt.clients.get(&name).unwrap_or(&state.http); - let upstream_headers = state - .headers_for(provider, http) - .await - .map_err(|e| GatewayError::config(crate::state::credential_failed(e, &name)))?; + // 凭据取不到是这一家的问题,和 HTTP 那条路一样记一次失败(见 `crate::health`) + let upstream_headers = match state.headers_for(provider, http).await { + Ok(h) => h, + Err(e) => { + let change = state.health.record_failure(&name); + super::note_health(&state.bus, &state.health, &name, change); + return Err(GatewayError::config(crate::state::credential_failed( + e, &name, + ))); + } + }; // Responses 的连接:每个 `response.create` 是一个请求(见 `crate::ws::turn`) let responses = crate::client_api::ClientApi::of_path(uri.path()) == Some(crate::client_api::ClientApi::OpenaiResponses) diff --git a/crates/tw-gateway/src/ws.rs b/crates/tw-gateway/src/ws.rs index e7edea17..47fab7b4 100644 --- a/crates/tw-gateway/src/ws.rs +++ b/crates/tw-gateway/src/ws.rs @@ -719,6 +719,31 @@ pub(crate) async fn proxy( (id, Some(ending), None) } }; + // 这一家接没接下这条连接,和 HTTP 那条路一跳的成败记在同一笔账上(见 `crate::health`): + // 连不上、回了 5xx 是失败,凭据被拒、限流按原因停用,请求本身的问题算这一家答上了。每一轮 + // 一行的连接接下了不记 —— 它的每一轮各记各的(见 `turn`) + let change = match &connected { + Ok(_) if turns.is_some() => None, + Ok(_) => state.health.record_success(name), + Err(NotConnected { + status: Some(s), .. + }) => match crate::failure::classify( + *s, + &axum::http::HeaderMap::new(), + &[], + crate::server::now_ms(), + ) { + crate::failure::Verdict::Failed(cause) => state.health.record_cause(name, cause), + crate::failure::Verdict::ClientError => state.health.record_success(name), + }, + // 地址、请求头写坏了:配置的事,不是这一家的 + Err(NotConnected { + source: tw_api::FailureSource::Config, + .. + }) => None, + Err(_) => state.health.record_failure(name), + }; + crate::server::note_health(&state.bus, &state.health, name, change); if ending.is_some() { // 这一跳记发出去的模型名,和 HTTP 那条路一样。一条连接跑好几轮、每一帧写的模型可能 // 不一样,升级时定得下来的只有每一帧都发的那个(指定模型、阶段一的改写,见 @@ -1292,6 +1317,10 @@ enum Flow { /// Responses 的连接上它属于上游此刻在回答的那一轮(见 [`turn`]):这一轮的结局按上游原话认 /// (用量、第一个 token、上游报的错),回答完了的那一帧交给客户端之后,这一轮收场、放掉它 /// 占着的。被工具墙切断的,这一轮记成拒绝。 +/// +/// **收了尾的那一轮又来的帧不算任何一轮的**(见 [`turn::Turns::late`]):上游先报 `error`、 +/// 再为同一次回答补一个 `response.failed` 时,后一帧照原样交给客户端,不让排在后面的那一轮 +/// 背上它的失败。被切掉的那一轮已经替它发过 `response.failed`,它补发的不再发。 async fn upstream_text( state: &AppState, p: &mut Pipes, @@ -1300,13 +1329,29 @@ async fn upstream_text( ending: &mut Option, ) -> Flow { let kind = frame_kind(t); - if let Some(turn) = p.turns.as_mut().and_then(turn::Turns::front) { - turn.upstream(t); + // 一次回答从开始到收尾的那几帧带着它的 id。只有它们要解第二遍 + let response = kind + .as_deref() + .filter(|k| lifecycle(k)) + .and_then(|_| response_id(t)); + let late = p + .turns + .as_mut() + .zip(response.as_deref().zip(kind.as_deref())) + .and_then(|(turns, (id, kind))| turns.late(id, kind)); + match late { + Some(true) => return Flow::Sent, + Some(false) => {} + None => { + if let Some(turn) = p.turns.as_mut().and_then(turn::Turns::front) { + turn.upstream(t, kind.as_deref(), response.as_deref()); + } + } } - let flow = relay(state, p, t, kind.as_deref(), c_tx, ending).await; + let flow = relay(state, p, t, kind.as_deref(), late.is_some(), c_tx, ending).await; if let Some(turns) = p.turns.as_mut() { match &flow { - Flow::Sent if ends_turn(kind.as_deref()) => turns.finish_front(), + Flow::Sent if late.is_none() && ends_turn(kind.as_deref()) => turns.finish_front(), Flow::End(End::Cut(why)) => { turns.fail_front(tw_api::FailureSource::Denied, why.clone()) } @@ -1325,12 +1370,28 @@ fn ends_turn(kind: Option<&str>) -> bool { ) } -/// [`upstream_text`] 的转发那一半。`kind` 是这一帧的 `type` +/// 一次回答从开始到收尾的那几帧(`type` 是 `kind`):它们带着这次回答(`response.id`) +fn lifecycle(kind: &str) -> bool { + matches!( + kind, + "response.created" + | "response.queued" + | "response.in_progress" + | "response.completed" + | "response.failed" + | "response.incomplete" + ) +} + +/// [`upstream_text`] 的转发那一半。`kind` 是这一帧的 `type`。`late` 是收了尾的那一次回答又来 +/// 的一帧(见 [`turn::Turns::late`]):照原样交给客户端(占位符照样还原、工具墙照样看), +/// 不碰此刻那一次回答的回答钩子和切掉的状态 async fn relay( state: &AppState, p: &mut Pipes, t: &str, kind: Option<&str>, + late: bool, c_tx: &mut ClientSink, ending: &mut Option, ) -> Flow { @@ -1350,7 +1411,10 @@ async fn relay( kind, Some("response.completed" | "response.failed" | "response.incomplete") ); - if kind == Some("response.created") { + // 收了尾的那一次回答又来的一帧不是此刻这一次的边界 + if late { + // 原样往下走 + } else if kind == Some("response.created") { p.response = response_id(&restored); p.dropping = false; // 回答钩子:这一次回答起一组实例 @@ -1364,7 +1428,7 @@ async fn relay( return Flow::Sent; } // 回答钩子:一帧可能变成几帧,也可能先扣着 - let (outgoing, failed) = match p.reply.as_mut() { + let (outgoing, failed) = match p.reply.as_mut().filter(|_| !late) { None => (vec![restored], None), Some(s) => { let (out, mut err) = s.feed(as_sse(&restored).as_bytes()).await; diff --git a/crates/tw-gateway/src/ws/turn.rs b/crates/tw-gateway/src/ws/turn.rs index c4977247..f7cf2841 100644 --- a/crates/tw-gateway/src/ws/turn.rs +++ b/crates/tw-gateway/src/ws/turn.rs @@ -16,10 +16,17 @@ //! **上限和并发按轮算,和 HTTP 那条路的一个请求同一套**([`admit`]):天、周、月用满了就拒, //! 密钥的并发上限等前面的结束,分钟、小时等得到就等;然后占这一家的一个位置 //! (`max_concurrent`)—— 这条连接只连着这一家,没有下一家可换:满着就等,等不到回 -//! `response.failed`。被拒的这一轮替它回一个 `response.failed`,连接照常。占着的(密钥的 +//! `response.failed`。等分钟、小时的空位和等这一家的位置**共用这一轮的一个期限**,和 HTTP +//! 那条路一样。被拒的这一轮替它回一个 `response.failed`,连接照常。占着的(密钥的 //! 并发通行证、这一家的位置)从发出去占到这一次回答完、或者连接断了;**闲着的连接什么都 //! 不占**。 //! +//! **快慢和成败也按轮记,和 HTTP 那条路的一跳同一笔账**:这一家的快慢样本从上游开始答这一轮 +//! 量到第一段内容([`crate::latency`],认第一段内容的是同一个),成败一轮记一次(熔断和 +//! `load-balance` 按成败分的成功率,[`crate::health`])—— 第一段内容到了是答上了;内容之前就 +//! 收了尾的,按上游最后那一帧报的错判断,和 HTTP 那条路开头报错的一跳同一个判据;上游断了是 +//! 失败。等位置等不到、被上限拒的、客户端走了的、被防护切断的不记:那不是这一家的错。 +//! //! 上限看的是这一轮开始那一刻的配置,和 HTTP 那条路每个请求一样:去向和防护按升级时的走 //! 到底,上限改了,下一轮就照新的数。 @@ -109,11 +116,16 @@ pub(crate) struct Line { pub(crate) billing: tw_config::Billing, } +/// 在 `error` 那一帧收尾的几轮,最多记住多少个回答的 id([`Turns::late`]) +const ENDED_MAX: usize = 8; + /// 这条连接上在跑的几轮,按发出去的先后。**上游按顺序回答**:上游来的帧都算头一轮的,头一轮 /// 回答完了,下一轮接上。 pub(crate) struct Turns { pub(crate) line: Arc, queue: VecDeque, + /// 最近在 `error` 那一帧收尾的几轮的回答 id,和那一轮是不是被切掉了(见 [`Self::late`]) + ended: VecDeque<(String, bool)>, } impl Turns { @@ -121,7 +133,50 @@ impl Turns { Self { line, queue: VecDeque::new(), + ended: VecDeque::new(), + } + } + + /// 上游的这一帧(`type` 是 `kind`、带着回答 `response`)是不是**已经收了尾的那一轮补发 + /// 的**,是的话那一轮是不是被切掉了。 + /// + /// 一轮在上游的 `error` 那一帧收尾,上游之后还可能为**同一次回答**补发一个 + /// `response.failed`。按「上游来的帧都算头一轮的」,那一帧会落到排在后面的那一轮头上: + /// 那一轮的结局成了上一轮的失败,它自己的回答再落到下一轮头上,一路错下去。所以带着回答 + /// id 的帧先对一下 id。认得保守 —— 认错了的那一轮等不到收尾,比记错一行更糟: + /// + /// - 只记在 `error` 那一帧收尾的几轮:在完成、失败、没答完那一帧收尾的,上游不会再补; + /// - 一次回答只补一个收尾帧,补过就忘掉它; + /// - `response.created` 和头一轮自己的回答 id 一律算头一轮的:有的上游每次回答都用同一个 id + pub(crate) fn late(&mut self, response: &str, kind: &str) -> Option { + let own = self.queue.front().and_then(|t| t.response.as_deref()) == Some(response); + if own || kind == "response.created" { + return None; + } + let at = self.ended.iter().position(|(id, _)| id == response)?; + let cut = self.ended[at].1; + if matches!( + kind, + "response.completed" | "response.failed" | "response.incomplete" + ) { + self.ended.remove(at); + } + Some(cut) + } + + /// 头一轮收尾了,下一轮接上。在 `error` 那一帧收尾的记下它的回答 id(见 [`Self::late`]) + fn pop_front(&mut self) -> Option { + let t = self.queue.pop_front()?; + if let (true, Some(id)) = (t.errored, t.response.clone()) { + if self.ended.len() == ENDED_MAX { + self.ended.pop_front(); + } + self.ended.push_back((id, t.cut.is_some())); + } + if let Some(next) = self.queue.front_mut() { + next.begin(); } + Some(t) } /// 上游此刻在回答的那一轮 @@ -133,21 +188,24 @@ impl Turns { self.queue.front().map(|t| t.id) } - /// 发出去了:排在后面等上游回答 - pub(crate) fn push(&mut self, t: Turn) { + /// 发出去了:排在后面等上游回答。前面没有在答的,上游这就开始答它 + pub(crate) fn push(&mut self, mut t: Turn) { + if self.queue.is_empty() { + t.begin(); + } self.queue.push_back(t); } /// 头一轮回答完了:报结局,放掉它占着的 pub(crate) fn finish_front(&mut self) { - if let Some(t) = self.queue.pop_front() { + if let Some(t) = self.pop_front() { t.finish(); } } /// 头一轮失败了:被防护切断了 pub(crate) fn fail_front(&mut self, source: tw_api::FailureSource, why: Msg) { - if let Some(t) = self.queue.pop_front() { + if let Some(t) = self.pop_front() { t.fail(source, why); } } @@ -187,7 +245,17 @@ pub(crate) struct Turn { pub(crate) id: u64, /// 报了就没有了 ending: Option, - bus: tw_observe::EventBus, + state: AppState, + /// 这条连接连着的那一家:快慢和成败记在它头上 + provider: String, + /// 上游这一次回答的 id(`response.created` 里的)。收尾之后靠它认出上游补发的帧(见 + /// [`Turns::late`]) + response: Option, + /// 上游这一轮来过一个 `error`:在它那一帧收尾的,上游之后还可能为同一次回答补一个收尾帧 + errored: bool, + /// 这一轮的成败给这一家记过了(见 [`Self::judge`])。**一轮只记一次**,和 HTTP 那条路一跳 + /// 只记一次一样 + judged: bool, /// 这一帧到的那一刻:首字节时间从它算,和 HTTP 那条路从请求进来算一样 started: Instant, /// 上游这一轮的第一帧到了没有。到了报响应头 @@ -210,13 +278,31 @@ impl Turn { } } - /// 上游这一轮来了一帧,**上游原话**(带占位符的那一版):头一帧报响应头,每一帧喂给结局 - /// 认用量、第一个 token、上游报的错 - pub(crate) fn upstream(&mut self, text: &str) { + /// 上游开始答这一轮了:发出去时前面没有在答的,或者前面那一轮刚收尾。这一家的快慢样本 + /// 从这一刻量到第一段内容(见 [`crate::latency`])—— 前一轮还在答的时候,这一轮排在后面 + /// 等,那一段不是这一家慢 + fn begin(&mut self) { + if let Some(e) = self.ending.as_mut() { + e.timed(crate::ending::Lap { + latency: self.state.latency.clone(), + provider: self.provider.clone(), + sent: Instant::now(), + }); + } + } + + /// 上游这一轮来了一帧,**上游原话**(带占位符的那一版),`type` 是 `kind`、带着的回答 id + /// 是 `response`:头一帧报响应头,每一帧喂给结局认用量、第一个 token、上游报的错。成败 + /// 在这里记:第一段内容到了是答上了;内容之前就收了尾的,看收尾的这一帧报没报错 + pub(crate) fn upstream(&mut self, text: &str, kind: Option<&str>, response: Option<&str>) { + if self.response.is_none() { + self.response = response.map(str::to_string); + } + self.errored |= kind == Some("error"); if !self.responded { self.responded = true; self.routed(); - self.bus.emit(tw_api::Event::RequestHeaders { + self.state.bus.emit(tw_api::Event::RequestHeaders { id: self.id, status: 200, ttfb_ms: self.started.elapsed().as_millis() as u64, @@ -228,13 +314,43 @@ impl Turn { if let Some(e) = self.ending.as_mut() { e.frame(text); } + if self.judged { + return; + } + if self.ending.as_ref().is_some_and(Ending::has_content) { + // 内容到了:这一家答上了。之后流里再报错也不改 —— HTTP 那条路的一跳也是开头 + // 一过就记成功 + self.judge(|h, p| h.record_success(p)); + } else if super::ends_turn(kind) { + // 一段内容都没有就收了尾:和 HTTP 那条路开头报错的一跳同一个判据 + let dialect = tw_dialect::ir::Dialect::Responses; + match crate::server::stream_fault(&self.state, &self.provider, dialect, text) { + Some(cause) => self.judge(|h, p| h.record_cause(p, cause)), + None => self.judge(|h, p| h.record_success(p)), + } + } + } + + /// 给这一家记这一轮的成败(见 [`crate::health`]):熔断和按成败分的 `load-balance` 看的是 + /// 同一笔账,状态变了照常报。**一轮只记一次** + fn judge( + &mut self, + record: impl FnOnce(&crate::health::Health, &str) -> Option, + ) { + if std::mem::replace(&mut self.judged, true) { + return; + } + let health = &self.state.health; + let change = record(health, &self.provider); + crate::server::note_health(&self.state.bus, health, &self.provider, change); } fn routed(&mut self) { let Some(r) = self.route.take() else { return }; let mut attempt = r.attempt; attempt.ms = r.since.elapsed().as_millis() as u64; - self.bus + self.state + .bus .emit(routed(self.id, r.choice, vec![attempt], r.billing)); } @@ -249,8 +365,12 @@ impl Turn { } } - /// 失败了:上游断了、被防护切断了。用量照样带着(上游已经计了费) + /// 失败了:上游断了、被防护切断了。用量照样带着(上游已经计了费)。上游断了的,没答上 + /// 之前断的给这一家记一次失败,和 HTTP 那条路流在第一段内容之前断了一样 pub(crate) fn fail(mut self, source: tw_api::FailureSource, why: Msg) { + if source == tw_api::FailureSource::Upstream { + self.judge(|h, p| h.record_failure(p)); + } self.routed(); if let Some(e) = self.ending.take() { e.failed(source, why); @@ -258,15 +378,19 @@ impl Turn { } /// 准入过了、这一帧却没发出去:插件拒绝了它,或者写不过去。尝试链上是 `attempts`(写不 - /// 过去的那一跳;插件拒绝的没有),没接下的不按那一家记账 + /// 过去的那一跳;插件拒绝的没有),没接下的不按那一家记账。写不过去是上游断了,给这一家 + /// 记一次失败,和 HTTP 那条路发不出去一样 pub(crate) fn unsent( mut self, attempts: Vec, source: tw_api::FailureSource, why: Msg, ) { + if source == tw_api::FailureSource::Upstream { + self.judge(|h, p| h.record_failure(p)); + } if let Some(r) = self.route.take() { - self.bus.emit(routed( + self.state.bus.emit(routed( self.id, r.choice, attempts, @@ -340,12 +464,15 @@ pub(crate) enum NotAdmitted { /// /// 1. 天、周、月的上限 —— 用满了就拒; /// 2. 密钥的并发上限 —— 等前面的结束; -/// 3. 分钟、小时的上限 —— 下一个空位在 `slot_wait_secs` 之内空出来就等,等不到就拒;过了 +/// 3. 分钟、小时的上限 —— 下一个空位在这一轮的等待期限之前空出来就等,等不到就拒;过了 /// 就按输入的估算占着,存储层记下这一行时换成实数; /// 4. 开始:发开始事件,内容过滤的结论挂在这一轮上报(拒绝的切断连接); -/// 5. 这一家的位置(`max_concurrent`)—— 满着就等,最多 `slot_wait_secs`;这条连接只连着 +/// 5. 这一家的位置(`max_concurrent`)—— 满着就等,等到同一个期限;这条连接只连着 /// 这一家,等不到就是这一家忙,回 429 那句话。 /// +/// **等待期限一轮只有一个**:过了密钥的并发上限那一刻起算 `failover.slot_wait_secs`,第 3 步、 +/// 第 5 步的等待都算在里面,和 HTTP 那条路一样(见 `server::pipeline::admission`)。 +/// /// 被上限拒的、等不到位置的**照样留一行**,流量里看得见它为什么没发出去。拿到的通行证和 /// 位置交给这一轮([`Turn`]),跟着它走。 pub(crate) async fn admit(a: Admit) -> Result { @@ -356,8 +483,6 @@ pub(crate) async fn admit(a: Admit) -> Result { let key = rt.config.clients.iter().find(|c| c.name == key_name); let max_concurrent = key.and_then(|c| c.max_concurrent); let limits: &[tw_config::KeyLimit] = key.map(|c| c.limits.as_slice()).unwrap_or_default(); - // 等滚动窗口的空位、等这一家的位置,各等最多这么久(见 `crate::key_limits::slot_wait`) - let wait = crate::key_limits::slot_wait(&rt.config); let choice = Choice { rewritten_by: a.rewritten_by.clone(), ..line.choice.clone() @@ -366,8 +491,15 @@ pub(crate) async fn admit(a: Admit) -> Result { return Err(refused(&a, &choice, r.error())); } let pass = state.gate.acquire(key_name, max_concurrent).await; + // 等待期限从这里起算:等滚动窗口的空位、等这一家的位置共用它(见 + // `crate::key_limits::slot_wait`) + let wait_until = tokio::time::Instant::now() + crate::key_limits::slot_wait(&rt.config); let ask = ask(state, &a, limits); - let hold = match state.key_limits.admit(key_name, limits, ask, wait).await { + let hold = match state + .key_limits + .admit_by(key_name, limits, ask, wait_until) + .await + { Ok(hold) => hold, Err(r) => return Err(refused(&a, &choice, r.error())), }; @@ -394,11 +526,9 @@ pub(crate) async fn admit(a: Admit) -> Result { None => { let mut waited = None; let mut slot = None; - if !wait.is_zero() { - slot = state - .slots - .take_by(&line.provider, tokio::time::Instant::now() + wait) - .await; + // 这一轮能等的已经等完了(或者配置的是不等):不再等 + if tokio::time::Instant::now() < wait_until { + slot = state.slots.take_by(&line.provider, wait_until).await; waited = Some(hop_started.elapsed().as_millis() as u64); } match slot { @@ -433,7 +563,11 @@ pub(crate) async fn admit(a: Admit) -> Result { Ok(Turn { id, ending: Some(ending), - bus: state.bus.clone(), + state: state.clone(), + provider: line.provider.clone(), + response: None, + errored: false, + judged: false, started: a.arrived, responded: false, route: Some(Route { diff --git a/crates/tw-gateway/tests/upstream_slots.rs b/crates/tw-gateway/tests/upstream_slots.rs index fc8ccfa2..8222ceb0 100644 --- a/crates/tw-gateway/tests/upstream_slots.rs +++ b/crates/tw-gateway/tests/upstream_slots.rs @@ -87,7 +87,11 @@ async fn upstream() -> Up { } } -/// 甲、乙两家按顺序(`fallback`)。`limits` 是两家各自的上限 +/// 甲、乙两家按顺序(`fallback`)的组「池」,默认路由「默认」把请求都交给它。`limits` 是两家 +/// 各自的上限。 +/// +/// **`default_route` 要写**:不写的话默认路由叫 `default`,「默认」这条路由谁也不走,请求走的 +/// 是内置的「全部上游」组 —— 测的就不是这个组了([`Log::routed`] 会说出来) fn config(a: &Up, b: &Up, limits: (Option, Option), wait_secs: u64) -> tw_config::Config { let limit = |n: Option| { n.map(|n| format!(" max_concurrent: {n}\n")) @@ -95,6 +99,7 @@ fn config(a: &Up, b: &Up, limits: (Option, Option), wait_secs: u64) -> }; let yaml = format!( "version: 1 +default_route: 默认 clients: - name: me key: tw-k @@ -153,21 +158,36 @@ impl Log { .await } - /// 那个请求的尝试链和对话留在哪一家的说明 + /// 那个请求的尝试链和对话留在哪一家的说明。**它走的是「默认」路由上的组「池」**(见 + /// [`config`]):走到内置的「全部上游」组上的话,测的就不是这个组了 async fn routed(&self, model: &str) -> (Vec, Option) { let id = self.id(model).await; - self.until(&format!(" {model} 的路由"), |evs| { - evs.iter().find_map(|e| match e { - Event::RequestRouted { - id: i, - attempts, - affinity, - .. - } if *i == id => Some((attempts.clone(), affinity.as_ref().and_then(|a| a.stayed))), - _ => None, + let (route, group, attempts, stayed) = self + .until(&format!(" {model} 的路由"), |evs| { + evs.iter().find_map(|e| match e { + Event::RequestRouted { + id: i, + route, + group, + attempts, + affinity, + .. + } if *i == id => Some(( + route.clone(), + group.clone(), + attempts.clone(), + affinity.as_ref().and_then(|a| a.stayed), + )), + _ => None, + }) }) - }) - .await + .await; + assert_eq!( + (route.as_str(), group.as_deref()), + ("默认", Some("池")), + "{model} 没走配置里的那个组" + ); + (attempts, stayed) } async fn finished(&self, model: &str) { @@ -627,7 +647,6 @@ async fn a_full_member_of_a_load_balance_group_sits_its_turns_out() { let (a, b) = (upstream().await, upstream().await); let mut cfg = config(&a, &b, (Some(1), None), 5); cfg.groups[0].kind = tw_engine::GroupType::LoadBalance; - cfg.default_route = Some("默认".into()); let (gw, state, log) = serve(cfg).await; let group = state.runtime().engine.groups()[0].clone(); // 头一个轮到甲,占着它的那个位置 diff --git a/crates/tw-gateway/tests/ws.rs b/crates/tw-gateway/tests/ws.rs index c649d2fe..72ab7285 100644 --- a/crates/tw-gateway/tests/ws.rs +++ b/crates/tw-gateway/tests/ws.rs @@ -1302,3 +1302,295 @@ async fn a_turn_waits_for_a_full_upstream_and_fails_as_busy_when_none_frees() { let frames = one_answer(&mut c).await; assert_eq!(frames.last().unwrap()["type"], "response.completed"); } + +// ---------------------------------------------------------------- 快慢和成败 + +/// 一轮的剧本:收到第 n 个 `response.create`(从 1 数)时回的每一帧,和发它之前先等的毫秒数 +type Script = fn(usize) -> Vec<(u64, serde_json::Value)>; + +/// 照剧本回答的 Responses 上游 +async fn scripted(script: Script) -> SocketAddr { + let app = Router::new().route( + "/backend-api/codex/responses", + axum::routing::any(move |ws: WebSocketUpgrade| async move { + ws.on_upgrade(move |mut sock| async move { + let mut n = 0; + while let Some(Ok(m)) = sock.recv().await { + if !matches!(m, Message::Text(_)) { + continue; + } + n += 1; + for (wait, f) in script(n) { + tokio::time::sleep(Duration::from_millis(wait)).await; + if sock + .send(Message::Text(f.to_string().into())) + .await + .is_err() + { + return; + } + } + } + }) + }), + ); + let l = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap(); + let addr = l.local_addr().unwrap(); + tokio::spawn(async move { axum::serve(l, app).await.unwrap() }); + addr +} + +fn created(id: &str) -> serde_json::Value { + serde_json::json!({"type":"response.created","response":{"id":id,"status":"in_progress","model":"gpt-5","output":[]}}) +} + +fn delta() -> serde_json::Value { + serde_json::json!({"type":"response.output_text.delta","output_index":0,"content_index":0,"item_id":"msg","delta":"hello"}) +} + +fn completed(id: &str) -> serde_json::Value { + let usage: serde_json::Value = serde_json::from_str(USAGE).unwrap(); + serde_json::json!({"type":"response.completed","response":{"id":id,"status":"completed","model":"gpt-5","output":[],"usage":usage}}) +} + +fn failed(id: &str, code: &str) -> serde_json::Value { + serde_json::json!({"type":"response.failed","response":{"id":id,"status":"failed","model":"gpt-5","output":[],"error":{"code":code,"message":"boom"}}}) +} + +fn error(code: &str) -> serde_json::Value { + serde_json::json!({"type":"error","error":{"type":code,"code":code,"message":"boom"}}) +} + +/// 收到这一轮收尾的那一帧为止(完成、失败、没答完,或者一个 `error`) +async fn until_end(c: &mut Socket) -> Vec { + let mut out = Vec::new(); + loop { + let m = tokio::time::timeout(Duration::from_secs(5), c.next()) + .await + .unwrap_or_else(|_| panic!("the answer did not end: {out:?}")) + .expect("the connection closed") + .unwrap(); + let tokio_tungstenite::tungstenite::Message::Text(t) = m else { + continue; + }; + let v: serde_json::Value = serde_json::from_str(&t).unwrap(); + let end = matches!( + v["type"].as_str(), + Some("response.completed" | "response.failed" | "response.incomplete" | "error") + ); + out.push(v); + if end { + return out; + } + } +} + +/// **每一轮和 HTTP 那条路的一跳记同一笔账**:快慢样本从上游开始答这一轮量到第一段内容;成败 +/// 一轮记一次 —— 内容到了是答上了(之后流里再报错也是),内容之前就收了尾的看收尾那一帧报的 +/// 错(上游的错是失败,请求本身的问题算答上了),客户端没等到内容就走了的不记 +#[tokio::test] +async fn each_turn_feeds_the_upstreams_speed_and_success_like_an_http_hop() { + fn script(n: usize) -> Vec<(u64, serde_json::Value)> { + let id = format!("resp_{n}"); + match n { + // 头三轮的第一段内容 200 毫秒之后才到 + 1..=3 => vec![(0, created(&id)), (200, delta()), (0, completed(&id))], + 4 | 5 => vec![(0, created(&id)), (0, delta()), (0, completed(&id))], + // 内容之前就失败了:这一家的错 + 6 => vec![(0, created(&id)), (0, failed(&id, "server_error"))], + // 内容到了之后才失败:和 HTTP 那条路一样算答上了 + 7 => vec![ + (0, created(&id)), + (0, delta()), + (0, failed(&id, "server_error")), + ], + // 请求本身的问题 + 8 => vec![(0, error("invalid_request_error"))], + // 客户端等不到内容就走了 + _ => vec![(0, created(&id)), (5_000, delta())], + } + } + let up = scripted(script).await; + let (gw, mut rx, state) = turns_gateway(up, |_| {}, None).await; + let rate = || { + state + .health + .success_rates(&["up".to_string()]) + .get("up") + .copied() + }; + let mut c = connect(gw).await; + for _ in 1..=5 { + c.send(create("hi")).await.unwrap(); + assert_eq!( + until_end(&mut c).await.last().unwrap()["type"], + "response.completed" + ); + } + let typical = state.latency.typical("up").expect("每一轮都该量到"); + assert!( + (200..2_000).contains(&typical), + "从这一轮开始到第一段内容是 200 多毫秒,量到的是 {typical}" + ); + assert_eq!(rate(), Some(1.0)); + + c.send(create("hi")).await.unwrap(); + until_end(&mut c).await; + assert_eq!( + rate(), + Some(5.0 / 6.0), + "内容之前的 server_error 是这一家的错" + ); + c.send(create("hi")).await.unwrap(); + until_end(&mut c).await; + assert_eq!(rate(), Some(6.0 / 7.0), "内容到了之后的失败不算"); + c.send(create("hi")).await.unwrap(); + until_end(&mut c).await; + assert_eq!(rate(), Some(7.0 / 8.0), "请求本身的问题算这一家答上了"); + + // 客户端没等到内容就走了:这一轮取消,不记 + c.send(create("hi")).await.unwrap(); + loop { + let m = c.next().await.unwrap().unwrap(); + if m.into_text().unwrap().contains("response.created") { + break; + } + } + drop(c); + loop { + let ev = tokio::time::timeout(Duration::from_secs(5), rx.recv()) + .await + .expect("the turn was not cancelled") + .unwrap(); + if matches!(ev, Event::RequestCancelled { .. }) { + break; + } + } + assert_eq!(rate(), Some(7.0 / 8.0), "客户端走了被算进了成败"); +} + +/// 一轮在上游的 `error` 那一帧收尾,上游又为**同一次回答**补发一个 `response.failed`:那一帧 +/// 照原样交给客户端,**不算排在后面的那一轮的** —— 下一轮照样有自己的回答、用量和成败。 +/// 每次回答都用同一个 id 的上游也照常:下一轮自己的回答不会被当成补发的 +#[tokio::test] +async fn a_late_frame_for_an_answer_that_ended_is_not_the_next_turns() { + fn script(n: usize) -> Vec<(u64, serde_json::Value)> { + match n { + 1 => vec![(0, created("resp_1")), (0, error("server_error"))], + 2 => vec![ + (0, failed("resp_1", "server_error")), + (0, created("resp_2")), + (0, delta()), + (0, completed("resp_2")), + ], + 3 => vec![(0, created("same")), (0, error("server_error"))], + _ => vec![(0, created("same")), (0, delta()), (0, completed("same"))], + } + } + let up = scripted(script).await; + let (gw, mut rx, state) = turns_gateway(up, |c| c.failover.failures_to_pause = 1, None).await; + let mut c = connect(gw).await; + c.send(create("one")).await.unwrap(); + assert_eq!(until_end(&mut c).await.last().unwrap()["type"], "error"); + assert_eq!( + state.health.state("up"), + tw_gateway::health::State::Open, + "内容之前的 server_error 是这一家的错" + ); + + c.send(create("two")).await.unwrap(); + // 补发的那一帧照原样到了客户端,然后是第二轮自己的回答 + let late = until_end(&mut c).await; + assert_eq!(late.last().unwrap()["response"]["id"], "resp_1", "{late:?}"); + let second = until_end(&mut c).await; + assert_eq!(second.last().unwrap()["type"], "response.completed"); + + let evs = requests(&mut rx, 2).await; + let ids = started_ids(&evs); + assert_eq!(ids.len(), 2, "{evs:?}"); + assert!( + matches!(of(&evs, ids[0]).last(), Some(Event::RequestFailed { source, .. }) if source.slug() == "upstream"), + "{evs:?}" + ); + match of(&evs, ids[1]).last() { + Some(Event::RequestFinished { usage: Some(u), .. }) => { + assert_eq!((u.input, u.cache_read, u.output), (200, 1000, 30)); + } + other => panic!("第二轮背上了第一轮的失败:{other:?}"), + } + assert_eq!( + state.health.state("up"), + tw_gateway::health::State::Closed, + "第二轮答上了,补发的那一帧不该记到它头上" + ); + + // 第三轮在 `error` 收尾,第四轮的回答用的还是那个 id:它是第四轮自己的 + c.send(create("three")).await.unwrap(); + assert_eq!(until_end(&mut c).await.last().unwrap()["type"], "error"); + c.send(create("four")).await.unwrap(); + let fourth = until_end(&mut c).await; + assert_eq!( + fourth.last().unwrap()["type"], + "response.completed", + "{fourth:?}" + ); + let evs = requests(&mut rx, 2).await; + let ids = started_ids(&evs); + assert!( + matches!( + of(&evs, ids[1]).last(), + Some(Event::RequestFinished { usage: Some(_), .. }) + ), + "{evs:?}" + ); + assert_eq!(state.health.state("up"), tw_gateway::health::State::Closed); +} + +/// 一轮只有一段可等的时间,和 HTTP 那条路的一个请求一样:等密钥的分钟上限用掉的,等这一家的 +/// 位置时就少等那么久。两段各给一份的话,这一轮要等两倍那么久才收到「忙」 +#[tokio::test] +async fn a_turns_key_limit_wait_and_slot_wait_share_one_budget() { + let (up, _seen) = responder(None).await; + let clock = JumpClock::new(); + let (gw, mut rx, state) = turns_gateway( + up, + |c| { + c.clients[0].limits = serde_yaml_ng::from_str("[{per: minute, requests: 1}]").unwrap(); + c.providers[0].max_concurrent = Some(1); + c.failover.slot_wait_secs = 2; + }, + Some(clock.clone()), + ) + .await; + let mut c = connect(gw).await; + c.send(create("one")).await.unwrap(); + one_answer(&mut c).await; + requests(&mut rx, 1).await; + // 这一分钟的一个用掉了,再过 1 秒滑出去;这一家的位置一直有人占着 + clock.jump(59_000); + let _other = state.slots.try_take("up").unwrap(); + + let t = std::time::Instant::now(); + c.send(create("two")).await.unwrap(); + let frames = one_answer(&mut c).await; + let took = t.elapsed(); + assert_eq!( + frames.last().unwrap()["type"], + "response.failed", + "{frames:?}" + ); + let evs = requests(&mut rx, 1).await; + match evs.last().unwrap() { + Event::RequestFailed { message, .. } => { + assert_eq!( + message.code, "gw.busy_all", + "该是等过了分钟上限、再等这一家的位置" + ) + } + other => panic!("{other:?}"), + } + assert!( + took >= Duration::from_millis(1_800) && took < Duration::from_millis(2_600), + "等了 {took:?}:两段该共用 2 秒" + ); +} diff --git a/docs/config.md b/docs/config.md index a0d049fb..7184ccc8 100644 --- a/docs/config.md +++ b/docs/config.md @@ -987,6 +987,10 @@ go straight to the next candidate. How long depends on the reason the upstream gives: an insufficient balance waits for a top-up, a used-up quota waits until the moment the upstream says it resets, and a rate limit usually passes within seconds. A request with a single candidate is never affected. +On a Responses WebSocket connection, each `response.create` counts here like +a request: one that fails because of the upstream before any content arrives +counts as a failure, and so does a connection the upstream refuses or that +cannot be made. Before the first content of a streamed answer reaches the client, an error the upstream sends in the stream moves the request to the next candidate, @@ -1151,7 +1155,10 @@ group shares out requests by the result in the same way as above. Speed is measured on streamed answers only, from the moment the request is sent to that upstream, so waiting and upstreams that failed before it do not count. An upstream given up on because its stream was slow to start -(`failover.next_on_slow_start`) counts as taking the whole wait. +(`failover.next_on_slow_start`) counts as taking the whole wait. On a +Responses WebSocket connection, each `response.create` counts as one request +for both speed and failures, its speed measured from the moment the upstream +starts answering it. An upstream without enough measurements yet counts as average. As with weights alone, conversations in progress stay where they are, and new diff --git a/docs/config.zh-CN.md b/docs/config.zh-CN.md index 7e0e5ced..ff7cab6f 100644 --- a/docs/config.zh-CN.md +++ b/docs/config.zh-CN.md @@ -769,7 +769,9 @@ security: 上游失败后会停用一段时间,接下来的请求直接交给下一个候选。停用多久取决于上游 给出的原因:余额不足要等充值,额度用完要等到上游说的重置时刻,限流通常几秒钟就 -过去。只有一个候选的请求不受影响。 +过去。只有一个候选的请求不受影响。Responses 的 WebSocket 连接上,每个 `response.create` +在这里和一个请求一样算:第一段内容到达之前因为上游出错而失败的,算作一次失败;上游 +拒绝或连不上的连接也一样。 流式回答在第一段内容交给客户端之前,上游在流里报的错误和错误状态码一样,会把 请求换到下一个候选。 @@ -875,7 +877,7 @@ groups: - `health`:越少失败的上游分得越多。依据是最近 30 分钟内的最近 50 次请求:服务器错误、限流、额度或余额用尽、凭据被拒、超时和连接失败算作失败;请求本身导致的错误不算,客户端取消、因开头太慢而换走、因并发数满了(`max_concurrent`)而跳过也不算。经常失败的上游至少保留权重的二十分之一,仍会偶尔分到新对话,以便发现它已经恢复;完全失败的上游照旧由 [`failover`](#cfg-failover) 暂停。 - `latency-health`:两个系数相乘。 -快慢只在流式回答上测,从请求发给这家上游的那一刻算起:之前的等待、之前失败的上游都不算在内。因开头太慢而被放弃的上游(`failover.next_on_slow_start`),按等满的那段时间计。 +快慢只在流式回答上测,从请求发给这家上游的那一刻算起:之前的等待、之前失败的上游都不算在内。因开头太慢而被放弃的上游(`failover.next_on_slow_start`),按等满的那段时间计。Responses 的 WebSocket 连接上,每个 `response.create` 在快慢和成败上都算一个请求,快慢从上游开始回答它的那一刻算起。 测量还不够的上游按中等对待。与只按权重时一样,进行中的对话留在原来的上游,差额由新对话补齐。 From 0620f53d0b178a9eadec511e3148f54dfef0464a Mon Sep 17 00:00:00 2001 From: fylorn <249551762+fylorn@users.noreply.github.com> Date: Mon, 5 Oct 2026 18:40:31 +0800 Subject: [PATCH 11/22] Control-plane protocol 39 for the routing round The routing round adds view, input and event fields across the control plane: upstream max_concurrent, failover slot_wait_secs and next_on_slow_start, attempt queued_ms / skipped / usage and the slow_start outcome, load-balance weights and balance_by with the dry-run factors, manual model specs and PUT /provider-model-spec, key usage limits with KeyLimitAlert, and one row per Responses WebSocket turn. A UI written for 38 would drop group weights and key limits when it saves, so the version moves to 39, with the notes above the constant describing it. Co-Authored-By: Claude Opus 5.5 --- crates/tw-api/src/lib.rs | 19 ++++++++++++++++++- 1 file changed, 18 insertions(+), 1 deletion(-) diff --git a/crates/tw-api/src/lib.rs b/crates/tw-api/src/lib.rs index bcbacaea..3037da0b 100644 --- a/crates/tw-api/src/lib.rs +++ b/crates/tw-api/src/lib.rs @@ -777,7 +777,24 @@ pub const MSG_CODES: &str = include_str!("../msg-codes.txt"); /// `preview` 的上游和叫 `test` 的代理改和删都是 405;插件 id 不再保留 `order`、`inspect`、 /// `rewrite`、`confirmed` 这几个词,消息码 `config.plugin.reserved_id`、 /// `control.plugin.reserved_id` 跟着删。 -pub const CONTROL_API_VERSION: u32 = 38; +/// +/// **39 起路由看得到上游忙不忙、快不快、可不可靠,密钥有用量上限**。上游能设并发上限 +/// ([`ProviderView`]、[`ProviderInput`] 的 `max_concurrent`);[`FailoverView`] 多了 +/// `slot_wait_secs`(一个请求合计最多等多久,等密钥的分钟、小时上限也算在里面)和 +/// `next_on_slow_start`(流开头太慢就换下一家)。尝试链([`AttemptView`])多了 `queued_ms`、 +/// `skipped`([`ServeSkip`] 的 `busy`)和 `usage`([`AttemptUsage`]:开头慢被放弃的那一跳上游 +/// 可能已经收了钱的输入),结果多了 `slow_start`([`AttemptOutcome`])。`load-balance` 组的 +/// 成员有权重、能按快慢和成败分([`GroupView`]、[`GroupInput`] 的 `weights`、`balance_by`, +/// [`BalanceBy`]),试算说得出每个成员的权重、快慢、成功率和系数([`DryRunCandidate`]、 +/// [`DryRunResult::balance_by`])。一家上游的模型规格可以手动设:新端点 +/// `PUT /provider-model-spec`([`ModelSpecSave`] → [`ConfigWritten`]),[`ModelRow`] 多了 +/// `context_window_source`、`max_output_tokens`、`max_output_tokens_source`([`SpecSource`])。 +/// 密钥的用量上限:[`ClientView`] 多了 `limits`([`KeyLimitView`])和 `unpriced_models`, +/// [`KeyInput`] 多了 `limits`([`KeyLimitInput`]),事件多了 [`Event::KeyLimitAlert`]。 +/// Responses 的 WebSocket 连接上每个 `response.create` 是一行请求(带用量和费用,按轮算上限、 +/// 并发、快慢和成败),连接本身不留行;Realtime 和别的路径照旧整条连接一行。照 38 写的界面 +/// 保存组和密钥时会把权重、上限丢掉。 +pub const CONTROL_API_VERSION: u32 = 39; #[derive(Debug, Clone, Serialize, Deserialize)] #[cfg_attr(feature = "ts", derive(ts_rs::TS))] From 07a94ea2c1cc427178f51d7013b30dc85ca9b11a Mon Sep 17 00:00:00 2001 From: fylorn <249551762+fylorn@users.noreply.github.com> Date: Mon, 5 Oct 2026 19:29:29 +0800 Subject: [PATCH 12/22] Keep a working upstream when the only next ones are full Two ways a full upstream (max_concurrent) made failover worse than no failover at all: - The slow-start switch counted a full candidate as "a later one that can take the request". The request gave up on an upstream that was about to answer, then waited for the full one; with the recommended settings the shared wait deadline had already passed, so the client got an immediate 429. At the deadline only candidates with a free slot now count; otherwise the slow upstream is treated as the last one and keeps streaming. - When the wait for full candidates ended without a slot, the busy 429 won over a real failure: an upstream that answered 401 was hidden behind "every upstream is busy", with x-should-retry: true. The busy 429 is now only for requests that never reached an upstream; otherwise the last real failure is answered, as it is when no candidate is full. Errors that failover cannot fix (request errors) were already passed through at once. Co-Authored-By: Claude Opus 5.5 --- crates/tw-gateway/src/server/pipeline/hop.rs | 40 ++++-- crates/tw-gateway/src/server/pipeline/slow.rs | 6 +- crates/tw-gateway/tests/slow_start.rs | 124 ++++++++++++++++++ crates/tw-gateway/tests/upstream_slots.rs | 96 +++++++++++++- docs/config.md | 7 +- docs/config.zh-CN.md | 6 +- 6 files changed, 260 insertions(+), 19 deletions(-) diff --git a/crates/tw-gateway/src/server/pipeline/hop.rs b/crates/tw-gateway/src/server/pipeline/hop.rs index db76d7fa..ece5c1de 100644 --- a/crates/tw-gateway/src/server/pipeline/hop.rs +++ b/crates/tw-gateway/src/server/pipeline/hop.rs @@ -156,6 +156,9 @@ pub(super) async fn try_upstreams<'a>( let mut queued: Option<(usize, u64)> = None; // 等到最后,剩下的候选还都满着 let mut stalled = false; + // 有一跳真的发到了上游、失败了(回了错误、流在内容之前断了、连不上、开头慢被放弃)。 + // 等满着的那几家等到最后的话,交出去的是这一家的错,不是「都满着」(见下面的 `stalled`) + let mut reached = false; loop { stamp_queued(&mut chain, queued.take()); @@ -505,10 +508,15 @@ pub(super) async fn try_upstreams<'a>( .clone() .unwrap_or_else(|| reading.facts.model.clone()); let (attempt, bridge) = (chain.len(), plugged.bridge); - // 后面还有没有接得下这个请求的(见 `successor`)。开头慢了才问。后面的是还没看的 - // 候选,加上满着、跳过了的那几家:它们还可能空出来(见 `crate::slots`) + // 后面还有没有此刻接得下这个请求的(见 `successor`)。开头慢了才问。后面的是还没看的 + // 候选,加上满着、跳过了的那几家:**此刻有空位的才算**(见 `crate::slots`)—— 满着的 + // 那一家要等,而这个请求能等的多半已经等完了,放弃了这一家,换来的是一个 429 let others = || { - let rest = queue.iter().chain(busy.iter()).copied(); + let rest = queue + .iter() + .chain(busy.iter()) + .copied() + .filter(|n| !state.slots.is_full(n)); successor(state, rt, req, reading, decision, &catalog, allow, rest) }; // 这一跳发出去的那一刻:这一家的快慢样本从这里算起(见 `crate::latency`) @@ -557,6 +565,7 @@ pub(super) async fn try_upstreams<'a>( // 响应头都还没来,后面又有接得下的:放弃这一家。不停用、不算失败(见 // `super::slow`) Err(_) if others() => { + reached = true; let waited = slow_wait.unwrap_or_default(); super::slow::timed_out(state, &provider.name, waited); chain.push(super::slow::abandoned( @@ -616,7 +625,9 @@ pub(super) async fn try_upstreams<'a>( match verdict { Verdict::Failed(cause) if !hand_on => { // 5xx、限流、没钱了、额度用完、凭据被拒、没有这个模型:换一家有 - // 意义,那边是另一把密钥、另一个账户。停用多久看原因。 + // 意义,那边是另一把密钥、另一个账户。停用多久看原因。后面只剩满着 + // 的那几家也一样等它们:空出来的那一家可能答得上 + reached = true; // // 交给客户端的那一跳,下面这几样由回程(`relay`)去记 // @@ -719,6 +730,7 @@ pub(super) async fn try_upstreams<'a>( { let status = response.status().as_u16(); drop(response); + reached = true; // 上游回了话,说明代理是通的 state.note_proxy_ok(&provider.proxy); let waited = slow_wait.unwrap_or_default(); @@ -751,6 +763,7 @@ pub(super) async fn try_upstreams<'a>( match crate::failure::classify(status, &headers, &body, now_ms()) { Verdict::ClientError => response, Verdict::Failed(cause) => { + reached = true; state.note_quota(id, &provider.name, &headers); state.note_proxy_ok(&provider.proxy); let cause = known_reset(state, &provider.name, cause); @@ -783,6 +796,7 @@ pub(super) async fn try_upstreams<'a>( } } super::opening::Opening::Broken(err) => { + reached = true; note_health( &state.bus, &state.health, @@ -837,6 +851,7 @@ pub(super) async fn try_upstreams<'a>( ); let err = match e { SendError::Http(e) => { + reached = true; // 连不上的可能是代理而不是上游 —— 检一次那个代理,说清是哪一件事 state.check_proxy(&provider.proxy); forward::map_reqwest_error(e) @@ -918,9 +933,13 @@ pub(super) async fn try_upstreams<'a>( )))); } - // 剩下的候选等过了还都满着:**429,带 `Retry-After`**。请求本身没问题,过一会儿再来就 - // 发得出去;报成哪一家的失败都不对,它们一个字节都没收到 - if stalled { + // 剩下的候选等过了还都满着,**一家都没发到过**:429,带 `Retry-After`。请求本身没问题, + // 过一会儿再来就发得出去;报成哪一家的失败都不对,它们一个字节都没收到。 + // + // 发到过的话,交出去的是最后那一家失败的原因,和没有满着的上游时一样:乙回了 401、甲一直 + // 满着,客户端该看到的是乙的错 —— 一个「都满着」的 429 会让它退避了再试,而该修的是乙的 + // 密钥 + if stalled && !reached { let upstreams = busy .iter() .map(|n| format!("`{n}`")) @@ -1092,9 +1111,10 @@ fn unsendable_tool( /// 没有的话,慢了的这一家就是最后一家**,照常等下去 —— 放弃了它,换来的是一个注定失败的 /// 请求。 /// -/// **满着的算接得下**(`max_concurrent`,见 [`crate::slots`]):满着只是此刻,等空位的 -/// 那一段(`failover.slot_wait_secs`)里它可能空出来。于是放弃了慢的这一家之后,请求可能 -/// 去等一家满着的、等不到时回 429 —— 不在这里猜它空不空得出来,「最后一家」只看接不接得下。 +/// **满着的不算**(`max_concurrent`,见 [`crate::slots`]),调用方先把它们滤掉:到点的这一刻 +/// 有空位的(没设上限的、跳过时满着、此刻空出来了的)才接得下。满着的那一家要等,而一个请求 +/// 只有一段等待期限(`failover.slot_wait_secs`),开头慢的这一段多半已经把它用完了:放弃了 +/// 一家正在答的,换来的是去等一家满着的、等不到回 429。 /// /// **只看不跑**:看的是客户端的原话,不跑插件、不取密钥(那两样在真发的那一跳才知道拒 /// 不拒),不发转换事件 diff --git a/crates/tw-gateway/src/server/pipeline/slow.rs b/crates/tw-gateway/src/server/pipeline/slow.rs index ce7e1830..85a88d9a 100644 --- a/crates/tw-gateway/src/server/pipeline/slow.rs +++ b/crates/tw-gateway/src/server/pipeline/slow.rs @@ -7,9 +7,9 @@ //! //! 几条规矩: //! -//! - **最后一家不换**,照常等下去。「最后」按后面还有没有接得下的算(见 -//! `hop::successor`):停用着的、这一跳发不出去的不算 —— 否则放弃了一个慢的,换来的是 -//! 一个注定失败的。 +//! - **最后一家不换**,照常等下去。「最后」按到点的那一刻后面还有没有接得下的算(见 +//! `hop::successor`):停用着的、这一跳发不出去的、并发数满着的(`max_concurrent`)不算 —— +//! 否则放弃了一个慢的,换来的是一个注定失败的,或者一个等不到空位的 429。 //! - **这一家不停用、不算失败**:慢不是坏,下一个请求它可能就快了。 //! - **它的快慢样本记它被给的那段时间**(见 [`timed_out`]):`url-test` 和按快慢分的 //! `load-balance` 照这个把它往后排。 diff --git a/crates/tw-gateway/tests/slow_start.rs b/crates/tw-gateway/tests/slow_start.rs index e06d71d1..3f17f7d1 100644 --- a/crates/tw-gateway/tests/slow_start.rs +++ b/crates/tw-gateway/tests/slow_start.rs @@ -604,3 +604,127 @@ async fn the_slot_of_an_upstream_given_up_on_is_free_at_once() { assert_eq!(status, 200, "{text}"); assert!(text.contains("hello"), "{text}"); } + +/// 起网关,交回数据面的状态(要占住某一家的位置)。`slot_wait_secs` 是等空位的期限 +async fn gateway_with_state( + providers: Vec, + slot_wait_secs: u64, +) -> ( + SocketAddr, + tokio::sync::broadcast::Receiver, + tw_gateway::AppState, +) { + let cfg = Config { + version: 1, + clients: vec![Client { + name: "c".into(), + key: "tw-k".into(), + ..Default::default() + }], + providers, + failover: Failover { + stream_start_wait_secs: 1, + next_on_slow_start: true, + slot_wait_secs, + ..Default::default() + }, + ..Default::default() + }; + let state = tw_gateway::AppState::new(cfg).unwrap(); + let rx = state.bus.subscribe(); + let addr = tw_gateway::serve(state.clone(), ([127, 0, 0, 1], 0).into()) + .await + .unwrap(); + tokio::time::sleep(Duration::from_millis(50)).await; + (addr, rx, state) +} + +/// 并发数满着的那一家(`max_concurrent`) +fn full(name: &str, up: &Upstream) -> Provider { + Provider { + max_concurrent: Some(1), + ..provider(name, up, Protocol::Anthropic) + } +} + +/// 后面只剩一家满着的(`max_concurrent`):到点时它接不下,慢的这一家就是最后一家,照常等它。 +/// 放弃了它,换来的是去等那一家空出来、等不到回 429 —— 一个本来答得上的请求就这样丢了 +#[tokio::test] +async fn a_next_upstream_that_is_full_at_the_deadline_does_not_count() { + let busy = upstream(prompt("busy")).await; + let slow = upstream(late(2_500, "patience")).await; + let (gw, mut rx, state) = gateway_with_state( + vec![ + full("busy", &busy), + provider("slow", &slow, Protocol::Anthropic), + ], + 3, + ) + .await; + // 满着的那一家排在前面:它被当场跳过,慢的那一家发出去时后面还「有」它 + let _held = state.slots.try_take("busy").expect("一个空位"); + let (status, text) = post(gw, "/v1/messages", &messages(true)).await; + assert_eq!(status, 200, "{text}"); + assert!(text.contains("patience"), "{text}"); + assert_eq!(busy.hits.load(Ordering::SeqCst), 0); + let (_, attempts) = estimate_and_attempts(&mut rx).await; + use tw_api::AttemptOutcome::{Error, Served}; + assert_eq!(outcomes(&attempts), [("busy", Error), ("slow", Served)]); + assert_eq!(attempts[0].skipped, Some(tw_api::ServeSkip::Busy)); +} + +/// 后面有一家此刻接得下:照常换过去,排在前面满着的那一家不挡它 +#[tokio::test] +async fn a_next_upstream_with_a_free_slot_still_takes_over() { + let busy = upstream(prompt("busy")).await; + let slow = upstream(stalled()).await; + let good = upstream(prompt("good")).await; + let (gw, mut rx, state) = gateway_with_state( + vec![ + full("busy", &busy), + provider("slow", &slow, Protocol::Anthropic), + provider("good", &good, Protocol::Anthropic), + ], + 3, + ) + .await; + let _held = state.slots.try_take("busy").expect("一个空位"); + let (status, text) = post(gw, "/v1/messages", &messages(true)).await; + assert_eq!(status, 200, "{text}"); + assert!(text.contains("good"), "{text}"); + assert_eq!(busy.hits.load(Ordering::SeqCst), 0); + let (_, attempts) = estimate_and_attempts(&mut rx).await; + use tw_api::AttemptOutcome::{Error, Served, SlowStart}; + assert_eq!( + outcomes(&attempts), + [("busy", Error), ("slow", SlowStart), ("good", Served)] + ); +} + +/// 跳过时满着的那一家,到点之前空出来了:它此刻接得下,换到它 +#[tokio::test] +async fn an_upstream_that_frees_before_the_deadline_takes_over() { + let busy = upstream(prompt("freed")).await; + let slow = upstream(stalled()).await; + let (gw, mut rx, state) = gateway_with_state( + vec![ + full("busy", &busy), + provider("slow", &slow, Protocol::Anthropic), + ], + 3, + ) + .await; + let held = state.slots.try_take("busy").expect("一个空位"); + let asking = tokio::spawn(async move { post(gw, "/v1/messages", &messages(true)).await }); + tokio::time::sleep(Duration::from_millis(400)).await; + drop(held); + let (status, text) = asking.await.unwrap(); + assert_eq!(status, 200, "{text}"); + assert!(text.contains("freed"), "{text}"); + let (_, attempts) = estimate_and_attempts(&mut rx).await; + use tw_api::AttemptOutcome::{Error, Served, SlowStart}; + assert_eq!( + outcomes(&attempts), + [("busy", Error), ("slow", SlowStart), ("busy", Served)] + ); +} diff --git a/crates/tw-gateway/tests/upstream_slots.rs b/crates/tw-gateway/tests/upstream_slots.rs index 8222ceb0..ecf43fd9 100644 --- a/crates/tw-gateway/tests/upstream_slots.rs +++ b/crates/tw-gateway/tests/upstream_slots.rs @@ -18,7 +18,8 @@ use tokio::sync::Semaphore; use tw_api::{AttemptOutcome, AttemptView, Event, FailureSource, ServeSkip, Stay}; /// 一家假的 Anthropic 上游。请求里写着 `HOLD` 的:先吐开头和一段内容,然后一直等到测试 -/// 放行([`Up::release`])才说完 —— 一个正在长篇作答的模型,占着一个位置。别的当场答完 +/// 放行([`Up::release`])才说完 —— 一个正在长篇作答的模型,占着一个位置。写着 `REFUSE401`、 +/// `REFUSE400` 的回那个状态码的错误。别的当场答完 struct Up { addr: SocketAddr, /// 收到的生成请求数 @@ -54,7 +55,23 @@ async fn upstream() -> Up { let (hits, release) = (h.clone(), r.clone()); async move { hits.fetch_add(1, Ordering::SeqCst); - if String::from_utf8_lossy(&body).contains("HOLD") { + let text = String::from_utf8_lossy(&body); + // 上游拒绝:凭据不对(换一家有意义),或者请求本身写错了(换一家也一样) + for (marker, status, kind) in [ + ("REFUSE401", 401, "authentication_error"), + ("REFUSE400", 400, "invalid_request_error"), + ] { + if text.contains(marker) { + return axum::response::Response::builder() + .status(status) + .header("content-type", "application/json") + .body(axum::body::Body::from(format!( + r#"{{"type":"error","error":{{"type":"{kind}","message":"upstream says {marker}"}}}}"# + ))) + .unwrap(); + } + } + if text.contains("HOLD") { let stream = async_stream::stream! { yield Ok::<_, std::io::Error>(Bytes::from_static(OPENING)); if let Ok(p) = release.acquire().await { @@ -733,3 +750,78 @@ async fn the_key_limit_wait_and_the_slot_wait_share_one_budget() { "等了 {took:?}:两段该共用 2 秒" ); } + +/// 满着的那一家等到最后也没空出来,而另一家真的收到了请求、回了错:交出去的是那一家的错 +/// (和没有满着的上游时最后一家失败一样),**不是「都满着」的 429** —— 那会让客户端退避了 +/// 再试,而该修的是乙的密钥 +#[tokio::test] +async fn a_real_failure_is_not_hidden_behind_a_full_upstream() { + let (a, b) = (upstream().await, upstream().await); + let (gw, _state, log) = serve(config(&a, &b, (Some(1), None), 1)).await; + let _held = hold(gw, &log, "占着").await; + + let t = std::time::Instant::now(); + let r = ask(gw, "被拒", None, &format!("[{}]", user("REFUSE401")), false) + .send() + .await + .unwrap(); + // 凭据被拒换一家有意义:等过满着的甲 + assert!(t.elapsed() >= Duration::from_millis(900), "没等甲就交了"); + assert_eq!(r.status(), 502); + assert_eq!(r.headers()["x-thinkwatch-error"], "upstream"); + assert!(r.headers().get("retry-after").is_none()); + let json: serde_json::Value = r.json().await.unwrap(); + let text = json["error"]["message"].as_str().unwrap(); + assert!(text.contains("`乙` answered 401"), "{text}"); + + let id = log.id("被拒").await; + let code = log + .until("失败", |evs| { + evs.iter().find_map(|e| match e { + Event::RequestFailed { id: i, message, .. } if *i == id => { + Some(message.code.clone()) + } + _ => None, + }) + }) + .await; + assert_eq!(code, "gw.upstream.status"); + let (chain, _) = log.routed("被拒").await; + assert_eq!(chain.len(), 2, "{chain:#?}"); + assert!(busy(&chain[0]), "{chain:#?}"); + assert_eq!( + ( + chain[1].provider.as_str(), + chain[1].outcome, + chain[1].status + ), + ("乙", AttemptOutcome::Status, Some(401)) + ); +} + +/// 请求本身的问题(换一家也一样被拒):当场原样交出去,不为满着的那一家等 +#[tokio::test] +async fn a_request_error_is_handed_on_without_waiting_for_a_full_upstream() { + let (a, b) = (upstream().await, upstream().await); + let (gw, _state, log) = serve(config(&a, &b, (Some(1), None), 5)).await; + let _held = hold(gw, &log, "占着").await; + + let t = std::time::Instant::now(); + let r = ask( + gw, + "写错了", + None, + &format!("[{}]", user("REFUSE400")), + false, + ) + .send() + .await + .unwrap(); + assert!(t.elapsed() < Duration::from_secs(2), "为满着的甲等了"); + assert_eq!(r.status(), 400); + let text = r.text().await.unwrap(); + assert!( + text.contains("upstream says REFUSE400"), + "上游的原话:{text}" + ); +} diff --git a/docs/config.md b/docs/config.md index 7184ccc8..95379990 100644 --- a/docs/config.md +++ b/docs/config.md @@ -1001,7 +1001,9 @@ nothing for a long time. With `next_on_slow_start`, the request moves on to the next candidate when no content has arrived `stream_start_wait_secs` after it was sent. It is off by default, because models that think before they write can take long to start; with it on, wait 30 seconds or more. The last -candidate always waits, and the upstream given up on is not set aside. +candidate always waits, and the upstream given up on is not set aside. A +candidate that is at its `max_concurrent` at that moment does not count as a +next one: the slow upstream keeps the request. ```yaml failover: @@ -1017,7 +1019,8 @@ upstream it stays on and, if no slot frees in time, moves on to the next one, where its cache starts over. A new conversation skips a full upstream at once. When every candidate is full, the request waits for whichever frees first; if none does, the client gets a 429 with `Retry-After` saying the upstreams are -busy. +busy. If an upstream did receive the request and failed, the client gets that +failure instead. diff --git a/docs/config.zh-CN.md b/docs/config.zh-CN.md index ff7cab6f..2cd6bd6d 100644 --- a/docs/config.zh-CN.md +++ b/docs/config.zh-CN.md @@ -779,7 +779,8 @@ security: 上游也可能开头很慢:收下请求之后很久都不发内容。开启 `next_on_slow_start` 后,请求 发出 `stream_start_wait_secs` 秒仍没有内容,就换到下一个候选。默认关闭,因为先思考 再输出的模型本来就可能很久才开始;开启时建议等 30 秒以上。最后一个候选总是等下去, -被放弃的上游不会停用。 +被放弃的上游不会停用。到点那一刻并发数已满(`max_concurrent`)的候选不算下一个:请求 +留在慢的那一家。 ```yaml failover: @@ -791,7 +792,8 @@ failover: 秒。等密钥的 `minute`、`hour` 上限也算在这段时间里,两样加起来不超过 `slot_wait_secs`。 进行中的对话等它留在的那一家,到时还没有空位就换下一家,缓存在那边从头建; 新的对话遇到满着的上游直接跳过。候选全满时,请求等先空出来的那一家;都没有空出来, -客户端收到 429 和 `Retry-After`,说明上游都忙。 +客户端收到 429 和 `Retry-After`,说明上游都忙。如果有上游收到过这个请求并且失败了, +客户端收到的是那次失败。 From d090016d18c9bb6dd557dd314b0d59070b8e127f Mon Sep 17 00:00:00 2001 From: fylorn <249551762+fylorn@users.noreply.github.com> Date: Mon, 5 Oct 2026 19:32:26 +0800 Subject: [PATCH 13/22] Weigh load-balance by real speed samples only balance_by: latency read the same latency snapshot as url-test, and that snapshot hands out the startup link-test seed until an upstream has three real samples. The seed measures a TCP/TLS handshake (tens of ms), the real samples measure time to first content (seconds), so a member with only a seed looked tens of times faster and took the 10x factor cap. The latency factor now uses real samples only; a member with just a seed is unmeasured and neutral (1.0), as documented. url-test keeps using the seed, which is what keeps it from acting like fallback right after startup. Co-Authored-By: Claude Opus 5.5 --- crates/tw-control/tests/dryrun.rs | 39 +++++++++++++++++++++++++++++++ crates/tw-engine/src/engine.rs | 5 ++-- crates/tw-gateway/src/latency.rs | 39 ++++++++++++++++++++++++------- crates/tw-gateway/src/state.rs | 7 ++++-- 4 files changed, 78 insertions(+), 12 deletions(-) diff --git a/crates/tw-control/tests/dryrun.rs b/crates/tw-control/tests/dryrun.rs index 47938064..4a187692 100644 --- a/crates/tw-control/tests/dryrun.rs +++ b/crates/tw-control/tests/dryrun.rs @@ -729,6 +729,45 @@ async fn a_balancing_group_shows_what_each_member_is_weighed_by() { ); } +/// 启动时握手垫的底(`url-test` 用的,见 `tw_gateway::latency`)**不是按快慢分的速度**:握手 +/// 量的是建连,几十毫秒;真实样本是从发出去到第一段内容,几秒。拿它们比,只垫过底的那一家 +/// 会被算成快十倍、拿走十倍的份额。只垫过底的算没测过:系数 1,不带首字节时间。`url-test` +/// 照旧用它(不然刚起来时它和 `fallback` 一样) +#[tokio::test] +async fn a_link_test_seed_is_no_speed_for_load_balance() { + let cfg = CFG.replace( + " type: load-balance\n", + " type: load-balance\n balance_by: latency\n", + ); + let (_d, app, gw) = app_and_gateway(&cfg); + gw.latency.seed("官方", 30); + for _ in 0..3 { + gw.latency.record("中转", 3_000); + } + let r = run(&app, r#"{"model":"claude-sonnet-4-5"}"#).await; + let official = candidate(&r, "官方"); + assert_eq!( + (official.ttfb_ms, official.balance_factor), + (None, Some(1.0)), + "{official:?}" + ); + let relay = candidate(&r, "中转"); + assert_eq!( + (relay.ttfb_ms, relay.balance_factor), + (Some(3_000), Some(1.0)), + "{relay:?}" + ); + + let (_d, app, gw) = app_and_gateway(&CFG.replace("type: load-balance", "type: url-test")); + gw.latency.seed("官方", 30); + for _ in 0..3 { + gw.latency.record("中转", 3_000); + } + let r = run(&app, r#"{"model":"claude-sonnet-4-5"}"#).await; + assert_eq!(r.candidates, ["官方", "中转"], "url-test 照旧用垫的底"); + assert_eq!(candidate(&r, "官方").ttfb_ms, Some(30)); +} + /// 只按权重分时没有系数可说;`url-test` 只说首字节时间 —— 它就按这个排 #[tokio::test] async fn only_the_numbers_the_order_uses_are_shown() { diff --git a/crates/tw-engine/src/engine.rs b/crates/tw-engine/src/engine.rs index 1f4eae7e..862fa01f 100644 --- a/crates/tw-engine/src/engine.rs +++ b/crates/tw-engine/src/engine.rs @@ -176,7 +176,7 @@ fn median_ttfb(members: &[String], f: &Facts) -> Option { }) } -/// 快慢系数。**不到 1 毫秒的按 1 毫秒算**:本机的假上游、或者垫底的握手计时可能是 0 +/// 快慢系数。**不到 1 毫秒的按 1 毫秒算**:本机的假上游可能是 0 fn latency_factor(median_ms: f64, ttfb_ms: u32) -> f64 { let ratio = median_ms.max(1.0) / f64::from(ttfb_ms.max(1)); (ratio * ratio).clamp(LATENCY_FACTOR_MIN, LATENCY_FACTOR_MAX) @@ -205,7 +205,8 @@ pub struct Facts { /// 此刻并发数满着的上游(`providers[].max_concurrent`)。`load-balance` 这一轮也不算它们, /// 和停着的一样:网关会当场跳过它们,轮到它们的那一次记了账却没答 pub busy: std::collections::HashSet, - /// 每家典型的快慢:从发出去到回答的第一段内容,毫秒。**缺席 = 样本不够**,不是「很快」 + /// 每家典型的快慢:从发出去到回答的第一段内容,毫秒。**缺席 = 样本不够**,不是「很快」。 + /// `url-test` 拿到的里面有启动时握手垫的底;`load-balance` 拿到的只有真实样本 pub ttfb_ms: std::collections::HashMap, /// 每家跑这个模型的单价,(输入, 输出),微分/百万 token。 /// **缺席 = 算不出价钱**,不是「免费」 diff --git a/crates/tw-gateway/src/latency.rs b/crates/tw-gateway/src/latency.rs index deb42165..2c8567ad 100644 --- a/crates/tw-gateway/src/latency.rs +++ b/crates/tw-gateway/src/latency.rs @@ -1,5 +1,6 @@ //! 每家上游典型的快慢:**从这一跳发出去,到回答的第一段内容**。`url-test` 按它挑最快的, -//! `load-balance` 按快慢分时(`balance_by: latency`)按它算系数。 +//! `load-balance` 按快慢分时(`balance_by: latency`)按它算系数 —— **只认真实样本**(见 +//! [`Latency::measured`])。 //! //! # 量的是哪一段 //! @@ -57,12 +58,17 @@ struct Window { impl Window { /// 这家的典型值:真实样本够数就是它们的中位数,不够时有垫底的用垫底的,都没有是 `None` fn typical(&self) -> Option { - if self.real.len() >= MIN_SAMPLES { - let mut v: Vec = self.real.iter().copied().collect(); - v.sort_unstable(); - return Some(v[v.len() / 2]); + self.measured().or(self.seed) + } + + /// 真实样本够数时它们的中位数。垫的底不算 + fn measured(&self) -> Option { + if self.real.len() < MIN_SAMPLES { + return None; } - self.seed + let mut v: Vec = self.real.iter().copied().collect(); + v.sort_unstable(); + Some(v[v.len() / 2]) } } @@ -117,12 +123,24 @@ impl Latency { g.get(provider)?.typical() } - /// 一次取一批 —— 排序时要用到,逐个取会连着锁好几次。 + /// 一次取一批 —— 排序时要用到,逐个取会连着锁好几次。`url-test` 用:垫过底的算测过 pub fn snapshot(&self, names: &[String]) -> HashMap { + self.batch(names, Window::typical) + } + + /// 一批里**真实样本够数的**那几家,`load-balance` 按快慢分时用。**垫的底不算**:L1 握手 + /// 量的是建连(几十毫秒),真实样本量到第一段内容(几秒),两样一比,只垫过底的那一家 + /// 被算成快几十倍、拿到十倍的份额。没测过的不在里面,系数算 1(中等)—— 和 `url-test` + /// 不同,这里没测过的照样分到请求,攒得到真实样本 + pub fn measured(&self, names: &[String]) -> HashMap { + self.batch(names, Window::measured) + } + + fn batch(&self, names: &[String], of: fn(&Window) -> Option) -> HashMap { let g = self.inner.lock().expect("lock not poisoned"); names .iter() - .filter_map(|n| Some((n.clone(), g.get(n)?.typical()?))) + .filter_map(|n| Some((n.clone(), of(g.get(n)?)?))) .collect() } } @@ -287,9 +305,14 @@ mod tests { "丙".to_string(), "丁".to_string(), ]; + // `url-test` 看的那一份:垫过底的算测过 let s = l.snapshot(&names); assert_eq!(s.len(), 2, "{s:?}"); assert_eq!(s.get("乙"), Some(&200)); assert_eq!(s.get("丁"), Some(&30), "垫过底的算测过"); + // 按快慢分的 `load-balance` 看的那一份:只有真实样本够数的 + let m = l.measured(&names); + assert_eq!(m.len(), 1, "{m:?}"); + assert_eq!(m.get("乙"), Some(&200)); } } diff --git a/crates/tw-gateway/src/state.rs b/crates/tw-gateway/src/state.rs index 75af7f45..74a624a1 100644 --- a/crates/tw-gateway/src/state.rs +++ b/crates/tw-gateway/src/state.rs @@ -426,9 +426,12 @@ impl AppState { } else { Default::default() }, - // `url-test` 选最快的;`load-balance` 按快慢分的也看它 - ttfb_ms: if g.kind == GroupType::UrlTest || (balanced && g.balance_by.uses_latency()) { + // `url-test` 选最快的,握手垫的底也算;`load-balance` 按快慢分的也看它,只认真实 + // 样本(见 `Latency::measured`) + ttfb_ms: if g.kind == GroupType::UrlTest { self.latency.snapshot(candidates) + } else if balanced && g.balance_by.uses_latency() { + self.latency.measured(candidates) } else { Default::default() }, From 923f841b8f63b8faa9fb8af251dad496baa04246 Mon Sep 17 00:00:00 2001 From: fylorn <249551762+fylorn@users.noreply.github.com> Date: Mon, 5 Oct 2026 19:34:12 +0800 Subject: [PATCH 14/22] Charge the sticky leader even when it is full or paused A load-balance member that is full (max_concurrent) or paused sits out of the pick for new conversations. advance() charged only members of that round, so when stickiness kept a conversation on a full member, nothing was charged: the member served after waiting for its slot, yet its share went unrecorded, and the busier it was the more it got beyond its weight. The member that actually leads after stickiness now always joins the round it is charged in, so new conversations make up the difference as they do for any other sticky request. Co-Authored-By: Claude Opus 5.5 --- crates/tw-engine/src/weighted.rs | 57 +++++++++++++++++++++++++++----- 1 file changed, 48 insertions(+), 9 deletions(-) diff --git a/crates/tw-engine/src/weighted.rs b/crates/tw-engine/src/weighted.rs index 7ee14091..ff5f3f91 100644 --- a/crates/tw-engine/src/weighted.rs +++ b/crates/tw-engine/src/weighted.rs @@ -177,12 +177,7 @@ pub(crate) mod members { /// 越满越被记空账,拿到的比它的权重少。停着的、满着的加起来是全部时都算 —— 那时网关 /// 照样一家家试(fail-open)、等先空出来的那一家,排头的还是按权重来。 fn round<'a>(g: &Group, members: &'a [String], f: &Facts) -> Vec<(&'a str, i64)> { - let factors = balance_factors(g.balance_by, members, f); - let all: Vec<(&'a str, i64)> = members - .iter() - .zip(factors) - .map(|(m, x)| (m.as_str(), effective(g.weight(m), x))) - .collect(); + let all = everyone(g, members, f); let up: Vec<(&'a str, i64)> = all .iter() .copied() @@ -191,6 +186,16 @@ fn round<'a>(g: &Group, members: &'a [String], f: &Facts) -> Vec<(&'a str, i64)> if up.is_empty() { all } else { up } } +/// 全部候选和它们的有效权重,停着的、满着的也在里面 +fn everyone<'a>(g: &Group, members: &'a [String], f: &Facts) -> Vec<(&'a str, i64)> { + let factors = balance_factors(g.balance_by, members, f); + members + .iter() + .zip(factors) + .map(|(m, x)| (m.as_str(), effective(g.weight(m), x))) + .collect() +} + /// 下一个排头:当前权重加上自己的有效权重,最大的那个;一样大取组里靠前的。 /// /// `members` 是这次的候选(服务不了这个请求的已经去掉了):不在里面的成员这一轮不参加, @@ -211,13 +216,18 @@ pub fn lead<'a>(g: &Group, members: &'a [String], f: &Facts) -> Option<&'a str> /// /// `leader` 是**会话粘性之后实际排头的那一家**,不一定是 [`lead`] 挑的那个:一段对话 /// 留在了上次回答它的那一家,这一次就记在那一家头上,之后的新对话把差的补回去。 +/// **它满着、停着也记**:满着、停着的不参加给新对话挑排头,可粘性留下的那一家是真要答的 +/// (满着的等到空位再答,见 `tw_gateway::slots`)—— 这一次它带着自己的权重加进这一轮。 /// `members` 和 `f` 要和排序时的一样(同一轮;系数也就是排序时的那一份)。`leader` -/// 不在这一轮里时什么都不记。 +/// 不在候选里时什么都不记。 pub fn advance(g: &Group, members: &[String], f: &Facts, leader: &str) -> HashMap { - let round = round(g, members, f); let mut out = f.current_weight.clone(); - if !round.iter().any(|(m, _)| *m == leader) { + let Some(&led) = everyone(g, members, f).iter().find(|(m, _)| *m == leader) else { return out; + }; + let mut round = round(g, members, f); + if !round.iter().any(|(m, _)| *m == leader) { + round.push(led); } let total: i64 = round.iter().map(|(_, w)| w).sum(); for (m, w) in &round { @@ -387,6 +397,35 @@ mod tests { assert_eq!(advance(&g, &members, &f, "别家"), before); } + /// 粘性留下的那一家此刻满着(或者停着):它不参加给新对话挑排头,但这段对话真的由它答 —— + /// 等到空位之后答的 —— 账照样记在它头上。不记的话,它答得越多越满、越满越不记账,拿到的 + /// 比它的权重多;记了,之后的新对话先去别家,把差的补回来 + #[test] + fn a_member_that_leads_by_stickiness_is_charged_even_when_it_sits_out() { + let g = group(&[("甲", 1), ("乙", 1)]); + let members = names(&g); + for sitting_out in ["busy", "paused"] { + let mut f = Facts::default(); + let set: std::collections::HashSet = ["甲".to_string()].into_iter().collect(); + match sitting_out { + "busy" => f.busy = set, + _ => f.paused = set, + } + // 新对话这一轮轮不到甲 + assert_eq!(lead(&g, &members, &f), Some("乙")); + // 一段对话留在了甲:这一轮甲、乙都加上自己的,甲再减去两家之和 + let next = advance(&g, &members, &f, "甲"); + assert_eq!(next.get("甲"), Some(&-1000), "{sitting_out}: {next:?}"); + assert_eq!(next.get("乙"), Some(&1000), "{sitting_out}: {next:?}"); + // 甲回来之后,下一个新对话先去乙,然后接着挨个轮 + f = Facts { + current_weight: next, + ..Default::default() + }; + assert_eq!(run(&g, &members, &mut f, 4), ["乙", "甲", "乙", "甲"]); + } + } + /// 权重 1 写回去是名字,写了别的权重写回去是映射;数、真假照名字读 #[test] fn a_member_is_a_name_or_a_name_with_a_weight() { From 8aa65f344a413ec8877b08ec9b7de578fce540d9 Mon Sep 17 00:00:00 2001 From: fylorn <249551762+fylorn@users.noreply.github.com> Date: Mon, 5 Oct 2026 19:37:37 +0800 Subject: [PATCH 15/22] Refuse key limits whose records or amounts cannot work - Day and week totals are rebuilt from the request records after a restart, just like the month's, but only month limits checked retention.row_days. A weekly limit with row_days 3 silently forgot the start of the week. Every calendar limit now needs records covering its period (day 1, week 7, month 31). The sentence names the period, so it gets a new code, config.key_limit_retention, replacing config.key_limit_month_retention. - A cost limit below $0.01 rounds to zero micro-dollars or is used up by a single request; it is refused with config.key_limit_cost_too_small. Co-Authored-By: Claude Opus 5.5 --- crates/tw-api/msg-codes.txt | 3 +- crates/tw-config/src/validate.rs | 114 +++++++++++++++++++----- crates/tw-config/tests/manual/schema.rs | 4 +- docs/config.md | 9 +- docs/config.zh-CN.md | 7 +- 5 files changed, 105 insertions(+), 32 deletions(-) diff --git a/crates/tw-api/msg-codes.txt b/crates/tw-api/msg-codes.txt index 19de9be7..7ef362df 100644 --- a/crates/tw-api/msg-codes.txt +++ b/crates/tw-api/msg-codes.txt @@ -68,10 +68,11 @@ config.empty_key config.empty_models_only config.failover_range config.key_limit_cache_reads +config.key_limit_cost_too_small config.key_limit_duplicate config.key_limit_empty -config.key_limit_month_retention config.key_limit_not_positive +config.key_limit_retention config.key_limit_two_measures config.model_spec_blank_model config.model_spec_empty diff --git a/crates/tw-config/src/validate.rs b/crates/tw-config/src/validate.rs index db69aa5c..4f877f25 100644 --- a/crates/tw-config/src/validate.rs +++ b/crates/tw-config/src/validate.rs @@ -136,7 +136,18 @@ pub enum ValidationError { #[error("{}", self.msg())] KeyLimitCacheReads { key: String, measure: &'static str }, #[error("{}", self.msg())] - KeyLimitMonthRetention { key: String, days: u64 }, + KeyLimitCostTooSmall { + key: String, + per: &'static str, + value: String, + }, + #[error("{}", self.msg())] + KeyLimitRetention { + key: String, + per: &'static str, + days: u64, + min: u64, + }, } impl ValidationError { @@ -390,11 +401,21 @@ impl ValidationError { "the {measure} limit of gateway key `{key}` sets cache_reads, which only a \ tokens limit takes" ), - KeyLimitMonthRetention { key, days } => msg!( - "config.key_limit_month_retention", key = key, days = days => - "gateway key `{key}` has a limit per month, and retention.row_days is {days}. \ - After a restart the month's total is added up again from the request records, \ - so row_days has to be at least 31" + KeyLimitCostTooSmall { key, per, value } => msg!( + "config.key_limit_cost_too_small", key = key, per = per, value = value => + "the cost limit per {per} of gateway key `{key}` is {value}; it has to be at \ + least 0.01" + ), + KeyLimitRetention { + key, + per, + days, + min, + } => msg!( + "config.key_limit_retention", key = key, per = per, days = days, min = min => + "gateway key `{key}` has a limit per {per}, and retention.row_days is {days}. \ + After a restart the total for the {per} is added up again from the request \ + records, so row_days has to be at least {min}" ), } } @@ -636,9 +657,20 @@ pub fn validate(cfg: &Config) -> Result<(), ValidationError> { Ok(()) } -/// 一个月的记录最少要留几天:月度上限的用量在重启之后从请求记录里重新加起来, -/// 留得比一个月短,月初那几天的就加不回来 -pub const MONTH_ROW_DAYS: u64 = 31; +/// 天、周、月的上限要请求记录最少留几天:这一期的用量在重启之后(和这一期的开头变了的 +/// 时候,比如时区改了)从请求记录里重新加起来,留得比一期短,这一期开头那几天的就加不回来。 +/// 分钟、小时从空的开始,不要记录 +fn row_days_needed(per: crate::LimitPer) -> Option { + match per { + crate::LimitPer::Day => Some(1), + crate::LimitPer::Week => Some(7), + crate::LimitPer::Month => Some(31), + crate::LimitPer::Minute | crate::LimitPer::Hour => None, + } +} + +/// 费用上限最小是多少美元:再小的话换成微分是 0(不到一微分),或者小到一个请求就用完 +const COST_MIN: f64 = 0.01; /// 一把密钥的用量上限写得对不对。 /// @@ -682,6 +714,13 @@ fn check_limits(c: &crate::Client, row_days: u64) -> Result<(), ValidationError> value, }); } + if let Some(x) = l.cost.filter(|x| *x < COST_MIN) { + return Err(ValidationError::KeyLimitCostTooSmall { + key: key(), + per: l.per.word(), + value: x.to_string(), + }); + } if l.cache_reads && measure != crate::LimitMeasure::Tokens { return Err(ValidationError::KeyLimitCacheReads { key: key(), @@ -695,10 +734,12 @@ fn check_limits(c: &crate::Client, row_days: u64) -> Result<(), ValidationError> measure: measure.word(), }); } - if l.per == crate::LimitPer::Month && row_days < MONTH_ROW_DAYS { - return Err(ValidationError::KeyLimitMonthRetention { + if let Some(min) = row_days_needed(l.per).filter(|min| row_days < *min) { + return Err(ValidationError::KeyLimitRetention { key: key(), + per: l.per.word(), days: row_days, + min, }); } } @@ -1488,8 +1529,8 @@ groups: } } - /// 用量上限:每一条恰好一种量、大于 0、`cache_reads` 只给 token、不重复;月度上限要 - /// 请求记录留够一个月。 + /// 用量上限:每一条恰好一种量、大于 0、费用至少一分、`cache_reads` 只给 token、不重复; + /// 天、周、月的上限要请求记录留够那一期。 #[test] fn key_limits_are_checked_one_entry_at_a_time() { let with = |limits: &str, row_days: u64| { @@ -1543,15 +1584,37 @@ groups: ), "config.key_limit_duplicate" ); - // 月度上限:记录要留够 31 天,用量在重启之后从记录里加回来 - let m = with("[{per: month, requests: 3}]", 30).unwrap_err(); - assert_eq!(m.code, "config.key_limit_month_retention"); - assert_eq!(m.arg("days"), "30"); - assert!(with("[{per: month, requests: 3}]", 31).is_ok()); - assert!( - with("[{per: week, requests: 3}]", 7).is_ok(), - "周以内的不受影响" + // 天、周、月的用量在重启之后从记录里加回来:记录要留够那一期 —— 月 31 天、周 7 天、 + // 天 1 天。分钟、小时从空的开始,不看 + for (per, need) in [("month", 31), ("week", 7), ("day", 1)] { + let limits = format!("[{{per: {per}, requests: 3}}]"); + let m = with(&limits, need - 1).unwrap_err(); + assert_eq!(m.code, "config.key_limit_retention", "{per}"); + assert_eq!( + (m.arg("per"), m.arg("days"), m.arg("min")), + ( + per, + (need - 1).to_string().as_str(), + need.to_string().as_str() + ) + ); + assert!(with(&limits, need).is_ok(), "{per}"); + } + assert!(with("[{per: hour, requests: 3}, {per: minute, cost: 1}]", 0).is_ok()); + // 费用不到一分:换成微分是 0,或者小到没有意义 + for bad in [ + "[{per: day, cost: 0.009}]", + "[{per: hour, cost: 0.0000001}]", + ] { + let m = with(bad, 90).unwrap_err(); + assert_eq!(m.code, "config.key_limit_cost_too_small", "{bad}"); + } + let m = with("[{per: day, cost: 0.005}]", 90).unwrap_err(); + assert_eq!( + (m.arg("key"), m.arg("per"), m.arg("value")), + ("k", "day", "0.005") ); + assert!(with("[{per: day, cost: 0.01}]", 90).is_ok()); } #[test] @@ -1786,9 +1849,16 @@ mod msg_codes { key: "k".into(), measure: "cost", }, - KeyLimitMonthRetention { + KeyLimitCostTooSmall { + key: "k".into(), + per: "day", + value: "0.005".into(), + }, + KeyLimitRetention { key: "k".into(), + per: "month", days: 30, + min: 31, }, ]; check( diff --git a/crates/tw-config/tests/manual/schema.rs b/crates/tw-config/tests/manual/schema.rs index b175a74b..30ed88c0 100644 --- a/crates/tw-config/tests/manual/schema.rs +++ b/crates/tw-config/tests/manual/schema.rs @@ -478,8 +478,8 @@ pub fn sections() -> Vec
{ Kind::Num, Def::Unset, t( - "At most this much, in US dollars, as recorded for each request. Models without a price and upstreams with `billing: free` count as 0.", - "最多花这么多美元,按每个请求记下的费用算。没有价格的模型、`billing: free` 的上游算 0。", + "At most this much, in US dollars, as recorded for each request; at least 0.01. Models without a price and upstreams with `billing: free` count as 0.", + "最多花这么多美元,按每个请求记下的费用算;至少 0.01。没有价格的模型、`billing: free` 的上游算 0。", ), ), row( diff --git a/docs/config.md b/docs/config.md index 95379990..ec4159e4 100644 --- a/docs/config.md +++ b/docs/config.md @@ -340,9 +340,10 @@ key, the limit, the amount used and when it resets, and it shows in the traffic list. Cost is what is recorded for each request, so a model without a price and an upstream with `billing: free` count as $0. A request still running counts with an estimate of its input until it is recorded. After a -restart, the day, week and month are added up again from the request records; -a monthly limit therefore needs `retention.row_days` of at least 31. Minute -and hour limits start empty. +restart, the day, week and month are added up again from the request records, +so the records have to cover the period: `retention.row_days` of at least 1 +for a daily limit, 7 for a weekly one and 31 for a monthly one. Minute and +hour limits start empty. On a Responses WebSocket connection, each `response.create` is a request of its own: it is recorded with its usage and cost and checked against these @@ -357,7 +358,7 @@ connection stays open. | `per` | `minute` \| `hour` \| `day` \| `week` \| `month` | **required** | The period. `minute` and `hour` are rolling (the last 60 seconds, the last 60 minutes); `day`, `week` and `month` start again at local midnight, on Monday and on the 1st. | | `requests` | integer | — | At most this many requests. Token counts and answers the gateway gives itself do not count. | | `tokens` | integer | — | At most this many tokens: uncached input, cache writes and output. | -| `cost` | number | — | At most this much, in US dollars, as recorded for each request. Models without a price and upstreams with `billing: free` count as 0. | +| `cost` | number | — | At most this much, in US dollars, as recorded for each request; at least 0.01. Models without a price and upstreams with `billing: free` count as 0. | | `cache_reads` | bool | `false` | Count cache reads too. Only for a `tokens` limit. | diff --git a/docs/config.zh-CN.md b/docs/config.zh-CN.md index 2cd6bd6d..30c6a176 100644 --- a/docs/config.zh-CN.md +++ b/docs/config.zh-CN.md @@ -244,8 +244,9 @@ clients: 被拒的请求收到 HTTP 429,错误格式和客户端自己的一致,写明是哪把密钥、哪一条上限、用了多少、 什么时候重置;流量列表里也有这一条。费用按每个请求记下的费用算,所以没有价格的模型、 `billing: free` 的上游算 0。还在进行的请求先按输入的估算计入,记下之后换成实际用量。 -重启之后,这一天、这一周、这个月的用量从请求记录里重新加起来,所以设了按月上限时 -`retention.row_days` 至少要 31;按分钟、按小时的上限从零开始。 +重启之后,这一天、这一周、这个月的用量从请求记录里重新加起来,所以请求记录要留够那一期: +设了按天的上限时 `retention.row_days` 至少 1,按周至少 7,按月至少 31;按分钟、按小时的 +上限从零开始。 Responses 的 WebSocket 连接上,每个 `response.create` 各算一个请求:带着各自的用量和费用记下, 按这些上限检查;被拒的那一个收到 `response.failed`,连接保持不断。 @@ -258,7 +259,7 @@ Responses 的 WebSocket 连接上,每个 `response.create` 各算一个请求 | `per` | `minute` \| `hour` \| `day` \| `week` \| `month` | **必填** | 按多长一段时间算。`minute`、`hour` 是滚动的(最近 60 秒、最近 60 分钟);`day`、`week`、`month` 在本地时间的零点、周一零点、每月一号零点重新算。 | | `requests` | 整数 | — | 最多这么多个请求。数 token 的请求和网关自己答的不算。 | | `tokens` | 整数 | — | 最多这么多 token:未命中缓存的输入、写入缓存的和输出。 | -| `cost` | 数字 | — | 最多花这么多美元,按每个请求记下的费用算。没有价格的模型、`billing: free` 的上游算 0。 | +| `cost` | 数字 | — | 最多花这么多美元,按每个请求记下的费用算;至少 0.01。没有价格的模型、`billing: free` 的上游算 0。 | | `cache_reads` | 布尔 | `false` | 把读取缓存的 token 也算进去。只有 `tokens` 上限能写。 | From 9255f472b072a716bf3535c8fea6a06c3039e739 Mon Sep 17 00:00:00 2001 From: fylorn <249551762+fylorn@users.noreply.github.com> Date: Mon, 5 Oct 2026 19:38:56 +0800 Subject: [PATCH 16/22] Say what a load-balance weight shares: requests, not new conversations The schedule charges every request to the member that leads it, sticky conversations included, so a weight is the long-run share of requests; conversations in progress stay where they are and count toward that member's share. The groups[].type cell, the prose around it and the API comments on GroupView.weights and BalanceBy still said the weights share out new conversations. Co-Authored-By: Claude Opus 5.5 --- crates/tw-api/src/lib.rs | 17 +++++++++-------- crates/tw-config/tests/manual/schema.rs | 4 ++-- crates/tw-engine/src/engine.rs | 11 ++++++----- crates/tw-engine/src/weighted.rs | 3 ++- docs/config.md | 9 +++++---- docs/config.zh-CN.md | 6 +++--- 6 files changed, 27 insertions(+), 23 deletions(-) diff --git a/crates/tw-api/src/lib.rs b/crates/tw-api/src/lib.rs index 3037da0b..914afad3 100644 --- a/crates/tw-api/src/lib.rs +++ b/crates/tw-api/src/lib.rs @@ -2243,10 +2243,11 @@ pub struct GroupView { #[serde(default, skip_serializing_if = "Option::is_none")] pub selected: Option, pub providers: Vec, - /// `load-balance` 组每个成员的权重,**每个成员都在**,没写权重的是 1:新对话按这个 - /// 比例分。别的类型不用权重,是空的 + /// `load-balance` 组每个成员的权重,**每个成员都在**,没写权重的是 1:长期看各家分到的 + /// 请求就是这个比例。进行中的对话留在回答它的那一家,那一轮记在那一家的份额里,新对话把 + /// 差的补回去。别的类型不用权重,是空的 pub weights: std::collections::BTreeMap, - /// `load-balance` 按什么分新对话。别的类型永远是 `weights` + /// `load-balance` 按什么分请求。别的类型永远是 `weights` pub balance_by: BalanceBy, } @@ -3414,11 +3415,11 @@ slug_enum! { } slug_enum! { - /// `load-balance` 组按什么分新对话:配置里 `balance_by` 写的那个词。 + /// `load-balance` 组按什么分请求:配置里 `balance_by` 写的那个词。 /// /// 成员的权重永远是底数,快慢、成败算出的系数乘在上面 - /// ([`DryRunCandidate::balance_factor`]);进行中的对话照旧留在回答它的那一家。 - /// 没有测到的上游算中等。 + /// ([`DryRunCandidate::balance_factor`]),长期看各家分到的请求是乘出来的比例;进行中的 + /// 对话照旧留在回答它的那一家,记在那一家的份额里。没有测到的上游算中等。 pub enum BalanceBy { /// 只按成员的权重 Weights = "weights", @@ -3446,7 +3447,7 @@ pub struct GroupInput { /// 别的类型只能不给、或者都是 1 #[serde(default, skip_serializing_if = "Option::is_none")] pub weights: Option>, - /// `load-balance` 组按什么分新对话。不给 = `weights`;别的类型只能是 `weights` + /// `load-balance` 组按什么分请求。不给 = `weights`;别的类型只能是 `weights` #[serde(default, skip_serializing_if = "Option::is_none")] pub balance_by: Option, } @@ -4936,7 +4937,7 @@ pub struct DryRunResult { /// 写在第一个」。直指 provider 时是 None。 #[serde(default, skip_serializing_if = "Option::is_none")] pub strategy: Option, - /// 经过的是 `load-balance` 组时,它按什么分新对话。别的时候没有 + /// 经过的是 `load-balance` 组时,它按什么分请求。别的时候没有 #[serde(default, skip_serializing_if = "Option::is_none")] pub balance_by: Option, /// `route` | `deny` | `no_match` | `unavailable`(选中的上游都服务不了, diff --git a/crates/tw-config/tests/manual/schema.rs b/crates/tw-config/tests/manual/schema.rs index 30ed88c0..44fca006 100644 --- a/crates/tw-config/tests/manual/schema.rs +++ b/crates/tw-config/tests/manual/schema.rs @@ -1355,8 +1355,8 @@ pub fn sections() -> Vec
{ Kind::Enum(group_types), Def::Is("fallback"), t( - "`fallback`: the first healthy member, in order. `select`: the member named in `selected`. `load-balance`: new conversations take turns, in proportion to the members' weights. `url-test`: the fastest by measured time from sending a request to the first content of the answer. `cheapest`: the lowest input price.", - "`fallback`:按顺序取第一个健康的。`select`:取 `selected` 指定的那个。`load-balance`:新对话按成员的权重轮流。`url-test`:按实测从发出请求到回答第一段内容的时间取最快的。`cheapest`:取输入单价最低的。", + "`fallback`: the first healthy member, in order. `select`: the member named in `selected`. `load-balance`: requests are shared out in proportion to the members' weights; a conversation in progress stays where it is. `url-test`: the fastest by measured time from sending a request to the first content of the answer. `cheapest`: the lowest input price.", + "`fallback`:按顺序取第一个健康的。`select`:取 `selected` 指定的那个。`load-balance`:请求按成员的权重分;进行中的对话留在原来那一家。`url-test`:按实测从发出请求到回答第一段内容的时间取最快的。`cheapest`:取输入单价最低的。", ), ), row( diff --git a/crates/tw-engine/src/engine.rs b/crates/tw-engine/src/engine.rs index 862fa01f..0946ee9a 100644 --- a/crates/tw-engine/src/engine.rs +++ b/crates/tw-engine/src/engine.rs @@ -28,7 +28,8 @@ pub enum GroupType { Select, /// 按成员的权重轮流(平滑加权轮询,见 [`crate::weighted`]),权重都是 1 就是挨个轮; /// `balance_by` 还可以按快慢、成败给权重乘一个系数([`BalanceBy`])。 - /// **轮的是新对话**:已经有人回答过、缓存还热着的对话留在那一家 + /// **权重是长期看各家分到的请求的比例**:已经有人回答过、缓存还热着的对话留在那一家, + /// 那一轮记在那一家的份额里,新对话把差的补回去 LoadBalance, /// 选最快的。判据是**真实流量测出来的快慢**:从发出去到回答的第一段内容(流式回答才有), /// 样本不够时用启动时那次零成本的 L1 握手计时补。 @@ -67,11 +68,11 @@ impl GroupType { } } -/// `load-balance` 按什么分新对话([`Group::balance_by`])。 +/// `load-balance` 按什么分请求([`Group::balance_by`])。 /// /// **成员的权重永远是底数**:快慢、成败算出一个系数([`balance_factors`]),乘在每一家 -/// 的权重上,平滑加权轮询按乘出来的数轮([`crate::weighted`])。分的只是新对话 —— -/// 进行中的对话照旧留在回答它的那一家(`tw_gateway::affinity`)。 +/// 的权重上,平滑加权轮询按乘出来的数轮([`crate::weighted`])。进行中的对话照旧留在 +/// 回答它的那一家(`tw_gateway::affinity`),记在那一家的份额里,新对话把差的补回去。 #[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize, Default)] #[serde(rename_all = "kebab-case")] pub enum BalanceBy { @@ -234,7 +235,7 @@ pub struct Group { /// `select` 用:当前选中的那个 #[serde(default, skip_serializing_if = "Option::is_none")] pub selected: Option, - /// `load-balance` 用:按什么分新对话(见 [`BalanceBy`])。默认只按成员的权重,不写回配置 + /// `load-balance` 用:按什么分请求(见 [`BalanceBy`])。默认只按成员的权重,不写回配置 #[serde(default, skip_serializing_if = "BalanceBy::is_weights")] pub balance_by: BalanceBy, } diff --git a/crates/tw-engine/src/weighted.rs b/crates/tw-engine/src/weighted.rs index ff5f3f91..951c5922 100644 --- a/crates/tw-engine/src/weighted.rs +++ b/crates/tw-engine/src/weighted.rs @@ -50,7 +50,8 @@ pub fn effective(weight: u32, factor: f64) -> i64 { pub struct Member { /// 上游的名字 pub name: String, - /// 新对话按权重的比例分给各个成员。不写是 1 + /// 长期看,各个成员分到的请求是权重的比例(进行中的对话留在原来那一家,也记在它的 + /// 份额里)。不写是 1 #[serde(default = "one")] pub weight: u32, } diff --git a/docs/config.md b/docs/config.md index ec4159e4..09b4f168 100644 --- a/docs/config.md +++ b/docs/config.md @@ -1084,7 +1084,7 @@ group with `to`. | Field | Type | Default | Description | |---|---|---|---| | `name` | string | **required** | Name of the group; unique, and not the name of an upstream. | -| `type` | `fallback` \| `select` \| `load-balance` \| `url-test` \| `cheapest` | `fallback` | `fallback`: the first healthy member, in order. `select`: the member named in `selected`. `load-balance`: new conversations take turns, in proportion to the members' weights. `url-test`: the fastest by measured time from sending a request to the first content of the answer. `cheapest`: the lowest input price. | +| `type` | `fallback` \| `select` \| `load-balance` \| `url-test` \| `cheapest` | `fallback` | `fallback`: the first healthy member, in order. `select`: the member named in `selected`. `load-balance`: requests are shared out in proportion to the members' weights; a conversation in progress stays where it is. `url-test`: the fastest by measured time from sending a request to the first content of the answer. `cheapest`: the lowest input price. | | `providers` | list of strings or [`groups[].providers[]`](#cfg-groups-providers) | **required** | Member upstreams, by name; not groups. Each upstream appears once in a group. In a `load-balance` group, a member can be written as `{name, weight}`. | | `selected` | string | — | For `select`: the chosen member. | | `balance_by` | `weights` \| `latency` \| `health` \| `latency-health` | `weights` | For `load-balance`: what the members' weights are multiplied by. `weights`: nothing; the weights alone. `latency`: faster upstreams get more. `health`: upstreams that fail less get more. `latency-health`: both. Other group types take only `weights`. | @@ -1131,8 +1131,9 @@ the conversation, and whichever upstream answered after a failover is the one it stays on. The rule a turn matched at its start also holds for the rest of that turn: rules keyed on input size or images do not move a turn halfway, unless its input no longer fits the context window of a model the rule sends -it to. `load-balance` therefore takes turns between new conversations, by -weight. +it to. A `load-balance` weight is therefore the long-run share of requests: +conversations in progress stay where they are and count toward that +upstream's share. `balance_by` lets a `load-balance` group also look at how each upstream has been doing lately. Each member's weight is multiplied by a factor, and the @@ -1151,7 +1152,7 @@ group shares out requests by the result in the same way as above. neither does a client that cancels, a switch away from a stream that is slow to start, or an upstream skipped because it is at its `max_concurrent`. An upstream that keeps failing keeps a twentieth of its - weight, so it still gets the occasional new conversation and its recovery + weight, so it still gets the occasional request and its recovery is noticed; one that fails outright is set aside by [`failover`](#cfg-failover) as before. - `latency-health`: both factors, multiplied. diff --git a/docs/config.zh-CN.md b/docs/config.zh-CN.md index 30c6a176..230d1cf7 100644 --- a/docs/config.zh-CN.md +++ b/docs/config.zh-CN.md @@ -843,7 +843,7 @@ aliases: | 字段 | 类型 | 默认值 | 说明 | |---|---|---|---| | `name` | 字符串 | **必填** | 策略组的名字,不能重复,也不能和上游同名。 | -| `type` | `fallback` \| `select` \| `load-balance` \| `url-test` \| `cheapest` | `fallback` | `fallback`:按顺序取第一个健康的。`select`:取 `selected` 指定的那个。`load-balance`:新对话按成员的权重轮流。`url-test`:按实测从发出请求到回答第一段内容的时间取最快的。`cheapest`:取输入单价最低的。 | +| `type` | `fallback` \| `select` \| `load-balance` \| `url-test` \| `cheapest` | `fallback` | `fallback`:按顺序取第一个健康的。`select`:取 `selected` 指定的那个。`load-balance`:请求按成员的权重分;进行中的对话留在原来那一家。`url-test`:按实测从发出请求到回答第一段内容的时间取最快的。`cheapest`:取输入单价最低的。 | | `providers` | 列表,每项是字符串或对象,对象见 [`groups[].providers[]`](#cfg-groups-providers) | **必填** | 成员上游的名字,不能是策略组。同一个上游在一个策略组中只出现一次。`load-balance` 组的成员可以写成 `{name, weight}`。 | | `selected` | 字符串 | — | `select` 类型选中的成员。 | | `balance_by` | `weights` \| `latency` \| `health` \| `latency-health` | `weights` | `load-balance` 类型用:成员的权重再乘上什么。`weights`:不乘,只按权重。`latency`:越快的上游分得越多。`health`:越少失败的上游分得越多。`latency-health`:两者都看。其他类型只能是 `weights`。 | @@ -871,13 +871,13 @@ groups: - relay ``` -无论哪种类型,一段对话都留在上次回答它的那一家上游,让上游缓存着的那部分被再次读取,而不是换一家全价重算。同一轮之内(客户端正在回传工具结果)一律不换;跨轮时,上一次回答读或写了至少 1024 个 token 的 prompt cache、且距今不到五分钟,才继续留下。上游因失败进入冷却时,对话随之放开;故障转移之后接下回答的那一家,就是之后留下的那一家。一轮开始时命中的规则也沿用到这一轮结束:按输入大小或图片分流的规则不会让一轮半路换家,除非输入已经超出规则所指模型的上下文窗口。因此 `load-balance` 按权重轮流的是新对话。 +无论哪种类型,一段对话都留在上次回答它的那一家上游,让上游缓存着的那部分被再次读取,而不是换一家全价重算。同一轮之内(客户端正在回传工具结果)一律不换;跨轮时,上一次回答读或写了至少 1024 个 token 的 prompt cache、且距今不到五分钟,才继续留下。上游因失败进入冷却时,对话随之放开;故障转移之后接下回答的那一家,就是之后留下的那一家。一轮开始时命中的规则也沿用到这一轮结束:按输入大小或图片分流的规则不会让一轮半路换家,除非输入已经超出规则所指模型的上下文窗口。因此 `load-balance` 的权重是长期看各家分到的请求的比例:进行中的对话留在原来的上游,也算进那一家的份额。 `balance_by` 让 `load-balance` 组再看各上游最近的表现:每个成员的权重乘上一个系数,组内请求按乘出来的结果照上文的方式分。 - `weights`(默认):只按权重。 - `latency`:越快的上游分得越多。快慢看典型的从发出请求到回答第一段内容的时间,与 `url-test` 使用同一份测量。比组内居中者快一倍的上游,权重乘以四;最多乘以十,最少乘以十分之一。 -- `health`:越少失败的上游分得越多。依据是最近 30 分钟内的最近 50 次请求:服务器错误、限流、额度或余额用尽、凭据被拒、超时和连接失败算作失败;请求本身导致的错误不算,客户端取消、因开头太慢而换走、因并发数满了(`max_concurrent`)而跳过也不算。经常失败的上游至少保留权重的二十分之一,仍会偶尔分到新对话,以便发现它已经恢复;完全失败的上游照旧由 [`failover`](#cfg-failover) 暂停。 +- `health`:越少失败的上游分得越多。依据是最近 30 分钟内的最近 50 次请求:服务器错误、限流、额度或余额用尽、凭据被拒、超时和连接失败算作失败;请求本身导致的错误不算,客户端取消、因开头太慢而换走、因并发数满了(`max_concurrent`)而跳过也不算。经常失败的上游至少保留权重的二十分之一,仍会偶尔分到请求,以便发现它已经恢复;完全失败的上游照旧由 [`failover`](#cfg-failover) 暂停。 - `latency-health`:两个系数相乘。 快慢只在流式回答上测,从请求发给这家上游的那一刻算起:之前的等待、之前失败的上游都不算在内。因开头太慢而被放弃的上游(`failover.next_on_slow_start`),按等满的那段时间计。Responses 的 WebSocket 连接上,每个 `response.create` 在快慢和成败上都算一个请求,快慢从上游开始回答它的那一刻算起。 From 6149de8df411bcd299f8ac3eeb92a4b7df045a33 Mon Sep 17 00:00:00 2001 From: fylorn <249551762+fylorn@users.noreply.github.com> Date: Mon, 5 Oct 2026 19:46:32 +0800 Subject: [PATCH 17/22] Keep the cache_reads refusal test on a valid cost The test saved a cost limit of one micro-dollar with cache_reads, which the new minimum of $0.01 now refuses first; use $1 so it still checks that cache_reads is only for token limits. Co-Authored-By: Claude Opus 5.5 --- crates/tw-control/tests/key_limits.rs | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/crates/tw-control/tests/key_limits.rs b/crates/tw-control/tests/key_limits.rs index d000fe1c..99c08055 100644 --- a/crates/tw-control/tests/key_limits.rs +++ b/crates/tw-control/tests/key_limits.rs @@ -326,7 +326,7 @@ async fn saving_a_key_writes_its_limits_and_a_rename_keeps_what_it_used() { "PUT", "/keys/k2", serde_json::json!({ "key": { "name": "k2", "limits": [ - { "per": "day", "measure": "cost", "max": 1, "cache_reads": true }, + { "per": "day", "measure": "cost", "max": 1_000_000, "cache_reads": true }, ] } }), ) .await; From f7415ba7a37b03af8bc729524ea075a3e0860c71 Mon Sep 17 00:00:00 2001 From: fylorn <249551762+fylorn@users.noreply.github.com> Date: Mon, 5 Oct 2026 19:48:18 +0800 Subject: [PATCH 18/22] Count only requests that reached an upstream against key limits Key limits counted every admitted request. A 429 because every upstream was at its max_concurrent (sent with x-should-retry: true), a content-filter denial or a phase-two rule refusal all counted, so a client retrying as told burned its requests-per-minute limit although nothing was sent anywhere and nothing was spent. A request now counts only if it may have reached an upstream: some attempt was sent (the upstream answered, or it failed after sending: timeout, broken on the way, error before content). Skipped, refused, unconvertible, credential-less and unreachable hops were never sent. A request with an empty chain that failed was refused by the gateway; one that ended without failing (the client left before the chain was reported) may still have been in flight and counts as before. The rule lives once in tw-api (RoutingView::reached_upstream); the recorder hands its verdict to the settle hook, and the restart rebuild applies the same function in SQL (tw_reached), so both totals agree. A request that does not count gives back its reservation and the request it took in the rolling window, at settlement and when its hold is dropped before it starts. Co-Authored-By: Claude Opus 5.5 --- crates/tw-api/src/lib.rs | 110 +++++++++++++++++++ crates/tw-config/tests/manual/schema.rs | 4 +- crates/tw-control/src/key_limits.rs | 6 +- crates/tw-control/tests/key_limits.rs | 122 ++++++++++++++++++++++ crates/tw-gateway/src/key_limits/mod.rs | 85 ++++++++------- crates/tw-gateway/src/key_limits/tests.rs | 66 +++++++----- crates/tw-store/src/db.rs | 95 ++++++++++------- crates/tw-store/src/recorder.rs | 20 ++-- docs/config.md | 9 +- docs/config.zh-CN.md | 5 +- 10 files changed, 394 insertions(+), 128 deletions(-) diff --git a/crates/tw-api/src/lib.rs b/crates/tw-api/src/lib.rs index 914afad3..3a662816 100644 --- a/crates/tw-api/src/lib.rs +++ b/crates/tw-api/src/lib.rs @@ -1543,6 +1543,52 @@ pub struct AttemptView { pub skipped: Option, } +/// 发出去之后才失败的一跳报的码(见 [`AttemptView::sent`]):上游没在时限内回话、请求在路上 +/// 断了、流在第一段内容之前报了错。**上游可能已经收下了它、在算了**。 +/// +/// 别的 `error` 都没发出去:这家满着(`skipped`)、规则拒绝、格式对不上、发给它的名字对不上、 +/// 插件拒绝、转换不了、凭据取不到、签不了名,还有连不上(地址不通、握手失败)—— 那时上游 +/// 一个字节都没收到。 +pub const FAILED_AFTER_SENDING: &[&str] = &[ + "gw.upstream.timeout", + "gw.upstream.forward_failed", + "gw.upstream.stream_opening_error", +]; + +impl AttemptView { + /// 这一跳发到了上游:上游回了话(`served`、`status`、`slow_start`,`estimated` 里带着 + /// 状态码的),或者发出去之后才失败([`FAILED_AFTER_SENDING`])。 + pub fn sent(&self) -> bool { + match self.outcome { + AttemptOutcome::Served | AttemptOutcome::Status | AttemptOutcome::SlowStart => true, + AttemptOutcome::Estimated => self.status.is_some(), + AttemptOutcome::Error => { + self.skipped.is_none() + && self + .error + .as_ref() + .is_some_and(|m| FAILED_AFTER_SENDING.contains(&m.code.as_str())) + } + } + } +} + +impl RoutingView { + /// 这个请求可能发到了上游:尝试链上有一跳发出去了([`AttemptView::sent`])。**尝试链是空的 + /// 时看它怎么收场**(`failed`):失败的是网关在发往哪一家之前就拒了(规则、内容过滤、 + /// 用量上限);没失败的(客户端走了)是还没等到路由事件 —— 那时请求可能正在上游那里, + /// 算它发到了。 + /// + /// 密钥的用量上限只数这样的请求(见 `tw_gateway::key_limits`):上游都满着回的 429、 + /// 被拒的请求不该用掉客户端的上限,它重试的时候什么都没花。 + pub fn reached_upstream(&self, failed: bool) -> bool { + if self.attempts.is_empty() { + return !failed; + } + self.attempts.iter().any(AttemptView::sent) + } +} + /// 放弃了的一跳([`AttemptOutcome::SlowStart`])上游可能已经收了钱的输入。 /// /// 上游在流开头报了的(Anthropic 的 `message_start`)是它报的数;没报的只有 `input`,是网关 @@ -5668,6 +5714,70 @@ pub struct PluginRunView { mod tests { use super::*; + fn attempt(outcome: AttemptOutcome, status: Option, code: Option<&str>) -> AttemptView { + AttemptView { + provider: "p".into(), + model: None, + outcome, + status, + error: code.map(|c| Msg { + code: c.into(), + args: Default::default(), + text: String::new(), + }), + ms: 0, + usage: None, + queued_ms: None, + skipped: None, + } + } + + /// 发到了上游的一跳:上游回了话,或者发出去之后才失败。没发出去的(满着、被拒、连不上) + /// 不算。认的码都是 core 真发得出的 + #[test] + fn an_attempt_was_sent_when_the_upstream_may_have_received_it() { + use AttemptOutcome::*; + for a in [ + attempt(Served, Some(200), None), + attempt(Status, Some(503), None), + attempt(SlowStart, None, Some("gw.slow_start")), + attempt(Estimated, Some(404), None), + attempt(Error, None, Some("gw.upstream.timeout")), + attempt(Error, None, Some("gw.upstream.stream_opening_error")), + ] { + assert!(a.sent(), "{a:?}"); + } + let mut busy = attempt(Error, None, Some("gw.busy_upstream")); + busy.skipped = Some(ServeSkip::Busy); + for a in [ + busy, + attempt(Estimated, None, None), + attempt(Error, None, Some("gw.route.denied")), + attempt(Error, None, Some("gw.upstream.unreachable")), + attempt(Error, None, Some("gw.upstream.sign_failed")), + attempt(Error, None, None), + ] { + assert!(!a.sent(), "{a:?}"); + } + for code in FAILED_AFTER_SENDING { + assert!( + MSG_CODES + .lines() + .any(|l| l.split_whitespace().next() == Some(code)), + "{code} 不是 core 发得出的码" + ); + } + // 尝试链是空的:失败的是网关先拒了,没失败的是还没等到路由事件 + let none = RoutingView::default(); + assert!(!none.reached_upstream(true)); + assert!(none.reached_upstream(false)); + let only_denied = RoutingView { + attempts: vec![attempt(Error, None, Some("gw.route.denied"))], + ..Default::default() + }; + assert!(!only_denied.reached_upstream(false)); + } + /// 版本号就是它上面的说明写到的最新一版(「N 起」)。两条分支各自加了一版、合到一起 /// 时,常量那一行两边都没动、不会冲突,很容易照旧留在合并之前的那个数上 —— 照着说明 /// 写的界面就按旧版去读新的协议了 diff --git a/crates/tw-config/tests/manual/schema.rs b/crates/tw-config/tests/manual/schema.rs index 44fca006..c4a7010e 100644 --- a/crates/tw-config/tests/manual/schema.rs +++ b/crates/tw-config/tests/manual/schema.rs @@ -460,8 +460,8 @@ pub fn sections() -> Vec
{ Kind::Int, Def::Unset, t( - "At most this many requests. Token counts and answers the gateway gives itself do not count.", - "最多这么多个请求。数 token 的请求和网关自己答的不算。", + "At most this many requests. Token counts, answers the gateway gives itself and requests that never reach an upstream do not count.", + "最多这么多个请求。数 token 的请求、网关自己答的、没有发到上游的不算。", ), ), row( diff --git a/crates/tw-control/src/key_limits.rs b/crates/tw-control/src/key_limits.rs index 7f88493f..58d0b993 100644 --- a/crates/tw-control/src/key_limits.rs +++ b/crates/tw-control/src/key_limits.rs @@ -19,8 +19,7 @@ pub fn settle_hook(gw: &tw_gateway::AppState) -> tw_store::SettleHook { client: s.client.clone(), path: s.path.clone(), local: s.local, - error_code: s.error_code.clone(), - attempted: s.attempted, + reached: s.reached, requests: 1, input: s.input, output: s.output, @@ -53,8 +52,7 @@ fn recorded(u: tw_store::KeyUsage) -> Recorded { client: u.client, path: u.path, local: false, - error_code: u.error_code, - attempted: u.attempted, + reached: u.reached, requests: n(u.requests), input: n(u.input_tokens), output: n(u.output_tokens), diff --git a/crates/tw-control/tests/key_limits.rs b/crates/tw-control/tests/key_limits.rs index 99c08055..2120c358 100644 --- a/crates/tw-control/tests/key_limits.rs +++ b/crates/tw-control/tests/key_limits.rs @@ -383,3 +383,125 @@ async fn reaching_a_limit_is_told_on_the_bus() { ] ); } + +/// 一个没发到上游的请求的几条事件:密钥 `key`,开始于 `at`。`attempts` 是尝试链(`None` 是 +/// 一直没有路由事件 —— 内容过滤在发往哪一家之前就拒了),最后是网关自己的一句拒绝 +fn refused( + id: u64, + key: &str, + at: i64, + attempts: Option>, + code: &str, + source: tw_api::FailureSource, +) -> Vec { + let mut evs = request(id, key, at, 0); + evs.pop(); + match attempts { + Some(a) => { + if let Event::RequestRouted { attempts, .. } = &mut evs[1] { + *attempts = a; + } + } + None => { + evs.remove(1); + } + } + evs.push(Event::RequestFailed { + id, + model: String::new(), + source, + message: tw_api::Msg { + code: code.into(), + args: Default::default(), + text: "refused".into(), + }, + bytes: None, + duration_ms: Some(1), + usage: None, + answered_model: None, + }); + evs +} + +fn hop(outcome: tw_api::AttemptOutcome, code: &str, busy: bool) -> tw_api::AttemptView { + tw_api::AttemptView { + provider: "中转".into(), + model: None, + outcome, + status: None, + error: Some(tw_api::Msg { + code: code.into(), + args: Default::default(), + text: "x".into(), + }), + ms: 1, + usage: None, + queued_ms: None, + skipped: busy.then_some(tw_api::ServeSkip::Busy), + } +} + +/// 没发到上游的请求**不算进用量**:上游都满着回的 429(尝试链上只有跳过的)、被内容过滤拒掉 +/// 的(一跳都没有)、在那一跳上被阶段二的规则拒了的、凭据取不到的。客户端照着 `Retry-After` +/// 重试,不该把自己的上限用光。发出去之后才失败的(超时)照样算;重启之后从记录里加回来时 +/// 是同一个判断 +#[tokio::test] +async fn a_request_that_never_reached_an_upstream_does_not_count() { + use tw_api::AttemptOutcome::Error; + use tw_api::FailureSource::{Config, Denied, RateLimited, Upstream}; + let d = tempfile::tempdir().unwrap(); + let file = d.path().join("data.db"); + let b = bed(Some(tw_store::Db::open(&file).unwrap())); + let at = ms(NOON); + let events = [ + refused( + 1, + "plain", + at, + Some(vec![hop(Error, "gw.busy_upstream", true)]), + "gw.busy_all", + RateLimited, + ), + refused(2, "plain", at, None, "gw.content.denied", Denied), + refused( + 3, + "plain", + at, + Some(vec![hop(Error, "gw.route.denied", false)]), + "gw.route.denied", + Denied, + ), + refused( + 4, + "plain", + at, + Some(vec![hop(Error, "gw.upstream.credential_failed", false)]), + "gw.upstream.credential_failed", + Config, + ), + // 发出去了、上游没在时限内回话:它可能已经在算了 + refused( + 5, + "plain", + at, + Some(vec![hop(Error, "gw.upstream.timeout", false)]), + "gw.upstream.timeout", + Upstream, + ), + request(6, "plain", at, 10), + ]; + for e in events.iter().flatten() { + b.rec.lock().await.on_event(e); + } + let list = keys(&b).await; + assert_eq!( + key(&list, "plain")["limits"][0]["used"], + 2, + "只有超时的和答上了的算" + ); + drop(b); + // 重启:从记录里加回来,同一个判断 + let b = bed(Some(tw_store::Db::open(&file).unwrap())); + let list = keys(&b).await; + assert_eq!(key(&list, "plain")["limits"][0]["used"], 2); +} diff --git a/crates/tw-gateway/src/key_limits/mod.rs b/crates/tw-gateway/src/key_limits/mod.rs index 34960a64..21321ce9 100644 --- a/crates/tw-gateway/src/key_limits/mod.rs +++ b/crates/tw-gateway/src/key_limits/mod.rs @@ -2,7 +2,9 @@ //! //! 写法在 `tw_config::limits`。这里回答四件事。 //! -//! **数什么。**请求数:一个准入的请求算一个,数 token 的请求、网关自己答的不算。token: +//! **数什么。**请求数:一个发到了上游的请求算一个([`Recorded::counts`])—— 数 token 的请求、 +//! 网关自己答的不算,没发到上游的也不算:上游都满着回的 429、被规则、内容过滤拒掉的,客户端 +//! 照着重试,不该把自己的上限用光。token: //! 没走缓存的输入 + 写进缓存的 + 输出,那一条开了 `cache_reads` 再加上从缓存读的。费用: //! 记下的费用(实测的、估算的都算),没有价格的模型、不计费的上游算 0。**都按存储层记下的 //! 那一行算**([`KeyLimits::settle`]):重启之后从库里加回来的([`KeyLimits::rebuild`])和 @@ -44,21 +46,6 @@ const CALENDAR: [LimitPer; 3] = [LimitPer::Day, LimitPer::Week, LimitPer::Month] /// 额度。存储层正常时几毫秒就到 const GRACE_MS: i64 = 60_000; -/// 准入之前就被拒的请求那一行的失败码:**路由**拒绝了它(规则拒绝、选中的上游都服务不了), -/// 一跳都没有。这样的请求没有经过准入,不算进请求数([`Recorded::counts`])。 -/// -/// 只看码不够:同样的码在尝试链的某一跳上也会出现(阶段二的规则拒绝),那时请求已经准入 -/// 过了 —— 所以还要看有没有一跳。 -const NOT_ADMITTED: &[&str] = &[ - "gw.route.denied", - "gw.route.all_selected_disabled", - "gw.model.no_upstream_available", -]; - -/// 上限本身拒绝的请求那一行的失败码前缀。它们也不算进请求数:不然一个被拒的客户端每重试 -/// 一次,窗口就往后推一次,永远等不到空位 -const REFUSED: &str = "gw.key_limit."; - /// 一个请求最多等多久:`failover.slot_wait_secs`,0 是不等。 /// /// **一个请求合起来算**:等滚动窗口的空位在准入时、发给哪一家之前,等上游空位 @@ -204,9 +191,9 @@ pub struct Recorded { pub path: String, /// 网关自己答的 pub local: bool, - pub error_code: Option, - /// 尝试链上有没有至少一跳 - pub attempted: bool, + /// 这个请求可能发到了上游(`tw_api::RoutingView::reached_upstream`):尝试链上有一跳发 + /// 出去了,或者客户端在路由事件之前就走了 + pub reached: bool, /// 几个请求:结算时是 1,重建时是这一组的行数 pub requests: u64, pub input: u64, @@ -218,19 +205,14 @@ pub struct Recorded { } impl Recorded { - /// 算不算进密钥的用量:**准入过的才算**。 + /// 算不算进密钥的用量:**发到了上游的才算**。 /// - /// 网关自己答的、数 token 的不经过准入;上限本身拒绝的、路由就拒绝了的没有准入。 - /// 结算一行和从库里加回来用的是这同一个判断,两边的数才对得上。 + /// 网关自己答的、数 token 的不算;一个字节都没发到上游的也不算 —— 上限本身拒绝的、路由 + /// 就拒绝了的、内容过滤拒掉的、上游都满着回了 429 的、每一跳都没发出去的。它们什么都没 + /// 花,客户端照着 `Retry-After` 重试时,不该一次次把自己的上限用掉。结算一行和从库里 + /// 加回来用的是这同一个判断,两边的数才对得上。 pub fn counts(&self) -> bool { - if self.local || uncounted(&self.path) { - return false; - } - match self.error_code.as_deref() { - Some(code) if code.starts_with(REFUSED) => false, - Some(code) if !self.attempted && NOT_ADMITTED.contains(&code) => false, - _ => true, - } + self.reached && !self.local && !uncounted(&self.path) } fn amount(&self) -> Amount { @@ -534,18 +516,22 @@ impl KeyLimits { } } - /// 存储层记下了一行:预留换成实数。`at_ms` 是请求开始的时刻,算在哪一期看它 + /// 存储层记下了一行:预留换成实数。`at_ms` 是请求开始的时刻,算在哪一期看它。 + /// + /// **不算的那一行把准入时记上的也还回去**([`Recorded::counts`]):预留,和滚动窗口里的那 + /// 一个请求 —— 上游都满着回了 429 的请求,客户端过几秒重试,窗口里不该还留着它 pub fn settle(&self, id: u64, at_ms: i64, rec: &Recorded) { let now = self.clock.now_ms(); let events = { let mut g = self.lock(); - let key = match g.by_request.remove(&id).and_then(|s| g.held.remove(&s)) { - Some(r) => r.key, - None => rec.client.clone(), - }; + let held = g.by_request.remove(&id).and_then(|s| g.held.remove(&s)); if !rec.counts() { + if let Some(r) = held { + unrecord(&mut g, &r); + } return; } + let key = held.map_or_else(|| rec.client.clone(), |r| r.key); let amount = rec.amount(); let rolling = g .limits @@ -828,12 +814,15 @@ impl KeyLimits { } } + /// 准入之后、开始之前就被丢掉的请求([`Hold`] 的 Drop):它一个字节都没发出去,预留和 + /// 滚动窗口里的那一个请求都还回去 fn release(&self, seq: u64) { let mut g = self.lock(); - if let Some(r) = g.held.remove(&seq) - && let Some(id) = r.request - { - g.by_request.remove(&id); + if let Some(r) = g.held.remove(&seq) { + if let Some(id) = r.request { + g.by_request.remove(&id); + } + unrecord(&mut g, &r); } } @@ -966,6 +955,24 @@ impl KeyLimits { } } +/// 撤掉准入时给这个请求在滚动窗口里记的那一个请求([`KeyLimits::reserve`] 记在准入的那一刻, +/// 和预留的 `at_ms` 是同一个数)。请求数一个一个都一样,撤哪一个都行。已经滑出窗口的不用撤 +fn unrecord(g: &mut Inner, r: &Reservation) { + let Some(b) = g.books.get_mut(&r.key) else { + return; + }; + if let Some(i) = b + .recent + .iter() + .position(|(t, a)| *t == r.at_ms && a.requests > 0) + { + b.recent[i].1.requests -= 1; + if b.recent[i].1 == Amount::default() { + b.recent.remove(i); + } + } +} + /// 一把密钥的分钟、小时上限里最长的那个窗口。没有就是 None fn longest_window(limits: &[KeyLimit]) -> Option { limits.iter().filter_map(|l| l.per.rolling_ms()).max() diff --git a/crates/tw-gateway/src/key_limits/tests.rs b/crates/tw-gateway/src/key_limits/tests.rs index b11821db..5d33191c 100644 --- a/crates/tw-gateway/src/key_limits/tests.rs +++ b/crates/tw-gateway/src/key_limits/tests.rs @@ -129,8 +129,7 @@ fn row(input: u64, output: u64, cache_read: u64, cost: i64) -> Recorded { client: "k".into(), path: "/v1/messages".into(), local: false, - error_code: None, - attempted: true, + reached: true, requests: 1, input, output, @@ -363,37 +362,47 @@ async fn what_does_not_count_does_not_count() { not(&|r| r.path = "/v1/messages/count_tokens".into()), not(&|r| r.path = "/v1beta/models/gemini-2.5-pro:countTokens".into()), not(&|r| r.path = "/v1/responses/input_tokens".into()), - // 上限自己拒的 - not(&|r| { - r.attempted = false; - r.error_code = Some("gw.key_limit.requests_per_period".into()); - }), - // 路由就拒了,一跳都没有 - not(&|r| { - r.attempted = false; - r.error_code = Some("gw.route.denied".into()); - }), - not(&|r| { - r.attempted = false; - r.error_code = Some("gw.model.no_upstream_available".into()); - }), + // 没发到上游的:上限自己拒的、路由拒的、内容过滤拒的、上游都满着回了 429 的 + not(&|r| r.reached = false), ]; for (i, r) in cases.iter().enumerate() { assert!(!r.counts(), "{r:?}"); b.limits.settle(100 + i as u64, b.now(), r); } assert_eq!(b.used(), [0]); - // 准入过了、在某一跳上被阶段二的规则拒了:算 - let mut hop_denied = row(0, 0, 0, 0); - hop_denied.error_code = Some("gw.route.denied".into()); - // 准入过了、第一跳还没回话客户端就走了:一跳都没记下,也算 - let mut gone = row(0, 0, 0, 0); - gone.attempted = false; - for r in [hop_denied, gone] { - assert!(r.counts(), "{r:?}"); - b.limits.settle(200, b.now(), &r); - } - assert_eq!(b.used(), [2]); + // 发到了上游,失败了、或者客户端走了:算 + let r = row(0, 0, 0, 0); + assert!(r.counts(), "{r:?}"); + b.limits.settle(200, b.now(), &r); + assert_eq!(b.used(), [1]); +} + +/// 准入过了、却没发到上游的请求(上游都满着回了 429、内容过滤拒了):结算时把准入时记上的 +/// 都还回去 —— 预留,和滚动窗口里的那一个请求。客户端照着 `Retry-After` 过几秒重试,窗口里 +/// 不该还留着上一次 +#[tokio::test(start_paused = true)] +async fn a_request_that_never_reached_an_upstream_gives_back_what_it_took() { + let b = bed( + "2026-10-05T10:00:00+08:00", + "[{per: minute, requests: 1}, {per: day, requests: 5}, {per: day, tokens: 1000}]", + ); + b.run(1, ask(600, 0)).await.unwrap(); + assert_eq!(b.used(), [1, 1, 600]); + assert!(b.try_admit(Ask::default()).await.is_err(), "这一分钟用满了"); + let mut busy = row(0, 0, 0, 0); + busy.reached = false; + b.done(1, busy); + assert_eq!(b.used(), [0, 0, 0]); + // 马上就能再来 + b.run(2, ask(600, 0)).await.unwrap(); + assert_eq!(b.used(), [1, 1, 600]); + // 开始之前就被丢掉的(准入之后出了岔子):一样都还回去 + b.done(2, row(1, 0, 0, 0)); + tokio::time::advance(Duration::from_secs(61)).await; + let hold = b.try_admit(ask(100, 0)).await.unwrap(); + assert_eq!(b.used(), [1, 2, 101]); + drop(hold); + assert_eq!(b.used(), [0, 1, 1]); } #[tokio::test(start_paused = true)] @@ -468,8 +477,7 @@ async fn a_restart_adds_the_periods_back_from_the_store() { } // 不算的那几种,加回来时一样不算 let mut refused = row(0, 0, 0, 999_000); - refused.attempted = false; - refused.error_code = Some("gw.route.denied".into()); + refused.reached = false; rows.push(refused); rows }); diff --git a/crates/tw-store/src/db.rs b/crates/tw-store/src/db.rs index 80641ccb..b809c39c 100644 --- a/crates/tw-store/src/db.rs +++ b/crates/tw-store/src/db.rs @@ -148,15 +148,14 @@ pub struct RequestRow { /// 一把密钥从某一刻起用了多少,按几样分开数(见 [`Db::key_usage_since`])。 /// -/// **分开的那几样就是「算不算」要看的**:数 token 的请求、准入之前就被拒的不算进密钥的 +/// **分开的那几样就是「算不算」要看的**:数 token 的请求、没发到上游的不算进密钥的 /// 用量上限,判断在网关那边(`tw_gateway::key_limits`),和它结算一行时是同一个判断。 #[derive(Debug, Clone, PartialEq)] pub struct KeyUsage { pub client: String, pub path: String, - pub error_code: Option, - /// 尝试链上有没有至少一跳 - pub attempted: bool, + /// 这些请求可能发到了上游([`tw_api::RoutingView::reached_upstream`]) + pub reached: bool, pub requests: i64, pub input_tokens: i64, pub output_tokens: i64, @@ -166,6 +165,28 @@ pub struct KeyUsage { pub cost_micros: i64, } +/// 装上 `tw_reached(路由那一列, 失败没有)`:这一行的请求可能发到了上游 +/// ([`tw_api::RoutingView::reached_upstream`])。**判断写在 tw-api 那一处**,记下一行时交给 +/// 网关的([`crate::Settled::reached`])和从库里加回来的是同一个。路由那一列读不出来的当作 +/// 尝试链是空的 +fn register_reached(conn: &Connection) -> rusqlite::Result<()> { + use rusqlite::functions::FunctionFlags; + use rusqlite::types::ValueRef; + conn.create_scalar_function( + "tw_reached", + 2, + FunctionFlags::SQLITE_UTF8 | FunctionFlags::SQLITE_DETERMINISTIC, + |ctx| { + let failed = matches!(ctx.get_raw(1), ValueRef::Integer(n) if n != 0); + let routing = match ctx.get_raw(0) { + ValueRef::Text(t) => serde_json::from_slice::(t).ok(), + _ => None, + }; + Ok(routing.unwrap_or_default().reached_upstream(failed)) + }, + ) +} + #[derive(Debug)] pub struct Db { /// 上游体检的查询在 `crate::health`,和这里共用一个连接 @@ -239,6 +260,7 @@ impl Db { conn.busy_timeout(std::time::Duration::from_secs(5))?; // 搜索用的 SQL 函数。装在连接上,不进库文件:每次打开都要装 crate::search::register(&conn)?; + register_reached(&conn)?; let found: i64 = conn.pragma_query_value(None, "user_version", |r| r.get(0))?; if found != 0 && found != SCHEMA { return Err(DbError::OtherVersion { @@ -1224,37 +1246,34 @@ impl Db { rows.collect::, _>>().map_err(Into::into) } - /// 每把密钥从 `since_ms` 起用了多少:重启之后,密钥的用量上限按它把这一天、这一周、 - /// 这个月的数加回来。 + /// 每把密钥从 `since_ms` 起用了多少:重启之后(和一期的开头变了的时候),密钥的用量 + /// 上限按它把这一天、这一周、这个月的数加回来。 /// - /// 网关自己答的不在里面。**按「算不算」要看的几样分组**(路径、失败的码、有没有发往 - /// 上游),组数和密钥、路径的个数相当,一把密钥一个月的记录也只有几十组。 + /// 网关自己答的不在里面。**按「算不算」要看的几样分组**(路径、有没有发到上游),组数和 + /// 密钥、路径的个数相当,一把密钥一个月的记录也只有几十组。有没有发到上游按路由那一列 + /// 和这一行失败没有判断(`tw_reached`,和记下这一行时交给网关的是同一个判断) pub fn key_usage_since(&self, since_ms: i64) -> Result, DbError> { let mut st = self.conn.prepare( - "SELECT client, path, error_code, - (CASE WHEN json_valid(routing) - THEN COALESCE(json_array_length(routing, '$.attempts'), 0) - ELSE 0 END) > 0 AS attempted, + "SELECT client, path, tw_reached(routing, error IS NOT NULL) AS reached, COUNT(*), COALESCE(SUM(input_tokens), 0), COALESCE(SUM(output_tokens), 0), COALESCE(SUM(cache_read_tokens), 0), COALESCE(SUM(cache_write_tokens), 0), COALESCE(SUM(cost_micros), 0) FROM requests WHERE at_ms >= ?1 AND local = 0 - GROUP BY client, path, error_code, attempted", + GROUP BY client, path, reached", )?; let rows = st.query_map(params![since_ms], |r| { Ok(KeyUsage { client: r.get(0)?, path: r.get(1)?, - error_code: r.get(2)?, - attempted: r.get(3)?, - requests: r.get(4)?, - input_tokens: r.get(5)?, - output_tokens: r.get(6)?, - cache_read_tokens: r.get(7)?, - cache_write_tokens: r.get(8)?, - cost_micros: r.get(9)?, + reached: r.get(2)?, + requests: r.get(3)?, + input_tokens: r.get(4)?, + output_tokens: r.get(5)?, + cache_read_tokens: r.get(6)?, + cache_write_tokens: r.get(7)?, + cost_micros: r.get(8)?, }) })?; rows.collect::, _>>().map_err(Into::into) @@ -1876,8 +1895,8 @@ pub(crate) mod tests { ); } - /// 密钥用量上限重启之后加回来的数:按密钥、路径、失败的码、有没有发往上游分组, - /// 从那一刻起,网关自己答的不算。 + /// 密钥用量上限重启之后加回来的数:按密钥、路径、有没有发到上游分组,从那一刻起,网关 + /// 自己答的不算。 #[test] fn key_usage_adds_up_each_key_from_a_moment_on() { let db = Db::in_memory().unwrap(); @@ -1904,19 +1923,28 @@ pub(crate) mod tests { refused.output_tokens = None; refused.cache_read_tokens = None; refused.cost_micros = None; + // 上游都满着:尝试链上只有跳过的那一跳 + let mut busy = refused.clone(); + busy.id = 7; + busy.routing = Some( + r#"{"route":"default","rule":"r","rewritten_by":[],"attempts":[{"provider":"官方", + "outcome":"error","error":{"code":"gw.busy_upstream","args":{},"text":"x"}, + "ms":0,"skipped":"busy"}]}"# + .into(), + ); let mut local = row(5, t0 + 7); local.local = true; let mut other = row(6, t0 + 8); other.client = "codex".into(); - for r in [&early, &a, &b, &refused, &local, &other] { + for r in [&early, &a, &b, &refused, &busy, &local, &other] { db.insert(r).unwrap(); } let mut got = db.key_usage_since(t0).unwrap(); - got.sort_by(|x, y| (&x.client, x.attempted).cmp(&(&y.client, y.attempted))); + got.sort_by(|x, y| (&x.client, x.reached).cmp(&(&y.client, y.reached))); assert_eq!(got.len(), 3, "{got:?}"); let (refused_g, served, codex) = (&got[0], &got[1], &got[2]); assert_eq!( - (served.client.as_str(), served.attempted, served.requests), + (served.client.as_str(), served.reached, served.requests), ("claude-code", true, 2), "早于那一刻的、本地答的都不算" ); @@ -1932,17 +1960,14 @@ pub(crate) mod tests { "算不出钱的那一行算 0" ); assert_eq!( - ( - refused_g.attempted, - refused_g.error_code.as_deref(), - refused_g.requests - ), - (false, Some("gw.route.denied"), 1) + (refused_g.reached, refused_g.requests), + (false, 2), + "被拒的、都满着的没发到上游" ); assert_eq!( - (codex.client.as_str(), codex.attempted), - ("codex", false), - "没有路由那一列的当作没发往上游" + (codex.client.as_str(), codex.reached), + ("codex", true), + "没有路由那一列、也没失败的:客户端在等上游时走了,可能发到了" ); } diff --git a/crates/tw-store/src/recorder.rs b/crates/tw-store/src/recorder.rs index 057819b6..42f97333 100644 --- a/crates/tw-store/src/recorder.rs +++ b/crates/tw-store/src/recorder.rs @@ -114,10 +114,9 @@ pub struct Settled { pub path: String, /// 网关自己答的(本地估的 token 数) pub local: bool, - /// 失败的原因的码。没失败的是 None - pub error_code: Option, - /// 尝试链上有没有至少一跳。路由就拒绝了的、准入没过的没有 - pub attempted: bool, + /// 这个请求可能发到了上游([`tw_api::RoutingView::reached_upstream`])。都满着回的 429、 + /// 被规则、内容过滤、用量上限拒的没有 + pub reached: bool, pub input: u64, pub output: u64, pub cache_read: u64, @@ -743,11 +742,7 @@ impl Recorder { client: p.client.clone(), path: p.path.clone(), local, - error_code: match &how { - Ending::Failed(message) => Some(message.code.clone()), - _ => None, - }, - attempted: !p.routing.attempts.is_empty(), + reached: p.routing.reached_upstream(matches!(how, Ending::Failed(_))), input: u.map_or(0, |u| u.input), output: u.map_or(0, |u| u.output), cache_read: u.map_or(0, |u| u.cache_read), @@ -2747,10 +2742,8 @@ mod settle_hook_tests { ("claude-code", "/v1/messages") ); assert_eq!(s.at_ms, 1_000_000); - assert!(s.attempted && !s.local); + assert!(s.reached && !s.local); } - assert_eq!(seen[2].error_code.as_deref(), Some("t.broke")); - assert_eq!(seen[0].error_code, None); } /// 路由就拒绝了的:没有一跳。没有用量的:费用是 None,不是 0 @@ -2773,8 +2766,7 @@ mod settle_hook_tests { answered_model: None, }); let s = seen.lock().unwrap()[0].clone(); - assert!(!s.attempted); + assert!(!s.reached, "网关在发往哪一家之前就拒了"); assert_eq!(s.cost_micros, None); - assert_eq!(s.error_code.as_deref(), Some("gw.route.denied")); } } diff --git a/docs/config.md b/docs/config.md index 09b4f168..0011235c 100644 --- a/docs/config.md +++ b/docs/config.md @@ -337,8 +337,11 @@ twcore runs on and start again at midnight, on Monday and on the 1st. When one i A refused request gets HTTP 429 in the client's own error format, naming the key, the limit, the amount used and when it resets, and it shows in the -traffic list. Cost is what is recorded for each request, so a model without a -price and an upstream with `billing: free` count as $0. A request still +traffic list. A request that never reaches an upstream counts toward no +limit: one refused by a rule, the content filter or a limit, or turned away +because every upstream was at its `max_concurrent`. Cost is what is recorded +for each request, so a model without a price and an upstream with +`billing: free` count as $0. A request still running counts with an estimate of its input until it is recorded. After a restart, the day, week and month are added up again from the request records, so the records have to cover the period: `retention.row_days` of at least 1 @@ -356,7 +359,7 @@ connection stays open. | Field | Type | Default | Description | |---|---|---|---| | `per` | `minute` \| `hour` \| `day` \| `week` \| `month` | **required** | The period. `minute` and `hour` are rolling (the last 60 seconds, the last 60 minutes); `day`, `week` and `month` start again at local midnight, on Monday and on the 1st. | -| `requests` | integer | — | At most this many requests. Token counts and answers the gateway gives itself do not count. | +| `requests` | integer | — | At most this many requests. Token counts, answers the gateway gives itself and requests that never reach an upstream do not count. | | `tokens` | integer | — | At most this many tokens: uncached input, cache writes and output. | | `cost` | number | — | At most this much, in US dollars, as recorded for each request; at least 0.01. Models without a price and upstreams with `billing: free` count as 0. | | `cache_reads` | bool | `false` | Count cache reads too. Only for a `tokens` limit. | diff --git a/docs/config.zh-CN.md b/docs/config.zh-CN.md index 230d1cf7..7a32278a 100644 --- a/docs/config.zh-CN.md +++ b/docs/config.zh-CN.md @@ -242,7 +242,8 @@ clients: 本地时区算,在零点、周一零点、每月一号零点重新开始;用满之后,到重新开始之前的请求一律拒绝。 被拒的请求收到 HTTP 429,错误格式和客户端自己的一致,写明是哪把密钥、哪一条上限、用了多少、 -什么时候重置;流量列表里也有这一条。费用按每个请求记下的费用算,所以没有价格的模型、 +什么时候重置;流量列表里也有这一条。没有发到上游的请求不计入任何上限:被规则、内容过滤或 +上限拒绝的,以及因为上游都已达到 `max_concurrent` 而被退回的。费用按每个请求记下的费用算,所以没有价格的模型、 `billing: free` 的上游算 0。还在进行的请求先按输入的估算计入,记下之后换成实际用量。 重启之后,这一天、这一周、这个月的用量从请求记录里重新加起来,所以请求记录要留够那一期: 设了按天的上限时 `retention.row_days` 至少 1,按周至少 7,按月至少 31;按分钟、按小时的 @@ -257,7 +258,7 @@ Responses 的 WebSocket 连接上,每个 `response.create` 各算一个请求 | 字段 | 类型 | 默认值 | 说明 | |---|---|---|---| | `per` | `minute` \| `hour` \| `day` \| `week` \| `month` | **必填** | 按多长一段时间算。`minute`、`hour` 是滚动的(最近 60 秒、最近 60 分钟);`day`、`week`、`month` 在本地时间的零点、周一零点、每月一号零点重新算。 | -| `requests` | 整数 | — | 最多这么多个请求。数 token 的请求和网关自己答的不算。 | +| `requests` | 整数 | — | 最多这么多个请求。数 token 的请求、网关自己答的、没有发到上游的不算。 | | `tokens` | 整数 | — | 最多这么多 token:未命中缓存的输入、写入缓存的和输出。 | | `cost` | 数字 | — | 最多花这么多美元,按每个请求记下的费用算;至少 0.01。没有价格的模型、`billing: free` 的上游算 0。 | | `cache_reads` | 布尔 | `false` | 把读取缓存的 token 也算进去。只有 `tokens` 上限能写。 | From dcddf6a7944ff28b9a4774c86a68779630858956 Mon Sep 17 00:00:00 2001 From: fylorn <249551762+fylorn@users.noreply.github.com> Date: Mon, 5 Oct 2026 19:54:34 +0800 Subject: [PATCH 19/22] Add a moved day, week or month up again from the records Books::period started a period's totals at zero whenever its computed start differed from the one kept. A rollover is fine that way, but a time zone change moves the start of today, this week and this month to a moment that already has requests in it, and those totals dropped to zero: changing the machine's time zone wiped out what a key had spent today. When a key's period start changes, its totals are now read back from the request store with the same query the startup rebuild uses. The gateway asks once per key and period start (KeyLimits::reread_with), outside its lock; tw_control::key_limits::follow does the read on a task, under the recorder's lock, where every settlement also happens, so the totals it hands back are exactly what has been settled. An answer for a period that has moved on again is dropped. twcore wires it up after the startup rebuild. Until the answer arrives the period counts from zero, as it did before. Co-Authored-By: Claude Opus 5.5 --- bin/twcore/src/main.rs | 2 + crates/tw-control/src/key_limits.rs | 41 ++++++- crates/tw-control/tests/key_limits.rs | 90 ++++++++++++++- crates/tw-gateway/src/key_limits/mod.rs | 132 ++++++++++++++++++++-- crates/tw-gateway/src/key_limits/tests.rs | 104 +++++++++++++++++ docs/config.md | 4 +- docs/config.zh-CN.md | 1 + 7 files changed, 356 insertions(+), 18 deletions(-) diff --git a/bin/twcore/src/main.rs b/bin/twcore/src/main.rs index 8f3e65bd..9cf231dc 100644 --- a/bin/twcore/src/main.rs +++ b/bin/twcore/src/main.rs @@ -828,6 +828,8 @@ fn cmd_serve(path: &Path, port: Option, safe: bool, parent: Option) -> // 密钥的用量上限:这一天、这一周、这个月已经用了多少,从记录里加回来。**在第一个 // 请求之前**(网关还没开始听) tw_control::key_limits::rebuild(&state, rec.lock().await.db()); + // 之后一期的开头变了(到了下一期、换了时区),从记录里把新的那一期加起来 + tw_control::key_limits::follow(&state, rec.clone()); } /* diff --git a/crates/tw-control/src/key_limits.rs b/crates/tw-control/src/key_limits.rs index 58d0b993..891d447e 100644 --- a/crates/tw-control/src/key_limits.rs +++ b/crates/tw-control/src/key_limits.rs @@ -1,5 +1,6 @@ -//! 密钥的用量上限和存储层之间的两根线:每记下一行请求就结算,启动时从请求记录里把这一天、 -//! 这一周、这个月用了多少加回来(见 `tw_gateway::key_limits`)。 +//! 密钥的用量上限和存储层之间的三根线:每记下一行请求就结算,启动时从请求记录里把这一天、 +//! 这一周、这个月用了多少加回来,一期的开头变了(到了下一期、机器换了时区)时把新的那一期 +//! 重新加起来(见 `tw_gateway::key_limits`)。 //! //! **网关和存储层互不依赖**(两件平级的事),线在这里接:除了 twcore,控制面是唯一同时 //! 看得见两边的地方。twcore 起来时接上,测试也从这里接。 @@ -46,6 +47,42 @@ pub fn rebuild(gw: &tw_gateway::AppState, db: &tw_store::Db) { }); } +/// 一期的开头变了:从请求记录里把新的这一期重新加起来(见 +/// `tw_gateway::key_limits::KeyLimits::reread_with`)。启动时 [`rebuild`] 之后接上。 +/// +/// **在存储层那把锁里读、在锁里交回去**:结算也在那把锁里(存储层记下一行时调 +/// [`settle_hook`]),拿着它读到的就是此刻结算过的全部,换上去一行不多、一行不少。读库 +/// 在另起的任务上:要它的那一刻可能正拿着这把锁(结算一行时发现到了下一期)。读不了库就 +/// 从 0 起,只记一行 +pub fn follow(gw: &tw_gateway::AppState, store: Arc>) { + // 弱引用:网关的账拿着这个办法,办法再拿着账就是一个圈,谁都放不掉 + let limits = Arc::downgrade(&gw.key_limits); + gw.key_limits + .reread_with(Arc::new(move |m: tw_gateway::key_limits::Moved| { + let (limits, store) = (limits.clone(), store.clone()); + let Ok(rt) = tokio::runtime::Handle::try_current() else { + return; + }; + rt.spawn(async move { + let rec = store.lock().await; + let Some(limits) = limits.upgrade() else { + return; + }; + match rec.db().key_usage_since(m.start) { + Ok(rows) => { + let rows: Vec = rows.into_iter().map(recorded).collect(); + limits.reread(&m, &rows); + } + Err(e) => tracing::warn!( + key = %m.key, + "the usage of a gateway key for its new period could not be read back, \ + so it counts from zero: {e}" + ), + } + }); + })); +} + fn recorded(u: tw_store::KeyUsage) -> Recorded { let n = |v: i64| v.max(0) as u64; Recorded { diff --git a/crates/tw-control/tests/key_limits.rs b/crates/tw-control/tests/key_limits.rs index 2120c358..4d2032b4 100644 --- a/crates/tw-control/tests/key_limits.rs +++ b/crates/tw-control/tests/key_limits.rs @@ -57,15 +57,20 @@ struct Bed { /// `db`:上一次运行留下的请求记录(重启);没有就是一个空库 fn bed(db: Option) -> Bed { + bed_with( + db, + Arc::new(tw_gateway::key_limits::TestClock::new(ms(NOON), 8 * 3600)), + ) +} + +/// [`bed`],用量上限看的是 `clock` +fn bed_with(db: Option, clock: Arc) -> Bed { let d = tempfile::tempdir().unwrap(); let p = d.path().join("config.yaml"); std::fs::write(&p, CONFIG).unwrap(); let cfg = tw_config::try_parse(CONFIG).unwrap(); let mut gw = tw_gateway::AppState::new(cfg).unwrap(); - gw.set_key_limits_clock(Arc::new(tw_gateway::key_limits::TestClock::new( - ms(NOON), - 8 * 3600, - ))); + gw.set_key_limits_clock(clock); let db = db.unwrap_or_else(|| tw_store::Db::in_memory().unwrap()); // 重启:第一个请求之前把这一期加回来,和 twcore 起来时一样 tw_control::key_limits::rebuild(&gw, &db); @@ -505,3 +510,80 @@ async fn a_request_that_never_reached_an_upstream_does_not_count() { let list = keys(&b).await; assert_eq!(key(&list, "plain")["limits"][0]["used"], 2); } + +/// 换得了时区的时钟:东八区和西五区各一只测试时钟,按开关取一只。此刻是同一刻 +struct Moving { + east: tw_gateway::key_limits::TestClock, + west: tw_gateway::key_limits::TestClock, + moved: std::sync::atomic::AtomicBool, +} + +impl Moving { + fn now(&self) -> &tw_gateway::key_limits::TestClock { + if self.moved.load(std::sync::atomic::Ordering::SeqCst) { + &self.west + } else { + &self.east + } + } +} + +impl tw_gateway::key_limits::Clock for Moving { + fn now_ms(&self) -> i64 { + self.now().now_ms() + } + fn period(&self, per: tw_config::LimitPer, at_ms: i64) -> (i64, i64) { + self.now().period(per, at_ms) + } + fn show(&self, at_ms: i64) -> String { + self.now().show(at_ms) + } +} + +/// 机器换了时区:这一天的开头跟着变了。**这一天的数从请求记录里重新加起来**,不是从 0 起 —— +/// 从 0 起的话,换一次时区就能把今天花掉的钱一笔勾销。东八区的 10 月 5 日上午换到西五区, +/// 「今天」成了 10 月 4 日(西五区),从东八区 10 月 4 日 13:00 起:前一晚那一个请求也在里面 +#[tokio::test] +async fn a_time_zone_change_adds_the_new_day_up_again_from_the_records() { + use tw_gateway::key_limits::Clock as _; + let clock = Arc::new(Moving { + east: tw_gateway::key_limits::TestClock::new(ms(NOON), 8 * 3600), + west: tw_gateway::key_limits::TestClock::new(ms(NOON), -5 * 3600), + moved: Default::default(), + }); + let b = bed_with(None, clock.clone()); + tw_control::key_limits::follow(&b.gw, b.rec.clone()); + let events = [ + request(1, "plain", ms("2026-10-04T23:00:00+08:00"), 10), + request(2, "plain", ms("2026-10-05T09:00:00+08:00"), 10), + request(3, "plain", ms("2026-10-05T10:00:00+08:00"), 10), + ]; + for e in events.iter().flatten() { + b.rec.lock().await.on_event(e); + } + let limits: Vec = + serde_yaml_ng::from_str("[{per: day, requests: 100}]").unwrap(); + let used = |b: &Bed| b.gw.key_limits.view("plain", &limits)[0].used; + assert_eq!(used(&b), 2, "东八区的今天"); + clock.moved.store(true, std::sync::atomic::Ordering::SeqCst); + assert_eq!( + clock.period(tw_config::LimitPer::Day, clock.now_ms()).0, + ms("2026-10-04T00:00:00-05:00") + ); + // 第一次看到这一期变了就去读;读回来之前可能还是 0 + let _ = used(&b); + let mut seen = 0; + for _ in 0..100 { + seen = used(&b); + if seen == 3 { + break; + } + tokio::time::sleep(std::time::Duration::from_millis(20)).await; + } + assert_eq!(seen, 3, "西五区的今天:三个请求都在里面"); + // 之后照常结算 + for e in request(4, "plain", ms(NOON), 10) { + b.rec.lock().await.on_event(&e); + } + assert_eq!(used(&b), 4); +} diff --git a/crates/tw-gateway/src/key_limits/mod.rs b/crates/tw-gateway/src/key_limits/mod.rs index 21321ce9..477b4719 100644 --- a/crates/tw-gateway/src/key_limits/mod.rs +++ b/crates/tw-gateway/src/key_limits/mod.rs @@ -20,7 +20,9 @@ //! 那些。 //! //! **什么时候算到哪一期。**天、周、月按请求开始的时刻(那一行的 `at_ms`)归期,和从库里 -//! 加回来时同一个口径。滚动窗口里请求数记在准入的那一刻,token 和费用记在结算的那一刻: +//! 加回来时同一个口径。**一期的开头变了**(到了下一期,或者机器换了时区)就从请求记录里把 +//! 新的这一期重新加起来([`KeyLimits::reread_with`]),不从 0 起 —— 从 0 起的话,换一次时区 +//! 就把今天花掉的一笔勾销。滚动窗口里请求数记在准入的那一刻,token 和费用记在结算的那一刻: //! 一个跑了三分钟的请求,它的输出要等它跑完才知道,按开始的时刻记的话,「每分钟多少 //! token」永远数不到它。 @@ -109,16 +111,24 @@ struct Books { /// 滚动窗口用:最近的每一笔(请求数在准入时,token 和费用在结算时),按记下的先后。 /// **只有设了分钟、小时上限的密钥才记**,留到最长的那个窗口为止 recent: VecDeque<(i64, Amount)>, + /// 开头变了、换了一本新的那几期(哪一种、新的开头),等着从请求记录里重新加(见 + /// [`KeyLimits::reread_with`]) + moved: Vec<(LimitPer, i64)>, } impl Books { - /// `now` 所在的那一期。到了下一期就从 0 起 + /// `now` 所在的那一期。开头和记着的不一样(到了下一期、换了时区)就换一本新的,记下它 + /// 要从请求记录里重新加;加回来之前先从 0 起。这把密钥从没记过这一种的,从 0 起就是对的 + /// —— 启动时从库里加回来过([`KeyLimits::rebuild`]),之后的每一行都结算过 fn period(&mut self, clock: &dyn Clock, per: LimitPer, now: i64) -> &mut Period { let (start, end) = clock.period(per, now); let slot = &mut self.periods[calendar_index(per)]; match slot { Some(p) if p.start == start => {} _ => { + if slot.is_some() { + self.moved.push((per, start)); + } *slot = Some(Period { start, end, @@ -130,6 +140,19 @@ impl Books { } } +/// 开头变了的一期:哪把密钥、哪一种、新的开头。从请求记录里重新加它([`KeyLimits::reread_with`]) +#[derive(Debug, Clone, PartialEq, Eq, Hash)] +pub struct Moved { + pub key: String, + pub per: LimitPer, + /// 这一期的开头,Unix 毫秒:从这一刻起记下的都算进来 + pub start: i64, +} + +/// 从请求记录里重新加一期的办法([`KeyLimits::reread_with`])。拿到要加的那一期,读完了交给 +/// [`KeyLimits::reread`] +pub type Reread = Arc; + fn calendar_index(per: LimitPer) -> usize { CALENDAR.iter().position(|p| *p == per).unwrap_or(0) } @@ -174,6 +197,9 @@ struct Inner { /// 此刻配置里每把密钥的上限。结算时用:要不要记滚动窗口、报不报到了八成 limits: HashMap>, told: HashSet, + /// 已经要过、还没读回来的几期。**一期只要一次**:读回来之前每个请求都会看到同一个开头, + /// 不该每个请求去库里查一遍 + asked: HashSet, } /// 准入时一个请求要占的:输入 token 的估算,和按头一个候选算的输入费用。 @@ -366,6 +392,9 @@ impl Drop for Hold { pub struct KeyLimits { inner: Mutex, clock: Arc, + /// 一期的开头变了时从请求记录里重新加([`Self::reread_with`])。没接上(没有存储层)是 None: + /// 那时从 0 起 + reread: Mutex>, /// 报「到了八成、到了上限」,和认请求有没有结束(见 [`GRACE_MS`]) bus: tw_observe::EventBus, } @@ -379,10 +408,83 @@ impl KeyLimits { Self { inner: Mutex::default(), clock, + reread: Mutex::default(), bus, } } + /// 接上请求记录:一期的开头变了(到了下一期、机器换了时区)时,用 `f` 把新的这一期 + /// 从记录里重新加起来,加好了交回 [`Self::reread`]。 + /// + /// **`f` 在锁外面调,不该等**:它去读库,读完了再交回来(`tw_control::key_limits::follow` + /// 起一个任务做这件事)。读回来之前这一期先从 0 起,同一期只要一次 + pub fn reread_with(&self, f: Reread) { + *self.reread.lock().unwrap_or_else(PoisonError::into_inner) = Some(f); + } + + /// 读回来了:`rows` 是从 `m.start` 起记下的那些(存储层的 `key_usage_since`,各把密钥的都 + /// 在里面),这一期的数换成它们加起来的。**这一期的开头又变了的话不换**:那是另一期了, + /// 它自己会再要一次。 + /// + /// 调用方要保证读的时候没有结算在进行(存储层在同一把锁里记下一行、结算、读库):读到的 + /// 就是此刻结算过的全部,换上去一行不多、一行不少。到了八成、到了顶不报,和启动时加回来 + /// 一样 —— 前一期多半报过同一件事 + pub fn reread(&self, m: &Moved, rows: &[Recorded]) { + let now = self.clock.now_ms(); + let mut sum = Amount::default(); + for r in rows.iter().filter(|r| r.client == m.key && r.counts()) { + sum.add(&r.amount()); + } + let mut g = self.lock(); + g.asked.remove(m); + let Some(p) = g + .books + .get_mut(&m.key) + .and_then(|b| b.periods[calendar_index(m.per)].as_mut()) + .filter(|p| p.start == m.start) + else { + return; + }; + p.sum = sum; + let _ = self.alerts(&mut g, &m.key, now); + } + + /// 这把密钥刚换了新的那几期([`Books::moved`])里还没要过的:记下要过了,交回去,**在锁外面** + /// 交给 [`Self::ask_reread`] + fn moved(&self, g: &mut Inner, key: &str) -> Vec { + let Some(b) = g.books.get_mut(key) else { + return Vec::new(); + }; + let moved: Vec = std::mem::take(&mut b.moved) + .into_iter() + .map(|(per, start)| Moved { + key: key.to_string(), + per, + start, + }) + .collect(); + moved + .into_iter() + .filter(|m| g.asked.insert(m.clone())) + .collect() + } + + /// 去请求记录里重新加这几期。没接上记录的,从 0 起 + fn ask_reread(&self, moved: Vec) { + if moved.is_empty() { + return; + } + let f = self + .reread + .lock() + .unwrap_or_else(PoisonError::into_inner) + .clone(); + match f { + Some(f) => moved.into_iter().for_each(|m| f(m)), + None => self.lock().asked.clear(), + } + } + fn lock(&self) -> std::sync::MutexGuard<'_, Inner> { self.inner.lock().unwrap_or_else(PoisonError::into_inner) } @@ -436,7 +538,7 @@ impl KeyLimits { return Ok(()); } let now = self.clock.now_ms(); - let (out, events) = { + let (out, events, moved) = { let mut g = self.lock(); self.sweep(&mut g, now); let out = self.calendar_locked(&mut g, key, limits, now); @@ -444,9 +546,10 @@ impl KeyLimits { Err(r) => self.refused_alert(&mut g, r, now), Ok(()) => Vec::new(), }; - (out, events) + (out, events, self.moved(&mut g, key)) }; self.emit(events); + self.ask_reread(moved); out } @@ -481,10 +584,10 @@ impl KeyLimits { } loop { let now = self.clock.now_ms(); - let step = { + let (step, moved) = { let mut g = self.lock(); self.sweep(&mut g, now); - match self.calendar_locked(&mut g, key, limits, now) { + let step = match self.calendar_locked(&mut g, key, limits, now) { Err(r) => { let events = self.refused_alert(&mut g, &r, now); Err((r, events)) @@ -493,8 +596,10 @@ impl KeyLimits { None => Ok(self.reserve(&mut g, key, limits, ask, now)), Some(w) => Err((Box::new(w), Vec::new())), }, - } + }; + (step, self.moved(&mut g, key)) }; + self.ask_reread(moved); let refusal = match step { Ok(seq) => { return Ok(Hold { @@ -522,7 +627,7 @@ impl KeyLimits { /// 一个请求 —— 上游都满着回了 429 的请求,客户端过几秒重试,窗口里不该还留着它 pub fn settle(&self, id: u64, at_ms: i64, rec: &Recorded) { let now = self.clock.now_ms(); - let events = { + let (events, moved) = { let mut g = self.lock(); let held = g.by_request.remove(&id).and_then(|s| g.held.remove(&s)); if !rec.counts() { @@ -556,9 +661,10 @@ impl KeyLimits { }, )); } - self.alerts(&mut g, &key, now) + (self.alerts(&mut g, &key, now), self.moved(&mut g, &key)) }; self.emit(events); + self.ask_reread(moved); } /// 重启之后把天、周、月的数从请求记录里加回来。`since(t)` 给出从 `t` 起每把密钥 @@ -606,7 +712,7 @@ impl KeyLimits { let now = self.clock.now_ms(); let mut g = self.lock(); self.sweep(&mut g, now); - limits + let view = limits .iter() .map(|l| { let used = self.used(&mut g, key, l, now); @@ -625,7 +731,11 @@ impl KeyLimits { reached: used >= max, } }) - .collect() + .collect(); + let moved = self.moved(&mut g, key); + drop(g); + self.ask_reread(moved); + view } // ------------------------------------------------------------ 锁里面的 diff --git a/crates/tw-gateway/src/key_limits/tests.rs b/crates/tw-gateway/src/key_limits/tests.rs index 5d33191c..3dfe0bc6 100644 --- a/crates/tw-gateway/src/key_limits/tests.rs +++ b/crates/tw-gateway/src/key_limits/tests.rs @@ -494,6 +494,110 @@ async fn a_restart_adds_the_periods_back_from_the_store() { assert_eq!(b.alerts(), [(150_000, true)]); } +/// 换得了时区的时钟:东八区和西五区各一只,按开关取一只。此刻是同一刻 +struct Moving { + east: TestClock, + west: TestClock, + moved: std::sync::atomic::AtomicBool, +} + +impl Moving { + fn now(&self) -> &TestClock { + if self.moved.load(std::sync::atomic::Ordering::SeqCst) { + &self.west + } else { + &self.east + } + } +} + +impl Clock for Moving { + fn now_ms(&self) -> i64 { + self.now().now_ms() + } + fn period(&self, per: LimitPer, at_ms: i64) -> (i64, i64) { + self.now().period(per, at_ms) + } + fn show(&self, at_ms: i64) -> String { + self.now().show(at_ms) + } +} + +/// 一期的开头变了(机器换了时区):去请求记录里把新的这一期重新加起来,**一期只要一次** —— +/// 读回来之前来的请求看到的都是同一个开头,不该每个都去库里查一遍。读回来的数换上去;读回来 +/// 时这一期的开头又变了的,不换 +#[tokio::test(start_paused = true)] +async fn a_period_whose_start_moved_is_read_back_once() { + let now = at("2026-10-05T10:00:00+08:00"); + let clock = Arc::new(Moving { + east: TestClock::new(now, CST), + west: TestClock::new(now, -5 * 3600), + moved: Default::default(), + }); + let limits = Arc::new(KeyLimits::with_clock( + tw_observe::EventBus::new(), + clock.clone(), + )); + let set = parse("[{per: day, requests: 100}]"); + limits.configure(&tw_config::Config { + clients: vec![tw_config::Client { + name: "k".into(), + key: "tw-k".into(), + limits: set.clone(), + ..Default::default() + }], + ..Default::default() + }); + let asked: Arc>> = Arc::default(); + let into = asked.clone(); + limits.reread_with(Arc::new(move |m| into.lock().unwrap().push(m))); + let used = || limits.view("k", &set)[0].used; + for id in 1..=2 { + limits.settle(id, now, &row(1, 1, 0, 0)); + } + assert_eq!(used(), 2); + assert!(asked.lock().unwrap().is_empty(), "开头没变,不读"); + + clock.moved.store(true, std::sync::atomic::Ordering::SeqCst); + let west_day = at("2026-10-04T00:00:00-05:00"); + for _ in 0..3 { + used(); + limits.calendar("k", &set).unwrap(); + } + limits.settle(3, now, &row(1, 1, 0, 0)); + // 天、周、月的开头都跟着时区变了:各要一次 + let moved = asked.lock().unwrap().clone(); + let starts: Vec<(LimitPer, i64)> = moved.iter().map(|m| (m.per, m.start)).collect(); + assert_eq!( + starts, + [ + (LimitPer::Day, west_day), + (LimitPer::Week, at("2026-09-28T00:00:00-05:00")), + (LimitPer::Month, at("2026-10-01T00:00:00-05:00")), + ], + "一期只要一次" + ); + assert!(moved.iter().all(|m| m.key == "k")); + // 读回来之前:从 0 起,之后结算的照记 + assert_eq!(used(), 1); + // 读回来了:西五区的今天有五个(刚结算的那一个已经在库里),别的密钥的不算 + let mut rows: Vec = (0..5).map(|_| row(1, 1, 0, 0)).collect(); + let mut other = row(1, 1, 0, 0); + other.client = "别的".into(); + rows.push(other); + limits.reread(&moved[0], &rows); + assert_eq!(used(), 5); + // 读回来的时候开头又变了:那是另一期,不换 + clock + .moved + .store(false, std::sync::atomic::Ordering::SeqCst); + let before = used(); + limits.reread(&moved[0], &rows[..1]); + assert_eq!(used(), before, "旧的那一期读回来的数换到了新的这一期上"); + // 换回来:看过的这一期(天)又是新的开头,再要一次。周、月等下一次结算碰到时再要 + assert_eq!(asked.lock().unwrap().len(), 4); +} + #[tokio::test(start_paused = true)] async fn a_renamed_key_keeps_what_it_used() { let b = bed("2026-10-05T10:00:00+08:00", "[{per: day, requests: 3}]"); diff --git a/docs/config.md b/docs/config.md index 0011235c..7c2a9e5b 100644 --- a/docs/config.md +++ b/docs/config.md @@ -333,7 +333,9 @@ When one is used up, a request waits for the next free slot if it frees within `failover.slot_wait_secs` (30 seconds by default), and is refused otherwise. Any later wait for a busy upstream comes out of the same time. `day`, `week` and `month` follow the calendar in the time zone of the machine -twcore runs on and start again at midnight, on Monday and on the 1st. When one is used up, requests are refused until it starts again. +twcore runs on and start again at midnight, on Monday and on the 1st. When one is used up, requests are refused until it starts again. If the +machine's time zone changes, the current day, week and month are added up +again from the request records. A refused request gets HTTP 429 in the client's own error format, naming the key, the limit, the amount used and when it resets, and it shows in the diff --git a/docs/config.zh-CN.md b/docs/config.zh-CN.md index 7a32278a..f2d96130 100644 --- a/docs/config.zh-CN.md +++ b/docs/config.zh-CN.md @@ -240,6 +240,7 @@ clients: `failover.slot_wait_secs`(默认 30 秒)之内空出来,请求就等它;等不到就拒绝。之后若还要等 上游空出位置,用的是同一段时间里剩下的部分。`day`、`week`、`month` 按 twcore 所在机器的 本地时区算,在零点、周一零点、每月一号零点重新开始;用满之后,到重新开始之前的请求一律拒绝。 +机器换了时区时,当前这一天、这一周、这个月的用量按新的时区从请求记录里重新加起来。 被拒的请求收到 HTTP 429,错误格式和客户端自己的一致,写明是哪把密钥、哪一条上限、用了多少、 什么时候重置;流量列表里也有这一条。没有发到上游的请求不计入任何上限:被规则、内容过滤或 From 9dc5baa8bba5dcb0e236b13083f14b41c3075900 Mon Sep 17 00:00:00 2001 From: fylorn <249551762+fylorn@users.noreply.github.com> Date: Mon, 5 Oct 2026 19:58:11 +0800 Subject: [PATCH 20/22] Fail a WebSocket turn when the upstream closes mid-answer On a Responses WebSocket, a close frame or end of stream from the upstream ended the pump the same way as the client leaving, so turns still in flight were recorded as cancelled by the client, and one broken off before any content counted nothing against the upstream's health. An upstream close is now its own ending: the turns in flight fail with the new code gw.ws.upstream_closed, and one without content yet counts as an upstream failure, as an HTTP stream broken before content does. A client close is still a cancellation with no health record. Whole-connection rows (Realtime and other paths) still finish normally when either side closes. Co-Authored-By: Claude Opus 5.5 --- crates/tw-api/msg-codes.txt | 1 + crates/tw-gateway/src/ws.rs | 23 +++++-- crates/tw-gateway/src/ws/turn.rs | 12 ++-- crates/tw-gateway/tests/ws.rs | 104 +++++++++++++++++++++++++++++++ 4 files changed, 128 insertions(+), 12 deletions(-) diff --git a/crates/tw-api/msg-codes.txt b/crates/tw-api/msg-codes.txt index 7ef362df..b9162218 100644 --- a/crates/tw-api/msg-codes.txt +++ b/crates/tw-api/msg-codes.txt @@ -403,6 +403,7 @@ gw.ws.connect_failed gw.ws.proxy_unsupported gw.ws.send_failed gw.ws.upstream_broke +gw.ws.upstream_closed l1.config.bad_url l1.config.no_host l1.config.proxy_addr_form diff --git a/crates/tw-gateway/src/ws.rs b/crates/tw-gateway/src/ws.rs index 47fab7b4..4b0644a2 100644 --- a/crates/tw-gateway/src/ws.rs +++ b/crates/tw-gateway/src/ws.rs @@ -928,9 +928,13 @@ type Stream = tokio_tungstenite::WebSocketStream>; /// 一条连接是怎么断的。 enum End { - /// 有一边收场了:发了关闭帧,或者把连接收掉了。**客户端那一边怎么走 + /// 客户端那一边收场了:发了关闭帧,或者走了(连接收掉了、写不过去了)。**客户端怎么走 /// 都算这一种** —— 一次会话就是由客户端结束的,那是正常收场 Closed, + /// 上游收了连接:发了关闭帧,或者没有关闭帧就断开了。整条连接一行的,这也是收场; + /// Responses 的连接上还没答完的那几轮是**失败**,上游没答完就走了 —— 记成客户端取消的话, + /// 这一家内容之前断掉的那一轮不算它的失败,界面上也像是用户自己停下的 + UpstreamClosed, /// 上游那边出错断了,或者写不过去了 Broke(Msg), /// 被防护切断了:回答里的工具调用命中了切断规则,或者客户端发来的一帧被内容过滤拒了 @@ -1003,8 +1007,8 @@ async fn pump( msg = u_rx.next() => { let m = match msg { Some(Ok(m)) => m, - // 上游把连接收掉了,没有关闭帧也算收场 - None => break End::Closed, + // 上游把连接收掉了,没有关闭帧也算它收了 + None => break End::UpstreamClosed, Some(Err(e)) => { break End::Broke(msg!( "gw.ws.upstream_broke", detail = e => @@ -1027,7 +1031,7 @@ async fn pump( } UpMsg::Ping(b) => Message::Ping(b), UpMsg::Pong(b) => Message::Pong(b), - UpMsg::Close(_) => break End::Closed, + UpMsg::Close(_) => break End::UpstreamClosed, UpMsg::Frame(_) => continue, }; // 发不给客户端,就是客户端已经走了 @@ -1036,17 +1040,24 @@ async fn pump( } }; // **先报结局,再关连接。**关连接要等对面回话,而对面可能早就不在了。在等准入的那一轮 - // 开始了的话记成取消;在跑的几轮,上游断了的是失败,别的是取消 + // 开始了的话记成取消;在跑的几轮,上游断了、收了连接的是失败,客户端走了的是取消 drop(waiting); if let Some(t) = p.turns.as_mut() { match &end { End::Broke(why) => t.fail_all(tw_api::FailureSource::Upstream, why.clone()), + End::UpstreamClosed => t.fail_all( + tw_api::FailureSource::Upstream, + msg!( + "gw.ws.upstream_closed" => + "The upstream closed the connection before the answer was complete." + ), + ), _ => t.clear(), } } if let Some(ending) = ending { match end { - End::Closed => ending.finished(101), + End::Closed | End::UpstreamClosed => ending.finished(101), End::Broke(why) => ending.failed(tw_api::FailureSource::Upstream, why), End::Cut(why) => ending.failed(tw_api::FailureSource::Denied, why), } diff --git a/crates/tw-gateway/src/ws/turn.rs b/crates/tw-gateway/src/ws/turn.rs index f7cf2841..0ef5e561 100644 --- a/crates/tw-gateway/src/ws/turn.rs +++ b/crates/tw-gateway/src/ws/turn.rs @@ -5,8 +5,8 @@ //! 结局三条事件,存储层记一行。用量是这一次回答里的 `usage`(输入含缓存读、输出),结局事件 //! 交给同一个记录器,按发给这一家的模型名、照 HTTP 那条路同一套查价:费用、密钥的用量、 //! 体检、流量看到的都是它。第一个 token 什么时候到、回答里写的是哪个模型,也和 HTTP 那条路 -//! 一样认(见 [`Ending::frame`])。连接半路断了,这一轮记成取消;上游断了,记成失败。尝试链 -//! 只有一跳:这条连接连着的那一家。 +//! 一样认(见 [`Ending::frame`])。客户端半路走了,这一轮记成取消;上游断了、收了连接,记成 +//! 失败。尝试链只有一跳:这条连接连着的那一家。 //! //! **连接本身不留行**:一条连接跑好几轮、中间可以闲着很久,流量里该看的是每一轮。会话照 //! HTTP 那条路按每一帧认(Codex 每段对话带着 `prompt_cache_key`),同一段对话的几轮归到同一 @@ -210,7 +210,7 @@ impl Turns { } } - /// 上游断了:在跑的几轮都没答完,一样失败 + /// 上游断了、收了连接:在跑的几轮都没答完,一样失败 pub(crate) fn fail_all(&mut self, source: tw_api::FailureSource, why: Msg) { while let Some(t) = self.queue.pop_front() { t.fail(source, why.clone()); @@ -224,7 +224,7 @@ impl Turns { } } - /// 连接断了:没答完的几轮记成取消(结局的 Drop) + /// 客户端走了:没答完的几轮记成取消(结局的 Drop) pub(crate) fn clear(&mut self) { self.queue.clear(); } @@ -365,8 +365,8 @@ impl Turn { } } - /// 失败了:上游断了、被防护切断了。用量照样带着(上游已经计了费)。上游断了的,没答上 - /// 之前断的给这一家记一次失败,和 HTTP 那条路流在第一段内容之前断了一样 + /// 失败了:上游断了、收了连接,被防护切断了。用量照样带着(上游已经计了费)。上游那边的, + /// 没答上之前断的给这一家记一次失败,和 HTTP 那条路流在第一段内容之前断了一样 pub(crate) fn fail(mut self, source: tw_api::FailureSource, why: Msg) { if source == tw_api::FailureSource::Upstream { self.judge(|h, p| h.record_failure(p)); diff --git a/crates/tw-gateway/tests/ws.rs b/crates/tw-gateway/tests/ws.rs index 72ab7285..121e3073 100644 --- a/crates/tw-gateway/tests/ws.rs +++ b/crates/tw-gateway/tests/ws.rs @@ -1469,6 +1469,110 @@ async fn each_turn_feeds_the_upstreams_speed_and_success_like_an_http_hop() { assert_eq!(rate(), Some(7.0 / 8.0), "客户端走了被算进了成败"); } +/// 一轮回到一半上游就收了连接:帧里写着 `CLOSE` 的,回了开头之后发关闭帧;写着 `DROP` 的, +/// 回了开头之后直接断开(没有关闭帧);写着 `HOLD` 的,回了开头之后一直不说完。别的照常答完 +async fn closing() -> SocketAddr { + let app = Router::new().route( + "/backend-api/codex/responses", + axum::routing::any(|ws: WebSocketUpgrade| async move { + ws.on_upgrade(|mut sock| async move { + let mut n = 0; + while let Some(Ok(m)) = sock.recv().await { + let Message::Text(t) = m else { continue }; + n += 1; + let id = format!("resp_{n}"); + let frames = if ["CLOSE", "DROP", "HOLD"].iter().any(|w| t.contains(w)) { + vec![created(&id)] + } else { + vec![created(&id), delta(), completed(&id)] + }; + for f in frames { + if sock + .send(Message::Text(f.to_string().into())) + .await + .is_err() + { + return; + } + } + if t.contains("CLOSE") { + let _ = sock.send(Message::Close(None)).await; + return; + } + if t.contains("DROP") { + return; + } + } + }) + }), + ); + let l = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap(); + let addr = l.local_addr().unwrap(); + tokio::spawn(async move { axum::serve(l, app).await.unwrap() }); + addr +} + +/// 一轮还没答完上游就收了连接(关闭帧,或者直接断开):**这一轮失败了,是上游的事**,不是 +/// 客户端取消 —— 内容之前就断的给这一家记一次失败,和 HTTP 那条路流在第一段内容之前断了 +/// 一样。客户端自己走的才是取消,不记成败 +#[tokio::test] +async fn an_upstream_that_closes_mid_turn_fails_the_turn() { + let up = closing().await; + let (gw, mut rx, state) = turns_gateway(up, |_| {}, None).await; + let rate = || { + state + .health + .success_rates(&["up".to_string()]) + .get("up") + .copied() + }; + let mut c = connect(gw).await; + for _ in 0..4 { + c.send(create("hi")).await.unwrap(); + until_end(&mut c).await; + } + requests(&mut rx, 4).await; + for (i, how) in ["CLOSE", "DROP"].into_iter().enumerate() { + if i > 0 { + c = connect(gw).await; + } + c.send(create(how)).await.unwrap(); + let evs = requests(&mut rx, 1).await; + match evs.last().unwrap() { + Event::RequestFailed { + source, message, .. + } => { + assert_eq!(*source, tw_api::FailureSource::Upstream, "{how}"); + // 没有关闭帧就断开的,有的时候读到的是一个读错误(连接被重置) + let codes: &[&str] = match how { + "CLOSE" => &["gw.ws.upstream_closed"], + _ => &["gw.ws.upstream_closed", "gw.ws.upstream_broke"], + }; + assert!(codes.contains(&message.code.as_str()), "{how}: {message:?}"); + } + other => panic!("{how}:该是一次上游的失败:{other:?}"), + } + } + assert_eq!(rate(), Some(4.0 / 6.0), "内容之前断的是这一家的失败"); + + // 客户端自己走的:取消,不记成败 + let mut c = connect(gw).await; + c.send(create("HOLD")).await.unwrap(); + loop { + let m = c.next().await.unwrap().unwrap(); + if m.into_text().unwrap().contains("response.created") { + break; + } + } + drop(c); + let evs = requests(&mut rx, 1).await; + assert!( + matches!(evs.last().unwrap(), Event::RequestCancelled { .. }), + "{evs:?}" + ); + assert_eq!(rate(), Some(4.0 / 6.0), "客户端走了被算进了成败"); +} + /// 一轮在上游的 `error` 那一帧收尾,上游又为**同一次回答**补发一个 `response.failed`:那一帧 /// 照原样交给客户端,**不算排在后面的那一轮的** —— 下一轮照样有自己的回答、用量和成败。 /// 每次回答都用同一个 id 的上游也照常:下一轮自己的回答不会被当成补发的 From e1586a74ae668f219b0d5d1bd51a313a8f2df7f8 Mon Sep 17 00:00:00 2001 From: fylorn <249551762+fylorn@users.noreply.github.com> Date: Mon, 5 Oct 2026 20:03:14 +0800 Subject: [PATCH 21/22] Order WebSocket upgrades like HTTP requests An upgrade connected to the first available candidate, so a load-balance group sent every new WebSocket connection to its first member, and url-test/cheapest orders were ignored too; the dry-run kept predicting a rotation that WebSocket traffic never followed. The ordering step of the HTTP route (group order with weights and balance_by, busy and paused members sitting out, stickiness, then charging the leader) is now one function, pipeline::arrange, used by both paths. The upgrade first drops candidates that cannot serve it (disabled ones; for a Realtime connection also the ones whose alias or key scope does not fit), orders the rest the same way and charges the member it actually connects to. An upgrade has no body to recognise a conversation by, so it has no stickiness. Co-Authored-By: Claude Opus 5.5 --- crates/tw-control/tests/dryrun.rs | 86 +++++++++++++++++++++++ crates/tw-gateway/src/server/pipeline.rs | 87 ++++++++++++++++++------ crates/tw-gateway/src/server/upgrade.rs | 33 +++++++-- docs/config.md | 4 +- docs/config.zh-CN.md | 2 +- 5 files changed, 186 insertions(+), 26 deletions(-) diff --git a/crates/tw-control/tests/dryrun.rs b/crates/tw-control/tests/dryrun.rs index 4a187692..7cdebd2d 100644 --- a/crates/tw-control/tests/dryrun.rs +++ b/crates/tw-control/tests/dryrun.rs @@ -914,3 +914,89 @@ routes: assert_eq!(seen.len(), 3, "{seen:?}"); assert!(seen["甲"] < seen["丙"], "{seen:?}"); } + +/// WebSocket 的连接和 HTTP 的请求**按同一个组的顺序走**:3:1 的 `load-balance` 组,新连接照着 +/// 权重轮到两家(甲甲乙甲),不是都连头一家;每连一次之前试算一次,试算说的排头就是这一次 +/// 真连上的那一家 —— 两条路记的是同一本账 +#[tokio::test] +async fn websocket_connections_take_turns_like_requests_and_the_dry_run_agrees() { + use futures::StreamExt as _; + use std::sync::atomic::{AtomicUsize, Ordering}; + /// 一家只接 WebSocket 的 Responses 上游:数它接了几条连接 + async fn upstream() -> (std::net::SocketAddr, Arc) { + let n = Arc::new(AtomicUsize::new(0)); + let m = n.clone(); + let app = axum::Router::new().route( + "/v1/responses", + axum::routing::any(move |ws: axum::extract::WebSocketUpgrade| { + let m = m.clone(); + async move { + m.fetch_add(1, Ordering::SeqCst); + ws.on_upgrade(|mut sock| async move { while sock.next().await.is_some() {} }) + } + }), + ); + let l = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap(); + let a = l.local_addr().unwrap(); + tokio::spawn(async move { axum::serve(l, app).await.unwrap() }); + (a, n) + } + let (a, hits_a) = upstream().await; + let (b, hits_b) = upstream().await; + let text = format!( + r#"version: 1 +listen: + control: + key: c0ffee00c0ffee00c0ffee00c0ffee00c0ffee00c0ffee00c0ffee00c0ffee00 +clients: + - name: 我 + key: tw-k +providers: + - {{ name: 甲, base_url: "http://{a}", key: sk-a, protocol: openai-responses }} + - {{ name: 乙, base_url: "http://{b}", key: sk-b, protocol: openai-responses }} +groups: + - name: 池 + type: load-balance + providers: [{{ name: 甲, weight: 3 }}, 乙] +routes: + - name: default + rules: + - name: 都去池子 + to: 池 +"# + ); + let (_d, app, gw) = app_and_gateway(&text); + let addr = tw_gateway::serve(gw.clone(), ([127, 0, 0, 1], 0).into()) + .await + .unwrap(); + let mut went = Vec::new(); + for i in 0..8 { + let predicted = run(&app, r#"{"model":"gpt-5"}"#).await.candidates[0].clone(); + let before = (hits_a.load(Ordering::SeqCst), hits_b.load(Ordering::SeqCst)); + use tokio_tungstenite::tungstenite::client::IntoClientRequest; + let mut req = format!("ws://{addr}/v1/responses") + .into_client_request() + .unwrap(); + req.headers_mut() + .insert("authorization", "Bearer tw-k".parse().unwrap()); + let (c, _) = tokio_tungstenite::connect_async(req).await.unwrap(); + // 网关在升级之后才去连上游 + let mut now = before; + for _ in 0..100 { + now = (hits_a.load(Ordering::SeqCst), hits_b.load(Ordering::SeqCst)); + if now != before { + break; + } + tokio::time::sleep(std::time::Duration::from_millis(20)).await; + } + let to = match (now.0 - before.0, now.1 - before.1) { + (1, 0) => "甲", + (0, 1) => "乙", + other => panic!("第 {i} 条连接:{other:?}"), + }; + assert_eq!(to, predicted, "第 {i} 条连接"); + went.push(to); + drop(c); + } + assert_eq!(went, ["甲", "甲", "乙", "甲", "甲", "甲", "乙", "甲"]); +} diff --git a/crates/tw-gateway/src/server/pipeline.rs b/crates/tw-gateway/src/server/pipeline.rs index bee7fb95..87b5e8ec 100644 --- a/crates/tw-gateway/src/server/pipeline.rs +++ b/crates/tw-gateway/src/server/pipeline.rs @@ -678,12 +678,70 @@ fn route( tracing::debug!(skipped = ?serving.skipped, %model, "skipping the candidates that cannot serve this request"); } decision.candidates = serving.usable; + let arranged = arrange( + state, + rt, + &mut decision, + &crate::sent::pairs(&sent), + conv, + now, + ); + if let Some(why) = arranged.stayed { + choice.affinity = Some(tw_api::AffinityView { + held_route, + stayed: Some(why), + }); + // 留下的那一家在头上。满着时等它,不当场跳过(见 `crate::slots`) + choice.stayed_on = decision.candidates.first().cloned(); + } + // `load-balance` 记账:**记粘性之后排头的那一家**,不是按权重轮到的那一家。一段对话 + // 留在了上次回答它的那一家,这一次就算那一家的;之后的新对话把差的补回去 + if let Some(leader) = decision.candidates.first() { + arranged.charge(&decision.candidates, leader); + } + Ok(Routed::Go(choice, decision)) +} + +/// 排好了的候选([`arrange`]):会话粘性留下了哪一家的理由,和 `load-balance` 这一次还没记的账。 +pub(super) struct Arranged<'a> { + /// 排头的是上次回答这段对话的那一家:留下的理由。没留的是 None + pub(super) stayed: Option, + /// `load-balance` 这一轮(拿着它的锁)和排序时的那一份事实。别的组没有 + ledger: Option<(crate::balance::Turn<'a>, tw_engine::Facts)>, +} + +impl Arranged<'_> { + /// `load-balance` 记账:这一次排头的是 `leader`,`members` 是排好的那一份候选(粘性只换了 + /// 次序,没换集合,还是同一轮)。HTTP 的请求记粘性之后排头的那一家,WebSocket 的升级记 + /// 真连上的那一家。不经过 `load-balance` 的什么都不记 + pub(super) fn charge(self, members: &[String], leader: &str) { + if let Some((turn, f)) = self.ledger { + turn.charge(members, &f, leader); + } + } +} + +/// 管线第 2 步的最后:给候选排序、会话粘性。**HTTP 的请求和 WebSocket 的升级共用这一个**(见 +/// `super::upgrade`):同一个组排出同一个顺序,`load-balance` 记同一本账,试算说的就是两条路 +/// 下一个新对话会去的那一家。 +/// +/// `decision.candidates` 进来时是能服务这个请求的那几家,出去时是排好的次序。`sent` 是每一家和 +/// 发给它的名字(比价按它算);`conv` 是这段对话,认不出来的(WebSocket 的升级)没有粘性。 +/// 记账交给调用方([`Arranged::charge`]):从排序到记账一直拿着 `load-balance` 那一组的锁, +/// 同时进来的几个一个接一个地排,后一个看到的是前一个记过的账 +pub(super) fn arrange<'a>( + state: &'a AppState, + rt: &'a Runtime, + decision: &mut tw_engine::Decision, + sent: &[(String, String)], + conv: Option<&crate::affinity::Conversation>, + now: u64, +) -> Arranged<'a> { let group = decision .via_group .as_deref() .and_then(|n| rt.engine.groups().iter().find(|g| g.name == n)); - // `load-balance` 这一次轮到谁(见 `crate::balance`)。**从排序到记账一直拿着锁**: - // 同时进来的几个请求一个接一个地排,后一个看到的是前一个记过的账 + // `load-balance` 这一次轮到谁(见 `crate::balance`) let turn = group .filter(|g| g.kind == tw_engine::GroupType::LoadBalance) .map(|g| state.balance.turn(g)); @@ -701,7 +759,7 @@ fn route( &rt.config.providers, g, &decision.candidates, - &crate::sent::pairs(&sent), + sent, turn.as_ref().map(crate::balance::Turn::current), ); decision.candidates = rt.engine.order(Some(&g.name), &decision.candidates, &f); @@ -709,30 +767,19 @@ fn route( } // 留在上次回答这段对话的那一家:同一轮里一律留,跨轮看缓存值不值得留。**排在 // 策略组排序之后** —— 该留的时候盖过策略,放开的时候策略照常说了算 - if let Some(c) = conv - && let Some(why) = state.affinity.stay( + let stayed = conv.and_then(|c| { + state.affinity.stay( c, decision.via_group.as_deref(), &mut decision.candidates, |p| state.health.is_available(p), now, ) - { - choice.affinity = Some(tw_api::AffinityView { - held_route, - stayed: Some(why), - }); - // 留下的那一家在头上。满着时等它,不当场跳过(见 `crate::slots`) - choice.stayed_on = decision.candidates.first().cloned(); - } - // `load-balance` 记账:**记粘性之后排头的那一家**,不是按权重轮到的那一家。一段对话 - // 留在了上次回答它的那一家,这一次就算那一家的;之后的新对话把差的补回去 - if let (Some(turn), Some(f), Some(leader)) = - (turn, facts_rt.as_ref(), decision.candidates.first()) - { - turn.charge(&decision.candidates, f, leader); + }); + Arranged { + stayed, + ledger: turn.zip(facts_rt), } - Ok(Routed::Go(choice, decision)) } /// 发出开始事件:熔断过滤、出站脱敏看一遍、`RequestStarted`、结局、脱敏的记录、请求体留档。 diff --git a/crates/tw-gateway/src/server/upgrade.rs b/crates/tw-gateway/src/server/upgrade.rs index dc615aa3..256b133e 100644 --- a/crates/tw-gateway/src/server/upgrade.rs +++ b/crates/tw-gateway/src/server/upgrade.rs @@ -13,8 +13,9 @@ use tw_types::msg; /// 接管一次 WebSocket 升级。 /// /// 路由照走一遍 —— **一次升级也是一次请求**,`deny` 规则、熔断对它一样 -/// 有效。之后把连接交给 [`crate::ws::proxy`],那里会在 -/// 每一帧上重新点一遍管线的保护。路由事件也由那边发:选中的那一家接没 +/// 有效,候选的次序和 HTTP 那条路是同一段([`super::pipeline::arrange`]):`load-balance` +/// 按权重轮、记同一本账,`url-test`、`cheapest` 同样排。之后把连接交给 [`crate::ws::proxy`], +/// 那里会在每一帧上重新点一遍管线的保护。路由事件也由那边发:选中的那一家接没 /// 接下,要和它握完手才知道。 /// /// **Responses 的连接上每个 `response.create` 是一个请求**(见 `crate::ws::turn`):连接 @@ -76,7 +77,7 @@ pub(super) async fn ws_upgrade( at_ms: now_ms(), }) }; - let decision = match rt.engine.route(&facts).map_err(|e| { + let mut decision = match rt.engine.route(&facts).map_err(|e| { GatewayError::config(msg!("gw.route.failed", detail = e => "Routing failed: {detail}")) })? { tw_engine::Outcome::Route(d) => d, @@ -100,6 +101,28 @@ pub(super) async fn ws_upgrade( return Err(err); } }; + // 候选的次序:**和 HTTP 那条路同一段**(见 `super::pipeline::arrange`)。先去掉服务不了的 + // (停用的;Realtime 的连接写了模型,还有别名对不上、密钥不让用的),再按组排 —— + // `load-balance` 照权重轮到谁就是谁,不是一律连头一家。升级请求没有正文,认不出是哪段 + // 对话,没有粘性。一家都服务不了的照旧交给下面一家家看,说得出为什么 + let catalog = state.catalog.load(); + let allow = crate::models::key_allow(&rt.config, &client_name); + let asked = rt + .engine + .asked(rt.engine.rules_for_client(&client_name), &facts, &decision); + let sent = crate::sent::plan(&rt.config, &catalog, &decision, &facts.model, &asked, allow); + let serving = crate::sent::serving(&rt.config, &catalog, &decision, &asked, allow); + if !serving.usable.is_empty() { + decision.candidates = serving.usable; + } + let arranged = super::pipeline::arrange( + &state, + &rt, + &mut decision, + &crate::sent::pairs(&sent), + None, + now_ms(), + ); // 发给哪一家、每个 `response.create` 发出去的模型名怎么定:和 HTTP 那条路的一跳同一套, // 按这次的决定定(指定模型、阶段一的改写),别名对到这一家(见 `crate::ws::Naming`) let naming_for = |provider: &tw_config::Provider| crate::ws::Naming { @@ -116,7 +139,6 @@ pub(super) async fn ws_upgrade( let picked = if facts.model.is_empty() { alive.first().map(|s| (s.to_string(), None)) } else { - let catalog = state.catalog.load(); let mut unserved = None; let mut barred = None; let mut picked = None; @@ -176,6 +198,9 @@ pub(super) async fn ws_upgrade( "gw.route.no_upstream_alive" => "No upstream is available." ))); }; + // `load-balance` 记账:记在真连的那一家头上(排在它前面的服务不了这个模型时,不是排头的 + // 那一家) + arranged.charge(&decision.candidates, &name); let Some(provider) = rt.config.providers.iter().find(|p| p.name == name) else { return Err(GatewayError::config(msg!( "gw.route.upstream_missing", upstream = name.clone() => diff --git a/docs/config.md b/docs/config.md index 7c2a9e5b..37678264 100644 --- a/docs/config.md +++ b/docs/config.md @@ -1106,7 +1106,9 @@ stay on the upstream that answers them (see below) and count toward its share, so the balance is kept by where new conversations start. An upstream that is cooling down after failures, is at its `max_concurrent`, or cannot serve a request, sits that request out, and the others share it by their -weights. Other group types take no weights. +weights. A new WebSocket connection is placed the same way and counts as one +request; everything sent on it then goes to the upstream it connected to. +Other group types take no weights. diff --git a/docs/config.zh-CN.md b/docs/config.zh-CN.md index f2d96130..5af24d71 100644 --- a/docs/config.zh-CN.md +++ b/docs/config.zh-CN.md @@ -853,7 +853,7 @@ aliases: 默认类型为 `fallback`:单个使用者的机器上没有需要分散的负载。 -`load-balance` 组的成员可以带权重,取值 1 到 100;只写名字的成员权重为 1。权重决定组内请求怎么分:写成 `{ name: anthropic, weight: 7 }` 和 `relay` 时,每十个请求有七个由官方 API 服务。进行中的对话留在回答它的那一家(见下文),也算进那一家的份额,因此份额靠新对话从哪一家开始来补齐。因失败处于冷却、并发数已满(`max_concurrent`)、或服务不了某个请求的上游不参与这一次分配,其余成员按各自的权重分。其他类型的策略组不用权重。 +`load-balance` 组的成员可以带权重,取值 1 到 100;只写名字的成员权重为 1。权重决定组内请求怎么分:写成 `{ name: anthropic, weight: 7 }` 和 `relay` 时,每十个请求有七个由官方 API 服务。进行中的对话留在回答它的那一家(见下文),也算进那一家的份额,因此份额靠新对话从哪一家开始来补齐。因失败处于冷却、并发数已满(`max_concurrent`)、或服务不了某个请求的上游不参与这一次分配,其余成员按各自的权重分。新的 WebSocket 连接也这样分配,算作一个请求;之后在这条连接上发的都交给它连上的那一家。其他类型的策略组不用权重。 From adadbf53cfe93c10e8e42ffdb94e16d4e5a6180e Mon Sep 17 00:00:00 2001 From: fylorn <249551762+fylorn@users.noreply.github.com> Date: Mon, 5 Oct 2026 20:07:14 +0800 Subject: [PATCH 22/22] Count a Realtime connection's tokens and cost when it closes A Realtime connection is one row, admitted with an empty estimate, and the row carried no usage, so token and cost limits never saw anything spent on it. The upstream reports each answer's usage in response.done (response.usage: input_tokens including cached_tokens under input_token_details, output_tokens). The connection now adds those up into its row; the recorder prices it like any other row and settles it into the key's limits when the connection closes. Limits are still checked when the connection opens. Audio and image tokens are not told apart: the row is priced at the model's per-token price, as other rows are. Other whole-connection paths have no known usage format and still count only as a request. Co-Authored-By: Claude Opus 5.5 --- crates/tw-control/tests/ws_turns.rs | 137 ++++++++++++++++++++++++ crates/tw-gateway/src/ending.rs | 21 +++- crates/tw-gateway/src/server/upgrade.rs | 8 +- crates/tw-gateway/src/ws.rs | 76 +++++++++++-- docs/config.md | 4 +- docs/config.zh-CN.md | 4 +- 6 files changed, 234 insertions(+), 16 deletions(-) diff --git a/crates/tw-control/tests/ws_turns.rs b/crates/tw-control/tests/ws_turns.rs index 77483c73..7553b31a 100644 --- a/crates/tw-control/tests/ws_turns.rs +++ b/crates/tw-control/tests/ws_turns.rs @@ -170,3 +170,140 @@ providers: ); drop(c); } + +/// Realtime 的每一次回答的用量:`response.done` 里的 `usage`。输入 1200(其中 1000 走了缓存, +/// 细分叫 `input_token_details`)、输出 30 +const REALTIME_USAGE: &str = r#"{"total_tokens":1230,"input_tokens":1200,"output_tokens":30,"input_token_details":{"text_tokens":1200,"audio_tokens":0,"cached_tokens":1000,"cached_tokens_details":{"text_tokens":1000,"audio_tokens":0}},"output_token_details":{"text_tokens":30,"audio_tokens":0}}"#; + +/// 像 Realtime 那样回答的上游:每个 `response.create` 回 created、done(带用量) +async fn realtime_upstream() -> SocketAddr { + let app = axum::Router::new().route( + "/v1/realtime", + axum::routing::any(|ws: WebSocketUpgrade| async move { + ws.on_upgrade(|mut sock| async move { + let mut n = 0; + while let Some(Ok(m)) = sock.recv().await { + let Message::Text(_) = m else { continue }; + n += 1; + let usage: serde_json::Value = serde_json::from_str(REALTIME_USAGE).unwrap(); + let id = format!("resp_{n}"); + let frames = [ + serde_json::json!({"type":"response.created","event_id":"e1","response":{"id":id,"object":"realtime.response","status":"in_progress","output":[]}}), + serde_json::json!({"type":"response.done","event_id":"e2","response":{"id":id,"object":"realtime.response","status":"completed","output":[],"usage":usage}}), + ]; + for f in frames { + if sock.send(Message::Text(f.to_string().into())).await.is_err() { + return; + } + } + } + }) + }), + ); + let l = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap(); + let addr = l.local_addr().unwrap(); + tokio::spawn(async move { axum::serve(l, app).await.unwrap() }); + addr +} + +/// Realtime 的连接**整条连接一行**:连上之前看一遍上限,这条连接上每一次回答的用量(上游的 +/// `response.done`)加起来记在这一行上,照 HTTP 那条路查价,**断开时**算进密钥的 token 和费用。 +/// 费用按 gpt-realtime 的价($4/M 输入、$0.4/M 缓存读、$16/M 输出):两次回答,每次 +/// 200 × 4 + 1000 × 0.4 + 30 × 16 = 1680 微美元 +#[tokio::test] +async fn a_realtime_connection_counts_its_usage_when_it_closes() { + let up = realtime_upstream().await; + let d = tempfile::tempdir().unwrap(); + let yaml = format!( + "version: 1 +listen: + control: + key: c0ffee00c0ffee00c0ffee00c0ffee00c0ffee00c0ffee00c0ffee00c0ffee00 +clients: + - name: voice + key: tw-k + limits: + - {{ per: day, tokens: 100000 }} + - {{ per: day, cost: 5 }} +providers: + - name: openai + base_url: http://{up} + key: sk-x + protocol: openai-responses +" + ); + let cfg = tw_config::try_parse(&yaml).unwrap(); + let limits = cfg.clients[0].limits.clone(); + let gw = tw_gateway::AppState::new(cfg).unwrap(); + let (_bodies, rx) = tokio::sync::mpsc::channel(1); + let store = tw_store::task::spawn( + tw_store::Recorder::new( + tw_store::Db::open(&d.path().join("data.db")).unwrap(), + tw_store::Blobs::new(d.path().join("blobs")), + gw.pricing.clone(), + ) + .settling_to(tw_control::key_limits::settle_hook(&gw)), + gw.bus.subscribe(), + rx, + ); + let addr = tw_gateway::serve(gw.clone(), ([127, 0, 0, 1], 0).into()) + .await + .unwrap(); + tokio::time::sleep(Duration::from_millis(50)).await; + + use tokio_tungstenite::tungstenite::client::IntoClientRequest; + let mut req = format!("ws://{addr}/v1/realtime?model=gpt-realtime") + .into_client_request() + .unwrap(); + req.headers_mut() + .insert("authorization", "Bearer tw-k".parse().unwrap()); + let (mut c, _) = tokio_tungstenite::connect_async(req).await.unwrap(); + for _ in 0..2 { + let frame = serde_json::json!({"type": "response.create", "response": {}}); + c.send(tokio_tungstenite::tungstenite::Message::Text( + frame.to_string().into(), + )) + .await + .unwrap(); + loop { + let m = tokio::time::timeout(Duration::from_secs(5), c.next()) + .await + .expect("the answer did not end") + .unwrap() + .unwrap(); + if m.into_text().unwrap().contains("response.done") { + break; + } + } + } + // 连着的时候还没有这一行,用量也还没算进去 + let used = || { + gw.key_limits + .view("voice", &limits) + .iter() + .map(|v| v.used) + .collect::>() + }; + assert_eq!(used(), [0, 0]); + c.close(None).await.unwrap(); + + let deadline = std::time::Instant::now() + Duration::from_secs(10); + let rows = loop { + let rows = store.lock().await.db().recent(None, 10).unwrap(); + if !rows.is_empty() { + break rows; + } + assert!(std::time::Instant::now() < deadline, "the row never came"); + tokio::time::sleep(Duration::from_millis(20)).await; + }; + assert_eq!(rows.len(), 1, "{rows:?}"); + let r = &rows[0]; + assert_eq!(r.path, "/v1/realtime"); + assert_eq!( + (r.input_tokens, r.cache_read_tokens, r.output_tokens), + (Some(400), Some(2000), Some(60)), + "{r:?}" + ); + assert_eq!(r.cost_micros, Some(3360), "{r:?}"); + assert_eq!(used(), [2 * (200 + 30), 3360]); +} diff --git a/crates/tw-gateway/src/ending.rs b/crates/tw-gateway/src/ending.rs index 06fdf16a..dc1e2245 100644 --- a/crates/tw-gateway/src/ending.rs +++ b/crates/tw-gateway/src/ending.rs @@ -57,6 +57,8 @@ pub struct Ending { bytes: u64, /// 旁路嗅探。客户端走掉那一刻手里有多少用量,靠的就是它 sniffer: Sniffer, + /// 一条连接上每一次回答报的用量加起来([`Ending::add_usage`])。有它就不看嗅探器 + total: Option, tap: ResponseTap, /// 认第一个 token 的。**只在上游回的是成功的流时才有**(见 [`Ending::streaming`]), /// 认出来就扔掉 —— 之后的字节不必再解析 @@ -211,6 +213,7 @@ impl Ending { status: None, bytes: 0, sniffer: Sniffer::new(), + total: None, tap: ResponseTap::new(), first: None, opened: None, @@ -363,12 +366,25 @@ impl Ending { /// /// 整条连接一行的 WebSocket 用它(Realtime 和别的路径,见 [`crate::ws`])。一条连接上 /// 跑着好几轮回答,每轮各报一次用量,而嗅探器是「每个字段取最大值」—— 喂给它,得到的 - /// 是其中某一轮的数,看起来却像整条连接的。**与其报一个错的数,不如说没有。** + /// 是其中某一轮的数,看起来却像整条连接的。**与其报一个错的数,不如说没有**:认得出 + /// 每一轮用量的(Realtime 的 `response.done`)由调用方一轮一轮加上([`Ending::add_usage`])。 /// Responses 的连接每一轮各是一个请求,用的是 [`Ending::frame`]。 pub fn count(&mut self, bytes: usize) { self.bytes += bytes as u64; } + /// 一条连接上又一次回答的用量:加到这一行上(Realtime 的连接,见 [`crate::ws`])。结局 + /// 报的是加起来的数,存储层照它查价、算进密钥的用量 + pub fn add_usage(&mut self, u: &Usage) { + let t = self.total.get_or_insert_with(Usage::default); + t.input = t.input.saturating_add(u.input); + t.cache_read = t.cache_read.saturating_add(u.cache_read); + t.cache_write = t.cache_write.saturating_add(u.cache_write); + t.cache_1h |= u.cache_1h; + t.output = t.output.saturating_add(u.output); + t.reasoning = t.reasoning.saturating_add(u.reasoning); + } + /// WebSocket 上上游的一帧文本:Responses 连接上的一轮(见 `crate::ws::turn`)。一条消息 /// 就是一个事件,**按 SSE 的一帧喂**给认第一个 token、嗅用量、看错误的那几样 —— 它们 /// 读的是 SSE;字节只数消息本身。已经是 SSE 形状的(桥接过来的)原样喂 @@ -465,7 +481,8 @@ impl Ending { } let sniffer = std::mem::take(&mut self.sniffer); let model = sniffer.model().map(str::to_string); - (sniffer.finish(), model) + let usage = self.total.take().or_else(|| sniffer.finish()); + (usage, model) } fn duration_ms(&self) -> u64 { diff --git a/crates/tw-gateway/src/server/upgrade.rs b/crates/tw-gateway/src/server/upgrade.rs index 256b133e..aafd5241 100644 --- a/crates/tw-gateway/src/server/upgrade.rs +++ b/crates/tw-gateway/src/server/upgrade.rs @@ -20,7 +20,8 @@ use tw_types::msg; /// /// **Responses 的连接上每个 `response.create` 是一个请求**(见 `crate::ws::turn`):连接 /// 本身不留行,密钥的用量上限、并发上限按轮算,升级时不看。Realtime 和别的路径的连接照旧 -/// 整条连接一行,升级时过一遍用量上限。 +/// 整条连接一行,升级时过一遍用量上限;Realtime 的连接用了多少 token、花了多少,断开时 +/// 那一行记下来才算进去(见 `crate::ws`)。 #[allow(clippy::too_many_arguments)] pub(super) async fn ws_upgrade( state: AppState, @@ -267,7 +268,9 @@ pub(super) async fn ws_upgrade( } } else { // 这把密钥的用量上限:**整条连接算一个请求**,连上之前看一遍,和 HTTP 那条路的准入 - // 同一套(见 `crate::key_limits`)。这一行不带用量,用量的上限只数得到它的请求数 + // 同一套(见 `crate::key_limits`)。Realtime 的连接用了多少 token、花了多少,断开时这 + // 一行记下来才算进去(每一次回答的用量加起来,见 `crate::ws`):连着的时候不占预留 —— + // 一条语音连接用多少,开头估不出来。别的路径的连接不带用量,只数得到它的请求数 let limits = rt .config .clients @@ -300,6 +303,7 @@ pub(super) async fn ws_upgrade( crate::ws::Rows::Connection { id, ending: Box::new(ending), + realtime, } }; // 插件:升级那一刻的那一份表,一条连接用到底。**插件只管 Responses 的 WebSocket**(每个 diff --git a/crates/tw-gateway/src/ws.rs b/crates/tw-gateway/src/ws.rs index 4b0644a2..ee8d69f0 100644 --- a/crates/tw-gateway/src/ws.rs +++ b/crates/tw-gateway/src/ws.rs @@ -32,8 +32,10 @@ //! **Responses 的连接上每个 `response.create` 是一个请求**(见 [`turn`]):从这一帧到这一次 //! 回答完,开始、路由、结局三条事件,存储层记一行,带着这一次回答的用量,照 HTTP 那条路查价; //! 密钥的用量上限、并发上限、这一家的位置(`max_concurrent`)都按轮算,闲着的连接什么都不占。 -//! 连接本身不留行,连不上上游的除外。Realtime 和别的路径的连接照旧**整条连接一行**、不带 -//! 用量:它们的回答不按 Responses 的事件收尾,分不出一轮一轮。 +//! 连接本身不留行,连不上上游的除外。Realtime 和别的路径的连接照旧**整条连接一行**:它们的 +//! 回答不按 Responses 的事件收尾,分不出一轮一轮。Realtime 的每一次回答在 `response.done` 里 +//! 报用量,这一行带着它们加起来的数([`realtime_usage`]),断开时照 HTTP 那条路查价、算进密钥 +//! 的用量;别的路径的连接不知道用量的写法,不带。 //! //! 脚本插件也在这条路上跑(见 [`crate::plugin`]):客户端发来的每个 //! `response.create` 是一次请求。**这条路只有一跳**(升级时就连定了那一家,不换), @@ -631,6 +633,8 @@ struct Pipes { id: u64, /// Responses 的连接上在跑的几轮(见 [`turn`])。别的连接没有 turns: Option, + /// 整条连接一行的 Realtime 连接:每一次回答的用量(`response.done`)加到这一行上 + realtime: bool, /// 范围里可能有插件时才有 plugins: Option, /// 每个 `response.create` 发出去的模型名怎么定。Responses 的连接才有 @@ -660,10 +664,12 @@ impl Pipes { /// 这条连接在流量里怎么记(见 [`turn`])。 pub(crate) enum Rows { - /// 整条连接一行:Realtime 和别的路径的连接。升级时已经开始了,`ending` 是它欠着的结局 + /// 整条连接一行:Realtime 和别的路径的连接。升级时已经开始了,`ending` 是它欠着的结局。 + /// `realtime`:是 Realtime 的连接,这一行带着每一次回答的用量加起来的数 Connection { id: u64, ending: Box, + realtime: bool, }, /// 每一轮一行:Responses 的连接。连接本身不留行 —— 连不上上游的除外,那时按升级的那一刻 /// (`upgraded`:用时从哪一刻算起、那一刻的 Unix 毫秒)补上这一行 @@ -697,14 +703,24 @@ pub(crate) async fn proxy( let name = &upstream.provider.name; // 每一轮一行的连接连上了:连接本身不留行,每一轮各有各的号(见 `turn`)。连不上的补上 // 这一行,和整条连接一行的一样报 - let (id, ending, turns) = match (rows, &connected) { - (Rows::Connection { id, mut ending }, _) => { + let (id, ending, turns, realtime) = match (rows, &connected) { + ( + Rows::Connection { + id, + mut ending, + realtime, + }, + _, + ) => { ending.responded(101); - (id, Some(*ending), None) - } - (Rows::Turns { line, .. }, Ok(_)) => { - (state.bus.next_id(), None, Some(turn::Turns::new(line))) + (id, Some(*ending), None, realtime) } + (Rows::Turns { line, .. }, Ok(_)) => ( + state.bus.next_id(), + None, + Some(turn::Turns::new(line)), + false, + ), (Rows::Turns { line, upgraded }, Err(_)) => { let (id, mut ending) = line.opener.open(turn::Opening { choice: &line.choice, @@ -716,7 +732,7 @@ pub(crate) async fn proxy( at_ms: upgraded.1, }); ending.responded(101); - (id, Some(ending), None) + (id, Some(ending), None, false) } }; // 这一家接没接下这条连接,和 HTTP 那条路一跳的成败记在同一笔账上(见 `crate::health`): @@ -806,6 +822,7 @@ pub(crate) async fn proxy( provider: upstream.provider.name, id, turns, + realtime, plugins, naming, requested_model: String::new(), @@ -1340,6 +1357,14 @@ async fn upstream_text( ending: &mut Option, ) -> Flow { let kind = frame_kind(t); + // Realtime 的一次回答收了尾:它的用量加到这条连接的那一行上。**看的是上游原话**,和 + // 别的路一样(占位符不影响数字) + if p.realtime + && kind.as_deref() == Some("response.done") + && let (Some(e), Some(u)) = (ending.as_mut(), realtime_usage(t)) + { + e.add_usage(&u); + } // 一次回答从开始到收尾的那几帧带着它的 id。只有它们要解第二遍 let response = kind .as_deref() @@ -1372,6 +1397,26 @@ async fn upstream_text( flow } +/// Realtime 的 `response.done` 里这一次回答的用量(`response.usage`)。和 Responses 一样, +/// `input_tokens` 里含着从缓存读的(`cached_tokens`),只是细分叫 `input_token_details`。 +/// 语音、图片的 token 不分开:查价和别的请求一样按 token 的单价算。没有 `usage` 的是 None +fn realtime_usage(frame: &str) -> Option { + let v: serde_json::Value = serde_json::from_str(frame).ok()?; + let u = v.get("response")?.get("usage")?; + let n = |p: &str| { + u.pointer(p) + .and_then(serde_json::Value::as_u64) + .unwrap_or(0) + }; + let cached = n("/input_token_details/cached_tokens"); + Some(tw_dialect::usage::Usage { + input: n("/input_tokens").saturating_sub(cached), + cache_read: cached, + output: n("/output_tokens"), + ..Default::default() + }) +} + /// 上游的这一帧(`type` 是 `kind`)是不是一次回答的结尾:完成、失败、没答完,或者一个错误 /// (没开始回答就出错的,上游只回一个 `error`) fn ends_turn(kind: Option<&str>) -> bool { @@ -1876,6 +1921,17 @@ mod tests { assert_eq!(upstream_url("http://h", "/x", Some("")), "ws://h/x"); } + #[test] + fn a_realtime_answer_reports_its_usage_with_the_cache_read_split_out() { + let done = r#"{"type":"response.done","response":{"id":"r","status":"completed","usage":{"total_tokens":253,"input_tokens":132,"output_tokens":121,"input_token_details":{"text_tokens":119,"audio_tokens":13,"cached_tokens":64},"output_token_details":{"text_tokens":30,"audio_tokens":91}}}}"#; + let u = realtime_usage(done).unwrap(); + assert_eq!((u.input, u.cache_read, u.output), (68, 64, 121)); + assert_eq!( + realtime_usage(r#"{"type":"response.done","response":{"id":"r"}}"#), + None + ); + } + #[test] fn the_realtime_model_is_read_from_and_written_into_the_query() { assert!(realtime("/v1/realtime")); diff --git a/docs/config.md b/docs/config.md index 37678264..a84134c1 100644 --- a/docs/config.md +++ b/docs/config.md @@ -353,7 +353,9 @@ hour limits start empty. On a Responses WebSocket connection, each `response.create` is a request of its own: it is recorded with its usage and cost and checked against these limits, and a refused one is answered with `response.failed` while the -connection stays open. +connection stays open. A Realtime connection (`/v1/realtime`) is one request: +it is checked against the limits when it opens, and the tokens and cost of +all its answers count when it closes. diff --git a/docs/config.zh-CN.md b/docs/config.zh-CN.md index 5af24d71..49507c7c 100644 --- a/docs/config.zh-CN.md +++ b/docs/config.zh-CN.md @@ -251,7 +251,9 @@ clients: 上限从零开始。 Responses 的 WebSocket 连接上,每个 `response.create` 各算一个请求:带着各自的用量和费用记下, -按这些上限检查;被拒的那一个收到 `response.failed`,连接保持不断。 +按这些上限检查;被拒的那一个收到 `response.failed`,连接保持不断。Realtime 的连接 +(`/v1/realtime`)整条算一个请求:连上时按这些上限检查,这条连接上所有回答的 token 和费用在 +断开时计入。