From 67e074c0a7c7fd37b1705b6295a797245dfa6f58 Mon Sep 17 00:00:00 2001 From: Dex Miller Date: Sun, 29 Mar 2026 20:26:15 +0800 Subject: [PATCH] refactor(proxy): transparent header forwarding via hyper client (#1714) MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit * style(frontend): reformat provider forms, constants and hooks Apply prettier formatting across 5 frontend files. No logic changes. Changed files: - AddProviderDialog.tsx: reformat generic type annotation and callback - ClaudeFormFields.tsx: consolidate multi-line useState and Collapsible props - CodexConfigSections.tsx: expand single-line React imports to multi-line, collapse removeCodexTopLevelField() call - constants.ts: merge TemplateType into single line - useSkills.ts: expand single-line TanStack Query imports to multi-line, reformat uninstallSkill mutationFn chain * deps(proxy): add hyper ecosystem crates and manual decompression libs reqwest internally normalizes all header names to lowercase and does not preserve insertion order, causing proxied requests to differ from the original client requests. To achieve transparent header forwarding with original casing and order, introduce lower-level hyper HTTP client libs. New dependencies: - hyper-util 0.1: TokioExecutor + legacy Client with preserve_header_case support for HTTP/1.1 - hyper-rustls 0.27: rustls-based TLS connector for hyper - http 1 / http-body 1 / http-body-util 0.1: HTTP type crates for hyper 1.x request/response construction - flate2 1: manual gzip/deflate decompression (replaces reqwest auto) - brotli 7: manual brotli decompression Changed dependencies: - serde_json: enable preserve_order feature to keep JSON field order - reqwest: drop gzip feature to prevent reqwest from overriding the client's original accept-encoding header * refactor(proxy): use hyper client for header-case preserving forwarding Previously the proxy used reqwest for all upstream requests. reqwest normalizes header names to lowercase and reorders them internally, making proxied requests distinguishable from direct CLI requests. Some upstream providers are sensitive to these differences. This commit replaces reqwest with a hyper-based HTTP client on the default (non-proxy) path, achieving wire-level header fidelity: Server layer (server.rs): - Replace axum::serve with a manual hyper HTTP/1.1 accept loop - Enable preserve_header_case(true) so incoming header casing is captured in a HeaderCaseMap extension on each request - Bridge hyper requests to axum Router via tower::Service New hyper client module (hyper_client.rs): - Lazy-initialized hyper-util Client with preserve_header_case - ProxyResponse enum wrapping both hyper::Response and reqwest::Response behind a unified interface (status, headers, bytes, bytes_stream) - send_request() builds requests with ordered HeaderMap + case map Request handlers (handlers.rs): - Switch from (HeaderMap, Json) extractors to raw axum::extract::Request to preserve Extensions (containing the HeaderCaseMap from the accept loop) - Pass extensions through the forwarding chain Forwarder (forwarder.rs): - Remove HEADER_BLACKLIST array; replace with ordered header iteration that preserves original header sequence and casing - Build ordered_headers by iterating client headers, skipping only auth/host/content-length, and inserting auth headers at the original authorization position to maintain order - Handle anthropic-beta (ensure claude-code-20250219 tag) and anthropic-version (passthrough or default) inline during iteration - Remove should_force_identity_encoding() — accept-encoding is now transparently forwarded to upstream - Use hyper client by default; fall back to reqwest only when an HTTP/SOCKS5 proxy tunnel is configured Provider adapters (adapter.rs, claude.rs, codex.rs, gemini.rs): - Replace add_auth_headers(RequestBuilder) -> RequestBuilder with get_auth_headers(AuthInfo) -> Vec<(HeaderName, HeaderValue)> - Adapters now return header pairs instead of mutating a reqwest builder - Claude adapter: merge Anthropic/ClaudeAuth/Bearer into single branch; move Copilot fingerprint headers into get_auth_headers Response processing (response_processor.rs): - Add manual decompression (gzip/deflate/brotli via flate2 + brotli) for non-streaming responses, since reqwest auto-decompression is now disabled to allow accept-encoding passthrough - Add compressed-SSE warning log for streaming responses - Accept ProxyResponse instead of reqwest::Response HTTP client (http_client.rs): - Disable reqwest auto-decompression (.no_gzip/.no_brotli/.no_deflate) on both global and per-provider clients Streaming adapters (streaming.rs, streaming_responses.rs): - Generalize stream error type from reqwest::Error to generic E: Error Misc: - log_codes.rs: add SRV-005 (ACCEPT_ERR) and SRV-006 (CONN_ERR) - stream_check.rs: reformat copilot header lines - transform.rs: fix trailing whitespace alignment * fix(lint): resolve 35 clippy warnings across Rust codebase Fix all clippy warnings reported by `cargo clippy --lib`: - codex_config.rs: fix doc_overindented_list_items (3 spaces -> 2) - commands/copilot.rs: inline format args in 2 log::error! calls - commands/provider.rs: inline format args in 3 map_err closures - proxy/hyper_client.rs: inline format arg in log::debug! call - proxy/providers/copilot_auth.rs: inline format args in 16 locations (log macros, format! in headers, error constructors) - proxy/thinking_optimizer.rs: inline format args in 2 log::info! calls - services/skill.rs: inline format args in log::debug! call - services/webdav_sync.rs: inline format args in 6 format! calls (version compat messages, download limit messages) - services/webdav_sync/archive.rs: inline format args in 2 format! calls - session_manager/providers/opencode.rs: inline format args in source_path format! All fixes use the clippy::uninlined_format_args suggestion pattern: format!("msg: {}", var) -> format!("msg: {var}") * deps(proxy): add raw HTTP write and native TLS cert dependencies Add crates required for the raw TCP/TLS write path that bypasses hyper's header encoder to preserve original header name casing: - httparse: parse raw TCP peek bytes to capture header casings - tokio-rustls + rustls: direct TLS connections for raw write path - webpki-roots: Mozilla CA bundle baseline - rustls-native-certs: load system keychain CAs (trusts proxy MITM certificates from Clash, mitmproxy, etc.) * fix(proxy): address code review feedback on response handling Fixes from PR #1714 code review: - Extract `read_decoded_body()` and `strip_entity_headers_for_rebuilt_body()` in response_processor to properly clean content-encoding/content-length headers after decompression - Reuse `read_decoded_body()` in handlers.rs for Claude transform path, ensuring compressed responses are decoded before format conversion - Make `build_proxy_url_from_config()` public so forwarder can pass proxy URL to the hyper raw write path - Add `has_system_proxy_env()` utility with test coverage - Add 50ms backoff after accept() failures in server.rs to prevent tight-loop CPU spin on transient socket errors * feat(proxy): implement raw TCP/TLS write with HTTP CONNECT tunnel Rewrite hyper_client with a two-tier strategy for header case preservation: Primary path (raw write): - Peek raw TCP bytes in server.rs to capture OriginalHeaderCases before hyper lowercases them - Build raw HTTP/1.1 request bytes with exact original header name casing - Write directly to TLS stream, then use WriteFilter to let hyper parse the response while discarding its duplicate request writes - Support HTTP CONNECT tunneling through upstream proxies, so header case is preserved even when a proxy (Clash, V2Ray) is configured Fallback path (hyper-util Client): - Used when OriginalHeaderCases is empty or raw write fails - Configured with title_case_headers(true) for best-effort casing TLS improvements: - Load native system certificates alongside webpki roots so proxy MITM CAs (installed in system keychain) are trusted through CONNECT tunnels Key types added: - OriginalHeaderCases: maps lowercase name → original wire-casing bytes - WriteFilter: AsyncRead+AsyncWrite wrapper that discards writes - connect_via_proxy(): HTTP CONNECT tunnel establishment - ExtensionDebugMarker: diagnostic marker for extension chain debugging * refactor(proxy): route requests through hyper with proxy-aware forwarding Rework forwarder request dispatch to always prefer the hyper raw write path (header case preservation) over reqwest: Request routing: - HTTP/HTTPS proxy: hyper raw write through CONNECT tunnel (case preserved) - SOCKS5 proxy: reqwest fallback (CONNECT not supported for SOCKS5) - No proxy: hyper raw write direct connection Header handling improvements: - Replace host header in-place at original position instead of skip-and-append, preserving client's header ordering - Preserve client's original accept-encoding for transparent passthrough; only force identity encoding when transform path needs decompression - Add should_force_identity_encoding() to centralize the decision - Remove hardcoded 'br, gzip, deflate' override that masked client values Proxy URL resolution (priority order): 1. Provider-specific proxy config (if enabled) 2. Global proxy URL configured in CC Switch 3. Direct connection (no proxy) * chore(proxy): remove dead code, redundant tests and debug scaffolding - Inline should_force_identity_encoding() (was just `needs_transform`) and delete its 5 test cases - Remove ExtensionDebugMarker diagnostic type - Remove unused has_system_proxy_env() and its test - Remove strip_entity_headers test - Simplify hyper path: remove redundant is_socks_proxy ternary - Update hyper_client module doc to reflect CONNECT tunnel support * fix(proxy): block direct-connect fallback and complete CONNECT tunnel support * feat(hooks): improve proxy requirement warnings with specific reasons - Remove redundant OpenAI format hint toast messages - Add detailed reason detection for proxy requirements (OpenAI Chat, OpenAI Responses, full URL mode) - Update i18n files with new reason-specific keys * style(*): format code with prettier - Remove extra whitespace in http_client.rs - Fix formatting issues in useProviderActions.ts * fix(proxy): post-merge fixes for forward return type and clippy warnings - Restore forward() return type to (ProxyResponse, Option) to pass claude_api_format through to callers - Inline format args in log::warn! macro (clippy::uninlined_format_args) - Suppress too_many_arguments for check_claude_stream * refactor(proxy): preserve original header wire order and add non-streaming body timeout - Rewrite build_raw_request to emit headers in original client-sent sequence instead of hash-map order - Remove unused OriginalHeaderCases::get_all method - Add body_timeout to read_decoded_body to prevent requests hanging when upstream stalls after headers --- src-tauri/Cargo.lock | 129 +++- src-tauri/Cargo.toml | 16 +- src-tauri/src/codex_config.rs | 2 +- src-tauri/src/commands/copilot.rs | 4 +- src-tauri/src/commands/provider.rs | 6 +- src-tauri/src/commands/stream_check.rs | 10 +- src-tauri/src/proxy/error.rs | 1 - src-tauri/src/proxy/forwarder.rs | 438 +++++++---- src-tauri/src/proxy/handlers.rs | 112 ++- src-tauri/src/proxy/http_client.rs | 12 +- src-tauri/src/proxy/hyper_client.rs | 685 ++++++++++++++++++ src-tauri/src/proxy/log_codes.rs | 2 + src-tauri/src/proxy/mod.rs | 1 + src-tauri/src/proxy/providers/adapter.rs | 89 +-- src-tauri/src/proxy/providers/claude.rs | 63 +- src-tauri/src/proxy/providers/codex.rs | 9 +- src-tauri/src/proxy/providers/copilot_auth.rs | 39 +- src-tauri/src/proxy/providers/gemini.rs | 24 +- src-tauri/src/proxy/providers/streaming.rs | 12 +- .../proxy/providers/streaming_responses.rs | 28 +- src-tauri/src/proxy/providers/transform.rs | 2 +- src-tauri/src/proxy/response_processor.rs | 160 +++- src-tauri/src/proxy/server.rs | 83 ++- src-tauri/src/proxy/thinking_optimizer.rs | 4 +- src-tauri/src/services/skill.rs | 2 +- src-tauri/src/services/stream_check.rs | 38 +- src-tauri/src/services/webdav_sync.rs | 19 +- src-tauri/src/services/webdav_sync/archive.rs | 4 +- .../src/session_manager/providers/opencode.rs | 2 +- .../providers/AddProviderDialog.tsx | 10 +- .../providers/forms/ClaudeFormFields.tsx | 8 +- .../providers/forms/CodexConfigSections.tsx | 13 +- src/config/constants.ts | 3 +- src/hooks/useProviderActions.ts | 46 +- src/hooks/useSkills.ts | 18 +- src/i18n/locales/en.json | 7 +- src/i18n/locales/ja.json | 7 +- src/i18n/locales/zh.json | 7 +- 38 files changed, 1606 insertions(+), 509 deletions(-) create mode 100644 src-tauri/src/proxy/hyper_client.rs 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 && ( - +