diff --git a/src-tauri/Cargo.lock b/src-tauri/Cargo.lock index 220682a30..ff0024d0c 100644 --- a/src-tauri/Cargo.lock +++ b/src-tauri/Cargo.lock @@ -137,18 +137,6 @@ dependencies = [ "pin-project-lite", ] -[[package]] -name = "async-compression" -version = "0.4.41" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "d0f9ee0f6e02ffd7ad5816e9464499fba7b3effd01123b515c41d1697c43dad1" -dependencies = [ - "compression-codecs", - "compression-core", - "pin-project-lite", - "tokio", -] - [[package]] name = "async-executor" version = "1.14.0" @@ -324,6 +312,28 @@ version = "1.5.0" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "c08606f8c3cbf4ce6ec8e28fb0014a2c086708fe954eaa885384a6165172e7e8" +[[package]] +name = "aws-lc-rs" +version = "1.15.3" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "e84ce723ab67259cfeb9877c6a639ee9eb7a27b28123abd71db7f0d5d0cc9d86" +dependencies = [ + "aws-lc-sys", + "zeroize", +] + +[[package]] +name = "aws-lc-sys" +version = "0.36.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "43a442ece363113bd4bd4c8b18977a7798dd4d3c3383f34fb61936960e8f4ad8" +dependencies = [ + "cc", + "cmake", + "dunce", + "fs_extra", +] + [[package]] name = "axum" version = "0.7.9" @@ -496,6 +506,17 @@ dependencies = [ "syn 2.0.117", ] +[[package]] +name = "brotli" +version = "7.0.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "cc97b8f16f944bba54f0433f07e30be199b6dc2bd25937444bbad560bcea29bd" +dependencies = [ + "alloc-no-stdlib", + "alloc-stdlib", + "brotli-decompressor 4.0.3", +] + [[package]] name = "brotli" version = "8.0.2" @@ -504,7 +525,17 @@ checksum = "4bd8b9603c7aa97359dbd97ecf258968c95f3adddd6db2f7e7a5bef101c84560" dependencies = [ "alloc-no-stdlib", "alloc-stdlib", - "brotli-decompressor", + "brotli-decompressor 5.0.0", +] + +[[package]] +name = "brotli-decompressor" +version = "4.0.3" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "a334ef7c9e23abf0ce748e8cd309037da93e606ad52eb372e4ce327a0dcfbdfd" +dependencies = [ + "alloc-no-stdlib", + "alloc-stdlib", ] [[package]] @@ -691,11 +722,19 @@ dependencies = [ "auto-launch", "axum", "base64 0.22.1", + "brotli 7.0.0", "bytes", "chrono", "dirs 5.0.1", + "flate2", "futures", + "http", + "http-body", + "http-body-util", + "httparse", "hyper", + "hyper-rustls", + "hyper-util", "indexmap 2.13.0", "json-five", "json5", @@ -708,6 +747,8 @@ dependencies = [ "rquickjs", "rusqlite", "rust_decimal", + "rustls", + "rustls-native-certs", "serde", "serde_json", "serde_yaml", @@ -726,6 +767,7 @@ dependencies = [ "tempfile", "thiserror 2.0.18", "tokio", + "tokio-rustls", "toml 0.8.23", "toml_edit 0.22.27", "tower 0.4.13", @@ -733,6 +775,7 @@ dependencies = [ "url", "uuid", "webkit2gtk", + "webpki-roots 0.26.11", "winreg 0.52.0", "zip 2.4.2", ] @@ -800,6 +843,15 @@ dependencies = [ "inout", ] +[[package]] +name = "cmake" +version = "0.1.57" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "75443c44cd6b379beb8c5b45d85d0773baf31cce901fe7bb252f4eff3008ef7d" +dependencies = [ + "cc", +] + [[package]] name = "combine" version = "4.6.7" @@ -810,23 +862,6 @@ dependencies = [ "memchr", ] -[[package]] -name = "compression-codecs" -version = "0.4.37" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "eb7b51a7d9c967fc26773061ba86150f19c50c0d65c887cb1fbe295fd16619b7" -dependencies = [ - "compression-core", - "flate2", - "memchr", -] - -[[package]] -name = "compression-core" -version = "0.4.31" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "75984efb6ed102a0d42db99afb6c1948f0380d1d91808d5529916e6c08b49d8d" - [[package]] name = "concurrent-queue" version = "2.5.0" @@ -1573,6 +1608,12 @@ dependencies = [ "percent-encoding", ] +[[package]] +name = "fs_extra" +version = "1.3.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "42703706b716c37f96a77aea830392ad231f44c9e9a67872fa5548707e11b11c" + [[package]] name = "funty" version = "2.0.0" @@ -2206,12 +2247,14 @@ dependencies = [ "http", "hyper", "hyper-util", + "log", "rustls", + "rustls-native-certs", "rustls-pki-types", "tokio", "tokio-rustls", "tower-service", - "webpki-roots", + "webpki-roots 1.0.6", ] [[package]] @@ -4200,7 +4243,7 @@ dependencies = [ "wasm-bindgen-futures", "wasm-streams 0.4.2", "web-sys", - "webpki-roots", + "webpki-roots 1.0.6", ] [[package]] @@ -4410,6 +4453,8 @@ version = "0.23.37" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "758025cb5fccfd3bc2fd74708fd4682be41d99e5dff73c377c0646c6012c73a4" dependencies = [ + "aws-lc-rs", + "log", "once_cell", "ring", "rustls-pki-types", @@ -4473,6 +4518,7 @@ version = "0.103.9" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "d7df23109aa6c1567d1c575b9952556388da57401e4ace1d15f79eedad0d8f53" dependencies = [ + "aws-lc-rs", "ring", "rustls-pki-types", "untrusted", @@ -4715,6 +4761,7 @@ version = "1.0.149" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "83fc039473c5595ace860d8c4fafa220ff474b3fc6bfdb4293327f1a37e94d86" dependencies = [ + "indexmap 2.13.0", "itoa", "memchr", "serde", @@ -5325,7 +5372,7 @@ source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "d4a24476afd977c5d5d169f72425868613d82747916dd29e0a357c84c4bd6d29" dependencies = [ "base64 0.22.1", - "brotli", + "brotli 8.0.2", "ico", "json-patch", "plist", @@ -5613,7 +5660,7 @@ source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "219a1f983a2af3653f75b5747f76733b0da7ff03069c7a41901a5eb3ace4557d" dependencies = [ "anyhow", - "brotli", + "brotli 8.0.2", "cargo_metadata", "ctor", "dunce", @@ -6017,18 +6064,13 @@ version = "0.6.8" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "d4e6559d53cc268e5031cd8429d05415bc4cb4aefc4aa5d6cc35fbf5b924a1f8" dependencies = [ - "async-compression", "bitflags 2.11.0", "bytes", - "futures-core", "futures-util", "http", "http-body", - "http-body-util", "iri-string", "pin-project-lite", - "tokio", - "tokio-util", "tower 0.5.3", "tower-layer", "tower-service", @@ -6564,6 +6606,15 @@ dependencies = [ "rustls-pki-types", ] +[[package]] +name = "webpki-roots" +version = "0.26.11" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "521bc38abb08001b01866da9f51eb7c5d647a19260e00054a8c7fd5f9e57f7a9" +dependencies = [ + "webpki-roots 1.0.6", +] + [[package]] name = "webpki-roots" version = "1.0.6" diff --git a/src-tauri/Cargo.toml b/src-tauri/Cargo.toml index 764e1cafd..f4ba03c6e 100644 --- a/src-tauri/Cargo.toml +++ b/src-tauri/Cargo.toml @@ -23,7 +23,7 @@ test-hooks = [] tauri-build = { version = "2.4.0", features = [] } [dependencies] -serde_json = "1.0" +serde_json = { version = "1.0", features = ["preserve_order"] } serde = { version = "1.0", features = ["derive"] } log = "0.4" chrono = { version = "0.4", features = ["serde"] } @@ -38,7 +38,9 @@ tauri-plugin-deep-link = "2" dirs = "5.0" toml = "0.8" toml_edit = "0.22" -reqwest = { version = "0.12", features = ["rustls-tls", "json", "stream", "socks", "gzip"] } +reqwest = { version = "0.12", features = ["rustls-tls", "json", "stream", "socks"] } +flate2 = "1" +brotli = "7" tokio = { version = "1", features = ["macros", "rt-multi-thread", "time", "sync"] } futures = "0.3" async-stream = "0.3" @@ -47,6 +49,16 @@ axum = "0.7" tower = "0.4" tower-http = { version = "0.5", features = ["cors"] } hyper = { version = "1.0", features = ["full"] } +hyper-util = { version = "0.1", features = ["tokio", "http1", "client-legacy"] } +hyper-rustls = { version = "0.27", features = ["http1", "tls12", "ring", "webpki-tokio"] } +http = "1" +http-body = "1" +http-body-util = "0.1" +httparse = "1" +tokio-rustls = "0.26" +rustls = "0.23" +webpki-roots = "0.26" +rustls-native-certs = "0.8" regex = "1.10" rquickjs = { version = "0.8", features = ["array-buffer", "classes"] } thiserror = "2.0" diff --git a/src-tauri/src/codex_config.rs b/src-tauri/src/codex_config.rs index b84caab7a..98f6d06a0 100644 --- a/src-tauri/src/codex_config.rs +++ b/src-tauri/src/codex_config.rs @@ -141,7 +141,7 @@ pub fn read_and_validate_codex_config_text() -> Result { /// /// Supported fields: /// - `"base_url"`: writes to `[model_providers.].base_url` if `model_provider` exists, -/// otherwise falls back to top-level `base_url`. +/// otherwise falls back to top-level `base_url`. /// - `"model"`: writes to top-level `model` field. /// /// Empty value removes the field. diff --git a/src-tauri/src/commands/copilot.rs b/src-tauri/src/commands/copilot.rs index 4e16ddd89..7e36e8460 100644 --- a/src-tauri/src/commands/copilot.rs +++ b/src-tauri/src/commands/copilot.rs @@ -49,7 +49,7 @@ pub async fn copilot_poll_for_auth( Ok(false) } Err(e) => { - log::error!("[CopilotAuth] 轮询失败: {}", e); + log::error!("[CopilotAuth] 轮询失败: {e}"); Err(e.to_string()) } } @@ -70,7 +70,7 @@ pub async fn copilot_poll_for_account( Ok(None) } Err(e) => { - log::error!("[CopilotAuth] 轮询失败: {}", e); + log::error!("[CopilotAuth] 轮询失败: {e}"); Err(e.to_string()) } } diff --git a/src-tauri/src/commands/provider.rs b/src-tauri/src/commands/provider.rs index 883f8d3b9..f186d150d 100644 --- a/src-tauri/src/commands/provider.rs +++ b/src-tauri/src/commands/provider.rs @@ -159,7 +159,7 @@ pub async fn queryProviderUsage( let providers = state .db .get_all_providers(app_type.as_str()) - .map_err(|e| format!("Failed to get providers: {}", e))?; + .map_err(|e| format!("Failed to get providers: {e}"))?; let provider = providers.get(&providerId); let is_copilot = provider @@ -182,11 +182,11 @@ pub async fn queryProviderUsage( Some(account_id) => auth_manager .fetch_usage_for_account(account_id) .await - .map_err(|e| format!("Failed to fetch Copilot usage: {}", e))?, + .map_err(|e| format!("Failed to fetch Copilot usage: {e}"))?, None => auth_manager .fetch_usage() .await - .map_err(|e| format!("Failed to fetch Copilot usage: {}", e))?, + .map_err(|e| format!("Failed to fetch Copilot usage: {e}"))?, }; let premium = &usage.quota_snapshots.premium_interactions; let used = premium.entitlement - premium.remaining; diff --git a/src-tauri/src/commands/stream_check.rs b/src-tauri/src/commands/stream_check.rs index 7afe21018..cc7e4d519 100644 --- a/src-tauri/src/commands/stream_check.rs +++ b/src-tauri/src/commands/stream_check.rs @@ -216,7 +216,11 @@ async fn resolve_claude_api_format_override( .and_then(|meta| meta.managed_account_id_for("github_copilot")); let vendor_result = match account_id.as_deref() { - Some(id) => auth_manager.get_model_vendor_for_account(id, &model_id).await, + Some(id) => { + auth_manager + .get_model_vendor_for_account(id, &model_id) + .await + } None => auth_manager.get_model_vendor(&model_id).await, }; @@ -225,9 +229,7 @@ async fn resolve_claude_api_format_override( Ok(Some(_)) | Ok(None) => "openai_chat", Err(err) => { log::warn!( - "[StreamCheck] Failed to resolve Copilot model vendor for {}: {}. Falling back to chat/completions", - model_id, - err + "[StreamCheck] Failed to resolve Copilot model vendor for {model_id}: {err}. Falling back to chat/completions" ); "openai_chat" } diff --git a/src-tauri/src/proxy/error.rs b/src-tauri/src/proxy/error.rs index cc66d68e6..5468c445e 100644 --- a/src-tauri/src/proxy/error.rs +++ b/src-tauri/src/proxy/error.rs @@ -68,7 +68,6 @@ pub enum ProxyError { StreamIdleTimeout(u64), /// 认证错误 - #[allow(dead_code)] #[error("认证失败: {0}")] AuthError(String), diff --git a/src-tauri/src/proxy/forwarder.rs b/src-tauri/src/proxy/forwarder.rs index ad14ae020..f9dd145ed 100644 --- a/src-tauri/src/proxy/forwarder.rs +++ b/src-tauri/src/proxy/forwarder.rs @@ -2,6 +2,7 @@ //! //! 负责将请求转发到上游Provider,支持故障转移 +use super::hyper_client::ProxyResponse; use super::{ body_filter::filter_private_params_with_whitelist, error::*, @@ -19,67 +20,14 @@ use super::{ use crate::commands::CopilotAuthState; use crate::proxy::providers::copilot_auth::CopilotAuthManager; use crate::{app_config::AppType, provider::Provider}; -use reqwest::Response; +use http::Extensions; use serde_json::Value; use std::sync::Arc; use tauri::Manager; use tokio::sync::RwLock; -/// Headers 黑名单 - 不透传到上游的 Headers -/// -/// 精简版黑名单,只过滤必须覆盖或可能导致问题的 header -/// 参考成功透传的请求,保留更多原始 header -/// -/// 注意:客户端 IP 类(x-forwarded-for, x-real-ip)默认透传 -const HEADER_BLACKLIST: &[&str] = &[ - // 认证类(会被覆盖) - "authorization", - "x-api-key", - "x-goog-api-key", - // 连接类(由 HTTP 客户端管理) - "host", - "content-length", - "transfer-encoding", - // 编码类(会被覆盖为 identity) - "accept-encoding", - // 代理转发类(保留 x-forwarded-for 和 x-real-ip) - "x-forwarded-host", - "x-forwarded-port", - "x-forwarded-proto", - "forwarded", - // CDN/云服务商特定头 - "cf-connecting-ip", - "cf-ipcountry", - "cf-ray", - "cf-visitor", - "true-client-ip", - "fastly-client-ip", - "x-azure-clientip", - "x-azure-fdid", - "x-azure-ref", - "akamai-origin-hop", - "x-akamai-config-log-detail", - // 请求追踪类 - "x-request-id", - "x-correlation-id", - "x-trace-id", - "x-amzn-trace-id", - "x-b3-traceid", - "x-b3-spanid", - "x-b3-parentspanid", - "x-b3-sampled", - "traceparent", - "tracestate", - // anthropic 特定头单独处理,避免重复 - "anthropic-beta", - "anthropic-version", - // 客户端 IP 单独处理(默认透传) - "x-forwarded-for", - "x-real-ip", -]; - pub struct ForwardResult { - pub response: Response, + pub response: ProxyResponse, pub provider: Provider, pub claude_api_format: Option, } @@ -150,6 +98,7 @@ impl RequestForwarder { endpoint: &str, body: Value, headers: axum::http::HeaderMap, + extensions: Extensions, providers: Vec, ) -> Result { // 获取适配器 @@ -226,6 +175,7 @@ impl RequestForwarder { endpoint, &provider_body, &headers, + &extensions, adapter.as_ref(), ) .await @@ -355,6 +305,7 @@ impl RequestForwarder { endpoint, &provider_body, &headers, + &extensions, adapter.as_ref(), ) .await @@ -553,6 +504,7 @@ impl RequestForwarder { endpoint, &provider_body, &headers, + &extensions, adapter.as_ref(), ) .await @@ -791,8 +743,9 @@ impl RequestForwarder { endpoint: &str, body: &Value, headers: &axum::http::HeaderMap, + extensions: &Extensions, adapter: &dyn ProviderAdapter, - ) -> Result<(Response, Option), ProxyError> { + ) -> Result<(ProxyResponse, Option), ProxyError> { // 使用适配器提取 base_url let base_url = adapter.extract_base_url(provider)?; @@ -871,87 +824,11 @@ impl RequestForwarder { // 过滤私有参数(以 `_` 开头的字段),防止内部信息泄露到上游 // 默认使用空白名单,过滤所有 _ 前缀字段 let filtered_body = filter_private_params_with_whitelist(request_body, &[]); + let force_identity_encoding = needs_transform + || should_force_identity_encoding(&effective_endpoint, &filtered_body, headers); - // 获取 HTTP 客户端:优先使用供应商单独代理配置,否则使用全局客户端 - let proxy_config = provider.meta.as_ref().and_then(|m| m.proxy_config.as_ref()); - let client = super::http_client::get_for_provider(proxy_config); - let mut request = client.post(&url); - - // 只有当 timeout > 0 时才设置请求超时 - // Duration::ZERO 在 reqwest 中表示"立刻超时"而不是"禁用超时" - // 故障转移关闭时会传入 0,此时应该使用 client 的默认超时(600秒) - if !self.non_streaming_timeout.is_zero() { - request = request.timeout(self.non_streaming_timeout); - } - - // 过滤黑名单 Headers,保护隐私并避免冲突 - for (key, value) in headers { - let key_str = key.as_str(); - if HEADER_BLACKLIST - .iter() - .any(|h| key_str.eq_ignore_ascii_case(h)) - { - continue; - } - // Copilot 请求:过滤会由 add_auth_headers 注入的固定指纹头, - // 防止客户端原始头与注入头重复(reqwest header() 是追加语义) - if is_copilot - && (key_str.eq_ignore_ascii_case("user-agent") - || key_str.eq_ignore_ascii_case("editor-version") - || key_str.eq_ignore_ascii_case("editor-plugin-version") - || key_str.eq_ignore_ascii_case("copilot-integration-id") - || key_str.eq_ignore_ascii_case("x-github-api-version") - || key_str.eq_ignore_ascii_case("openai-intent")) - { - continue; - } - request = request.header(key, value); - } - - // 处理 anthropic-beta Header(仅 Claude) - // 关键:确保包含 claude-code-20250219 标记,这是上游服务验证请求来源的依据 - // 如果客户端发送的 beta 标记中没有包含 claude-code-20250219,需要补充 - if adapter.name() == "Claude" { - const CLAUDE_CODE_BETA: &str = "claude-code-20250219"; - let beta_value = if let Some(beta) = headers.get("anthropic-beta") { - if let Ok(beta_str) = beta.to_str() { - // 检查是否已包含 claude-code-20250219 - if beta_str.contains(CLAUDE_CODE_BETA) { - beta_str.to_string() - } else { - // 补充 claude-code-20250219 - format!("{CLAUDE_CODE_BETA},{beta_str}") - } - } else { - CLAUDE_CODE_BETA.to_string() - } - } else { - // 如果客户端没有发送,使用默认值 - CLAUDE_CODE_BETA.to_string() - }; - request = request.header("anthropic-beta", &beta_value); - } - - // 客户端 IP 透传(默认开启) - if let Some(xff) = headers.get("x-forwarded-for") { - if let Ok(xff_str) = xff.to_str() { - request = request.header("x-forwarded-for", xff_str); - } - } - if let Some(real_ip) = headers.get("x-real-ip") { - if let Ok(real_ip_str) = real_ip.to_str() { - request = request.header("x-real-ip", real_ip_str); - } - } - - // 流式请求保守禁用压缩,避免上游压缩 SSE 在连接中断时触发解压错误。 - // 非流式请求不显式设置 Accept-Encoding,让 reqwest 自动协商压缩并透明解压。 - if should_force_identity_encoding(&effective_endpoint, &filtered_body, headers) { - request = request.header("accept-encoding", "identity"); - } - - // 使用适配器添加认证头 - if let Some(mut auth) = adapter.extract_auth(provider) { + // 获取认证头(提前准备,用于内联替换) + let auth_headers = if let Some(mut auth) = adapter.extract_auth(provider) { // GitHub Copilot 特殊处理:从 CopilotAuthManager 获取真实 token if auth.strategy == AuthStrategy::GitHubCopilot { if let Some(app_handle) = &self.app_handle { @@ -1002,17 +879,211 @@ impl RequestForwarder { )); } } - request = adapter.add_auth_headers(request, &auth); + adapter.get_auth_headers(&auth) + } else { + Vec::new() + }; + + // Copilot 指纹头名(由 get_auth_headers 注入,需在原始头中去重) + let copilot_fingerprint_headers: &[&str] = if is_copilot { + &[ + "user-agent", + "editor-version", + "editor-plugin-version", + "copilot-integration-id", + "x-github-api-version", + "openai-intent", + ] + } else { + &[] + }; + + // 预计算上游 host 值(用于在原位替换 host header) + let upstream_host = url + .parse::() + .ok() + .and_then(|u| u.authority().map(|a| a.to_string())); + + // 预计算 anthropic-beta 值(仅 Claude) + let anthropic_beta_value = if adapter.name() == "Claude" { + const CLAUDE_CODE_BETA: &str = "claude-code-20250219"; + Some(if let Some(beta) = headers.get("anthropic-beta") { + if let Ok(beta_str) = beta.to_str() { + if beta_str.contains(CLAUDE_CODE_BETA) { + beta_str.to_string() + } else { + format!("{CLAUDE_CODE_BETA},{beta_str}") + } + } else { + CLAUDE_CODE_BETA.to_string() + } + } else { + CLAUDE_CODE_BETA.to_string() + }) + } else { + None + }; + + // ============================================================ + // 构建有序 HeaderMap — 内联替换,保持客户端原始顺序 + // ============================================================ + let mut ordered_headers = http::HeaderMap::new(); + let mut saw_auth = false; + let mut saw_accept_encoding = false; + let mut saw_anthropic_beta = false; + let mut saw_anthropic_version = false; + + for (key, value) in headers { + let key_str = key.as_str(); + + // --- host — 原位替换为上游 host(保持客户端原始位置) --- + if key_str.eq_ignore_ascii_case("host") { + if let Some(ref host_val) = upstream_host { + if let Ok(hv) = http::HeaderValue::from_str(host_val) { + ordered_headers.append(key.clone(), hv); + } + } + continue; + } + + // --- 连接 / 追踪 / CDN 类 — 无条件跳过 --- + if matches!( + key_str, + "content-length" + | "transfer-encoding" + | "x-forwarded-host" + | "x-forwarded-port" + | "x-forwarded-proto" + | "forwarded" + | "cf-connecting-ip" + | "cf-ipcountry" + | "cf-ray" + | "cf-visitor" + | "true-client-ip" + | "fastly-client-ip" + | "x-azure-clientip" + | "x-azure-fdid" + | "x-azure-ref" + | "akamai-origin-hop" + | "x-akamai-config-log-detail" + | "x-request-id" + | "x-correlation-id" + | "x-trace-id" + | "x-amzn-trace-id" + | "x-b3-traceid" + | "x-b3-spanid" + | "x-b3-parentspanid" + | "x-b3-sampled" + | "traceparent" + | "tracestate" + ) { + continue; + } + + // --- 认证类 — 用 adapter 提供的认证头替换(在原始位置) --- + if key_str.eq_ignore_ascii_case("authorization") + || key_str.eq_ignore_ascii_case("x-api-key") + || key_str.eq_ignore_ascii_case("x-goog-api-key") + { + if !saw_auth { + saw_auth = true; + for (ah_name, ah_value) in &auth_headers { + ordered_headers.append(ah_name.clone(), ah_value.clone()); + } + } + continue; + } + + // --- accept-encoding — transform / SSE 路径强制 identity,其余保留原值 --- + if key_str.eq_ignore_ascii_case("accept-encoding") { + if !saw_accept_encoding { + saw_accept_encoding = true; + if force_identity_encoding { + ordered_headers.append( + http::header::ACCEPT_ENCODING, + http::HeaderValue::from_static("identity"), + ); + } else { + ordered_headers.append(key.clone(), value.clone()); + } + } + continue; + } + + // --- anthropic-beta — 用重建值替换(确保含 claude-code 标记) --- + if key_str.eq_ignore_ascii_case("anthropic-beta") { + if !saw_anthropic_beta { + saw_anthropic_beta = true; + if let Some(ref beta_val) = anthropic_beta_value { + if let Ok(hv) = http::HeaderValue::from_str(beta_val) { + ordered_headers.append("anthropic-beta", hv); + } + } + } + continue; + } + + // --- anthropic-version — 透传客户端值 --- + if key_str.eq_ignore_ascii_case("anthropic-version") { + saw_anthropic_version = true; + ordered_headers.append(key.clone(), value.clone()); + continue; + } + + // --- Copilot 指纹头 — 跳过(由 auth_headers 提供) --- + if copilot_fingerprint_headers + .iter() + .any(|h| key_str.eq_ignore_ascii_case(h)) + { + continue; + } + + // --- 默认:透传 --- + ordered_headers.append(key.clone(), value.clone()); } - // anthropic-version 统一处理(仅 Claude):优先使用客户端的版本号,否则使用默认值 - // 注意:只设置一次,避免重复 - if adapter.name() == "Claude" { - let version_str = headers - .get("anthropic-version") - .and_then(|v| v.to_str().ok()) - .unwrap_or("2023-06-01"); - request = request.header("anthropic-version", version_str); + // 如果原始请求中没有认证头,在末尾追加 + if !saw_auth && !auth_headers.is_empty() { + for (ah_name, ah_value) in &auth_headers { + ordered_headers.append(ah_name.clone(), ah_value.clone()); + } + } + + // transform / SSE 路径在缺失时补 identity;普通透传不主动补 accept-encoding + if !saw_accept_encoding && force_identity_encoding { + ordered_headers.append( + http::header::ACCEPT_ENCODING, + http::HeaderValue::from_static("identity"), + ); + } + + // 如果原始请求中没有 anthropic-beta 且有值需要添加,追加 + if !saw_anthropic_beta { + if let Some(ref beta_val) = anthropic_beta_value { + if let Ok(hv) = http::HeaderValue::from_str(beta_val) { + ordered_headers.append("anthropic-beta", hv); + } + } + } + + // anthropic-version:仅在缺失时补充默认值 + if adapter.name() == "Claude" && !saw_anthropic_version { + ordered_headers.append( + "anthropic-version", + http::HeaderValue::from_static("2023-06-01"), + ); + } + + // 序列化请求体 + let body_bytes = serde_json::to_vec(&filtered_body) + .map_err(|e| ProxyError::Internal(format!("Failed to serialize request body: {e}")))?; + + // 确保 content-type 存在 + if !ordered_headers.contains_key(http::header::CONTENT_TYPE) { + ordered_headers.insert( + http::header::CONTENT_TYPE, + http::HeaderValue::from_static("application/json"), + ); } // 输出请求信息日志 @@ -1030,16 +1101,66 @@ impl RequestForwarder { ); } + // 确定超时 + let timeout = if self.non_streaming_timeout.is_zero() { + std::time::Duration::from_secs(600) // 默认 600 秒 + } else { + self.non_streaming_timeout + }; + + // 解析上游代理 URL(供应商单独代理 > 全局代理 > 无) + let proxy_config = provider.meta.as_ref().and_then(|m| m.proxy_config.as_ref()); + let upstream_proxy_url: Option = proxy_config + .filter(|c| c.enabled) + .and_then(super::http_client::build_proxy_url_from_config) + .or_else(super::http_client::get_current_proxy_url); + + // SOCKS5 代理不支持 CONNECT 隧道,需要用 reqwest + let is_socks_proxy = upstream_proxy_url + .as_deref() + .map(|u| u.starts_with("socks5")) + .unwrap_or(false); + + let uri: http::Uri = url + .parse() + .map_err(|e| ProxyError::ForwardFailed(format!("Invalid URL '{url}': {e}")))?; + // 发送请求 - let response = request.json(&filtered_body).send().await.map_err(|e| { - if e.is_timeout() { - ProxyError::Timeout(format!("请求超时: {e}")) - } else if e.is_connect() { - ProxyError::ForwardFailed(format!("连接失败: {e}")) - } else { - ProxyError::ForwardFailed(e.to_string()) + let response = if is_socks_proxy { + // SOCKS5 代理:只能走 reqwest(不支持 header case 保留) + log::debug!("[Forwarder] Using reqwest for SOCKS5 proxy"); + let client = super::http_client::get_for_provider(proxy_config); + let mut request = client.post(&url); + if !self.non_streaming_timeout.is_zero() { + request = request.timeout(self.non_streaming_timeout); } - })?; + for (key, value) in &ordered_headers { + request = request.header(key, value); + } + let reqwest_resp = request.body(body_bytes).send().await.map_err(|e| { + if e.is_timeout() { + ProxyError::Timeout(format!("请求超时: {e}")) + } else if e.is_connect() { + ProxyError::ForwardFailed(format!("连接失败: {e}")) + } else { + ProxyError::ForwardFailed(e.to_string()) + } + })?; + ProxyResponse::Reqwest(reqwest_resp) + } else { + // HTTP 代理或直连:走 hyper raw write(保持 header 大小写) + // 如果有 HTTP 代理,hyper_client 会用 CONNECT 隧道穿过代理 + super::hyper_client::send_request( + uri, + http::Method::POST, + ordered_headers, + extensions.clone(), + body_bytes, + timeout, + upstream_proxy_url.as_deref(), + ) + .await? + }; // 检查响应状态 let status = response.status(); @@ -1048,7 +1169,7 @@ impl RequestForwarder { Ok((response, resolved_claude_api_format)) } else { let status_code = status.as_u16(); - let body_text = response.text().await.ok(); + let body_text = String::from_utf8(response.bytes().await?.to_vec()).ok(); Err(ProxyError::UpstreamError { status: status_code, @@ -1094,7 +1215,11 @@ impl RequestForwarder { .and_then(|m| m.managed_account_id_for("github_copilot")); let vendor_result = match account_id.as_deref() { - Some(id) => copilot_auth.get_model_vendor_for_account(id, model_id).await, + Some(id) => { + copilot_auth + .get_model_vendor_for_account(id, model_id) + .await + } None => copilot_auth.get_model_vendor(model_id).await, }; @@ -1373,7 +1498,8 @@ fn summarize_text_for_log(text: &str, max_chars: usize) -> String { #[cfg(test)] mod tests { use super::*; - use axum::http::{header::ACCEPT, HeaderMap, HeaderValue}; + use axum::http::header::{HeaderValue, ACCEPT}; + use axum::http::HeaderMap; use serde_json::json; #[test] diff --git a/src-tauri/src/proxy/handlers.rs b/src-tauri/src/proxy/handlers.rs index 2210df314..08816246a 100644 --- a/src-tauri/src/proxy/handlers.rs +++ b/src-tauri/src/proxy/handlers.rs @@ -18,7 +18,10 @@ use super::{ streaming_responses::create_anthropic_sse_stream_from_responses, transform, transform_responses, }, - response_processor::{create_logged_passthrough_stream, process_response, SseUsageCollector}, + response_processor::{ + create_logged_passthrough_stream, process_response, read_decoded_body, + strip_entity_headers_for_rebuilt_body, SseUsageCollector, + }, server::ProxyState, types::*, usage::parser::TokenUsage, @@ -27,6 +30,7 @@ use super::{ use crate::app_config::AppType; use axum::{extract::State, http::StatusCode, response::IntoResponse, Json}; use bytes::Bytes; +use http_body_util::BodyExt; use serde_json::{json, Value}; // ============================================================================ @@ -61,10 +65,20 @@ pub async fn get_status(State(state): State) -> Result, - uri: axum::http::Uri, - headers: axum::http::HeaderMap, - Json(body): Json, + request: axum::extract::Request, ) -> Result { + let (parts, body) = request.into_parts(); + let uri = parts.uri; + let headers = parts.headers; + let extensions = parts.extensions; + let body_bytes = body + .collect() + .await + .map_err(|e| ProxyError::Internal(format!("Failed to read request body: {e}")))? + .to_bytes(); + let body: Value = serde_json::from_slice(&body_bytes) + .map_err(|e| ProxyError::Internal(format!("Failed to parse request body: {e}")))?; + let mut ctx = RequestContext::new(&state, &body, &headers, AppType::Claude, "Claude", "claude").await?; @@ -86,6 +100,7 @@ pub async fn handle_messages( endpoint, body.clone(), headers, + extensions, ctx.get_providers(), ) .await @@ -126,7 +141,7 @@ pub async fn handle_messages( /// /// 支持 OpenAI Chat Completions 和 Responses API 两种格式的转换 async fn handle_claude_transform( - response: reqwest::Response, + response: super::hyper_client::ProxyResponse, ctx: &RequestContext, state: &ProxyState, _original_body: &Value, @@ -211,12 +226,14 @@ async fn handle_claude_transform( } // 非流式响应转换 (OpenAI/Responses → Anthropic) - let response_headers = response.headers().clone(); - - let body_bytes = response.bytes().await.map_err(|e| { - log::error!("[Claude] 读取响应体失败: {e}"); - ProxyError::ForwardFailed(format!("Failed to read response body: {e}")) - })?; + let body_timeout = + if ctx.app_config.auto_failover_enabled && ctx.app_config.non_streaming_timeout > 0 { + std::time::Duration::from_secs(ctx.app_config.non_streaming_timeout as u64) + } else { + std::time::Duration::ZERO + }; + let (mut response_headers, _status, body_bytes) = + read_decoded_body(response, ctx.tag, body_timeout).await?; let body_str = String::from_utf8_lossy(&body_bytes); @@ -269,13 +286,10 @@ async fn handle_claude_transform( // 构建响应 let mut builder = axum::response::Response::builder().status(status); + strip_entity_headers_for_rebuilt_body(&mut response_headers); for (key, value) in response_headers.iter() { - if key.as_str().to_lowercase() != "content-length" - && key.as_str().to_lowercase() != "transfer-encoding" - { - builder = builder.header(key, value); - } + builder = builder.header(key, value); } builder = builder.header("content-type", "application/json"); @@ -306,10 +320,20 @@ fn endpoint_with_query(uri: &axum::http::Uri, endpoint: &str) -> String { /// 处理 /v1/chat/completions 请求(OpenAI Chat Completions API - Codex CLI) pub async fn handle_chat_completions( State(state): State, - uri: axum::http::Uri, - headers: axum::http::HeaderMap, - Json(body): Json, + request: axum::extract::Request, ) -> Result { + let (parts, req_body) = request.into_parts(); + let uri = parts.uri; + let headers = parts.headers; + let extensions = parts.extensions; + let body_bytes = req_body + .collect() + .await + .map_err(|e| ProxyError::Internal(format!("Failed to read request body: {e}")))? + .to_bytes(); + let body: Value = serde_json::from_slice(&body_bytes) + .map_err(|e| ProxyError::Internal(format!("Failed to parse request body: {e}")))?; + let mut ctx = RequestContext::new(&state, &body, &headers, AppType::Codex, "Codex", "codex").await?; let endpoint = endpoint_with_query(&uri, "/chat/completions"); @@ -326,6 +350,7 @@ pub async fn handle_chat_completions( &endpoint, body, headers, + extensions, ctx.get_providers(), ) .await @@ -349,10 +374,20 @@ pub async fn handle_chat_completions( /// 处理 /v1/responses 请求(OpenAI Responses API - Codex CLI 透传) pub async fn handle_responses( State(state): State, - uri: axum::http::Uri, - headers: axum::http::HeaderMap, - Json(body): Json, + request: axum::extract::Request, ) -> Result { + let (parts, req_body) = request.into_parts(); + let uri = parts.uri; + let headers = parts.headers; + let extensions = parts.extensions; + let body_bytes = req_body + .collect() + .await + .map_err(|e| ProxyError::Internal(format!("Failed to read request body: {e}")))? + .to_bytes(); + let body: Value = serde_json::from_slice(&body_bytes) + .map_err(|e| ProxyError::Internal(format!("Failed to parse request body: {e}")))?; + let mut ctx = RequestContext::new(&state, &body, &headers, AppType::Codex, "Codex", "codex").await?; let endpoint = endpoint_with_query(&uri, "/responses"); @@ -369,6 +404,7 @@ pub async fn handle_responses( &endpoint, body, headers, + extensions, ctx.get_providers(), ) .await @@ -392,10 +428,20 @@ pub async fn handle_responses( /// 处理 /v1/responses/compact 请求(OpenAI Responses Compact API - Codex CLI 透传) pub async fn handle_responses_compact( State(state): State, - uri: axum::http::Uri, - headers: axum::http::HeaderMap, - Json(body): Json, + request: axum::extract::Request, ) -> Result { + let (parts, req_body) = request.into_parts(); + let uri = parts.uri; + let headers = parts.headers; + let extensions = parts.extensions; + let body_bytes = req_body + .collect() + .await + .map_err(|e| ProxyError::Internal(format!("Failed to read request body: {e}")))? + .to_bytes(); + let body: Value = serde_json::from_slice(&body_bytes) + .map_err(|e| ProxyError::Internal(format!("Failed to parse request body: {e}")))?; + let mut ctx = RequestContext::new(&state, &body, &headers, AppType::Codex, "Codex", "codex").await?; let endpoint = endpoint_with_query(&uri, "/responses/compact"); @@ -412,6 +458,7 @@ pub async fn handle_responses_compact( &endpoint, body, headers, + extensions, ctx.get_providers(), ) .await @@ -440,9 +487,19 @@ pub async fn handle_responses_compact( pub async fn handle_gemini( State(state): State, uri: axum::http::Uri, - headers: axum::http::HeaderMap, - Json(body): Json, + request: axum::extract::Request, ) -> Result { + let (parts, req_body) = request.into_parts(); + let headers = parts.headers; + let extensions = parts.extensions; + let body_bytes = req_body + .collect() + .await + .map_err(|e| ProxyError::Internal(format!("Failed to read request body: {e}")))? + .to_bytes(); + let body: Value = serde_json::from_slice(&body_bytes) + .map_err(|e| ProxyError::Internal(format!("Failed to parse request body: {e}")))?; + // Gemini 的模型名称在 URI 中 let mut ctx = RequestContext::new(&state, &body, &headers, AppType::Gemini, "Gemini", "gemini") .await? @@ -466,6 +523,7 @@ pub async fn handle_gemini( endpoint, body, headers, + extensions, ctx.get_providers(), ) .await diff --git a/src-tauri/src/proxy/http_client.rs b/src-tauri/src/proxy/http_client.rs index c2123c31f..897da7fe8 100644 --- a/src-tauri/src/proxy/http_client.rs +++ b/src-tauri/src/proxy/http_client.rs @@ -219,7 +219,12 @@ fn build_client(proxy_url: Option<&str>) -> Result { .timeout(Duration::from_secs(600)) .connect_timeout(Duration::from_secs(30)) .pool_max_idle_per_host(10) - .tcp_keepalive(Duration::from_secs(60)); + .tcp_keepalive(Duration::from_secs(60)) + // 禁用 reqwest 自动解压:防止 reqwest 覆盖客户端原始 accept-encoding header。 + // 响应解压由 response_processor 根据 content-encoding 手动处理。 + .no_gzip() + .no_brotli() + .no_deflate(); // 有代理地址则使用代理,否则跟随系统代理 if let Some(url) = proxy_url { @@ -332,7 +337,7 @@ pub fn mask_url(url: &str) -> String { /// 根据供应商单独代理配置构建代理 URL /// /// 将 ProviderProxyConfig 转换为代理 URL 字符串 -fn build_proxy_url_from_config(config: &ProviderProxyConfig) -> Option { +pub fn build_proxy_url_from_config(config: &ProviderProxyConfig) -> Option { let proxy_type = config.proxy_type.as_deref().unwrap_or("http"); let host = config.proxy_host.as_deref()?; let port = config.proxy_port?; @@ -387,6 +392,9 @@ pub fn build_client_for_provider(proxy_config: Option<&ProviderProxyConfig>) -> .connect_timeout(Duration::from_secs(30)) .pool_max_idle_per_host(10) .tcp_keepalive(Duration::from_secs(60)) + .no_gzip() + .no_brotli() + .no_deflate() .proxy(proxy) .build() { diff --git a/src-tauri/src/proxy/hyper_client.rs b/src-tauri/src/proxy/hyper_client.rs new file mode 100644 index 000000000..ecd69ae61 --- /dev/null +++ b/src-tauri/src/proxy/hyper_client.rs @@ -0,0 +1,685 @@ +//! Hyper-based HTTP client for proxy forwarding +//! +//! Uses raw TCP/TLS writes to preserve exact original header name casing. +//! Supports HTTP CONNECT tunneling through upstream proxies. +//! Falls back to hyper-util Client (title-case headers) when raw write is not feasible. + +use super::ProxyError; +use bytes::Bytes; +use futures::stream::Stream; +use http_body_util::BodyExt; +use hyper_rustls::HttpsConnectorBuilder; +use hyper_util::{client::legacy::Client, rt::TokioExecutor}; +use std::sync::OnceLock; + +/// Our own header case map: maps lowercase header name → original wire-casing bytes. +/// +/// This is a backup mechanism independent of hyper's internal `HeaderCaseMap` (which is +/// `pub(crate)` and cannot be directly inspected or constructed from outside hyper). +/// +/// Populated in `server.rs` by peeking at raw TCP bytes before hyper parses them. +/// Used in `send_request` to manually write headers with original casing when hyper's +/// own mechanism fails. +#[derive(Clone, Debug, Default)] +pub(crate) struct OriginalHeaderCases { + /// Ordered list of (lowercase_name, original_wire_bytes) pairs. + /// Multiple entries with the same name are allowed (for repeated headers). + pub cases: Vec<(String, Vec)>, +} + +impl OriginalHeaderCases { + /// Parse raw HTTP request bytes (from TcpStream::peek) to extract original header casings. + pub fn from_raw_bytes(buf: &[u8]) -> Self { + let mut headers_buf = [httparse::EMPTY_HEADER; 128]; + let mut req = httparse::Request::new(&mut headers_buf); + // We don't care if parsing is partial — we just want the header names we can get + let _ = req.parse(buf); + + let mut cases = Vec::new(); + for header in req.headers.iter() { + if header.name.is_empty() { + break; + } + cases.push(( + header.name.to_ascii_lowercase(), + header.name.as_bytes().to_vec(), + )); + } + + Self { cases } + } +} + +type HyperClient = Client< + hyper_rustls::HttpsConnector, + http_body_util::Full, +>; + +/// Lazily-initialized hyper client with header-case preservation enabled. +fn global_hyper_client() -> &'static HyperClient { + static CLIENT: OnceLock = OnceLock::new(); + CLIENT.get_or_init(|| { + let connector = HttpsConnectorBuilder::new() + .with_webpki_roots() + .https_or_http() + .enable_http1() + .build(); + + Client::builder(TokioExecutor::new()) + .http1_preserve_header_case(true) + .http1_title_case_headers(true) + .build(connector) + }) +} + +/// Unified response wrapper that can hold either a hyper or reqwest response. +/// +/// The hyper variant is used for the main (direct) path with header-case preservation. +/// The reqwest variant is the fallback when an upstream HTTP/SOCKS5 proxy is configured. +pub enum ProxyResponse { + Hyper(hyper::Response), + Reqwest(reqwest::Response), +} + +impl ProxyResponse { + pub fn status(&self) -> http::StatusCode { + match self { + Self::Hyper(r) => r.status(), + Self::Reqwest(r) => r.status(), + } + } + + pub fn headers(&self) -> &http::HeaderMap { + match self { + Self::Hyper(r) => r.headers(), + Self::Reqwest(r) => r.headers(), + } + } + + /// Shortcut: extract `content-type` header value as `&str`. + pub fn content_type(&self) -> Option<&str> { + self.headers() + .get("content-type") + .and_then(|v| v.to_str().ok()) + } + + /// Check if the response is an SSE stream. + pub fn is_sse(&self) -> bool { + self.content_type() + .map(|ct| ct.contains("text/event-stream")) + .unwrap_or(false) + } + + /// Consume the response and collect the full body into `Bytes`. + pub async fn bytes(self) -> Result { + match self { + Self::Hyper(r) => { + let collected = r.into_body().collect().await.map_err(|e| { + ProxyError::ForwardFailed(format!("Failed to read response body: {e}")) + })?; + Ok(collected.to_bytes()) + } + Self::Reqwest(r) => r.bytes().await.map_err(|e| { + ProxyError::ForwardFailed(format!("Failed to read response body: {e}")) + }), + } + } + + /// Consume the response and return a byte-chunk stream (for SSE pass-through). + pub fn bytes_stream(self) -> impl Stream> + Send { + use futures::StreamExt; + + match self { + Self::Hyper(r) => { + let body = r.into_body(); + let stream = futures::stream::unfold(body, |mut body| async { + match body.frame().await { + Some(Ok(frame)) => { + if let Ok(data) = frame.into_data() { + if data.is_empty() { + Some((Ok(Bytes::new()), body)) + } else { + Some((Ok(data), body)) + } + } else { + Some((Ok(Bytes::new()), body)) + } + } + Some(Err(e)) => Some((Err(std::io::Error::other(e.to_string())), body)), + None => None, + } + }) + .filter(|result| { + futures::future::ready(!matches!(result, Ok(ref b) if b.is_empty())) + }); + Box::pin(stream) + as std::pin::Pin> + Send>> + } + Self::Reqwest(r) => { + let stream = r + .bytes_stream() + .map(|r| r.map_err(|e| std::io::Error::other(e.to_string()))); + Box::pin(stream) + } + } + } +} + +/// Send an HTTP request with header-case preservation. +/// +/// Uses a two-tier strategy: +/// 1. Primary: raw HTTP/1.1 write via TLS stream with exact original header casing +/// (from `OriginalHeaderCases` captured by peek in server.rs), then hand off to +/// hyper for response parsing. +/// 2. Fallback: hyper-util Client with `title_case_headers(true)` when raw write +/// isn't feasible (e.g., missing original cases). +/// +/// The caller is expected to include `Host` in the supplied `headers` at the +/// correct position. +/// +/// `proxy_url`: optional upstream HTTP proxy URL (e.g. `http://127.0.0.1:7890`). +/// When set, the raw write path uses HTTP CONNECT tunneling through the proxy, +/// so header-case preservation works even when an upstream proxy is configured. +pub async fn send_request( + uri: http::Uri, + method: http::Method, + headers: http::HeaderMap, + original_extensions: http::Extensions, + body: Vec, + timeout: std::time::Duration, + proxy_url: Option<&str>, +) -> Result { + // Extract our own OriginalHeaderCases if available + let original_cases = original_extensions.get::().cloned(); + let has_cases = original_cases + .as_ref() + .map(|c| !c.cases.is_empty()) + .unwrap_or(false); + + log::debug!( + "[HyperClient] Sending request: uri={uri}, header_count={}, \ + has_host={}, has_original_cases={has_cases}, proxy={:?}", + headers.len(), + headers.contains_key(http::header::HOST), + proxy_url, + ); + + if has_cases { + // Primary path: use raw write + hyper handshake for exact header casing + let result = tokio::time::timeout( + timeout, + send_raw_request( + &uri, + &method, + &headers, + original_cases.as_ref().unwrap(), + &body, + proxy_url, + ), + ) + .await + .map_err(|_| ProxyError::Timeout(format!("请求超时: {}s", timeout.as_secs())))?; + + match result { + Ok(resp) => return Ok(resp), + Err(e) => { + if proxy_url.is_some() { + // Don't bypass configured proxy with direct connect fallback + return Err(e); + } + log::warn!("[HyperClient] Raw write failed, falling back to hyper-util: {e}"); + // Fall through to hyper-util Client + } + } + } + + // Fallback: hyper-util Client (title-case headers, no proxy support) + let mut req = http::Request::builder() + .method(method) + .uri(&uri) + .body(http_body_util::Full::new(Bytes::from(body))) + .map_err(|e| ProxyError::ForwardFailed(format!("Failed to build request: {e}")))?; + + *req.headers_mut() = headers; + *req.extensions_mut() = original_extensions; + + let client = global_hyper_client(); + let resp = tokio::time::timeout(timeout, client.request(req)) + .await + .map_err(|_| ProxyError::Timeout(format!("请求超时: {}s", timeout.as_secs())))? + .map_err(|e| ProxyError::ForwardFailed(format!("上游请求失败: {e}")))?; + + Ok(ProxyResponse::Hyper(resp)) +} + +/// TCP or TLS stream returned by `connect_via_proxy`. +/// +/// When the proxy URL uses `https://`, the connection to the proxy itself is +/// TLS-wrapped before sending the CONNECT request. The enum lets +/// `send_raw_request` work with either variant generically. +enum ProxyStream { + Tcp(tokio::net::TcpStream), + Tls(Box>), +} + +impl tokio::io::AsyncRead for ProxyStream { + fn poll_read( + self: std::pin::Pin<&mut Self>, + cx: &mut std::task::Context<'_>, + buf: &mut tokio::io::ReadBuf<'_>, + ) -> std::task::Poll> { + match self.get_mut() { + ProxyStream::Tcp(s) => std::pin::Pin::new(s).poll_read(cx, buf), + ProxyStream::Tls(s) => std::pin::Pin::new(s).poll_read(cx, buf), + } + } +} + +impl tokio::io::AsyncWrite for ProxyStream { + fn poll_write( + self: std::pin::Pin<&mut Self>, + cx: &mut std::task::Context<'_>, + buf: &[u8], + ) -> std::task::Poll> { + match self.get_mut() { + ProxyStream::Tcp(s) => std::pin::Pin::new(s).poll_write(cx, buf), + ProxyStream::Tls(s) => std::pin::Pin::new(s).poll_write(cx, buf), + } + } + + fn poll_flush( + self: std::pin::Pin<&mut Self>, + cx: &mut std::task::Context<'_>, + ) -> std::task::Poll> { + match self.get_mut() { + ProxyStream::Tcp(s) => std::pin::Pin::new(s).poll_flush(cx), + ProxyStream::Tls(s) => std::pin::Pin::new(s).poll_flush(cx), + } + } + + fn poll_shutdown( + self: std::pin::Pin<&mut Self>, + cx: &mut std::task::Context<'_>, + ) -> std::task::Poll> { + match self.get_mut() { + ProxyStream::Tcp(s) => std::pin::Pin::new(s).poll_shutdown(cx), + ProxyStream::Tls(s) => std::pin::Pin::new(s).poll_shutdown(cx), + } + } +} + +/// Send request via raw TCP/TLS with exact original header casing. +/// +/// When `proxy_url` is provided, establishes an HTTP CONNECT tunnel through +/// the proxy first, then performs TLS + raw write through the tunnel. +/// This preserves header casing even when an upstream proxy is configured. +async fn send_raw_request( + uri: &http::Uri, + method: &http::Method, + headers: &http::HeaderMap, + original_cases: &OriginalHeaderCases, + body: &[u8], + proxy_url: Option<&str>, +) -> Result { + use tokio::io::AsyncWriteExt; + + let scheme = uri.scheme_str().unwrap_or("https"); + let host = uri + .host() + .ok_or_else(|| ProxyError::ForwardFailed("URI has no host".into()))?; + let port = uri + .port_u16() + .unwrap_or(if scheme == "https" { 443 } else { 80 }); + let path_and_query = uri.path_and_query().map(|pq| pq.as_str()).unwrap_or("/"); + + // Build raw HTTP request bytes + let raw = build_raw_request(method, path_and_query, headers, original_cases, body); + + // Establish TCP connection — either direct or through HTTP CONNECT proxy + let stream = if let Some(proxy) = proxy_url { + connect_via_proxy(proxy, host, port).await? + } else { + ProxyStream::Tcp( + tokio::net::TcpStream::connect((host, port)) + .await + .map_err(|e| ProxyError::ForwardFailed(format!("TCP connect failed: {e}")))?, + ) + }; + + if scheme == "https" { + let tls_connector = global_tls_connector(); + let server_name = rustls::pki_types::ServerName::try_from(host.to_string()) + .map_err(|e| ProxyError::ForwardFailed(format!("Invalid server name: {e}")))?; + let mut tls_stream = tls_connector + .connect(server_name, stream) + .await + .map_err(|e| ProxyError::ForwardFailed(format!("TLS handshake failed: {e}")))?; + + tls_stream + .write_all(&raw) + .await + .map_err(|e| ProxyError::ForwardFailed(format!("Write failed: {e}")))?; + + let filtered = WriteFilter::new(tls_stream); + do_hyper_response(filtered, method.clone()).await + } else { + let mut stream = stream; + stream + .write_all(&raw) + .await + .map_err(|e| ProxyError::ForwardFailed(format!("Write failed: {e}")))?; + + let filtered = WriteFilter::new(stream); + do_hyper_response(filtered, method.clone()).await + } +} + +/// Establish a connection through an HTTP CONNECT proxy tunnel. +/// +/// 1. Connect TCP to the proxy server (TLS-wrapped when `https://` proxy) +/// 2. Send `CONNECT host:port HTTP/1.1` with optional `Proxy-Authorization` +/// 3. Read the proxy's 200 response (407 → `AuthError`) +/// 4. Return the tunneled stream (ready for target TLS handshake + raw write) +async fn connect_via_proxy( + proxy_url: &str, + target_host: &str, + target_port: u16, +) -> Result { + use base64::Engine; + use tokio::io::{AsyncBufReadExt, AsyncWriteExt, BufReader}; + + let parsed = url::Url::parse(proxy_url) + .map_err(|e| ProxyError::ForwardFailed(format!("Invalid proxy URL: {e}")))?; + + let proxy_host = parsed + .host_str() + .ok_or_else(|| ProxyError::ForwardFailed("Proxy URL has no host".into()))?; + let proxy_port = parsed + .port() + .unwrap_or(if parsed.scheme() == "https" { 443 } else { 80 }); + + // Build Proxy-Authorization header if credentials are present + let proxy_auth = if !parsed.username().is_empty() { + let password = parsed.password().unwrap_or(""); + let credentials = format!("{}:{}", parsed.username(), password); + let encoded = base64::engine::general_purpose::STANDARD.encode(credentials); + Some(format!("Proxy-Authorization: Basic {encoded}\r\n")) + } else { + None + }; + + // Connect to the proxy + let tcp = tokio::net::TcpStream::connect((proxy_host, proxy_port)) + .await + .map_err(|e| ProxyError::ForwardFailed(format!("Proxy TCP connect failed: {e}")))?; + + // Wrap with TLS if the proxy URL uses https:// + let mut stream: ProxyStream = if parsed.scheme() == "https" { + let tls_connector = global_tls_connector(); + let server_name = rustls::pki_types::ServerName::try_from(proxy_host.to_string()) + .map_err(|e| ProxyError::ForwardFailed(format!("Invalid proxy server name: {e}")))?; + let tls_stream = tls_connector + .connect(server_name, tcp) + .await + .map_err(|e| ProxyError::ForwardFailed(format!("Proxy TLS handshake failed: {e}")))?; + ProxyStream::Tls(Box::new(tls_stream)) + } else { + ProxyStream::Tcp(tcp) + }; + + // Send CONNECT request + let mut connect_req = format!( + "CONNECT {target_host}:{target_port} HTTP/1.1\r\n\ + Host: {target_host}:{target_port}\r\n" + ); + if let Some(auth) = &proxy_auth { + connect_req.push_str(auth); + } + connect_req.push_str("\r\n"); + + stream + .write_all(connect_req.as_bytes()) + .await + .map_err(|e| ProxyError::ForwardFailed(format!("CONNECT write failed: {e}")))?; + + // Read the proxy's response status line + let mut reader = BufReader::new(&mut stream); + let mut status_line = String::new(); + reader + .read_line(&mut status_line) + .await + .map_err(|e| ProxyError::ForwardFailed(format!("CONNECT read failed: {e}")))?; + + // Expect "HTTP/1.1 200 ..." or "HTTP/1.0 200 ..." + if !status_line.contains(" 200 ") { + if status_line.contains(" 407 ") { + return Err(ProxyError::AuthError(format!( + "Proxy authentication required (407): {}", + status_line.trim() + ))); + } + return Err(ProxyError::ForwardFailed(format!( + "Proxy CONNECT rejected: {}", + status_line.trim() + ))); + } + + // Drain remaining response headers (until empty line) + loop { + let mut line = String::new(); + reader + .read_line(&mut line) + .await + .map_err(|e| ProxyError::ForwardFailed(format!("CONNECT header read: {e}")))?; + if line.trim().is_empty() { + break; + } + } + // BufReader might have buffered data; drop it to get raw stream back. + // Since CONNECT response is headers-only (no body), and we read until \r\n\r\n, + // the BufReader buffer should be empty at this point. + drop(reader); + + log::debug!( + "[HyperClient] CONNECT tunnel established via {proxy_host}:{proxy_port} -> {target_host}:{target_port}" + ); + + Ok(stream) +} + +/// Lazily-initialized TLS connector for raw connections. +/// +/// Loads both webpki roots AND native system certificates so that +/// proxy MITM CAs (e.g. Clash, mitmproxy) installed in the system +/// keychain are trusted through the CONNECT tunnel. +fn global_tls_connector() -> &'static tokio_rustls::TlsConnector { + static CONNECTOR: OnceLock = OnceLock::new(); + CONNECTOR.get_or_init(|| { + let mut root_store = rustls::RootCertStore::empty(); + // Baseline: Mozilla/webpki roots + root_store.extend(webpki_roots::TLS_SERVER_ROOTS.iter().cloned()); + // Native system certs (includes user-installed proxy CAs) + let native = rustls_native_certs::load_native_certs(); + let (added, _errors) = root_store.add_parsable_certificates(native.certs); + log::debug!("[HyperClient] TLS root store: webpki + {added} native certs"); + let config = rustls::ClientConfig::builder() + .with_root_certificates(root_store) + .with_no_client_auth(); + tokio_rustls::TlsConnector::from(std::sync::Arc::new(config)) + }) +} + +/// Build raw HTTP/1.1 request bytes with original header casing. +fn build_raw_request( + method: &http::Method, + path_and_query: &str, + headers: &http::HeaderMap, + original_cases: &OriginalHeaderCases, + body: &[u8], +) -> Vec { + let mut raw = Vec::with_capacity(4096 + body.len()); + + // Request line + raw.extend_from_slice(method.as_str().as_bytes()); + raw.extend_from_slice(b" "); + raw.extend_from_slice(path_and_query.as_bytes()); + raw.extend_from_slice(b" HTTP/1.1\r\n"); + + // Headers with original casing, emitted in original wire order. + // + // Strategy: + // 1. Walk `original_cases.cases` in order — this preserves the exact + // header sequence the client sent. For each entry, emit the stored + // original-casing name plus the current value from `headers` (the + // proxy may have rewritten the value, e.g. Authorization). + // Repeated headers with the same name are handled by tracking a + // per-name value cursor so we step through `get_all()` in order. + // 2. After the original headers, append any headers that exist in + // `headers` but were not present in the original request (i.e. added + // by the proxy). These are emitted in lowercase. + // + // This replaces the old `for name in headers.keys()` loop which iterated + // in hash-map order, destroying the original header sequence. + let mut emitted: std::collections::HashSet = + std::collections::HashSet::with_capacity(original_cases.cases.len()); + // Per-name cursor: how many values we have already emitted for each name. + let mut value_cursor: std::collections::HashMap = + std::collections::HashMap::with_capacity(original_cases.cases.len()); + + for (lower_name, orig_name_bytes) in &original_cases.cases { + if let Ok(header_name) = http::header::HeaderName::from_bytes(lower_name.as_bytes()) { + let all_values: Vec<_> = headers.get_all(&header_name).iter().collect(); + let cursor = value_cursor.entry(lower_name.clone()).or_insert(0); + if let Some(value) = all_values.get(*cursor) { + raw.extend_from_slice(orig_name_bytes); + raw.extend_from_slice(b": "); + raw.extend_from_slice(value.as_bytes()); + raw.extend_from_slice(b"\r\n"); + *cursor += 1; + emitted.insert(lower_name.clone()); + } + } + } + + // Append proxy-added headers (not present in the original request). + for name in headers.keys() { + let lower = name.as_str().to_ascii_lowercase(); + if !emitted.contains(&lower) { + for value in headers.get_all(name) { + raw.extend_from_slice(name.as_str().as_bytes()); + raw.extend_from_slice(b": "); + raw.extend_from_slice(value.as_bytes()); + raw.extend_from_slice(b"\r\n"); + } + emitted.insert(lower); + } + } + + // Add Content-Length if not already present + if !headers.contains_key(http::header::CONTENT_LENGTH) { + raw.extend_from_slice(b"Content-Length: "); + raw.extend_from_slice(body.len().to_string().as_bytes()); + raw.extend_from_slice(b"\r\n"); + } + + // End of headers + body + raw.extend_from_slice(b"\r\n"); + raw.extend_from_slice(body); + + raw +} + +/// Use hyper's low-level client to parse the response on a stream where we've +/// already written the request. +/// +/// `WriteFilter` discards any writes from hyper (it would try to send its own +/// request encoding), while passing reads through transparently. +async fn do_hyper_response( + stream: WriteFilter, + method: http::Method, +) -> Result +where + S: tokio::io::AsyncRead + tokio::io::AsyncWrite + Unpin + Send + 'static, +{ + let io = hyper_util::rt::TokioIo::new(stream); + + let (mut sender, conn) = hyper::client::conn::http1::Builder::new() + .preserve_header_case(true) + .handshake::<_, http_body_util::Full>(io) + .await + .map_err(|e| ProxyError::ForwardFailed(format!("Handshake failed: {e}")))?; + + // Spawn the connection driver (reads responses from the stream) + tokio::spawn(async move { + if let Err(e) = conn.await { + log::debug!("[HyperClient] raw conn driver error: {e}"); + } + }); + + // Send a dummy request through hyper — hyper will encode this and try to write it, + // but WriteFilter discards all writes. Hyper will then read the response from the stream. + let dummy_req = http::Request::builder() + .method(method) + .uri("/") + .body(http_body_util::Full::new(Bytes::new())) + .map_err(|e| ProxyError::ForwardFailed(format!("Build dummy request: {e}")))?; + + let resp = sender + .send_request(dummy_req) + .await + .map_err(|e| ProxyError::ForwardFailed(format!("Response parse failed: {e}")))?; + + Ok(ProxyResponse::Hyper(resp)) +} + +/// A stream wrapper that discards all writes but passes reads through. +/// +/// This lets hyper's connection driver think it sent a request (its encoded bytes +/// go to /dev/null), while correctly parsing the response that the upstream server +/// sends in reply to our raw-written request. +struct WriteFilter { + inner: S, +} + +impl WriteFilter { + fn new(inner: S) -> Self { + Self { inner } + } +} + +impl tokio::io::AsyncRead for WriteFilter { + fn poll_read( + self: std::pin::Pin<&mut Self>, + cx: &mut std::task::Context<'_>, + buf: &mut tokio::io::ReadBuf<'_>, + ) -> std::task::Poll> { + // Pass reads through to the underlying stream + let inner = std::pin::Pin::new(&mut self.get_mut().inner); + inner.poll_read(cx, buf) + } +} + +impl tokio::io::AsyncWrite for WriteFilter { + fn poll_write( + self: std::pin::Pin<&mut Self>, + _cx: &mut std::task::Context<'_>, + buf: &[u8], + ) -> std::task::Poll> { + // Discard all writes — pretend they succeeded + std::task::Poll::Ready(Ok(buf.len())) + } + + fn poll_flush( + self: std::pin::Pin<&mut Self>, + _cx: &mut std::task::Context<'_>, + ) -> std::task::Poll> { + std::task::Poll::Ready(Ok(())) + } + + fn poll_shutdown( + self: std::pin::Pin<&mut Self>, + _cx: &mut std::task::Context<'_>, + ) -> std::task::Poll> { + std::task::Poll::Ready(Ok(())) + } +} diff --git a/src-tauri/src/proxy/log_codes.rs b/src-tauri/src/proxy/log_codes.rs index 0190dce52..ed4ae0183 100644 --- a/src-tauri/src/proxy/log_codes.rs +++ b/src-tauri/src/proxy/log_codes.rs @@ -26,6 +26,8 @@ pub mod srv { pub const STOPPED: &str = "SRV-002"; pub const STOP_TIMEOUT: &str = "SRV-003"; pub const TASK_ERROR: &str = "SRV-004"; + pub const ACCEPT_ERR: &str = "SRV-005"; + pub const CONN_ERR: &str = "SRV-006"; } /// 转发器日志码 diff --git a/src-tauri/src/proxy/mod.rs b/src-tauri/src/proxy/mod.rs index 4386230a7..984f04ceb 100644 --- a/src-tauri/src/proxy/mod.rs +++ b/src-tauri/src/proxy/mod.rs @@ -14,6 +14,7 @@ pub mod handler_context; mod handlers; mod health; pub mod http_client; +pub mod hyper_client; pub mod log_codes; pub mod model_mapper; pub mod provider_router; diff --git a/src-tauri/src/proxy/providers/adapter.rs b/src-tauri/src/proxy/providers/adapter.rs index 6393cf62a..d764d3eb2 100644 --- a/src-tauri/src/proxy/providers/adapter.rs +++ b/src-tauri/src/proxy/providers/adapter.rs @@ -5,7 +5,6 @@ use super::auth::AuthInfo; use crate::provider::Provider; use crate::proxy::error::ProxyError; -use reqwest::RequestBuilder; use serde_json::Value; /// 供应商适配器 Trait @@ -14,116 +13,36 @@ use serde_json::Value; /// - URL 构建 /// - 认证信息提取和头部注入 /// - 请求/响应格式转换(可选) -/// -/// # 示例 -/// -/// ```ignore -/// pub struct ClaudeAdapter; -/// -/// impl ProviderAdapter for ClaudeAdapter { -/// fn name(&self) -> &'static str { "Claude" } -/// -/// fn extract_base_url(&self, provider: &Provider) -> Result { -/// // 从 provider 配置中提取 base_url -/// } -/// -/// fn extract_auth(&self, provider: &Provider) -> Option { -/// // 从 provider 配置中提取认证信息 -/// } -/// -/// fn build_url(&self, base_url: &str, endpoint: &str) -> String { -/// format!("{}{}", base_url.trim_end_matches('/'), endpoint) -/// } -/// -/// fn add_auth_headers(&self, request: RequestBuilder, auth: &AuthInfo) -> RequestBuilder { -/// // 添加认证头 -/// } -/// } -/// ``` pub trait ProviderAdapter: Send + Sync { /// 适配器名称(用于日志和调试) fn name(&self) -> &'static str; /// 从 Provider 配置中提取 base_url - /// - /// # Arguments - /// * `provider` - Provider 配置 - /// - /// # Returns - /// * `Ok(String)` - 提取到的 base_url(已去除尾部斜杠) - /// * `Err(ProxyError)` - 提取失败 fn extract_base_url(&self, provider: &Provider) -> Result; /// 从 Provider 配置中提取认证信息 - /// - /// # Arguments - /// * `provider` - Provider 配置 - /// - /// # Returns - /// * `Some(AuthInfo)` - 提取到的认证信息 - /// * `None` - 未找到认证信息 fn extract_auth(&self, provider: &Provider) -> Option; /// 构建请求 URL - /// - /// # Arguments - /// * `base_url` - 基础 URL - /// * `endpoint` - 请求端点(如 `/v1/messages`) - /// - /// # Returns - /// 完整的请求 URL fn build_url(&self, base_url: &str, endpoint: &str) -> String; - /// 添加认证头到请求 + /// Return auth headers as `(name, value)` pairs. /// - /// # Arguments - /// * `request` - reqwest RequestBuilder - /// * `auth` - 认证信息 - /// - /// # Returns - /// 添加了认证头的 RequestBuilder - fn add_auth_headers(&self, request: RequestBuilder, auth: &AuthInfo) -> RequestBuilder; + /// The forwarder inserts these at the position of the original auth header + /// so that header order is preserved. + fn get_auth_headers(&self, auth: &AuthInfo) -> Vec<(http::HeaderName, http::HeaderValue)>; /// 是否需要格式转换 - /// - /// 默认返回 `false`(透传模式)。 - /// 仅当供应商需要格式转换时(如 Claude + OpenRouter 旧 OpenAI 兼容接口)才返回 `true`。 - /// - /// # Arguments - /// * `provider` - Provider 配置 fn needs_transform(&self, _provider: &Provider) -> bool { false } /// 转换请求体 - /// - /// 将请求体从一种格式转换为另一种格式(如 Anthropic → OpenAI)。 - /// 默认实现直接返回原始请求体(透传)。 - /// - /// # Arguments - /// * `body` - 原始请求体 - /// * `provider` - Provider 配置(用于获取模型映射等) - /// - /// # Returns - /// * `Ok(Value)` - 转换后的请求体 - /// * `Err(ProxyError)` - 转换失败 fn transform_request(&self, body: Value, _provider: &Provider) -> Result { Ok(body) } /// 转换响应体 - /// - /// 将响应体从一种格式转换为另一种格式(如 OpenAI → Anthropic)。 - /// 默认实现直接返回原始响应体(透传)。 - /// - /// # Arguments - /// * `body` - 原始响应体 - /// - /// # Returns - /// * `Ok(Value)` - 转换后的响应体 - /// * `Err(ProxyError)` - 转换失败 - /// - /// Note: 响应转换将在 handler 层集成,目前预留接口 #[allow(dead_code)] fn transform_response(&self, body: Value) -> Result { Ok(body) diff --git a/src-tauri/src/proxy/providers/claude.rs b/src-tauri/src/proxy/providers/claude.rs index 5cecdf7ce..8d59dc2a7 100644 --- a/src-tauri/src/proxy/providers/claude.rs +++ b/src-tauri/src/proxy/providers/claude.rs @@ -16,7 +16,6 @@ use super::{AuthInfo, AuthStrategy, ProviderAdapter, ProviderType}; use crate::provider::Provider; use crate::proxy::error::ProxyError; -use reqwest::RequestBuilder; /// 获取 Claude 供应商的 API 格式 /// @@ -337,32 +336,50 @@ impl ProviderAdapter for ClaudeAdapter { base } - fn add_auth_headers(&self, request: RequestBuilder, auth: &AuthInfo) -> RequestBuilder { + fn get_auth_headers(&self, auth: &AuthInfo) -> Vec<(http::HeaderName, http::HeaderValue)> { + use http::{HeaderName, HeaderValue}; // 注意:anthropic-version 由 forwarder.rs 统一处理(透传客户端值或设置默认值) - // 这里不再设置 anthropic-version,避免 header 重复 + let bearer = format!("Bearer {}", auth.api_key); match auth.strategy { - // Anthropic 官方: Authorization Bearer + x-api-key - AuthStrategy::Anthropic => request - .header("Authorization", format!("Bearer {}", auth.api_key)) - .header("x-api-key", &auth.api_key), - // ClaudeAuth 中转服务: 仅 Bearer,无 x-api-key - AuthStrategy::ClaudeAuth => { - request.header("Authorization", format!("Bearer {}", auth.api_key)) + AuthStrategy::Anthropic | AuthStrategy::ClaudeAuth | AuthStrategy::Bearer => { + vec![( + HeaderName::from_static("authorization"), + HeaderValue::from_str(&bearer).unwrap(), + )] } - // OpenRouter: Bearer - AuthStrategy::Bearer => { - request.header("Authorization", format!("Bearer {}", auth.api_key)) + AuthStrategy::GitHubCopilot => { + vec![ + ( + HeaderName::from_static("authorization"), + HeaderValue::from_str(&bearer).unwrap(), + ), + ( + HeaderName::from_static("editor-version"), + HeaderValue::from_static(super::copilot_auth::COPILOT_EDITOR_VERSION), + ), + ( + HeaderName::from_static("editor-plugin-version"), + HeaderValue::from_static(super::copilot_auth::COPILOT_PLUGIN_VERSION), + ), + ( + HeaderName::from_static("copilot-integration-id"), + HeaderValue::from_static(super::copilot_auth::COPILOT_INTEGRATION_ID), + ), + ( + HeaderName::from_static("user-agent"), + HeaderValue::from_static(super::copilot_auth::COPILOT_USER_AGENT), + ), + ( + HeaderName::from_static("x-github-api-version"), + HeaderValue::from_static(super::copilot_auth::COPILOT_API_VERSION), + ), + ( + HeaderName::from_static("openai-intent"), + HeaderValue::from_static("conversation-panel"), + ), + ] } - // GitHub Copilot: Bearer + 统一指纹头 - AuthStrategy::GitHubCopilot => request - .header("Authorization", format!("Bearer {}", auth.api_key)) - .header("editor-version", super::copilot_auth::COPILOT_EDITOR_VERSION) - .header("editor-plugin-version", super::copilot_auth::COPILOT_PLUGIN_VERSION) - .header("copilot-integration-id", super::copilot_auth::COPILOT_INTEGRATION_ID) - .header("user-agent", super::copilot_auth::COPILOT_USER_AGENT) - .header("x-github-api-version", super::copilot_auth::COPILOT_API_VERSION) - .header("openai-intent", "conversation-panel"), - _ => request, + _ => vec![], } } diff --git a/src-tauri/src/proxy/providers/codex.rs b/src-tauri/src/proxy/providers/codex.rs index 02c24bcc2..eebd33704 100644 --- a/src-tauri/src/proxy/providers/codex.rs +++ b/src-tauri/src/proxy/providers/codex.rs @@ -9,7 +9,6 @@ use super::{AuthInfo, AuthStrategy, ProviderAdapter}; use crate::provider::Provider; use crate::proxy::error::ProxyError; use regex::Regex; -use reqwest::RequestBuilder; use std::sync::LazyLock; /// 官方 Codex 客户端 User-Agent 正则 @@ -174,8 +173,12 @@ impl ProviderAdapter for CodexAdapter { url } - fn add_auth_headers(&self, request: RequestBuilder, auth: &AuthInfo) -> RequestBuilder { - request.header("Authorization", format!("Bearer {}", auth.api_key)) + fn get_auth_headers(&self, auth: &AuthInfo) -> Vec<(http::HeaderName, http::HeaderValue)> { + let bearer = format!("Bearer {}", auth.api_key); + vec![( + http::HeaderName::from_static("authorization"), + http::HeaderValue::from_str(&bearer).unwrap(), + )] } } diff --git a/src-tauri/src/proxy/providers/copilot_auth.rs b/src-tauri/src/proxy/providers/copilot_auth.rs index 11ee38789..04f3f9f33 100644 --- a/src-tauri/src/proxy/providers/copilot_auth.rs +++ b/src-tauri/src/proxy/providers/copilot_auth.rs @@ -341,7 +341,7 @@ impl CopilotAuthManager { // 尝试从磁盘加载(同步,不发起网络请求) if let Err(e) = manager.load_from_disk_sync() { - log::warn!("[CopilotAuth] 加载存储失败: {}", e); + log::warn!("[CopilotAuth] 加载存储失败: {e}"); } manager @@ -364,7 +364,7 @@ impl CopilotAuthManager { /// 移除指定账号 pub async fn remove_account(&self, account_id: &str) -> Result<(), CopilotAuthError> { - log::info!("[CopilotAuth] 移除账号: {}", account_id); + log::info!("[CopilotAuth] 移除账号: {account_id}"); { let mut accounts = self.accounts.write().await; @@ -482,8 +482,7 @@ impl CopilotAuthManager { let status = response.status(); let text = response.text().await.unwrap_or_default(); return Err(CopilotAuthError::NetworkError(format!( - "GitHub 设备码请求失败: {} - {}", - status, text + "GitHub 设备码请求失败: {status} - {text}" ))); } @@ -581,10 +580,7 @@ impl CopilotAuthManager { } // 需要刷新 - log::info!( - "[CopilotAuth] 账号 {} 的 Copilot Token 需要刷新", - account_id - ); + log::info!("[CopilotAuth] 账号 {account_id} 的 Copilot Token 需要刷新"); let refresh_lock = self.get_refresh_lock(account_id).await; let _refresh_guard = refresh_lock.lock().await; @@ -660,12 +656,12 @@ impl CopilotAuthManager { ) -> Result, CopilotAuthError> { let copilot_token = self.get_valid_token_for_account(account_id).await?; - log::info!("[CopilotAuth] 获取账号 {} 的 Copilot 可用模型", account_id); + log::info!("[CopilotAuth] 获取账号 {account_id} 的 Copilot 可用模型"); let response = self .http_client .get(COPILOT_MODELS_URL) - .header("Authorization", format!("Bearer {}", copilot_token)) + .header("Authorization", format!("Bearer {copilot_token}")) .header("Content-Type", "application/json") .header("copilot-integration-id", "vscode-chat") .header("editor-version", COPILOT_EDITOR_VERSION) @@ -679,8 +675,7 @@ impl CopilotAuthManager { let status = response.status(); let text = response.text().await.unwrap_or_default(); return Err(CopilotAuthError::CopilotTokenFetchFailed(format!( - "获取模型列表失败: {} - {}", - status, text + "获取模型列表失败: {status} - {text}" ))); } @@ -749,12 +744,12 @@ impl CopilotAuthManager { .ok_or_else(|| CopilotAuthError::AccountNotFound(account_id.to_string()))? }; - log::info!("[CopilotAuth] 获取账号 {} 的 Copilot 使用量", account_id); + log::info!("[CopilotAuth] 获取账号 {account_id} 的 Copilot 使用量"); let response = self .http_client .get(COPILOT_USAGE_URL) - .header("Authorization", format!("token {}", github_token)) + .header("Authorization", format!("token {github_token}")) .header("Content-Type", "application/json") .header("editor-version", COPILOT_EDITOR_VERSION) .header("editor-plugin-version", COPILOT_PLUGIN_VERSION) @@ -771,8 +766,7 @@ impl CopilotAuthManager { let status = response.status(); let text = response.text().await.unwrap_or_default(); return Err(CopilotAuthError::CopilotTokenFetchFailed(format!( - "获取使用量失败: {} - {}", - status, text + "获取使用量失败: {status} - {text}" ))); } @@ -1000,7 +994,7 @@ impl CopilotAuthManager { let response = self .http_client .get(GITHUB_USER_URL) - .header("Authorization", format!("token {}", github_token)) + .header("Authorization", format!("token {github_token}")) .header("User-Agent", COPILOT_USER_AGENT) .header("Editor-Version", COPILOT_EDITOR_VERSION) .header("Editor-Plugin-Version", COPILOT_PLUGIN_VERSION) @@ -1027,12 +1021,12 @@ impl CopilotAuthManager { github_token: &str, account_id: &str, ) -> Result<(), CopilotAuthError> { - log::debug!("[CopilotAuth] 获取账号 {} 的 Copilot Token", account_id); + log::debug!("[CopilotAuth] 获取账号 {account_id} 的 Copilot Token"); let response = self .http_client .get(COPILOT_TOKEN_URL) - .header("Authorization", format!("token {}", github_token)) + .header("Authorization", format!("token {github_token}")) .header("User-Agent", COPILOT_USER_AGENT) .header("Editor-Version", COPILOT_EDITOR_VERSION) .header("Editor-Plugin-Version", COPILOT_PLUGIN_VERSION) @@ -1051,8 +1045,7 @@ impl CopilotAuthManager { let status = response.status(); let text = response.text().await.unwrap_or_default(); return Err(CopilotAuthError::CopilotTokenFetchFailed(format!( - "{}: {}", - status, text + "{status}: {text}" ))); } @@ -1135,7 +1128,7 @@ impl CopilotAuthManager { .fetch_copilot_token_with_github_token(&legacy_token, &account_id) .await { - log::warn!("[CopilotAuth] 迁移时验证 Copilot 订阅失败: {}", e); + log::warn!("[CopilotAuth] 迁移时验证 Copilot 订阅失败: {e}"); } // 添加账号 @@ -1149,7 +1142,7 @@ impl CopilotAuthManager { "Legacy Copilot auth migration failed: {e}" ))) .await; - log::warn!("[CopilotAuth] 迁移失败,旧 token 可能已失效: {}", e); + log::warn!("[CopilotAuth] 迁移失败,旧 token 可能已失效: {e}"); } } diff --git a/src-tauri/src/proxy/providers/gemini.rs b/src-tauri/src/proxy/providers/gemini.rs index 646d237f7..e0e88feb8 100644 --- a/src-tauri/src/proxy/providers/gemini.rs +++ b/src-tauri/src/proxy/providers/gemini.rs @@ -9,7 +9,6 @@ use super::{AuthInfo, AuthStrategy, ProviderAdapter, ProviderType}; use crate::provider::Provider; use crate::proxy::error::ProxyError; -use reqwest::RequestBuilder; /// Gemini 适配器 pub struct GeminiAdapter; @@ -217,17 +216,26 @@ impl ProviderAdapter for GeminiAdapter { url } - fn add_auth_headers(&self, request: RequestBuilder, auth: &AuthInfo) -> RequestBuilder { + fn get_auth_headers(&self, auth: &AuthInfo) -> Vec<(http::HeaderName, http::HeaderValue)> { + use http::{HeaderName, HeaderValue}; match auth.strategy { - // OAuth Bearer 认证 AuthStrategy::GoogleOAuth => { let token = auth.access_token.as_ref().unwrap_or(&auth.api_key); - request - .header("Authorization", format!("Bearer {token}")) - .header("x-goog-api-client", "GeminiCLI/1.0") + vec![ + ( + HeaderName::from_static("authorization"), + HeaderValue::from_str(&format!("Bearer {token}")).unwrap(), + ), + ( + HeaderName::from_static("x-goog-api-client"), + HeaderValue::from_static("GeminiCLI/1.0"), + ), + ] } - // API Key 认证 - _ => request.header("x-goog-api-key", &auth.api_key), + _ => vec![( + HeaderName::from_static("x-goog-api-key"), + HeaderValue::from_str(&auth.api_key).unwrap(), + )], } } } diff --git a/src-tauri/src/proxy/providers/streaming.rs b/src-tauri/src/proxy/providers/streaming.rs index 599499016..65a10260e 100644 --- a/src-tauri/src/proxy/providers/streaming.rs +++ b/src-tauri/src/proxy/providers/streaming.rs @@ -88,8 +88,8 @@ struct ToolBlockState { } /// 创建 Anthropic SSE 流 -pub fn create_anthropic_sse_stream( - stream: impl Stream> + Send + 'static, +pub fn create_anthropic_sse_stream( + stream: impl Stream> + Send + 'static, ) -> impl Stream> + Send { async_stream::stream! { let mut buffer = String::new(); @@ -598,7 +598,9 @@ mod tests { "data: [DONE]\n\n" ); - let upstream = stream::iter(vec![Ok(Bytes::from(input.as_bytes().to_vec()))]); + let upstream = stream::iter(vec![Ok::<_, std::io::Error>(Bytes::from( + input.as_bytes().to_vec(), + ))]); let converted = create_anthropic_sse_stream(upstream); let chunks: Vec<_> = converted.collect().await; @@ -686,7 +688,9 @@ mod tests { "data: [DONE]\n\n" ); - let upstream = stream::iter(vec![Ok(Bytes::from(input.as_bytes().to_vec()))]); + let upstream = stream::iter(vec![Ok::<_, std::io::Error>(Bytes::from( + input.as_bytes().to_vec(), + ))]); let converted = create_anthropic_sse_stream(upstream); let chunks: Vec<_> = converted.collect().await; let merged = chunks diff --git a/src-tauri/src/proxy/providers/streaming_responses.rs b/src-tauri/src/proxy/providers/streaming_responses.rs index 6b7e250f3..4ad941f09 100644 --- a/src-tauri/src/proxy/providers/streaming_responses.rs +++ b/src-tauri/src/proxy/providers/streaming_responses.rs @@ -96,8 +96,8 @@ fn resolve_content_index( /// /// 状态机跟踪: message_id, current_model, has_sent_message_start, item/content index map /// SSE 解析支持 named events (event: + data: 行) -pub fn create_anthropic_sse_stream_from_responses( - stream: impl Stream> + Send + 'static, +pub fn create_anthropic_sse_stream_from_responses( + stream: impl Stream> + Send + 'static, ) -> impl Stream> + Send { async_stream::stream! { let mut buffer = String::new(); @@ -800,7 +800,9 @@ mod tests { "data: {\"type\":\"response.completed\",\"response\":{\"status\":\"completed\",\"usage\":{\"input_tokens\":12,\"output_tokens\":3}}}\n\n" ); - let upstream = stream::iter(vec![Ok(Bytes::from(input.as_bytes().to_vec()))]); + let upstream = stream::iter(vec![Ok::<_, std::io::Error>(Bytes::from( + input.as_bytes().to_vec(), + ))]); let converted = create_anthropic_sse_stream_from_responses(upstream); let chunks: Vec<_> = converted.collect().await; @@ -842,7 +844,9 @@ mod tests { "data: {\"type\":\"response.completed\",\"response\":{\"status\":\"completed\",\"usage\":{\"input_tokens\":8,\"output_tokens\":4}}}\n\n" ); - let upstream = stream::iter(vec![Ok(Bytes::from(input.as_bytes().to_vec()))]); + let upstream = stream::iter(vec![Ok::<_, std::io::Error>(Bytes::from( + input.as_bytes().to_vec(), + ))]); let converted = create_anthropic_sse_stream_from_responses(upstream); let chunks: Vec<_> = converted.collect().await; let merged = chunks @@ -913,7 +917,9 @@ mod tests { "data: {\"type\":\"response.completed\",\"response\":{\"status\":\"completed\",\"usage\":{\"input_tokens\":5,\"output_tokens\":10}}}\n\n" ); - let upstream = stream::iter(vec![Ok(Bytes::from(input.as_bytes().to_vec()))]); + let upstream = stream::iter(vec![Ok::<_, std::io::Error>(Bytes::from( + input.as_bytes().to_vec(), + ))]); let converted = create_anthropic_sse_stream_from_responses(upstream); let chunks: Vec<_> = converted.collect().await; let merged = chunks @@ -991,12 +997,17 @@ mod tests { .iter() .filter(|event| { event.get("type").and_then(|v| v.as_str()) == Some("content_block_start") - && event.pointer("/content_block/type").and_then(|v| v.as_str()) == Some("text") + && event + .pointer("/content_block/type") + .and_then(|v| v.as_str()) + == Some("text") }) .count(); let text_stops = events .iter() - .filter(|event| event.get("type").and_then(|v| v.as_str()) == Some("content_block_stop")) + .filter(|event| { + event.get("type").and_then(|v| v.as_str()) == Some("content_block_stop") + }) .count(); let text_deltas: Vec = events .iter() @@ -1005,7 +1016,8 @@ mod tests { && event.pointer("/delta/type").and_then(|v| v.as_str()) == Some("text_delta") }) .filter_map(|event| { - event.pointer("/delta/text") + event + .pointer("/delta/text") .and_then(|v| v.as_str()) .map(ToString::to_string) }) diff --git a/src-tauri/src/proxy/providers/transform.rs b/src-tauri/src/proxy/providers/transform.rs index b2296c3c7..2cdc93c79 100644 --- a/src-tauri/src/proxy/providers/transform.rs +++ b/src-tauri/src/proxy/providers/transform.rs @@ -50,7 +50,7 @@ pub fn resolve_reasoning_effort(body: &Value) -> Option<&'static str> { "medium" => Some("medium"), "high" => Some("high"), "max" => Some("xhigh"), // OpenAI xhigh = maximum reasoning effort - _ => None, // unknown value — do not inject + _ => None, // unknown value — do not inject }; } diff --git a/src-tauri/src/proxy/response_processor.rs b/src-tauri/src/proxy/response_processor.rs index 9f329b678..609ff393a 100644 --- a/src-tauri/src/proxy/response_processor.rs +++ b/src-tauri/src/proxy/response_processor.rs @@ -5,17 +5,19 @@ use super::{ handler_config::UsageParserConfig, handler_context::{RequestContext, StreamingTimeoutConfig}, + hyper_client::ProxyResponse, server::ProxyState, sse::strip_sse_field, usage::parser::TokenUsage, ProxyError, }; +use axum::http::header::HeaderMap; use axum::response::{IntoResponse, Response}; use bytes::Bytes; use futures::stream::{Stream, StreamExt}; -use reqwest::header::HeaderMap; use serde_json::Value; use std::{ + io::Read, sync::{ atomic::{AtomicBool, Ordering}, Arc, @@ -24,24 +26,123 @@ use std::{ }; use tokio::sync::Mutex; +// ============================================================================ +// 响应解压 +// ============================================================================ + +/// 根据 content-encoding 解压响应体字节 +/// +/// reqwest 自动解压已禁用(为了透传 accept-encoding),需要手动解压。 +fn decompress_body(content_encoding: &str, body: &[u8]) -> Result, std::io::Error> { + match content_encoding { + "gzip" | "x-gzip" => { + let mut decoder = flate2::read::GzDecoder::new(body); + let mut decompressed = Vec::new(); + decoder.read_to_end(&mut decompressed)?; + Ok(decompressed) + } + "deflate" => { + let mut decoder = flate2::read::DeflateDecoder::new(body); + let mut decompressed = Vec::new(); + decoder.read_to_end(&mut decompressed)?; + Ok(decompressed) + } + "br" => { + let mut decompressed = Vec::new(); + brotli::BrotliDecompress(&mut std::io::Cursor::new(body), &mut decompressed)?; + Ok(decompressed) + } + _ => { + log::warn!("未知的 content-encoding: {content_encoding},跳过解压"); + Ok(body.to_vec()) + } + } +} + +/// 从响应头提取 content-encoding(忽略 identity 和 chunked) +fn get_content_encoding(headers: &HeaderMap) -> Option { + headers + .get("content-encoding") + .and_then(|v| v.to_str().ok()) + .map(|s| s.trim().to_lowercase()) + .filter(|s| !s.is_empty() && s != "identity") +} + +/// 移除在重建响应体后会失真的实体头。 +pub(crate) fn strip_entity_headers_for_rebuilt_body(headers: &mut HeaderMap) { + headers.remove(axum::http::header::CONTENT_ENCODING); + headers.remove(axum::http::header::CONTENT_LENGTH); + headers.remove(axum::http::header::TRANSFER_ENCODING); +} + +/// 读取响应体并在需要时解压,确保 headers 与返回 body 一致。 +/// +/// `body_timeout`: 整包超时。当非零时用 `tokio::time::timeout` 包住 `.bytes()` 调用, +/// 防止上游发完响应头后卡住 body 导致请求永远挂住。 +/// 传入 `Duration::ZERO` 表示不启用超时(故障转移关闭时)。 +pub(crate) async fn read_decoded_body( + response: ProxyResponse, + tag: &str, + body_timeout: Duration, +) -> Result<(HeaderMap, http::StatusCode, Bytes), ProxyError> { + let mut headers = response.headers().clone(); + let status = response.status(); + let raw_bytes = if body_timeout.is_zero() { + response.bytes().await? + } else { + tokio::time::timeout(body_timeout, response.bytes()) + .await + .map_err(|_| { + ProxyError::Timeout(format!( + "响应体读取超时: {}s(上游发完响应头后 body 未到达)", + body_timeout.as_secs() + )) + })?? + }; + + log::debug!( + "[{tag}] 已接收上游响应体: status={}, bytes={}, headers={}", + status.as_u16(), + raw_bytes.len(), + format_headers(&headers) + ); + + let mut body_bytes = raw_bytes.clone(); + let mut decoded = false; + + if let Some(encoding) = get_content_encoding(&headers) { + log::debug!("[{tag}] 解压非流式响应: content-encoding={encoding}"); + match decompress_body(&encoding, &raw_bytes) { + Ok(decompressed) => { + body_bytes = Bytes::from(decompressed); + decoded = true; + } + Err(e) => { + log::warn!("[{tag}] 解压失败 ({encoding}): {e},使用原始数据"); + } + } + } + + if decoded { + strip_entity_headers_for_rebuilt_body(&mut headers); + } + + Ok((headers, status, body_bytes)) +} + // ============================================================================ // 公共接口 // ============================================================================ /// 检测响应是否为 SSE 流式响应 #[inline] -pub fn is_sse_response(response: &reqwest::Response) -> bool { - response - .headers() - .get("content-type") - .and_then(|v| v.to_str().ok()) - .map(|ct| ct.contains("text/event-stream")) - .unwrap_or(false) +pub fn is_sse_response(response: &ProxyResponse) -> bool { + response.is_sse() } /// 处理流式响应 pub async fn handle_streaming( - response: reqwest::Response, + response: ProxyResponse, ctx: &RequestContext, state: &ProxyState, parser_config: &UsageParserConfig, @@ -53,6 +154,15 @@ pub async fn handle_streaming( status.as_u16(), format_headers(response.headers()) ); + // 检查流式响应是否被压缩(SSE 通常不压缩,如果压缩则 SSE 解析会失败) + if let Some(encoding) = get_content_encoding(response.headers()) { + log::warn!( + "[{}] 流式响应含 content-encoding={encoding},SSE 解析可能失败。\ + 上游在 accept-encoding 透传后压缩了 SSE 流。", + ctx.tag + ); + } + let mut builder = axum::response::Response::builder().status(status); // 复制响应头 @@ -61,9 +171,7 @@ pub async fn handle_streaming( } // 创建字节流 - let stream = response - .bytes_stream() - .map(|chunk| chunk.map_err(|e| std::io::Error::other(e.to_string()))); + let stream = response.bytes_stream(); // 创建使用量收集器 let usage_collector = create_usage_collector(ctx, state, status.as_u16(), parser_config); @@ -87,26 +195,20 @@ pub async fn handle_streaming( /// 处理非流式响应 pub async fn handle_non_streaming( - response: reqwest::Response, + response: ProxyResponse, ctx: &RequestContext, state: &ProxyState, parser_config: &UsageParserConfig, ) -> Result { - let response_headers = response.headers().clone(); - let status = response.status(); - - // 读取响应体 - let body_bytes = response.bytes().await.map_err(|e| { - log::error!("[{}] 读取响应失败: {e}", ctx.tag); - ProxyError::ForwardFailed(format!("Failed to read response body: {e}")) - })?; - log::debug!( - "[{}] 已接收上游响应体: status={}, bytes={}, headers={}", - ctx.tag, - status.as_u16(), - body_bytes.len(), - format_headers(&response_headers) - ); + // 整包超时:仅在故障转移开启且配置值非零时生效 + let body_timeout = + if ctx.app_config.auto_failover_enabled && ctx.app_config.non_streaming_timeout > 0 { + Duration::from_secs(ctx.app_config.non_streaming_timeout as u64) + } else { + Duration::ZERO + }; + let (response_headers, status, body_bytes) = + read_decoded_body(response, ctx.tag, body_timeout).await?; log::debug!( "[{}] 上游响应体内容: {}", @@ -190,7 +292,7 @@ pub async fn handle_non_streaming( /// /// 根据响应类型自动选择流式或非流式处理 pub async fn process_response( - response: reqwest::Response, + response: ProxyResponse, ctx: &RequestContext, state: &ProxyState, parser_config: &UsageParserConfig, diff --git a/src-tauri/src/proxy/server.rs b/src-tauri/src/proxy/server.rs index 46e3635f5..8fa37edf6 100644 --- a/src-tauri/src/proxy/server.rs +++ b/src-tauri/src/proxy/server.rs @@ -1,6 +1,12 @@ //! HTTP代理服务器 //! //! 基于Axum的HTTP服务器,处理代理请求 +//! +//! Uses a manual hyper HTTP/1.1 accept loop with `preserve_header_case(true)` so +//! that the original header-name casing from the CLI client is captured in a +//! `HeaderCaseMap` extension. This map is later forwarded to the upstream via +//! the hyper-based HTTP client, producing wire-level header casing identical to +//! a direct (non-proxied) CLI request. use super::{ failover_switch::FailoverSwitchManager, handlers, log_codes::srv as log_srv, @@ -12,6 +18,7 @@ use axum::{ routing::{get, post}, Router, }; +use hyper_util::rt::TokioIo; use std::net::SocketAddr; use std::sync::Arc; use tokio::sync::{oneshot, RwLock}; @@ -114,15 +121,77 @@ impl ProxyServer { // 记录启动时间 *self.state.start_time.write().await = Some(std::time::Instant::now()); - // 启动服务器 + // 启动服务器 — 使用手动 hyper HTTP/1.1 accept loop + // 开启 preserve_header_case 以捕获客户端请求头的原始大小写 let state = self.state.clone(); let handle = tokio::spawn(async move { - axum::serve(listener, app) - .with_graceful_shutdown(async { - shutdown_rx.await.ok(); - }) - .await - .ok(); + let mut shutdown_rx = shutdown_rx; + loop { + tokio::select! { + result = listener.accept() => { + let (stream, _remote_addr) = match result { + Ok(v) => v, + Err(e) => { + log::error!("[{SRV}] accept 失败: {e}", SRV = log_srv::ACCEPT_ERR); + tokio::time::sleep(std::time::Duration::from_millis(50)).await; + continue; + } + }; + + let app = app.clone(); + tokio::spawn(async move { + // Peek raw TCP bytes to capture original header casing + // before hyper parses (and lowercases) the header names. + let original_cases = { + let mut peek_buf = vec![0u8; 8192]; + match stream.peek(&mut peek_buf).await { + Ok(n) => { + let cases = super::hyper_client::OriginalHeaderCases::from_raw_bytes(&peek_buf[..n]); + log::debug!( + "[ProxyServer] Peeked {} bytes, captured {} header casings", + n, cases.cases.len() + ); + cases + } + Err(e) => { + log::debug!("[ProxyServer] peek failed (non-fatal): {e}"); + super::hyper_client::OriginalHeaderCases::default() + } + } + }; + + // service_fn 将 axum Router(tower::Service)桥接到 hyper + let service = hyper::service::service_fn(move |req: hyper::Request| { + let mut router = app.clone(); + let cases = original_cases.clone(); + async move { + // 将 hyper::body::Incoming 转为 axum::body::Body,保留 extensions + let (mut parts, body) = req.into_parts(); + + // Insert our own header case map alongside hyper's internal one + parts.extensions.insert(cases); + + let body = axum::body::Body::new(body); + let axum_req = http::Request::from_parts(parts, body); + >>::call(&mut router, axum_req).await + } + }); + + if let Err(e) = hyper::server::conn::http1::Builder::new() + .preserve_header_case(true) + .serve_connection(TokioIo::new(stream), service) + .await + { + // Connection reset / broken pipe 等在代理场景下很常见,debug 级别 + log::debug!("[{SRV}] connection error: {e}", SRV = log_srv::CONN_ERR); + } + }); + } + _ = &mut shutdown_rx => { + break; + } + } + } // 服务器停止后更新状态 state.status.write().await.running = false; diff --git a/src-tauri/src/proxy/thinking_optimizer.rs b/src-tauri/src/proxy/thinking_optimizer.rs index c9b770c0b..f29968e04 100644 --- a/src-tauri/src/proxy/thinking_optimizer.rs +++ b/src-tauri/src/proxy/thinking_optimizer.rs @@ -25,7 +25,7 @@ pub fn optimize(body: &mut Value, config: &OptimizerConfig) { } if model.contains("opus-4-6") || model.contains("sonnet-4-6") { - log::info!("[OPT] thinking: adaptive({})", model); + log::info!("[OPT] thinking: adaptive({model})"); body["thinking"] = json!({"type": "adaptive"}); body["output_config"] = json!({"effort": "max"}); append_beta(body, "context-1m-2025-08-07"); @@ -33,7 +33,7 @@ pub fn optimize(body: &mut Value, config: &OptimizerConfig) { } // legacy path - log::info!("[OPT] thinking: legacy({})", model); + log::info!("[OPT] thinking: legacy({model})"); let max_tokens = body .get("max_tokens") diff --git a/src-tauri/src/services/skill.rs b/src-tauri/src/services/skill.rs index 289396e25..7a36b0424 100644 --- a/src-tauri/src/services/skill.rs +++ b/src-tauri/src/services/skill.rs @@ -957,7 +957,7 @@ impl SkillService { if source_path.is_none() { source_path = Some(skill_path); } - log::debug!("Skill '{}' found in source '{}'", dir_name, label); + log::debug!("Skill '{dir_name}' found in source '{label}'"); } } diff --git a/src-tauri/src/services/stream_check.rs b/src-tauri/src/services/stream_check.rs index 504faff8b..39e21d054 100644 --- a/src-tauri/src/services/stream_check.rs +++ b/src-tauri/src/services/stream_check.rs @@ -12,8 +12,8 @@ use std::time::Instant; use crate::app_config::AppType; use crate::error::AppError; use crate::provider::Provider; -use crate::proxy::providers::transform::anthropic_to_openai; use crate::proxy::providers::copilot_auth; +use crate::proxy::providers::transform::anthropic_to_openai; use crate::proxy::providers::transform_responses::anthropic_to_responses; use crate::proxy::providers::{get_adapter, AuthInfo, AuthStrategy}; @@ -95,15 +95,14 @@ impl StreamCheckService { let mut last_result = None; for attempt in 0..=effective_config.max_retries { - let result = - Self::check_once( - app_type, - provider, - &effective_config, - auth_override.clone(), - claude_api_format_override.clone(), - ) - .await; + let result = Self::check_once( + app_type, + provider, + &effective_config, + auth_override.clone(), + claude_api_format_override.clone(), + ) + .await; match &result { Ok(r) if r.success => { @@ -304,6 +303,7 @@ impl StreamCheckService { /// 根据供应商的 api_format 选择请求格式: /// - "anthropic" (默认): Anthropic Messages API (/v1/messages) /// - "openai_chat": OpenAI Chat Completions API (/v1/chat/completions) + #[allow(clippy::too_many_arguments)] async fn check_claude_stream( client: &Client, base_url: &str, @@ -339,12 +339,8 @@ impl StreamCheckService { .unwrap_or(false); let is_openai_chat = effective_api_format == "openai_chat"; let is_openai_responses = effective_api_format == "openai_responses"; - let url = Self::resolve_claude_stream_url( - base, - auth.strategy, - effective_api_format, - is_full_url, - ); + let url = + Self::resolve_claude_stream_url(base, auth.strategy, effective_api_format, is_full_url); let max_tokens = if is_openai_responses { 16 } else { 1 }; @@ -375,8 +371,14 @@ impl StreamCheckService { .header("accept-encoding", "identity") .header("user-agent", copilot_auth::COPILOT_USER_AGENT) .header("editor-version", copilot_auth::COPILOT_EDITOR_VERSION) - .header("editor-plugin-version", copilot_auth::COPILOT_PLUGIN_VERSION) - .header("copilot-integration-id", copilot_auth::COPILOT_INTEGRATION_ID) + .header( + "editor-plugin-version", + copilot_auth::COPILOT_PLUGIN_VERSION, + ) + .header( + "copilot-integration-id", + copilot_auth::COPILOT_INTEGRATION_ID, + ) .header("x-github-api-version", copilot_auth::COPILOT_API_VERSION) .header("openai-intent", "conversation-panel"); } else if is_openai_chat || is_openai_responses { diff --git a/src-tauri/src/services/webdav_sync.rs b/src-tauri/src/services/webdav_sync.rs index db1ca70f6..951785b62 100644 --- a/src-tauri/src/services/webdav_sync.rs +++ b/src-tauri/src/services/webdav_sync.rs @@ -463,12 +463,10 @@ fn validate_manifest_compat(manifest: &SyncManifest, layout: RemoteLayout) -> Re return Err(localized( "webdav.sync.manifest_db_version_incompatible", format!( - "远端数据库快照版本不兼容: db-v{} (本地 db-v{DB_COMPAT_VERSION})", - db_compat_version + "远端数据库快照版本不兼容: db-v{db_compat_version} (本地 db-v{DB_COMPAT_VERSION})" ), format!( - "Remote database snapshot version is incompatible: db-v{} (local db-v{DB_COMPAT_VERSION})", - db_compat_version + "Remote database snapshot version is incompatible: db-v{db_compat_version} (local db-v{DB_COMPAT_VERSION})" ), )); } @@ -476,12 +474,10 @@ fn validate_manifest_compat(manifest: &SyncManifest, layout: RemoteLayout) -> Re return Err(localized( "webdav.sync.manifest_db_version_incompatible", format!( - "远端数据库快照版本不兼容: db-v{} (本地最高支持 db-v{DB_COMPAT_VERSION})", - db_compat_version + "远端数据库快照版本不兼容: db-v{db_compat_version} (本地最高支持 db-v{DB_COMPAT_VERSION})" ), format!( - "Remote database snapshot version is incompatible: db-v{} (local supports up to db-v{DB_COMPAT_VERSION})", - db_compat_version + "Remote database snapshot version is incompatible: db-v{db_compat_version} (local supports up to db-v{DB_COMPAT_VERSION})" ), )); } @@ -661,11 +657,8 @@ fn validate_artifact_size_limit(artifact_name: &str, size: u64) -> Result<(), Ap let max_mb = MAX_SYNC_ARTIFACT_BYTES / 1024 / 1024; return Err(localized( "webdav.sync.artifact_too_large", - format!("artifact {artifact_name} 超过下载上限({} MB)", max_mb), - format!( - "Artifact {artifact_name} exceeds download limit ({} MB)", - max_mb - ), + format!("artifact {artifact_name} 超过下载上限({max_mb} MB)"), + format!("Artifact {artifact_name} exceeds download limit ({max_mb} MB)"), )); } Ok(()) diff --git a/src-tauri/src/services/webdav_sync/archive.rs b/src-tauri/src/services/webdav_sync/archive.rs index aae5de1a8..6741c3832 100644 --- a/src-tauri/src/services/webdav_sync/archive.rs +++ b/src-tauri/src/services/webdav_sync/archive.rs @@ -350,8 +350,8 @@ fn copy_entry_with_total_limit( let max_mb = max_total_bytes / 1024 / 1024; return Err(localized( "webdav.sync.skills_zip_too_large", - format!("skills.zip 解压后体积超过上限({} MB)", max_mb), - format!("skills.zip extracted size exceeds limit ({} MB)", max_mb), + format!("skills.zip 解压后体积超过上限({max_mb} MB)"), + format!("skills.zip extracted size exceeds limit ({max_mb} MB)"), )); } diff --git a/src-tauri/src/session_manager/providers/opencode.rs b/src-tauri/src/session_manager/providers/opencode.rs index 69e4ea76a..2cfad00f8 100644 --- a/src-tauri/src/session_manager/providers/opencode.rs +++ b/src-tauri/src/session_manager/providers/opencode.rs @@ -148,7 +148,7 @@ fn scan_sessions_sqlite() -> Vec { }, created_at: Some(created), last_active_at: Some(updated), - source_path: Some(format!("sqlite:{}:{}", db_display, session_id)), + source_path: Some(format!("sqlite:{db_display}:{session_id}")), resume_command: Some(format!("opencode session resume {session_id}")), }); } diff --git a/src/components/providers/AddProviderDialog.tsx b/src/components/providers/AddProviderDialog.tsx index a4f223d7c..b6e9075a7 100644 --- a/src/components/providers/AddProviderDialog.tsx +++ b/src/components/providers/AddProviderDialog.tsx @@ -42,9 +42,9 @@ export function AddProviderDialog({ const { t } = useTranslation(); // OpenCode and OpenClaw don't support universal providers const showUniversalTab = appId !== "opencode" && appId !== "openclaw"; - const [activeTab, setActiveTab] = useState< - "app-specific" | "universal" - >("app-specific"); + const [activeTab, setActiveTab] = useState<"app-specific" | "universal">( + "app-specific", + ); const [universalFormOpen, setUniversalFormOpen] = useState(false); const [selectedUniversalPreset, setSelectedUniversalPreset] = useState(null); @@ -284,9 +284,7 @@ export function AddProviderDialog({ {showUniversalTab ? ( - setActiveTab(v as "app-specific" | "universal") - } + onValueChange={(v) => setActiveTab(v as "app-specific" | "universal")} > diff --git a/src/components/providers/forms/ClaudeFormFields.tsx b/src/components/providers/forms/ClaudeFormFields.tsx index 851a7f7e9..bdcfb8ebb 100644 --- a/src/components/providers/forms/ClaudeFormFields.tsx +++ b/src/components/providers/forms/ClaudeFormFields.tsx @@ -166,8 +166,7 @@ export function ClaudeFormFields({ apiFormat !== "anthropic" || apiKeyField !== "ANTHROPIC_AUTH_TOKEN" ); - const [advancedExpanded, setAdvancedExpanded] = - useState(hasAnyAdvancedValue); + const [advancedExpanded, setAdvancedExpanded] = useState(hasAnyAdvancedValue); // 预设填充高级值后自动展开(仅从折叠→展开,不会自动折叠) useEffect(() => { @@ -409,10 +408,7 @@ export function ClaudeFormFields({ {/* 高级选项(API 格式 + 认证字段 + 模型映射) */} {shouldShowModelSelector && ( - +