diff --git a/crates/tw-gateway/src/lib.rs b/crates/tw-gateway/src/lib.rs index f3b64680..293977b7 100644 --- a/crates/tw-gateway/src/lib.rs +++ b/crates/tw-gateway/src/lib.rs @@ -55,6 +55,14 @@ pub use server::{router, serve}; pub use state::credential_failed; pub use state::{AppState, Runtime}; +/// Anthropic 流里上游静默多久补一个 `ping`。 +/// +/// Claude Code(和内嵌它的 Claude Desktop)数的是网关发来的每一个字节:一条流静默 +/// 五分钟就被放弃,只有 ping 在来的话还肯多等一段。Anthropic 上游思考时自己会发 +/// ping,**转换别的上游时没有人发**,长时间的推理就会被客户端当成断线。十五秒远低于 +/// 那个上限,又不至于让一条正常的流塞满心跳 +pub const PING_EVERY: std::time::Duration = std::time::Duration::from_secs(15); + /// 请求来自谁。**如实写 ThinkWatch** —— 我们从不把自己报成别的客户端。 pub const ORIGINATOR: &str = "thinkwatch"; diff --git a/crates/tw-gateway/src/server.rs b/crates/tw-gateway/src/server.rs index 7a5e2ac8..016e16cf 100644 --- a/crates/tw-gateway/src/server.rs +++ b/crates/tw-gateway/src/server.rs @@ -22,6 +22,10 @@ mod upgrade; pub fn router(state: AppState) -> Router { Router::new() .route("/healthz", get(|| async { "ok" })) + // Claude Code(和内嵌它的 Claude Desktop)启动时发一个 `HEAD /api/hello` 预热 + // 连接。**不鉴权、不转发**:它不带密钥也不需要上游,掉进透传的话会被当成一次 + // 没有密钥的请求拒成 401。`get` 也接 HEAD + .route("/api/hello", get(|| async {})) // **和准入共用同一个函数** —— 列表和准入不可能不一致。 .route("/v1/models", get(list_models)) // **单点查询要走同一道准入**。不接这条的话它掉进 diff --git a/crates/tw-gateway/src/server/listing.rs b/crates/tw-gateway/src/server/listing.rs index e05ff162..c989b94a 100644 --- a/crates/tw-gateway/src/server/listing.rs +++ b/crates/tw-gateway/src/server/listing.rs @@ -90,8 +90,14 @@ pub(super) async fn list_models( "name": format!("models/{m}"), })).collect::>() }), - // Anthropic 和 OpenAI 的 /v1/models 形状一样 - _ => serde_json::json!({ + ListingShape::Anthropic => serde_json::json!({ + "object": "list", + "data": models.iter().map(|m| anthropic_model(m, now)).collect::>(), + "has_more": false, + "first_id": models.first(), + "last_id": models.last(), + }), + ListingShape::Openai => serde_json::json!({ "object": "list", "data": models.iter().map(|m| serde_json::json!({ "id": m, "object": "model", "created": now, @@ -101,6 +107,88 @@ pub(super) async fn list_models( Ok(axum::Json(body).into_response()) } +/// Anthropic 格式的一个模型对象。 +/// +/// **是 OpenAI 那个对象的超集**:Anthropic 的字段(`type`、`display_name`、 +/// `created_at`)之外,`object` 和 `created` 照样在。只放 `x-api-key` 的客户端也被 +/// 认成 Anthropic,其中有按 OpenAI 的形状读列表的,不能让它们读不出来。 +/// +/// Claude 的模型再带上 `anthropic_family_tier`:Claude Desktop 按它把模型归到 +/// opus / sonnet / haiku,配置里写的 `sonnet` 这样的简称靠它解析。**只看模型名**, +/// 名字里看不出是 Claude 的一律不标 —— 把别家的模型标成 Claude 是在替客户端撒谎。 +fn anthropic_model(id: &str, now: u64) -> serde_json::Value { + let created_at = chrono::DateTime::from_timestamp(now as i64, 0) + .unwrap_or_default() + .to_rfc3339_opts(chrono::SecondsFormat::Secs, true); + let mut m = serde_json::json!({ + "type": "model", + "id": id, + "display_name": display_name(id).unwrap_or_else(|| id.to_string()), + "created_at": created_at, + "object": "model", + "created": now, + }); + if let Some(tier) = family_tier(id) { + m["anthropic_family_tier"] = tier.into(); + } + m +} + +/// Claude 模型的名字:`claude-sonnet-4-5-20250929` → `Claude Sonnet 4.5`。 +/// +/// 只认最后一段(`/` 之后)以 `claude-` 开头、每一节都是字母或数字的;日期那一节 +/// 去掉,相邻的数字用点连起来。**认不出就是 `None`**,调用方用模型 ID 本身 —— +/// 客户端看到和 ID 一样的名字时会自己想办法,一个猜错的名字它却会照着显示。 +fn display_name(id: &str) -> Option { + let last = id.rsplit('/').next()?; + let rest = last.strip_prefix("claude-")?; + let mut words: Vec = vec!["Claude".into()]; + let mut number = false; + for part in rest.split('-') { + if part.is_empty() { + return None; + } + if part.len() == 8 && part.bytes().all(|b| b.is_ascii_digit()) { + // 发布日期,不是名字的一部分 + continue; + } + if part.bytes().all(|b| b.is_ascii_digit() || b == b'.') { + match words.last_mut() { + Some(w) if number => { + w.push('.'); + w.push_str(part); + } + _ => words.push(part.to_string()), + } + number = true; + } else if part.bytes().all(|b| b.is_ascii_alphanumeric()) { + let mut c = part.chars(); + let first = c.next()?.to_ascii_uppercase(); + words.push(std::iter::once(first).chain(c).collect()); + number = false; + } else { + return None; + } + } + (words.len() > 1).then(|| words.join(" ")) +} + +/// 名字里看得出是哪一档的 Claude 模型:`opus`、`sonnet` 或 `haiku`。 +fn family_tier(id: &str) -> Option<&'static str> { + let lower = id.to_ascii_lowercase(); + if !lower.contains("claude") && !lower.contains("anthropic") { + return None; + } + let mut tiers = ["opus", "sonnet", "haiku"] + .into_iter() + .filter(|t| lower.contains(t)); + // 名字里同时出现两档的(一个路由别名)不猜 + match (tiers.next(), tiers.next()) { + (Some(t), None) => Some(t), + _ => None, + } +} + /// `GET /v1/models/:model`。 /// /// **不许可就当它不存在(404),不是 403。**回 403 等于告诉对方 @@ -146,7 +234,56 @@ pub(super) async fn get_model( ListingShape::Gemini => { serde_json::json!({ "name": format!("models/{model}") }) } - _ => serde_json::json!({ "id": model, "object": "model", "created": now }), + ListingShape::Anthropic => anthropic_model(&model, now), + ListingShape::Openai => { + serde_json::json!({ "id": model, "object": "model", "created": now }) + } }; Ok(axum::Json(body).into_response()) } + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn claude_model_ids_get_their_names() { + for (id, name) in [ + ("claude-sonnet-4-5", "Claude Sonnet 4.5"), + ("claude-sonnet-4-5-20250929", "Claude Sonnet 4.5"), + ("claude-opus-4-1-20250805", "Claude Opus 4.1"), + ("claude-3-5-haiku-20241022", "Claude 3.5 Haiku"), + ("claude-fable-5-1", "Claude Fable 5.1"), + ("anthropic/claude-sonnet-4.5", "Claude Sonnet 4.5"), + ] { + assert_eq!(display_name(id).as_deref(), Some(name), "{id}"); + } + } + + #[test] + fn a_name_that_cannot_be_read_is_not_guessed() { + for id in [ + "deepseek-chat", + "gpt-5", + "claude-", + "claude--x", + "us.anthropic.claude-sonnet-4-5-20250929-v1:0", + ] { + assert_eq!(display_name(id), None, "{id}"); + } + } + + #[test] + fn only_claude_models_are_given_a_family_tier() { + assert_eq!(family_tier("claude-opus-4-1"), Some("opus")); + assert_eq!(family_tier("claude-3-5-haiku-20241022"), Some("haiku")); + assert_eq!( + family_tier("us.anthropic.claude-sonnet-4-5-20250929-v1:0"), + Some("sonnet") + ); + assert_eq!(family_tier("claude-fable-5-1"), None); + // 名字里有 sonnet,但看不出是 Claude + assert_eq!(family_tier("my-sonnet-alias"), None); + assert_eq!(family_tier("claude-opus-or-sonnet"), None); + } +} diff --git a/crates/tw-gateway/src/server/pipeline/relay.rs b/crates/tw-gateway/src/server/pipeline/relay.rs index 1b8aa7c5..7169bdd3 100644 --- a/crates/tw-gateway/src/server/pipeline/relay.rs +++ b/crates/tw-gateway/src/server/pipeline/relay.rs @@ -89,6 +89,12 @@ pub(super) fn respond( let mut relay = Relay::new(state, rt, req, plan, &ledger, session, provider, id); let chunks = upstream.bytes_stream(); let dialect = req.dialect; + // 上游静默时补心跳:**只给 Anthropic Messages 的流**。那是客户端按字节计时、 + // 认得 `ping` 事件的格式;别的格式没有这个事件,塞进去就是一帧解析不了的东西 + let ping_every = (req.api == Some(crate::client_api::ClientApi::AnthropicMessages) + && plan.client_sse + && status.is_success()) + .then_some(state.ping_every); let stream = async_stream::stream! { // **通行证跟着响应体走。**这个流被丢掉的时候它才还回去:正常 // 发完是一种,客户端中途断开、hyper 丢掉响应体是另一种 —— 两种 @@ -98,7 +104,24 @@ pub(super) fn respond( let mut ending = ending; let mut chunks = std::pin::pin!(chunks); let mut broke: Option = None; - while let Some(item) = chunks.next().await { + loop { + let next = match ping_every { + None => chunks.next().await, + // `next()` 被超时丢掉不丢数据:它只是去问一次流,没拿走任何东西 + Some(every) => match tokio::time::timeout(every, chunks.next()).await { + Ok(next) => next, + Err(_) => { + // **不经过留档、计量和审查**:心跳不是上游说的话,不进请求记录, + // 也不算输出。**只在帧的边界上插**,上游停在一帧中间时插进去 + // 会把那一帧拆坏 —— 那时宁可不补 + if relay.between_frames() { + yield Ok::(Bytes::from_static(PING)); + } + continue; + } + }, + }; + let Some(item) = next else { break }; match item { Ok(chunk) => { // **旁路嗅探和留档,不缓冲**:字节照常流向客户端,同时 @@ -161,6 +184,9 @@ pub(super) fn respond( resp } +/// Anthropic 的心跳帧,和它自己的 API 发的一样。 +const PING: &[u8] = b"event: ping\ndata: {\"type\": \"ping\"}\n\n"; + /// 响应头到手时就定下的处理方式。 #[derive(Clone, Copy)] struct Plan { @@ -317,6 +343,8 @@ struct Relay { /// **切断时要把数组收好**(见 [`Relay::error_tail`]) array_opened: bool, array_element: bool, + /// 发给客户端的最后一段停在帧的边界上(或者还什么都没发)。心跳只能插在这里 + at_boundary: bool, bus: tw_observe::EventBus, id: u64, provider: String, @@ -396,6 +424,7 @@ impl Relay { client_dialect, array_opened: false, array_element: false, + at_boundary: true, bus: state.bus.clone(), id, provider: provider.name.clone(), @@ -601,9 +630,13 @@ impl Relay { (tail, None) } - /// 记下发给客户端的这一段。**只有直通的 JSON 数组流要记**:切断的位置总在元素 - /// 边界上(分隔符算在后面那个元素上),所以只要知道 `[` 之后有没有过 `{` + /// 记下发给客户端的这一段:停没停在帧的边界上(心跳要看)。直通的 JSON 数组流 + /// 还要记数组发到哪儿了:切断的位置总在元素边界上(分隔符算在后面那个元素上), + /// 所以只要知道 `[` 之后有没有过 `{` fn sent(&mut self, out: &[u8]) { + if !out.is_empty() { + self.at_boundary = out.ends_with(b"\n\n") || out.ends_with(b"\r\n\r\n"); + } if !self.plan.client_json_stream || self.session.is_some() { return; } @@ -619,6 +652,11 @@ impl Relay { } } + /// 现在插一帧会不会拆坏客户端正在收的那一帧。 + fn between_frames(&self) -> bool { + self.at_boundary + } + /// 流断了之后还能对客户端说的最后一句:按它收到的格式收尾。 fn error_tail(&mut self, err: &GatewayError) -> Option> { if let Some(c) = self.back.as_mut() { diff --git a/crates/tw-gateway/src/state.rs b/crates/tw-gateway/src/state.rs index ae102f17..07ccaa66 100644 --- a/crates/tw-gateway/src/state.rs +++ b/crates/tw-gateway/src/state.rs @@ -204,6 +204,9 @@ pub struct AppState { pub live: crate::live::Live, /// 每段对话此刻归到哪一次会话(见 [`crate::session::Sessions`])。**跨重载存活** pub sessions: Arc, + /// Anthropic 流里上游静默多久就补一个 `ping`(见 `relay`)。**测试会把它调短**, + /// 否则一条心跳的测试要干等十五秒 + pub ping_every: std::time::Duration, } impl AppState { @@ -252,6 +255,7 @@ impl AppState { proxies: Arc::new(std::sync::Mutex::new(Default::default())), live: crate::live::Live::default(), sessions: Default::default(), + ping_every: crate::PING_EVERY, }; // 手写的清单马上可用;向上游问是后台的事,不挡启动 state.publish_catalog(); diff --git a/crates/tw-gateway/tests/claude_desktop.rs b/crates/tw-gateway/tests/claude_desktop.rs new file mode 100644 index 00000000..a566ccfd --- /dev/null +++ b/crates/tw-gateway/tests/claude_desktop.rs @@ -0,0 +1,389 @@ +//! Claude Desktop(第三方推理模式)和它内嵌的 Claude Code 对网关的要求,端到端。 +//! +//! - 启动时的 `HEAD /api/hello` 预热:不要密钥,不打上游 +//! - 推理请求打到 `/v1/messages?beta=true`,`anthropic-beta` 和 `cache_control` 原样到上游 +//! - `/v1/messages/count_tokens` 照样转给 Anthropic 上游 +//! - `/v1/models` 回 Anthropic 的列表格式,Claude 模型带上名字和档位 +//! - 上游静默时,Anthropic 流里补 `ping`:客户端按字节计时,五分钟没有字节就放弃 + +use std::net::SocketAddr; +use std::sync::{Arc, Mutex}; +use std::time::Duration; + +use axum::Router; +use axum::extract::{OriginalUri, State}; +use axum::http::{HeaderMap, Method}; +use serde_json::{Value, json}; +use tw_config::{Client, Config, Listen, Protocol, Provider}; + +#[derive(Default, Debug)] +struct Seen { + method: Option, + uri: String, + headers: HeaderMap, + body: Vec, +} + +/// 记下收到的请求、回一个 JSON 的假上游。 +async fn recording_upstream() -> (SocketAddr, Arc>) { + let seen = Arc::new(Mutex::new(Seen::default())); + let app = Router::new() + .fallback( + |State(s): State>>, + method: Method, + OriginalUri(uri): OriginalUri, + headers: HeaderMap, + body: bytes::Bytes| async move { + *s.lock().unwrap() = Seen { + method: Some(method), + uri: uri.to_string(), + headers, + body: body.to_vec(), + }; + axum::response::Response::builder() + .header("content-type", "application/json") + .body(axum::body::Body::from(r#"{"input_tokens":3}"#)) + .unwrap() + }, + ) + .with_state(seen.clone()); + (serve(app).await, seen) +} + +/// 按顺序发出 `parts` 的 SSE 上游。`None` 是一段静默。 +async fn stalling_upstream(parts: Vec>, pause: Duration) -> SocketAddr { + let app = Router::new().fallback(move || { + let parts = parts.clone(); + async move { + let body = async_stream::stream! { + for p in parts { + match p { + Some(text) => yield Ok::<_, std::io::Error>(bytes::Bytes::from_static(text.as_bytes())), + None => tokio::time::sleep(pause).await, + } + } + }; + axum::response::Response::builder() + .header("content-type", "text/event-stream") + .body(axum::body::Body::from_stream(body)) + .unwrap() + } + }); + serve(app).await +} + +async fn serve(app: Router) -> SocketAddr { + let l = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap(); + let addr = l.local_addr().unwrap(); + tokio::spawn(async move { axum::serve(l, app).await.unwrap() }); + addr +} + +fn provider(up: SocketAddr, protocol: Protocol, models: &[&str]) -> Provider { + Provider { + name: "up".into(), + base_url: format!("http://{up}"), + key: Some("sk-upstream".into()), + protocol: Some(protocol), + models: models.iter().map(|m| m.to_string()).collect(), + ..Default::default() + } +} + +struct Gateway { + addr: SocketAddr, + events: tokio::sync::broadcast::Receiver, + bodies: tokio::sync::mpsc::Receiver, +} + +async fn gateway(p: Provider) -> Gateway { + let cfg = Config { + version: 1, + listen: Listen::default(), + clients: vec![Client { + name: "desktop".into(), + key: "tw-k".into(), + ..Default::default() + }], + providers: vec![p], + ..Default::default() + }; + let mut state = tw_gateway::AppState::new(cfg).unwrap(); + // 心跳的间隔调短,一条测试不必干等十五秒 + state.ping_every = Duration::from_millis(100); + let events = state.bus.subscribe(); + let (tx, bodies) = tokio::sync::mpsc::channel(16); + state.set_body_sink(tx); + let addr = tw_gateway::serve(state, ([127, 0, 0, 1], 0).into()) + .await + .unwrap(); + tokio::time::sleep(Duration::from_millis(50)).await; + Gateway { + addr, + events, + bodies, + } +} + +#[tokio::test] +async fn the_startup_probe_is_answered_without_a_key_or_the_upstream() { + let (up, seen) = recording_upstream().await; + let gw = gateway(provider(up, Protocol::Anthropic, &[])).await; + let http = reqwest::Client::new(); + for method in [Method::HEAD, Method::GET] { + let resp = http + .request(method.clone(), format!("http://{}/api/hello", gw.addr)) + .send() + .await + .unwrap(); + assert_eq!(resp.status(), 200, "{method}"); + } + assert!( + seen.lock().unwrap().method.is_none(), + "预热探测被转给了上游" + ); +} + +#[tokio::test] +async fn a_desktop_request_reaches_the_upstream_as_it_was_sent() { + let (up, seen) = recording_upstream().await; + let gw = gateway(provider(up, Protocol::Anthropic, &["claude-sonnet-4-5"])).await; + let body = json!({ + "model": "claude-sonnet-4-5", + "max_tokens": 16, + "system": [{"type": "text", "text": "be brief", "cache_control": {"type": "ephemeral"}}], + "messages": [{"role": "user", "content": [ + {"type": "text", "text": "hi", "cache_control": {"type": "ephemeral", "ttl": "1h"}} + ]}], + }); + let beta = "extended-cache-ttl-2025-04-11,some-future-beta-2099-01-01"; + let resp = reqwest::Client::new() + .post(format!("http://{}/v1/messages?beta=true", gw.addr)) + .header("authorization", "Bearer tw-k") + .header("anthropic-version", "2023-06-01") + .header("anthropic-beta", beta) + .json(&body) + .send() + .await + .unwrap(); + assert_eq!(resp.status(), 200, "{}", resp.text().await.unwrap()); + let g = seen.lock().unwrap(); + assert_eq!(g.uri, "/v1/messages?beta=true"); + assert_eq!(g.headers.get("anthropic-beta").unwrap(), beta); + assert_eq!(g.headers.get("anthropic-version").unwrap(), "2023-06-01"); + let sent: Value = serde_json::from_slice(&g.body).unwrap(); + assert_eq!(sent["system"], body["system"], "cache_control 被动过"); + assert_eq!(sent["messages"], body["messages"], "cache_control 被动过"); +} + +#[tokio::test] +async fn counting_tokens_goes_to_the_anthropic_upstream() { + let (up, seen) = recording_upstream().await; + let gw = gateway(provider(up, Protocol::Anthropic, &["claude-sonnet-4-5"])).await; + let resp = reqwest::Client::new() + .post(format!( + "http://{}/v1/messages/count_tokens?beta=true", + gw.addr + )) + .header("x-api-key", "tw-k") + .header("anthropic-version", "2023-06-01") + .header("anthropic-beta", "token-counting-2024-11-01") + .json( + &json!({"model": "claude-sonnet-4-5", "messages": [{"role": "user", "content": "hi"}]}), + ) + .send() + .await + .unwrap(); + assert_eq!(resp.status(), 200); + let v: Value = resp.json().await.unwrap(); + assert_eq!(v["input_tokens"], 3); + let g = seen.lock().unwrap(); + assert_eq!(g.uri, "/v1/messages/count_tokens?beta=true"); + assert_eq!( + g.headers.get("anthropic-beta").unwrap(), + "token-counting-2024-11-01" + ); +} + +#[tokio::test] +async fn models_are_listed_in_the_anthropic_shape() { + let (up, _) = recording_upstream().await; + let gw = gateway(provider( + up, + Protocol::Anthropic, + &["claude-sonnet-4-5-20250929", "my-alias"], + )) + .await; + let get = |path: &'static str| { + reqwest::Client::new() + .get(format!("http://{}{path}", gw.addr)) + .header("x-api-key", "tw-k") + .header("anthropic-version", "2023-06-01") + .send() + }; + let list: Value = get("/v1/models?limit=1000") + .await + .unwrap() + .json() + .await + .unwrap(); + assert_eq!(list["has_more"], false); + assert_eq!(list["first_id"], "claude-sonnet-4-5-20250929"); + assert_eq!(list["last_id"], "my-alias"); + let data = list["data"].as_array().unwrap(); + assert_eq!(data.len(), 2, "{list}"); + let claude = &data[0]; + assert_eq!(claude["type"], "model"); + assert_eq!(claude["id"], "claude-sonnet-4-5-20250929"); + assert_eq!(claude["display_name"], "Claude Sonnet 4.5"); + assert_eq!(claude["anthropic_family_tier"], "sonnet"); + assert!( + chrono::DateTime::parse_from_rfc3339(claude["created_at"].as_str().unwrap()).is_ok(), + "{claude}" + ); + // OpenAI 那几个字段照样在 + assert_eq!(claude["object"], "model"); + assert!(claude["created"].is_u64()); + // 看不出是 Claude 的:名字就是 ID,也不标档位 + let alias = &data[1]; + assert_eq!(alias["display_name"], "my-alias"); + assert!(alias.get("anthropic_family_tier").is_none(), "{alias}"); + + let one: Value = get("/v1/models/claude-sonnet-4-5-20250929") + .await + .unwrap() + .json() + .await + .unwrap(); + assert_eq!(one, *claude, "单点查询和列表里的应该是同一个对象"); +} + +#[tokio::test] +async fn an_openai_client_still_gets_the_openai_listing() { + let (up, _) = recording_upstream().await; + let gw = gateway(provider(up, Protocol::Anthropic, &["claude-sonnet-4-5"])).await; + let list: Value = reqwest::Client::new() + .get(format!("http://{}/v1/models", gw.addr)) + .header("authorization", "Bearer tw-k") + .send() + .await + .unwrap() + .json() + .await + .unwrap(); + assert_eq!(list["object"], "list"); + let m = &list["data"][0]; + assert_eq!(m["id"], "claude-sonnet-4-5"); + assert!(m.get("type").is_none(), "{m}"); + assert!(list.get("has_more").is_none(), "{list}"); +} + +const CHAT_FIRST: &str = "data: {\"id\":\"c1\",\"object\":\"chat.completion.chunk\",\"model\":\"gpt-x\",\"choices\":[{\"index\":0,\"delta\":{\"role\":\"assistant\",\"content\":\"hel\"}}]}\n\n"; +const CHAT_REST: &str = "data: {\"id\":\"c1\",\"object\":\"chat.completion.chunk\",\"model\":\"gpt-x\",\"choices\":[{\"index\":0,\"delta\":{\"content\":\"lo\"},\"finish_reason\":\"stop\"}]}\n\n\ +data: {\"id\":\"c1\",\"object\":\"chat.completion.chunk\",\"model\":\"gpt-x\",\"choices\":[],\"usage\":{\"prompt_tokens\":5,\"completion_tokens\":2,\"total_tokens\":7}}\n\n\ +data: [DONE]\n\n"; + +async fn stream_from(gw: &Gateway, path: &str, body: Value) -> String { + let resp = reqwest::Client::new() + .post(format!("http://{}{path}", gw.addr)) + .header("authorization", "Bearer tw-k") + .json(&body) + .send() + .await + .unwrap(); + assert_eq!(resp.status(), 200); + resp.text().await.unwrap() +} + +#[tokio::test] +async fn a_silent_converted_upstream_is_covered_with_pings() { + let up = stalling_upstream( + vec![Some(CHAT_FIRST), None, Some(CHAT_REST)], + Duration::from_millis(700), + ) + .await; + let mut gw = gateway(provider(up, Protocol::OpenaiChat, &[])).await; + let text = stream_from( + &gw, + "/v1/messages", + json!({"model": "gpt-x", "max_tokens": 16, "stream": true, + "messages": [{"role": "user", "content": "hi"}]}), + ) + .await; + let ping = "event: ping\ndata: {\"type\": \"ping\"}\n\n"; + assert!(text.matches(ping).count() >= 2, "{text}"); + // 心跳插在帧之间:去掉它们,剩下的是一条完整的 Anthropic 流 + let rest = text.replace(ping, ""); + assert!(rest.starts_with("event: message_start"), "{rest}"); + assert!(rest.contains("\"text\":\"hel\""), "{rest}"); + assert!(rest.contains("\"text\":\"lo\""), "{rest}"); + assert!( + rest.trim_end().ends_with("\"type\":\"message_stop\"}"), + "{rest}" + ); + assert!(!rest.contains("ping"), "心跳拆进了别的帧:{rest}"); + + // 用量照上游报的记;请求记录里存的是上游的原话,没有心跳 + let usage = loop { + match gw.events.recv().await.unwrap() { + tw_api::Event::RequestFinished { usage, .. } => break usage, + tw_api::Event::RequestFailed { message, .. } => panic!("{message:?}"), + _ => {} + } + }; + let usage = usage.expect("上游报了用量"); + assert_eq!((usage.input, usage.output), (5, 2)); + let mut recorded = String::new(); + while let Ok(Some(rec)) = + tokio::time::timeout(Duration::from_millis(200), gw.bodies.recv()).await + { + recorded.push_str(&String::from_utf8_lossy(&rec.body)); + } + assert!(recorded.contains("\"content\":\"lo\""), "{recorded}"); + assert!(!recorded.contains("ping"), "心跳进了请求记录:{recorded}"); +} + +#[tokio::test] +async fn a_ping_never_splits_a_frame_the_upstream_left_half_sent() { + const START: &str = "event: message_start\ndata: {\"type\":\"message_start\",\"message\":{\"id\":\"m1\",\"type\":\"message\",\"role\":\"assistant\",\"model\":\"claude-sonnet-4-5\",\"content\":[],\"usage\":{\"input_tokens\":5,\"output_tokens\":1}}}\n\n"; + const HALF: &str = "event: content_block_start\ndata: {\"type\":\"content_bl"; + const REST: &str = "ock_start\",\"index\":0,\"content_block\":{\"type\":\"text\",\"text\":\"\"}}\n\n\ +event: message_stop\ndata: {\"type\":\"message_stop\"}\n\n"; + let up = stalling_upstream( + vec![Some(START), None, Some(HALF), None, Some(REST)], + Duration::from_millis(500), + ) + .await; + let gw = gateway(provider(up, Protocol::Anthropic, &[])).await; + let text = stream_from( + &gw, + "/v1/messages", + json!({"model": "claude-sonnet-4-5", "max_tokens": 16, "stream": true, + "messages": [{"role": "user", "content": "hi"}]}), + ) + .await; + let ping = "event: ping\ndata: {\"type\": \"ping\"}\n\n"; + // 第一段静默在帧之间,补了;第二段停在一帧中间,没有补 + assert!(text.contains(&format!("{START}{ping}")), "{text}"); + assert!(text.contains(&format!("{HALF}{REST}")), "{text}"); + assert_eq!(text.replace(ping, ""), format!("{START}{HALF}{REST}")); +} + +#[tokio::test] +async fn other_formats_get_no_anthropic_pings() { + let up = stalling_upstream( + vec![Some(CHAT_FIRST), None, Some(CHAT_REST)], + Duration::from_millis(500), + ) + .await; + let gw = gateway(provider(up, Protocol::OpenaiChat, &[])).await; + let text = stream_from( + &gw, + "/v1/chat/completions", + json!({"model": "gpt-x", "stream": true, + "messages": [{"role": "user", "content": "hi"}]}), + ) + .await; + assert!(!text.contains("ping"), "{text}"); + assert_eq!(text, format!("{CHAT_FIRST}{CHAT_REST}")); +}