mirror of
https://github.com/farion1231/cc-switch.git
synced 2026-07-26 14:35:22 +08:00
Compare commits
17 Commits
| Author | SHA1 | Date | |
|---|---|---|---|
| 5a433b0ff7 | |||
| 1555dbc55e | |||
| 3551e3c496 | |||
| 08014f99e6 | |||
| 3006c6a23d | |||
| 9b14721d4c | |||
| 65c96db0d1 | |||
| e5867ca2d1 | |||
| 854f19d0e1 | |||
| 905f7ccbfe | |||
| 3e4c87278f | |||
| f4e960253e | |||
| 324a1da8e6 | |||
| 49f66bcc9a | |||
| 4084b53834 | |||
| 2c2c72271a | |||
| 6a1ba46f2a |
Generated
+90
-39
@@ -137,18 +137,6 @@ dependencies = [
|
|||||||
"pin-project-lite",
|
"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]]
|
[[package]]
|
||||||
name = "async-executor"
|
name = "async-executor"
|
||||||
version = "1.14.0"
|
version = "1.14.0"
|
||||||
@@ -324,6 +312,28 @@ version = "1.5.0"
|
|||||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||||
checksum = "c08606f8c3cbf4ce6ec8e28fb0014a2c086708fe954eaa885384a6165172e7e8"
|
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]]
|
[[package]]
|
||||||
name = "axum"
|
name = "axum"
|
||||||
version = "0.7.9"
|
version = "0.7.9"
|
||||||
@@ -496,6 +506,17 @@ dependencies = [
|
|||||||
"syn 2.0.117",
|
"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]]
|
[[package]]
|
||||||
name = "brotli"
|
name = "brotli"
|
||||||
version = "8.0.2"
|
version = "8.0.2"
|
||||||
@@ -504,7 +525,17 @@ checksum = "4bd8b9603c7aa97359dbd97ecf258968c95f3adddd6db2f7e7a5bef101c84560"
|
|||||||
dependencies = [
|
dependencies = [
|
||||||
"alloc-no-stdlib",
|
"alloc-no-stdlib",
|
||||||
"alloc-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]]
|
[[package]]
|
||||||
@@ -691,11 +722,19 @@ dependencies = [
|
|||||||
"auto-launch",
|
"auto-launch",
|
||||||
"axum",
|
"axum",
|
||||||
"base64 0.22.1",
|
"base64 0.22.1",
|
||||||
|
"brotli 7.0.0",
|
||||||
"bytes",
|
"bytes",
|
||||||
"chrono",
|
"chrono",
|
||||||
"dirs 5.0.1",
|
"dirs 5.0.1",
|
||||||
|
"flate2",
|
||||||
"futures",
|
"futures",
|
||||||
|
"http",
|
||||||
|
"http-body",
|
||||||
|
"http-body-util",
|
||||||
|
"httparse",
|
||||||
"hyper",
|
"hyper",
|
||||||
|
"hyper-rustls",
|
||||||
|
"hyper-util",
|
||||||
"indexmap 2.13.0",
|
"indexmap 2.13.0",
|
||||||
"json-five",
|
"json-five",
|
||||||
"json5",
|
"json5",
|
||||||
@@ -708,6 +747,8 @@ dependencies = [
|
|||||||
"rquickjs",
|
"rquickjs",
|
||||||
"rusqlite",
|
"rusqlite",
|
||||||
"rust_decimal",
|
"rust_decimal",
|
||||||
|
"rustls",
|
||||||
|
"rustls-native-certs",
|
||||||
"serde",
|
"serde",
|
||||||
"serde_json",
|
"serde_json",
|
||||||
"serde_yaml",
|
"serde_yaml",
|
||||||
@@ -726,6 +767,7 @@ dependencies = [
|
|||||||
"tempfile",
|
"tempfile",
|
||||||
"thiserror 2.0.18",
|
"thiserror 2.0.18",
|
||||||
"tokio",
|
"tokio",
|
||||||
|
"tokio-rustls",
|
||||||
"toml 0.8.23",
|
"toml 0.8.23",
|
||||||
"toml_edit 0.22.27",
|
"toml_edit 0.22.27",
|
||||||
"tower 0.4.13",
|
"tower 0.4.13",
|
||||||
@@ -733,6 +775,7 @@ dependencies = [
|
|||||||
"url",
|
"url",
|
||||||
"uuid",
|
"uuid",
|
||||||
"webkit2gtk",
|
"webkit2gtk",
|
||||||
|
"webpki-roots 0.26.11",
|
||||||
"winreg 0.52.0",
|
"winreg 0.52.0",
|
||||||
"zip 2.4.2",
|
"zip 2.4.2",
|
||||||
]
|
]
|
||||||
@@ -800,6 +843,15 @@ dependencies = [
|
|||||||
"inout",
|
"inout",
|
||||||
]
|
]
|
||||||
|
|
||||||
|
[[package]]
|
||||||
|
name = "cmake"
|
||||||
|
version = "0.1.57"
|
||||||
|
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||||
|
checksum = "75443c44cd6b379beb8c5b45d85d0773baf31cce901fe7bb252f4eff3008ef7d"
|
||||||
|
dependencies = [
|
||||||
|
"cc",
|
||||||
|
]
|
||||||
|
|
||||||
[[package]]
|
[[package]]
|
||||||
name = "combine"
|
name = "combine"
|
||||||
version = "4.6.7"
|
version = "4.6.7"
|
||||||
@@ -810,23 +862,6 @@ dependencies = [
|
|||||||
"memchr",
|
"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]]
|
[[package]]
|
||||||
name = "concurrent-queue"
|
name = "concurrent-queue"
|
||||||
version = "2.5.0"
|
version = "2.5.0"
|
||||||
@@ -1573,6 +1608,12 @@ dependencies = [
|
|||||||
"percent-encoding",
|
"percent-encoding",
|
||||||
]
|
]
|
||||||
|
|
||||||
|
[[package]]
|
||||||
|
name = "fs_extra"
|
||||||
|
version = "1.3.0"
|
||||||
|
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||||
|
checksum = "42703706b716c37f96a77aea830392ad231f44c9e9a67872fa5548707e11b11c"
|
||||||
|
|
||||||
[[package]]
|
[[package]]
|
||||||
name = "funty"
|
name = "funty"
|
||||||
version = "2.0.0"
|
version = "2.0.0"
|
||||||
@@ -2206,12 +2247,14 @@ dependencies = [
|
|||||||
"http",
|
"http",
|
||||||
"hyper",
|
"hyper",
|
||||||
"hyper-util",
|
"hyper-util",
|
||||||
|
"log",
|
||||||
"rustls",
|
"rustls",
|
||||||
|
"rustls-native-certs",
|
||||||
"rustls-pki-types",
|
"rustls-pki-types",
|
||||||
"tokio",
|
"tokio",
|
||||||
"tokio-rustls",
|
"tokio-rustls",
|
||||||
"tower-service",
|
"tower-service",
|
||||||
"webpki-roots",
|
"webpki-roots 1.0.6",
|
||||||
]
|
]
|
||||||
|
|
||||||
[[package]]
|
[[package]]
|
||||||
@@ -4200,7 +4243,7 @@ dependencies = [
|
|||||||
"wasm-bindgen-futures",
|
"wasm-bindgen-futures",
|
||||||
"wasm-streams 0.4.2",
|
"wasm-streams 0.4.2",
|
||||||
"web-sys",
|
"web-sys",
|
||||||
"webpki-roots",
|
"webpki-roots 1.0.6",
|
||||||
]
|
]
|
||||||
|
|
||||||
[[package]]
|
[[package]]
|
||||||
@@ -4410,6 +4453,8 @@ version = "0.23.37"
|
|||||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||||
checksum = "758025cb5fccfd3bc2fd74708fd4682be41d99e5dff73c377c0646c6012c73a4"
|
checksum = "758025cb5fccfd3bc2fd74708fd4682be41d99e5dff73c377c0646c6012c73a4"
|
||||||
dependencies = [
|
dependencies = [
|
||||||
|
"aws-lc-rs",
|
||||||
|
"log",
|
||||||
"once_cell",
|
"once_cell",
|
||||||
"ring",
|
"ring",
|
||||||
"rustls-pki-types",
|
"rustls-pki-types",
|
||||||
@@ -4473,6 +4518,7 @@ version = "0.103.9"
|
|||||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||||
checksum = "d7df23109aa6c1567d1c575b9952556388da57401e4ace1d15f79eedad0d8f53"
|
checksum = "d7df23109aa6c1567d1c575b9952556388da57401e4ace1d15f79eedad0d8f53"
|
||||||
dependencies = [
|
dependencies = [
|
||||||
|
"aws-lc-rs",
|
||||||
"ring",
|
"ring",
|
||||||
"rustls-pki-types",
|
"rustls-pki-types",
|
||||||
"untrusted",
|
"untrusted",
|
||||||
@@ -4715,6 +4761,7 @@ version = "1.0.149"
|
|||||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||||
checksum = "83fc039473c5595ace860d8c4fafa220ff474b3fc6bfdb4293327f1a37e94d86"
|
checksum = "83fc039473c5595ace860d8c4fafa220ff474b3fc6bfdb4293327f1a37e94d86"
|
||||||
dependencies = [
|
dependencies = [
|
||||||
|
"indexmap 2.13.0",
|
||||||
"itoa",
|
"itoa",
|
||||||
"memchr",
|
"memchr",
|
||||||
"serde",
|
"serde",
|
||||||
@@ -5325,7 +5372,7 @@ source = "registry+https://github.com/rust-lang/crates.io-index"
|
|||||||
checksum = "d4a24476afd977c5d5d169f72425868613d82747916dd29e0a357c84c4bd6d29"
|
checksum = "d4a24476afd977c5d5d169f72425868613d82747916dd29e0a357c84c4bd6d29"
|
||||||
dependencies = [
|
dependencies = [
|
||||||
"base64 0.22.1",
|
"base64 0.22.1",
|
||||||
"brotli",
|
"brotli 8.0.2",
|
||||||
"ico",
|
"ico",
|
||||||
"json-patch",
|
"json-patch",
|
||||||
"plist",
|
"plist",
|
||||||
@@ -5613,7 +5660,7 @@ source = "registry+https://github.com/rust-lang/crates.io-index"
|
|||||||
checksum = "219a1f983a2af3653f75b5747f76733b0da7ff03069c7a41901a5eb3ace4557d"
|
checksum = "219a1f983a2af3653f75b5747f76733b0da7ff03069c7a41901a5eb3ace4557d"
|
||||||
dependencies = [
|
dependencies = [
|
||||||
"anyhow",
|
"anyhow",
|
||||||
"brotli",
|
"brotli 8.0.2",
|
||||||
"cargo_metadata",
|
"cargo_metadata",
|
||||||
"ctor",
|
"ctor",
|
||||||
"dunce",
|
"dunce",
|
||||||
@@ -6017,18 +6064,13 @@ version = "0.6.8"
|
|||||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||||
checksum = "d4e6559d53cc268e5031cd8429d05415bc4cb4aefc4aa5d6cc35fbf5b924a1f8"
|
checksum = "d4e6559d53cc268e5031cd8429d05415bc4cb4aefc4aa5d6cc35fbf5b924a1f8"
|
||||||
dependencies = [
|
dependencies = [
|
||||||
"async-compression",
|
|
||||||
"bitflags 2.11.0",
|
"bitflags 2.11.0",
|
||||||
"bytes",
|
"bytes",
|
||||||
"futures-core",
|
|
||||||
"futures-util",
|
"futures-util",
|
||||||
"http",
|
"http",
|
||||||
"http-body",
|
"http-body",
|
||||||
"http-body-util",
|
|
||||||
"iri-string",
|
"iri-string",
|
||||||
"pin-project-lite",
|
"pin-project-lite",
|
||||||
"tokio",
|
|
||||||
"tokio-util",
|
|
||||||
"tower 0.5.3",
|
"tower 0.5.3",
|
||||||
"tower-layer",
|
"tower-layer",
|
||||||
"tower-service",
|
"tower-service",
|
||||||
@@ -6564,6 +6606,15 @@ dependencies = [
|
|||||||
"rustls-pki-types",
|
"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]]
|
[[package]]
|
||||||
name = "webpki-roots"
|
name = "webpki-roots"
|
||||||
version = "1.0.6"
|
version = "1.0.6"
|
||||||
|
|||||||
+14
-2
@@ -23,7 +23,7 @@ test-hooks = []
|
|||||||
tauri-build = { version = "2.4.0", features = [] }
|
tauri-build = { version = "2.4.0", features = [] }
|
||||||
|
|
||||||
[dependencies]
|
[dependencies]
|
||||||
serde_json = "1.0"
|
serde_json = { version = "1.0", features = ["preserve_order"] }
|
||||||
serde = { version = "1.0", features = ["derive"] }
|
serde = { version = "1.0", features = ["derive"] }
|
||||||
log = "0.4"
|
log = "0.4"
|
||||||
chrono = { version = "0.4", features = ["serde"] }
|
chrono = { version = "0.4", features = ["serde"] }
|
||||||
@@ -38,7 +38,9 @@ tauri-plugin-deep-link = "2"
|
|||||||
dirs = "5.0"
|
dirs = "5.0"
|
||||||
toml = "0.8"
|
toml = "0.8"
|
||||||
toml_edit = "0.22"
|
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"] }
|
tokio = { version = "1", features = ["macros", "rt-multi-thread", "time", "sync"] }
|
||||||
futures = "0.3"
|
futures = "0.3"
|
||||||
async-stream = "0.3"
|
async-stream = "0.3"
|
||||||
@@ -47,6 +49,16 @@ axum = "0.7"
|
|||||||
tower = "0.4"
|
tower = "0.4"
|
||||||
tower-http = { version = "0.5", features = ["cors"] }
|
tower-http = { version = "0.5", features = ["cors"] }
|
||||||
hyper = { version = "1.0", features = ["full"] }
|
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"
|
regex = "1.10"
|
||||||
rquickjs = { version = "0.8", features = ["array-buffer", "classes"] }
|
rquickjs = { version = "0.8", features = ["array-buffer", "classes"] }
|
||||||
thiserror = "2.0"
|
thiserror = "2.0"
|
||||||
|
|||||||
@@ -49,7 +49,7 @@ pub async fn copilot_poll_for_auth(
|
|||||||
Ok(false)
|
Ok(false)
|
||||||
}
|
}
|
||||||
Err(e) => {
|
Err(e) => {
|
||||||
log::error!("[CopilotAuth] 轮询失败: {}", e);
|
log::error!("[CopilotAuth] 轮询失败: {e}");
|
||||||
Err(e.to_string())
|
Err(e.to_string())
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
@@ -70,7 +70,7 @@ pub async fn copilot_poll_for_account(
|
|||||||
Ok(None)
|
Ok(None)
|
||||||
}
|
}
|
||||||
Err(e) => {
|
Err(e) => {
|
||||||
log::error!("[CopilotAuth] 轮询失败: {}", e);
|
log::error!("[CopilotAuth] 轮询失败: {e}");
|
||||||
Err(e.to_string())
|
Err(e.to_string())
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -159,7 +159,7 @@ pub async fn queryProviderUsage(
|
|||||||
let providers = state
|
let providers = state
|
||||||
.db
|
.db
|
||||||
.get_all_providers(app_type.as_str())
|
.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 provider = providers.get(&providerId);
|
||||||
let is_copilot = provider
|
let is_copilot = provider
|
||||||
@@ -182,11 +182,11 @@ pub async fn queryProviderUsage(
|
|||||||
Some(account_id) => auth_manager
|
Some(account_id) => auth_manager
|
||||||
.fetch_usage_for_account(account_id)
|
.fetch_usage_for_account(account_id)
|
||||||
.await
|
.await
|
||||||
.map_err(|e| format!("Failed to fetch Copilot usage: {}", e))?,
|
.map_err(|e| format!("Failed to fetch Copilot usage: {e}"))?,
|
||||||
None => auth_manager
|
None => auth_manager
|
||||||
.fetch_usage()
|
.fetch_usage()
|
||||||
.await
|
.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 premium = &usage.quota_snapshots.premium_interactions;
|
||||||
let used = premium.entitlement - premium.remaining;
|
let used = premium.entitlement - premium.remaining;
|
||||||
|
|||||||
@@ -216,7 +216,11 @@ async fn resolve_claude_api_format_override(
|
|||||||
.and_then(|meta| meta.managed_account_id_for("github_copilot"));
|
.and_then(|meta| meta.managed_account_id_for("github_copilot"));
|
||||||
|
|
||||||
let vendor_result = match account_id.as_deref() {
|
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,
|
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",
|
Ok(Some(_)) | Ok(None) => "openai_chat",
|
||||||
Err(err) => {
|
Err(err) => {
|
||||||
log::warn!(
|
log::warn!(
|
||||||
"[StreamCheck] Failed to resolve Copilot model vendor for {}: {}. Falling back to chat/completions",
|
"[StreamCheck] Failed to resolve Copilot model vendor for {model_id}: {err}. Falling back to chat/completions"
|
||||||
model_id,
|
|
||||||
err
|
|
||||||
);
|
);
|
||||||
"openai_chat"
|
"openai_chat"
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -68,7 +68,6 @@ pub enum ProxyError {
|
|||||||
StreamIdleTimeout(u64),
|
StreamIdleTimeout(u64),
|
||||||
|
|
||||||
/// 认证错误
|
/// 认证错误
|
||||||
#[allow(dead_code)]
|
|
||||||
#[error("认证失败: {0}")]
|
#[error("认证失败: {0}")]
|
||||||
AuthError(String),
|
AuthError(String),
|
||||||
|
|
||||||
|
|||||||
+275
-149
@@ -2,6 +2,7 @@
|
|||||||
//!
|
//!
|
||||||
//! 负责将请求转发到上游Provider,支持故障转移
|
//! 负责将请求转发到上游Provider,支持故障转移
|
||||||
|
|
||||||
|
use super::hyper_client::ProxyResponse;
|
||||||
use super::{
|
use super::{
|
||||||
body_filter::filter_private_params_with_whitelist,
|
body_filter::filter_private_params_with_whitelist,
|
||||||
error::*,
|
error::*,
|
||||||
@@ -19,67 +20,14 @@ use super::{
|
|||||||
use crate::commands::CopilotAuthState;
|
use crate::commands::CopilotAuthState;
|
||||||
use crate::proxy::providers::copilot_auth::CopilotAuthManager;
|
use crate::proxy::providers::copilot_auth::CopilotAuthManager;
|
||||||
use crate::{app_config::AppType, provider::Provider};
|
use crate::{app_config::AppType, provider::Provider};
|
||||||
use reqwest::Response;
|
use http::Extensions;
|
||||||
use serde_json::Value;
|
use serde_json::Value;
|
||||||
use std::sync::Arc;
|
use std::sync::Arc;
|
||||||
use tauri::Manager;
|
use tauri::Manager;
|
||||||
use tokio::sync::RwLock;
|
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 struct ForwardResult {
|
||||||
pub response: Response,
|
pub response: ProxyResponse,
|
||||||
pub provider: Provider,
|
pub provider: Provider,
|
||||||
pub claude_api_format: Option<String>,
|
pub claude_api_format: Option<String>,
|
||||||
}
|
}
|
||||||
@@ -150,6 +98,7 @@ impl RequestForwarder {
|
|||||||
endpoint: &str,
|
endpoint: &str,
|
||||||
body: Value,
|
body: Value,
|
||||||
headers: axum::http::HeaderMap,
|
headers: axum::http::HeaderMap,
|
||||||
|
extensions: Extensions,
|
||||||
providers: Vec<Provider>,
|
providers: Vec<Provider>,
|
||||||
) -> Result<ForwardResult, ForwardError> {
|
) -> Result<ForwardResult, ForwardError> {
|
||||||
// 获取适配器
|
// 获取适配器
|
||||||
@@ -226,6 +175,7 @@ impl RequestForwarder {
|
|||||||
endpoint,
|
endpoint,
|
||||||
&provider_body,
|
&provider_body,
|
||||||
&headers,
|
&headers,
|
||||||
|
&extensions,
|
||||||
adapter.as_ref(),
|
adapter.as_ref(),
|
||||||
)
|
)
|
||||||
.await
|
.await
|
||||||
@@ -355,6 +305,7 @@ impl RequestForwarder {
|
|||||||
endpoint,
|
endpoint,
|
||||||
&provider_body,
|
&provider_body,
|
||||||
&headers,
|
&headers,
|
||||||
|
&extensions,
|
||||||
adapter.as_ref(),
|
adapter.as_ref(),
|
||||||
)
|
)
|
||||||
.await
|
.await
|
||||||
@@ -553,6 +504,7 @@ impl RequestForwarder {
|
|||||||
endpoint,
|
endpoint,
|
||||||
&provider_body,
|
&provider_body,
|
||||||
&headers,
|
&headers,
|
||||||
|
&extensions,
|
||||||
adapter.as_ref(),
|
adapter.as_ref(),
|
||||||
)
|
)
|
||||||
.await
|
.await
|
||||||
@@ -791,8 +743,9 @@ impl RequestForwarder {
|
|||||||
endpoint: &str,
|
endpoint: &str,
|
||||||
body: &Value,
|
body: &Value,
|
||||||
headers: &axum::http::HeaderMap,
|
headers: &axum::http::HeaderMap,
|
||||||
|
extensions: &Extensions,
|
||||||
adapter: &dyn ProviderAdapter,
|
adapter: &dyn ProviderAdapter,
|
||||||
) -> Result<(Response, Option<String>), ProxyError> {
|
) -> Result<(ProxyResponse, Option<String>), ProxyError> {
|
||||||
// 使用适配器提取 base_url
|
// 使用适配器提取 base_url
|
||||||
let base_url = adapter.extract_base_url(provider)?;
|
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 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 auth_headers = if let Some(mut auth) = adapter.extract_auth(provider) {
|
||||||
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) {
|
|
||||||
// GitHub Copilot 特殊处理:从 CopilotAuthManager 获取真实 token
|
// GitHub Copilot 特殊处理:从 CopilotAuthManager 获取真实 token
|
||||||
if auth.strategy == AuthStrategy::GitHubCopilot {
|
if auth.strategy == AuthStrategy::GitHubCopilot {
|
||||||
if let Some(app_handle) = &self.app_handle {
|
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::<http::Uri>()
|
||||||
|
.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;
|
||||||
}
|
}
|
||||||
|
|
||||||
// anthropic-version 统一处理(仅 Claude):优先使用客户端的版本号,否则使用默认值
|
// --- 连接 / 追踪 / CDN 类 — 无条件跳过 ---
|
||||||
// 注意:只设置一次,避免重复
|
if matches!(
|
||||||
if adapter.name() == "Claude" {
|
key_str,
|
||||||
let version_str = headers
|
"content-length"
|
||||||
.get("anthropic-version")
|
| "transfer-encoding"
|
||||||
.and_then(|v| v.to_str().ok())
|
| "x-forwarded-host"
|
||||||
.unwrap_or("2023-06-01");
|
| "x-forwarded-port"
|
||||||
request = request.header("anthropic-version", version_str);
|
| "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());
|
||||||
|
}
|
||||||
|
|
||||||
|
// 如果原始请求中没有认证头,在末尾追加
|
||||||
|
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,8 +1101,43 @@ 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<String> = 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| {
|
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() {
|
if e.is_timeout() {
|
||||||
ProxyError::Timeout(format!("请求超时: {e}"))
|
ProxyError::Timeout(format!("请求超时: {e}"))
|
||||||
} else if e.is_connect() {
|
} else if e.is_connect() {
|
||||||
@@ -1040,6 +1146,21 @@ impl RequestForwarder {
|
|||||||
ProxyError::ForwardFailed(e.to_string())
|
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();
|
let status = response.status();
|
||||||
@@ -1048,7 +1169,7 @@ impl RequestForwarder {
|
|||||||
Ok((response, resolved_claude_api_format))
|
Ok((response, resolved_claude_api_format))
|
||||||
} else {
|
} else {
|
||||||
let status_code = status.as_u16();
|
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 {
|
Err(ProxyError::UpstreamError {
|
||||||
status: status_code,
|
status: status_code,
|
||||||
@@ -1094,7 +1215,11 @@ impl RequestForwarder {
|
|||||||
.and_then(|m| m.managed_account_id_for("github_copilot"));
|
.and_then(|m| m.managed_account_id_for("github_copilot"));
|
||||||
|
|
||||||
let vendor_result = match account_id.as_deref() {
|
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,
|
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)]
|
#[cfg(test)]
|
||||||
mod tests {
|
mod tests {
|
||||||
use super::*;
|
use super::*;
|
||||||
use axum::http::{header::ACCEPT, HeaderMap, HeaderValue};
|
use axum::http::header::{HeaderValue, ACCEPT};
|
||||||
|
use axum::http::HeaderMap;
|
||||||
use serde_json::json;
|
use serde_json::json;
|
||||||
|
|
||||||
#[test]
|
#[test]
|
||||||
|
|||||||
@@ -18,7 +18,10 @@ use super::{
|
|||||||
streaming_responses::create_anthropic_sse_stream_from_responses, transform,
|
streaming_responses::create_anthropic_sse_stream_from_responses, transform,
|
||||||
transform_responses,
|
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,
|
server::ProxyState,
|
||||||
types::*,
|
types::*,
|
||||||
usage::parser::TokenUsage,
|
usage::parser::TokenUsage,
|
||||||
@@ -27,6 +30,7 @@ use super::{
|
|||||||
use crate::app_config::AppType;
|
use crate::app_config::AppType;
|
||||||
use axum::{extract::State, http::StatusCode, response::IntoResponse, Json};
|
use axum::{extract::State, http::StatusCode, response::IntoResponse, Json};
|
||||||
use bytes::Bytes;
|
use bytes::Bytes;
|
||||||
|
use http_body_util::BodyExt;
|
||||||
use serde_json::{json, Value};
|
use serde_json::{json, Value};
|
||||||
|
|
||||||
// ============================================================================
|
// ============================================================================
|
||||||
@@ -61,10 +65,20 @@ pub async fn get_status(State(state): State<ProxyState>) -> Result<Json<ProxySta
|
|||||||
/// - 现在 OpenRouter 已推出 Claude Code 兼容接口,默认不再启用该转换(逻辑保留以备回退)
|
/// - 现在 OpenRouter 已推出 Claude Code 兼容接口,默认不再启用该转换(逻辑保留以备回退)
|
||||||
pub async fn handle_messages(
|
pub async fn handle_messages(
|
||||||
State(state): State<ProxyState>,
|
State(state): State<ProxyState>,
|
||||||
uri: axum::http::Uri,
|
request: axum::extract::Request,
|
||||||
headers: axum::http::HeaderMap,
|
|
||||||
Json(body): Json<Value>,
|
|
||||||
) -> Result<axum::response::Response, ProxyError> {
|
) -> Result<axum::response::Response, ProxyError> {
|
||||||
|
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 =
|
let mut ctx =
|
||||||
RequestContext::new(&state, &body, &headers, AppType::Claude, "Claude", "claude").await?;
|
RequestContext::new(&state, &body, &headers, AppType::Claude, "Claude", "claude").await?;
|
||||||
|
|
||||||
@@ -86,6 +100,7 @@ pub async fn handle_messages(
|
|||||||
endpoint,
|
endpoint,
|
||||||
body.clone(),
|
body.clone(),
|
||||||
headers,
|
headers,
|
||||||
|
extensions,
|
||||||
ctx.get_providers(),
|
ctx.get_providers(),
|
||||||
)
|
)
|
||||||
.await
|
.await
|
||||||
@@ -126,7 +141,7 @@ pub async fn handle_messages(
|
|||||||
///
|
///
|
||||||
/// 支持 OpenAI Chat Completions 和 Responses API 两种格式的转换
|
/// 支持 OpenAI Chat Completions 和 Responses API 两种格式的转换
|
||||||
async fn handle_claude_transform(
|
async fn handle_claude_transform(
|
||||||
response: reqwest::Response,
|
response: super::hyper_client::ProxyResponse,
|
||||||
ctx: &RequestContext,
|
ctx: &RequestContext,
|
||||||
state: &ProxyState,
|
state: &ProxyState,
|
||||||
_original_body: &Value,
|
_original_body: &Value,
|
||||||
@@ -211,12 +226,14 @@ async fn handle_claude_transform(
|
|||||||
}
|
}
|
||||||
|
|
||||||
// 非流式响应转换 (OpenAI/Responses → Anthropic)
|
// 非流式响应转换 (OpenAI/Responses → Anthropic)
|
||||||
let response_headers = response.headers().clone();
|
let body_timeout =
|
||||||
|
if ctx.app_config.auto_failover_enabled && ctx.app_config.non_streaming_timeout > 0 {
|
||||||
let body_bytes = response.bytes().await.map_err(|e| {
|
std::time::Duration::from_secs(ctx.app_config.non_streaming_timeout as u64)
|
||||||
log::error!("[Claude] 读取响应体失败: {e}");
|
} else {
|
||||||
ProxyError::ForwardFailed(format!("Failed to read response body: {e}"))
|
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);
|
let body_str = String::from_utf8_lossy(&body_bytes);
|
||||||
|
|
||||||
@@ -269,14 +286,11 @@ async fn handle_claude_transform(
|
|||||||
|
|
||||||
// 构建响应
|
// 构建响应
|
||||||
let mut builder = axum::response::Response::builder().status(status);
|
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() {
|
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");
|
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)
|
/// 处理 /v1/chat/completions 请求(OpenAI Chat Completions API - Codex CLI)
|
||||||
pub async fn handle_chat_completions(
|
pub async fn handle_chat_completions(
|
||||||
State(state): State<ProxyState>,
|
State(state): State<ProxyState>,
|
||||||
uri: axum::http::Uri,
|
request: axum::extract::Request,
|
||||||
headers: axum::http::HeaderMap,
|
|
||||||
Json(body): Json<Value>,
|
|
||||||
) -> Result<axum::response::Response, ProxyError> {
|
) -> Result<axum::response::Response, ProxyError> {
|
||||||
|
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 =
|
let mut ctx =
|
||||||
RequestContext::new(&state, &body, &headers, AppType::Codex, "Codex", "codex").await?;
|
RequestContext::new(&state, &body, &headers, AppType::Codex, "Codex", "codex").await?;
|
||||||
let endpoint = endpoint_with_query(&uri, "/chat/completions");
|
let endpoint = endpoint_with_query(&uri, "/chat/completions");
|
||||||
@@ -326,6 +350,7 @@ pub async fn handle_chat_completions(
|
|||||||
&endpoint,
|
&endpoint,
|
||||||
body,
|
body,
|
||||||
headers,
|
headers,
|
||||||
|
extensions,
|
||||||
ctx.get_providers(),
|
ctx.get_providers(),
|
||||||
)
|
)
|
||||||
.await
|
.await
|
||||||
@@ -349,10 +374,20 @@ pub async fn handle_chat_completions(
|
|||||||
/// 处理 /v1/responses 请求(OpenAI Responses API - Codex CLI 透传)
|
/// 处理 /v1/responses 请求(OpenAI Responses API - Codex CLI 透传)
|
||||||
pub async fn handle_responses(
|
pub async fn handle_responses(
|
||||||
State(state): State<ProxyState>,
|
State(state): State<ProxyState>,
|
||||||
uri: axum::http::Uri,
|
request: axum::extract::Request,
|
||||||
headers: axum::http::HeaderMap,
|
|
||||||
Json(body): Json<Value>,
|
|
||||||
) -> Result<axum::response::Response, ProxyError> {
|
) -> Result<axum::response::Response, ProxyError> {
|
||||||
|
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 =
|
let mut ctx =
|
||||||
RequestContext::new(&state, &body, &headers, AppType::Codex, "Codex", "codex").await?;
|
RequestContext::new(&state, &body, &headers, AppType::Codex, "Codex", "codex").await?;
|
||||||
let endpoint = endpoint_with_query(&uri, "/responses");
|
let endpoint = endpoint_with_query(&uri, "/responses");
|
||||||
@@ -369,6 +404,7 @@ pub async fn handle_responses(
|
|||||||
&endpoint,
|
&endpoint,
|
||||||
body,
|
body,
|
||||||
headers,
|
headers,
|
||||||
|
extensions,
|
||||||
ctx.get_providers(),
|
ctx.get_providers(),
|
||||||
)
|
)
|
||||||
.await
|
.await
|
||||||
@@ -392,10 +428,20 @@ pub async fn handle_responses(
|
|||||||
/// 处理 /v1/responses/compact 请求(OpenAI Responses Compact API - Codex CLI 透传)
|
/// 处理 /v1/responses/compact 请求(OpenAI Responses Compact API - Codex CLI 透传)
|
||||||
pub async fn handle_responses_compact(
|
pub async fn handle_responses_compact(
|
||||||
State(state): State<ProxyState>,
|
State(state): State<ProxyState>,
|
||||||
uri: axum::http::Uri,
|
request: axum::extract::Request,
|
||||||
headers: axum::http::HeaderMap,
|
|
||||||
Json(body): Json<Value>,
|
|
||||||
) -> Result<axum::response::Response, ProxyError> {
|
) -> Result<axum::response::Response, ProxyError> {
|
||||||
|
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 =
|
let mut ctx =
|
||||||
RequestContext::new(&state, &body, &headers, AppType::Codex, "Codex", "codex").await?;
|
RequestContext::new(&state, &body, &headers, AppType::Codex, "Codex", "codex").await?;
|
||||||
let endpoint = endpoint_with_query(&uri, "/responses/compact");
|
let endpoint = endpoint_with_query(&uri, "/responses/compact");
|
||||||
@@ -412,6 +458,7 @@ pub async fn handle_responses_compact(
|
|||||||
&endpoint,
|
&endpoint,
|
||||||
body,
|
body,
|
||||||
headers,
|
headers,
|
||||||
|
extensions,
|
||||||
ctx.get_providers(),
|
ctx.get_providers(),
|
||||||
)
|
)
|
||||||
.await
|
.await
|
||||||
@@ -440,9 +487,19 @@ pub async fn handle_responses_compact(
|
|||||||
pub async fn handle_gemini(
|
pub async fn handle_gemini(
|
||||||
State(state): State<ProxyState>,
|
State(state): State<ProxyState>,
|
||||||
uri: axum::http::Uri,
|
uri: axum::http::Uri,
|
||||||
headers: axum::http::HeaderMap,
|
request: axum::extract::Request,
|
||||||
Json(body): Json<Value>,
|
|
||||||
) -> Result<axum::response::Response, ProxyError> {
|
) -> Result<axum::response::Response, ProxyError> {
|
||||||
|
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 中
|
// Gemini 的模型名称在 URI 中
|
||||||
let mut ctx = RequestContext::new(&state, &body, &headers, AppType::Gemini, "Gemini", "gemini")
|
let mut ctx = RequestContext::new(&state, &body, &headers, AppType::Gemini, "Gemini", "gemini")
|
||||||
.await?
|
.await?
|
||||||
@@ -466,6 +523,7 @@ pub async fn handle_gemini(
|
|||||||
endpoint,
|
endpoint,
|
||||||
body,
|
body,
|
||||||
headers,
|
headers,
|
||||||
|
extensions,
|
||||||
ctx.get_providers(),
|
ctx.get_providers(),
|
||||||
)
|
)
|
||||||
.await
|
.await
|
||||||
|
|||||||
@@ -219,7 +219,12 @@ fn build_client(proxy_url: Option<&str>) -> Result<Client, String> {
|
|||||||
.timeout(Duration::from_secs(600))
|
.timeout(Duration::from_secs(600))
|
||||||
.connect_timeout(Duration::from_secs(30))
|
.connect_timeout(Duration::from_secs(30))
|
||||||
.pool_max_idle_per_host(10)
|
.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 {
|
if let Some(url) = proxy_url {
|
||||||
@@ -332,7 +337,7 @@ pub fn mask_url(url: &str) -> String {
|
|||||||
/// 根据供应商单独代理配置构建代理 URL
|
/// 根据供应商单独代理配置构建代理 URL
|
||||||
///
|
///
|
||||||
/// 将 ProviderProxyConfig 转换为代理 URL 字符串
|
/// 将 ProviderProxyConfig 转换为代理 URL 字符串
|
||||||
fn build_proxy_url_from_config(config: &ProviderProxyConfig) -> Option<String> {
|
pub fn build_proxy_url_from_config(config: &ProviderProxyConfig) -> Option<String> {
|
||||||
let proxy_type = config.proxy_type.as_deref().unwrap_or("http");
|
let proxy_type = config.proxy_type.as_deref().unwrap_or("http");
|
||||||
let host = config.proxy_host.as_deref()?;
|
let host = config.proxy_host.as_deref()?;
|
||||||
let port = config.proxy_port?;
|
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))
|
.connect_timeout(Duration::from_secs(30))
|
||||||
.pool_max_idle_per_host(10)
|
.pool_max_idle_per_host(10)
|
||||||
.tcp_keepalive(Duration::from_secs(60))
|
.tcp_keepalive(Duration::from_secs(60))
|
||||||
|
.no_gzip()
|
||||||
|
.no_brotli()
|
||||||
|
.no_deflate()
|
||||||
.proxy(proxy)
|
.proxy(proxy)
|
||||||
.build()
|
.build()
|
||||||
{
|
{
|
||||||
|
|||||||
@@ -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<u8>)>,
|
||||||
|
}
|
||||||
|
|
||||||
|
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<hyper_util::client::legacy::connect::HttpConnector>,
|
||||||
|
http_body_util::Full<Bytes>,
|
||||||
|
>;
|
||||||
|
|
||||||
|
/// Lazily-initialized hyper client with header-case preservation enabled.
|
||||||
|
fn global_hyper_client() -> &'static HyperClient {
|
||||||
|
static CLIENT: OnceLock<HyperClient> = 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<hyper::body::Incoming>),
|
||||||
|
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<Bytes, ProxyError> {
|
||||||
|
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<Item = Result<Bytes, std::io::Error>> + 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<Box<dyn Stream<Item = Result<Bytes, std::io::Error>> + 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<u8>,
|
||||||
|
timeout: std::time::Duration,
|
||||||
|
proxy_url: Option<&str>,
|
||||||
|
) -> Result<ProxyResponse, ProxyError> {
|
||||||
|
// Extract our own OriginalHeaderCases if available
|
||||||
|
let original_cases = original_extensions.get::<OriginalHeaderCases>().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<tokio_rustls::client::TlsStream<tokio::net::TcpStream>>),
|
||||||
|
}
|
||||||
|
|
||||||
|
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<std::io::Result<()>> {
|
||||||
|
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<std::io::Result<usize>> {
|
||||||
|
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<std::io::Result<()>> {
|
||||||
|
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<std::io::Result<()>> {
|
||||||
|
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<ProxyResponse, ProxyError> {
|
||||||
|
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<ProxyStream, ProxyError> {
|
||||||
|
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<tokio_rustls::TlsConnector> = 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<u8> {
|
||||||
|
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<String> =
|
||||||
|
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<String, usize> =
|
||||||
|
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<S>(
|
||||||
|
stream: WriteFilter<S>,
|
||||||
|
method: http::Method,
|
||||||
|
) -> Result<ProxyResponse, ProxyError>
|
||||||
|
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<Bytes>>(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<S> {
|
||||||
|
inner: S,
|
||||||
|
}
|
||||||
|
|
||||||
|
impl<S> WriteFilter<S> {
|
||||||
|
fn new(inner: S) -> Self {
|
||||||
|
Self { inner }
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
impl<S: tokio::io::AsyncRead + Unpin> tokio::io::AsyncRead for WriteFilter<S> {
|
||||||
|
fn poll_read(
|
||||||
|
self: std::pin::Pin<&mut Self>,
|
||||||
|
cx: &mut std::task::Context<'_>,
|
||||||
|
buf: &mut tokio::io::ReadBuf<'_>,
|
||||||
|
) -> std::task::Poll<std::io::Result<()>> {
|
||||||
|
// Pass reads through to the underlying stream
|
||||||
|
let inner = std::pin::Pin::new(&mut self.get_mut().inner);
|
||||||
|
inner.poll_read(cx, buf)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
impl<S: Unpin> tokio::io::AsyncWrite for WriteFilter<S> {
|
||||||
|
fn poll_write(
|
||||||
|
self: std::pin::Pin<&mut Self>,
|
||||||
|
_cx: &mut std::task::Context<'_>,
|
||||||
|
buf: &[u8],
|
||||||
|
) -> std::task::Poll<std::io::Result<usize>> {
|
||||||
|
// 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::io::Result<()>> {
|
||||||
|
std::task::Poll::Ready(Ok(()))
|
||||||
|
}
|
||||||
|
|
||||||
|
fn poll_shutdown(
|
||||||
|
self: std::pin::Pin<&mut Self>,
|
||||||
|
_cx: &mut std::task::Context<'_>,
|
||||||
|
) -> std::task::Poll<std::io::Result<()>> {
|
||||||
|
std::task::Poll::Ready(Ok(()))
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -26,6 +26,8 @@ pub mod srv {
|
|||||||
pub const STOPPED: &str = "SRV-002";
|
pub const STOPPED: &str = "SRV-002";
|
||||||
pub const STOP_TIMEOUT: &str = "SRV-003";
|
pub const STOP_TIMEOUT: &str = "SRV-003";
|
||||||
pub const TASK_ERROR: &str = "SRV-004";
|
pub const TASK_ERROR: &str = "SRV-004";
|
||||||
|
pub const ACCEPT_ERR: &str = "SRV-005";
|
||||||
|
pub const CONN_ERR: &str = "SRV-006";
|
||||||
}
|
}
|
||||||
|
|
||||||
/// 转发器日志码
|
/// 转发器日志码
|
||||||
|
|||||||
@@ -14,6 +14,7 @@ pub mod handler_context;
|
|||||||
mod handlers;
|
mod handlers;
|
||||||
mod health;
|
mod health;
|
||||||
pub mod http_client;
|
pub mod http_client;
|
||||||
|
pub mod hyper_client;
|
||||||
pub mod log_codes;
|
pub mod log_codes;
|
||||||
pub mod model_mapper;
|
pub mod model_mapper;
|
||||||
pub mod provider_router;
|
pub mod provider_router;
|
||||||
|
|||||||
@@ -5,7 +5,6 @@
|
|||||||
use super::auth::AuthInfo;
|
use super::auth::AuthInfo;
|
||||||
use crate::provider::Provider;
|
use crate::provider::Provider;
|
||||||
use crate::proxy::error::ProxyError;
|
use crate::proxy::error::ProxyError;
|
||||||
use reqwest::RequestBuilder;
|
|
||||||
use serde_json::Value;
|
use serde_json::Value;
|
||||||
|
|
||||||
/// 供应商适配器 Trait
|
/// 供应商适配器 Trait
|
||||||
@@ -14,116 +13,36 @@ use serde_json::Value;
|
|||||||
/// - URL 构建
|
/// - URL 构建
|
||||||
/// - 认证信息提取和头部注入
|
/// - 认证信息提取和头部注入
|
||||||
/// - 请求/响应格式转换(可选)
|
/// - 请求/响应格式转换(可选)
|
||||||
///
|
|
||||||
/// # 示例
|
|
||||||
///
|
|
||||||
/// ```ignore
|
|
||||||
/// pub struct ClaudeAdapter;
|
|
||||||
///
|
|
||||||
/// impl ProviderAdapter for ClaudeAdapter {
|
|
||||||
/// fn name(&self) -> &'static str { "Claude" }
|
|
||||||
///
|
|
||||||
/// fn extract_base_url(&self, provider: &Provider) -> Result<String, ProxyError> {
|
|
||||||
/// // 从 provider 配置中提取 base_url
|
|
||||||
/// }
|
|
||||||
///
|
|
||||||
/// fn extract_auth(&self, provider: &Provider) -> Option<AuthInfo> {
|
|
||||||
/// // 从 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 {
|
pub trait ProviderAdapter: Send + Sync {
|
||||||
/// 适配器名称(用于日志和调试)
|
/// 适配器名称(用于日志和调试)
|
||||||
fn name(&self) -> &'static str;
|
fn name(&self) -> &'static str;
|
||||||
|
|
||||||
/// 从 Provider 配置中提取 base_url
|
/// 从 Provider 配置中提取 base_url
|
||||||
///
|
|
||||||
/// # Arguments
|
|
||||||
/// * `provider` - Provider 配置
|
|
||||||
///
|
|
||||||
/// # Returns
|
|
||||||
/// * `Ok(String)` - 提取到的 base_url(已去除尾部斜杠)
|
|
||||||
/// * `Err(ProxyError)` - 提取失败
|
|
||||||
fn extract_base_url(&self, provider: &Provider) -> Result<String, ProxyError>;
|
fn extract_base_url(&self, provider: &Provider) -> Result<String, ProxyError>;
|
||||||
|
|
||||||
/// 从 Provider 配置中提取认证信息
|
/// 从 Provider 配置中提取认证信息
|
||||||
///
|
|
||||||
/// # Arguments
|
|
||||||
/// * `provider` - Provider 配置
|
|
||||||
///
|
|
||||||
/// # Returns
|
|
||||||
/// * `Some(AuthInfo)` - 提取到的认证信息
|
|
||||||
/// * `None` - 未找到认证信息
|
|
||||||
fn extract_auth(&self, provider: &Provider) -> Option<AuthInfo>;
|
fn extract_auth(&self, provider: &Provider) -> Option<AuthInfo>;
|
||||||
|
|
||||||
/// 构建请求 URL
|
/// 构建请求 URL
|
||||||
///
|
|
||||||
/// # Arguments
|
|
||||||
/// * `base_url` - 基础 URL
|
|
||||||
/// * `endpoint` - 请求端点(如 `/v1/messages`)
|
|
||||||
///
|
|
||||||
/// # Returns
|
|
||||||
/// 完整的请求 URL
|
|
||||||
fn build_url(&self, base_url: &str, endpoint: &str) -> String;
|
fn build_url(&self, base_url: &str, endpoint: &str) -> String;
|
||||||
|
|
||||||
/// 添加认证头到请求
|
/// Return auth headers as `(name, value)` pairs.
|
||||||
///
|
///
|
||||||
/// # Arguments
|
/// The forwarder inserts these at the position of the original auth header
|
||||||
/// * `request` - reqwest RequestBuilder
|
/// so that header order is preserved.
|
||||||
/// * `auth` - 认证信息
|
fn get_auth_headers(&self, auth: &AuthInfo) -> Vec<(http::HeaderName, http::HeaderValue)>;
|
||||||
///
|
|
||||||
/// # Returns
|
|
||||||
/// 添加了认证头的 RequestBuilder
|
|
||||||
fn add_auth_headers(&self, request: RequestBuilder, auth: &AuthInfo) -> RequestBuilder;
|
|
||||||
|
|
||||||
/// 是否需要格式转换
|
/// 是否需要格式转换
|
||||||
///
|
|
||||||
/// 默认返回 `false`(透传模式)。
|
|
||||||
/// 仅当供应商需要格式转换时(如 Claude + OpenRouter 旧 OpenAI 兼容接口)才返回 `true`。
|
|
||||||
///
|
|
||||||
/// # Arguments
|
|
||||||
/// * `provider` - Provider 配置
|
|
||||||
fn needs_transform(&self, _provider: &Provider) -> bool {
|
fn needs_transform(&self, _provider: &Provider) -> bool {
|
||||||
false
|
false
|
||||||
}
|
}
|
||||||
|
|
||||||
/// 转换请求体
|
/// 转换请求体
|
||||||
///
|
|
||||||
/// 将请求体从一种格式转换为另一种格式(如 Anthropic → OpenAI)。
|
|
||||||
/// 默认实现直接返回原始请求体(透传)。
|
|
||||||
///
|
|
||||||
/// # Arguments
|
|
||||||
/// * `body` - 原始请求体
|
|
||||||
/// * `provider` - Provider 配置(用于获取模型映射等)
|
|
||||||
///
|
|
||||||
/// # Returns
|
|
||||||
/// * `Ok(Value)` - 转换后的请求体
|
|
||||||
/// * `Err(ProxyError)` - 转换失败
|
|
||||||
fn transform_request(&self, body: Value, _provider: &Provider) -> Result<Value, ProxyError> {
|
fn transform_request(&self, body: Value, _provider: &Provider) -> Result<Value, ProxyError> {
|
||||||
Ok(body)
|
Ok(body)
|
||||||
}
|
}
|
||||||
|
|
||||||
/// 转换响应体
|
/// 转换响应体
|
||||||
///
|
|
||||||
/// 将响应体从一种格式转换为另一种格式(如 OpenAI → Anthropic)。
|
|
||||||
/// 默认实现直接返回原始响应体(透传)。
|
|
||||||
///
|
|
||||||
/// # Arguments
|
|
||||||
/// * `body` - 原始响应体
|
|
||||||
///
|
|
||||||
/// # Returns
|
|
||||||
/// * `Ok(Value)` - 转换后的响应体
|
|
||||||
/// * `Err(ProxyError)` - 转换失败
|
|
||||||
///
|
|
||||||
/// Note: 响应转换将在 handler 层集成,目前预留接口
|
|
||||||
#[allow(dead_code)]
|
#[allow(dead_code)]
|
||||||
fn transform_response(&self, body: Value) -> Result<Value, ProxyError> {
|
fn transform_response(&self, body: Value) -> Result<Value, ProxyError> {
|
||||||
Ok(body)
|
Ok(body)
|
||||||
|
|||||||
@@ -16,7 +16,6 @@
|
|||||||
use super::{AuthInfo, AuthStrategy, ProviderAdapter, ProviderType};
|
use super::{AuthInfo, AuthStrategy, ProviderAdapter, ProviderType};
|
||||||
use crate::provider::Provider;
|
use crate::provider::Provider;
|
||||||
use crate::proxy::error::ProxyError;
|
use crate::proxy::error::ProxyError;
|
||||||
use reqwest::RequestBuilder;
|
|
||||||
|
|
||||||
/// 获取 Claude 供应商的 API 格式
|
/// 获取 Claude 供应商的 API 格式
|
||||||
///
|
///
|
||||||
@@ -337,32 +336,50 @@ impl ProviderAdapter for ClaudeAdapter {
|
|||||||
base
|
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 由 forwarder.rs 统一处理(透传客户端值或设置默认值)
|
||||||
// 这里不再设置 anthropic-version,避免 header 重复
|
let bearer = format!("Bearer {}", auth.api_key);
|
||||||
match auth.strategy {
|
match auth.strategy {
|
||||||
// Anthropic 官方: Authorization Bearer + x-api-key
|
AuthStrategy::Anthropic | AuthStrategy::ClaudeAuth | AuthStrategy::Bearer => {
|
||||||
AuthStrategy::Anthropic => request
|
vec![(
|
||||||
.header("Authorization", format!("Bearer {}", auth.api_key))
|
HeaderName::from_static("authorization"),
|
||||||
.header("x-api-key", &auth.api_key),
|
HeaderValue::from_str(&bearer).unwrap(),
|
||||||
// ClaudeAuth 中转服务: 仅 Bearer,无 x-api-key
|
)]
|
||||||
AuthStrategy::ClaudeAuth => {
|
|
||||||
request.header("Authorization", format!("Bearer {}", auth.api_key))
|
|
||||||
}
|
}
|
||||||
// OpenRouter: Bearer
|
AuthStrategy::GitHubCopilot => {
|
||||||
AuthStrategy::Bearer => {
|
vec![
|
||||||
request.header("Authorization", format!("Bearer {}", auth.api_key))
|
(
|
||||||
|
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 + 统一指纹头
|
_ => vec![],
|
||||||
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,
|
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
@@ -9,7 +9,6 @@ use super::{AuthInfo, AuthStrategy, ProviderAdapter};
|
|||||||
use crate::provider::Provider;
|
use crate::provider::Provider;
|
||||||
use crate::proxy::error::ProxyError;
|
use crate::proxy::error::ProxyError;
|
||||||
use regex::Regex;
|
use regex::Regex;
|
||||||
use reqwest::RequestBuilder;
|
|
||||||
use std::sync::LazyLock;
|
use std::sync::LazyLock;
|
||||||
|
|
||||||
/// 官方 Codex 客户端 User-Agent 正则
|
/// 官方 Codex 客户端 User-Agent 正则
|
||||||
@@ -174,8 +173,12 @@ impl ProviderAdapter for CodexAdapter {
|
|||||||
url
|
url
|
||||||
}
|
}
|
||||||
|
|
||||||
fn add_auth_headers(&self, request: RequestBuilder, auth: &AuthInfo) -> RequestBuilder {
|
fn get_auth_headers(&self, auth: &AuthInfo) -> Vec<(http::HeaderName, http::HeaderValue)> {
|
||||||
request.header("Authorization", format!("Bearer {}", auth.api_key))
|
let bearer = format!("Bearer {}", auth.api_key);
|
||||||
|
vec![(
|
||||||
|
http::HeaderName::from_static("authorization"),
|
||||||
|
http::HeaderValue::from_str(&bearer).unwrap(),
|
||||||
|
)]
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
@@ -341,7 +341,7 @@ impl CopilotAuthManager {
|
|||||||
|
|
||||||
// 尝试从磁盘加载(同步,不发起网络请求)
|
// 尝试从磁盘加载(同步,不发起网络请求)
|
||||||
if let Err(e) = manager.load_from_disk_sync() {
|
if let Err(e) = manager.load_from_disk_sync() {
|
||||||
log::warn!("[CopilotAuth] 加载存储失败: {}", e);
|
log::warn!("[CopilotAuth] 加载存储失败: {e}");
|
||||||
}
|
}
|
||||||
|
|
||||||
manager
|
manager
|
||||||
@@ -364,7 +364,7 @@ impl CopilotAuthManager {
|
|||||||
|
|
||||||
/// 移除指定账号
|
/// 移除指定账号
|
||||||
pub async fn remove_account(&self, account_id: &str) -> Result<(), CopilotAuthError> {
|
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;
|
let mut accounts = self.accounts.write().await;
|
||||||
@@ -482,8 +482,7 @@ impl CopilotAuthManager {
|
|||||||
let status = response.status();
|
let status = response.status();
|
||||||
let text = response.text().await.unwrap_or_default();
|
let text = response.text().await.unwrap_or_default();
|
||||||
return Err(CopilotAuthError::NetworkError(format!(
|
return Err(CopilotAuthError::NetworkError(format!(
|
||||||
"GitHub 设备码请求失败: {} - {}",
|
"GitHub 设备码请求失败: {status} - {text}"
|
||||||
status, text
|
|
||||||
)));
|
)));
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -581,10 +580,7 @@ impl CopilotAuthManager {
|
|||||||
}
|
}
|
||||||
|
|
||||||
// 需要刷新
|
// 需要刷新
|
||||||
log::info!(
|
log::info!("[CopilotAuth] 账号 {account_id} 的 Copilot Token 需要刷新");
|
||||||
"[CopilotAuth] 账号 {} 的 Copilot Token 需要刷新",
|
|
||||||
account_id
|
|
||||||
);
|
|
||||||
|
|
||||||
let refresh_lock = self.get_refresh_lock(account_id).await;
|
let refresh_lock = self.get_refresh_lock(account_id).await;
|
||||||
let _refresh_guard = refresh_lock.lock().await;
|
let _refresh_guard = refresh_lock.lock().await;
|
||||||
@@ -660,12 +656,12 @@ impl CopilotAuthManager {
|
|||||||
) -> Result<Vec<CopilotModel>, CopilotAuthError> {
|
) -> Result<Vec<CopilotModel>, CopilotAuthError> {
|
||||||
let copilot_token = self.get_valid_token_for_account(account_id).await?;
|
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
|
let response = self
|
||||||
.http_client
|
.http_client
|
||||||
.get(COPILOT_MODELS_URL)
|
.get(COPILOT_MODELS_URL)
|
||||||
.header("Authorization", format!("Bearer {}", copilot_token))
|
.header("Authorization", format!("Bearer {copilot_token}"))
|
||||||
.header("Content-Type", "application/json")
|
.header("Content-Type", "application/json")
|
||||||
.header("copilot-integration-id", "vscode-chat")
|
.header("copilot-integration-id", "vscode-chat")
|
||||||
.header("editor-version", COPILOT_EDITOR_VERSION)
|
.header("editor-version", COPILOT_EDITOR_VERSION)
|
||||||
@@ -679,8 +675,7 @@ impl CopilotAuthManager {
|
|||||||
let status = response.status();
|
let status = response.status();
|
||||||
let text = response.text().await.unwrap_or_default();
|
let text = response.text().await.unwrap_or_default();
|
||||||
return Err(CopilotAuthError::CopilotTokenFetchFailed(format!(
|
return Err(CopilotAuthError::CopilotTokenFetchFailed(format!(
|
||||||
"获取模型列表失败: {} - {}",
|
"获取模型列表失败: {status} - {text}"
|
||||||
status, text
|
|
||||||
)));
|
)));
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -749,12 +744,12 @@ impl CopilotAuthManager {
|
|||||||
.ok_or_else(|| CopilotAuthError::AccountNotFound(account_id.to_string()))?
|
.ok_or_else(|| CopilotAuthError::AccountNotFound(account_id.to_string()))?
|
||||||
};
|
};
|
||||||
|
|
||||||
log::info!("[CopilotAuth] 获取账号 {} 的 Copilot 使用量", account_id);
|
log::info!("[CopilotAuth] 获取账号 {account_id} 的 Copilot 使用量");
|
||||||
|
|
||||||
let response = self
|
let response = self
|
||||||
.http_client
|
.http_client
|
||||||
.get(COPILOT_USAGE_URL)
|
.get(COPILOT_USAGE_URL)
|
||||||
.header("Authorization", format!("token {}", github_token))
|
.header("Authorization", format!("token {github_token}"))
|
||||||
.header("Content-Type", "application/json")
|
.header("Content-Type", "application/json")
|
||||||
.header("editor-version", COPILOT_EDITOR_VERSION)
|
.header("editor-version", COPILOT_EDITOR_VERSION)
|
||||||
.header("editor-plugin-version", COPILOT_PLUGIN_VERSION)
|
.header("editor-plugin-version", COPILOT_PLUGIN_VERSION)
|
||||||
@@ -771,8 +766,7 @@ impl CopilotAuthManager {
|
|||||||
let status = response.status();
|
let status = response.status();
|
||||||
let text = response.text().await.unwrap_or_default();
|
let text = response.text().await.unwrap_or_default();
|
||||||
return Err(CopilotAuthError::CopilotTokenFetchFailed(format!(
|
return Err(CopilotAuthError::CopilotTokenFetchFailed(format!(
|
||||||
"获取使用量失败: {} - {}",
|
"获取使用量失败: {status} - {text}"
|
||||||
status, text
|
|
||||||
)));
|
)));
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -1000,7 +994,7 @@ impl CopilotAuthManager {
|
|||||||
let response = self
|
let response = self
|
||||||
.http_client
|
.http_client
|
||||||
.get(GITHUB_USER_URL)
|
.get(GITHUB_USER_URL)
|
||||||
.header("Authorization", format!("token {}", github_token))
|
.header("Authorization", format!("token {github_token}"))
|
||||||
.header("User-Agent", COPILOT_USER_AGENT)
|
.header("User-Agent", COPILOT_USER_AGENT)
|
||||||
.header("Editor-Version", COPILOT_EDITOR_VERSION)
|
.header("Editor-Version", COPILOT_EDITOR_VERSION)
|
||||||
.header("Editor-Plugin-Version", COPILOT_PLUGIN_VERSION)
|
.header("Editor-Plugin-Version", COPILOT_PLUGIN_VERSION)
|
||||||
@@ -1027,12 +1021,12 @@ impl CopilotAuthManager {
|
|||||||
github_token: &str,
|
github_token: &str,
|
||||||
account_id: &str,
|
account_id: &str,
|
||||||
) -> Result<(), CopilotAuthError> {
|
) -> Result<(), CopilotAuthError> {
|
||||||
log::debug!("[CopilotAuth] 获取账号 {} 的 Copilot Token", account_id);
|
log::debug!("[CopilotAuth] 获取账号 {account_id} 的 Copilot Token");
|
||||||
|
|
||||||
let response = self
|
let response = self
|
||||||
.http_client
|
.http_client
|
||||||
.get(COPILOT_TOKEN_URL)
|
.get(COPILOT_TOKEN_URL)
|
||||||
.header("Authorization", format!("token {}", github_token))
|
.header("Authorization", format!("token {github_token}"))
|
||||||
.header("User-Agent", COPILOT_USER_AGENT)
|
.header("User-Agent", COPILOT_USER_AGENT)
|
||||||
.header("Editor-Version", COPILOT_EDITOR_VERSION)
|
.header("Editor-Version", COPILOT_EDITOR_VERSION)
|
||||||
.header("Editor-Plugin-Version", COPILOT_PLUGIN_VERSION)
|
.header("Editor-Plugin-Version", COPILOT_PLUGIN_VERSION)
|
||||||
@@ -1051,8 +1045,7 @@ impl CopilotAuthManager {
|
|||||||
let status = response.status();
|
let status = response.status();
|
||||||
let text = response.text().await.unwrap_or_default();
|
let text = response.text().await.unwrap_or_default();
|
||||||
return Err(CopilotAuthError::CopilotTokenFetchFailed(format!(
|
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)
|
.fetch_copilot_token_with_github_token(&legacy_token, &account_id)
|
||||||
.await
|
.await
|
||||||
{
|
{
|
||||||
log::warn!("[CopilotAuth] 迁移时验证 Copilot 订阅失败: {}", e);
|
log::warn!("[CopilotAuth] 迁移时验证 Copilot 订阅失败: {e}");
|
||||||
}
|
}
|
||||||
|
|
||||||
// 添加账号
|
// 添加账号
|
||||||
@@ -1149,7 +1142,7 @@ impl CopilotAuthManager {
|
|||||||
"Legacy Copilot auth migration failed: {e}"
|
"Legacy Copilot auth migration failed: {e}"
|
||||||
)))
|
)))
|
||||||
.await;
|
.await;
|
||||||
log::warn!("[CopilotAuth] 迁移失败,旧 token 可能已失效: {}", e);
|
log::warn!("[CopilotAuth] 迁移失败,旧 token 可能已失效: {e}");
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
@@ -9,7 +9,6 @@
|
|||||||
use super::{AuthInfo, AuthStrategy, ProviderAdapter, ProviderType};
|
use super::{AuthInfo, AuthStrategy, ProviderAdapter, ProviderType};
|
||||||
use crate::provider::Provider;
|
use crate::provider::Provider;
|
||||||
use crate::proxy::error::ProxyError;
|
use crate::proxy::error::ProxyError;
|
||||||
use reqwest::RequestBuilder;
|
|
||||||
|
|
||||||
/// Gemini 适配器
|
/// Gemini 适配器
|
||||||
pub struct GeminiAdapter;
|
pub struct GeminiAdapter;
|
||||||
@@ -217,17 +216,26 @@ impl ProviderAdapter for GeminiAdapter {
|
|||||||
url
|
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 {
|
match auth.strategy {
|
||||||
// OAuth Bearer 认证
|
|
||||||
AuthStrategy::GoogleOAuth => {
|
AuthStrategy::GoogleOAuth => {
|
||||||
let token = auth.access_token.as_ref().unwrap_or(&auth.api_key);
|
let token = auth.access_token.as_ref().unwrap_or(&auth.api_key);
|
||||||
request
|
vec![
|
||||||
.header("Authorization", format!("Bearer {token}"))
|
(
|
||||||
.header("x-goog-api-client", "GeminiCLI/1.0")
|
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 认证
|
_ => vec![(
|
||||||
_ => request.header("x-goog-api-key", &auth.api_key),
|
HeaderName::from_static("x-goog-api-key"),
|
||||||
|
HeaderValue::from_str(&auth.api_key).unwrap(),
|
||||||
|
)],
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -88,8 +88,8 @@ struct ToolBlockState {
|
|||||||
}
|
}
|
||||||
|
|
||||||
/// 创建 Anthropic SSE 流
|
/// 创建 Anthropic SSE 流
|
||||||
pub fn create_anthropic_sse_stream(
|
pub fn create_anthropic_sse_stream<E: std::error::Error + Send + 'static>(
|
||||||
stream: impl Stream<Item = Result<Bytes, reqwest::Error>> + Send + 'static,
|
stream: impl Stream<Item = Result<Bytes, E>> + Send + 'static,
|
||||||
) -> impl Stream<Item = Result<Bytes, std::io::Error>> + Send {
|
) -> impl Stream<Item = Result<Bytes, std::io::Error>> + Send {
|
||||||
async_stream::stream! {
|
async_stream::stream! {
|
||||||
let mut buffer = String::new();
|
let mut buffer = String::new();
|
||||||
@@ -598,7 +598,9 @@ mod tests {
|
|||||||
"data: [DONE]\n\n"
|
"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 converted = create_anthropic_sse_stream(upstream);
|
||||||
let chunks: Vec<_> = converted.collect().await;
|
let chunks: Vec<_> = converted.collect().await;
|
||||||
|
|
||||||
@@ -686,7 +688,9 @@ mod tests {
|
|||||||
"data: [DONE]\n\n"
|
"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 converted = create_anthropic_sse_stream(upstream);
|
||||||
let chunks: Vec<_> = converted.collect().await;
|
let chunks: Vec<_> = converted.collect().await;
|
||||||
let merged = chunks
|
let merged = chunks
|
||||||
|
|||||||
@@ -96,8 +96,8 @@ fn resolve_content_index(
|
|||||||
///
|
///
|
||||||
/// 状态机跟踪: message_id, current_model, has_sent_message_start, item/content index map
|
/// 状态机跟踪: message_id, current_model, has_sent_message_start, item/content index map
|
||||||
/// SSE 解析支持 named events (event: + data: 行)
|
/// SSE 解析支持 named events (event: + data: 行)
|
||||||
pub fn create_anthropic_sse_stream_from_responses(
|
pub fn create_anthropic_sse_stream_from_responses<E: std::error::Error + Send + 'static>(
|
||||||
stream: impl Stream<Item = Result<Bytes, reqwest::Error>> + Send + 'static,
|
stream: impl Stream<Item = Result<Bytes, E>> + Send + 'static,
|
||||||
) -> impl Stream<Item = Result<Bytes, std::io::Error>> + Send {
|
) -> impl Stream<Item = Result<Bytes, std::io::Error>> + Send {
|
||||||
async_stream::stream! {
|
async_stream::stream! {
|
||||||
let mut buffer = String::new();
|
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"
|
"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 converted = create_anthropic_sse_stream_from_responses(upstream);
|
||||||
let chunks: Vec<_> = converted.collect().await;
|
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"
|
"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 converted = create_anthropic_sse_stream_from_responses(upstream);
|
||||||
let chunks: Vec<_> = converted.collect().await;
|
let chunks: Vec<_> = converted.collect().await;
|
||||||
let merged = chunks
|
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"
|
"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 converted = create_anthropic_sse_stream_from_responses(upstream);
|
||||||
let chunks: Vec<_> = converted.collect().await;
|
let chunks: Vec<_> = converted.collect().await;
|
||||||
let merged = chunks
|
let merged = chunks
|
||||||
@@ -991,12 +997,17 @@ mod tests {
|
|||||||
.iter()
|
.iter()
|
||||||
.filter(|event| {
|
.filter(|event| {
|
||||||
event.get("type").and_then(|v| v.as_str()) == Some("content_block_start")
|
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();
|
.count();
|
||||||
let text_stops = events
|
let text_stops = events
|
||||||
.iter()
|
.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();
|
.count();
|
||||||
let text_deltas: Vec<String> = events
|
let text_deltas: Vec<String> = events
|
||||||
.iter()
|
.iter()
|
||||||
@@ -1005,7 +1016,8 @@ mod tests {
|
|||||||
&& event.pointer("/delta/type").and_then(|v| v.as_str()) == Some("text_delta")
|
&& event.pointer("/delta/type").and_then(|v| v.as_str()) == Some("text_delta")
|
||||||
})
|
})
|
||||||
.filter_map(|event| {
|
.filter_map(|event| {
|
||||||
event.pointer("/delta/text")
|
event
|
||||||
|
.pointer("/delta/text")
|
||||||
.and_then(|v| v.as_str())
|
.and_then(|v| v.as_str())
|
||||||
.map(ToString::to_string)
|
.map(ToString::to_string)
|
||||||
})
|
})
|
||||||
|
|||||||
@@ -5,17 +5,19 @@
|
|||||||
use super::{
|
use super::{
|
||||||
handler_config::UsageParserConfig,
|
handler_config::UsageParserConfig,
|
||||||
handler_context::{RequestContext, StreamingTimeoutConfig},
|
handler_context::{RequestContext, StreamingTimeoutConfig},
|
||||||
|
hyper_client::ProxyResponse,
|
||||||
server::ProxyState,
|
server::ProxyState,
|
||||||
sse::strip_sse_field,
|
sse::strip_sse_field,
|
||||||
usage::parser::TokenUsage,
|
usage::parser::TokenUsage,
|
||||||
ProxyError,
|
ProxyError,
|
||||||
};
|
};
|
||||||
|
use axum::http::header::HeaderMap;
|
||||||
use axum::response::{IntoResponse, Response};
|
use axum::response::{IntoResponse, Response};
|
||||||
use bytes::Bytes;
|
use bytes::Bytes;
|
||||||
use futures::stream::{Stream, StreamExt};
|
use futures::stream::{Stream, StreamExt};
|
||||||
use reqwest::header::HeaderMap;
|
|
||||||
use serde_json::Value;
|
use serde_json::Value;
|
||||||
use std::{
|
use std::{
|
||||||
|
io::Read,
|
||||||
sync::{
|
sync::{
|
||||||
atomic::{AtomicBool, Ordering},
|
atomic::{AtomicBool, Ordering},
|
||||||
Arc,
|
Arc,
|
||||||
@@ -24,24 +26,123 @@ use std::{
|
|||||||
};
|
};
|
||||||
use tokio::sync::Mutex;
|
use tokio::sync::Mutex;
|
||||||
|
|
||||||
|
// ============================================================================
|
||||||
|
// 响应解压
|
||||||
|
// ============================================================================
|
||||||
|
|
||||||
|
/// 根据 content-encoding 解压响应体字节
|
||||||
|
///
|
||||||
|
/// reqwest 自动解压已禁用(为了透传 accept-encoding),需要手动解压。
|
||||||
|
fn decompress_body(content_encoding: &str, body: &[u8]) -> Result<Vec<u8>, 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<String> {
|
||||||
|
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 流式响应
|
/// 检测响应是否为 SSE 流式响应
|
||||||
#[inline]
|
#[inline]
|
||||||
pub fn is_sse_response(response: &reqwest::Response) -> bool {
|
pub fn is_sse_response(response: &ProxyResponse) -> bool {
|
||||||
response
|
response.is_sse()
|
||||||
.headers()
|
|
||||||
.get("content-type")
|
|
||||||
.and_then(|v| v.to_str().ok())
|
|
||||||
.map(|ct| ct.contains("text/event-stream"))
|
|
||||||
.unwrap_or(false)
|
|
||||||
}
|
}
|
||||||
|
|
||||||
/// 处理流式响应
|
/// 处理流式响应
|
||||||
pub async fn handle_streaming(
|
pub async fn handle_streaming(
|
||||||
response: reqwest::Response,
|
response: ProxyResponse,
|
||||||
ctx: &RequestContext,
|
ctx: &RequestContext,
|
||||||
state: &ProxyState,
|
state: &ProxyState,
|
||||||
parser_config: &UsageParserConfig,
|
parser_config: &UsageParserConfig,
|
||||||
@@ -53,6 +154,15 @@ pub async fn handle_streaming(
|
|||||||
status.as_u16(),
|
status.as_u16(),
|
||||||
format_headers(response.headers())
|
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);
|
let mut builder = axum::response::Response::builder().status(status);
|
||||||
|
|
||||||
// 复制响应头
|
// 复制响应头
|
||||||
@@ -61,9 +171,7 @@ pub async fn handle_streaming(
|
|||||||
}
|
}
|
||||||
|
|
||||||
// 创建字节流
|
// 创建字节流
|
||||||
let stream = response
|
let stream = response.bytes_stream();
|
||||||
.bytes_stream()
|
|
||||||
.map(|chunk| chunk.map_err(|e| std::io::Error::other(e.to_string())));
|
|
||||||
|
|
||||||
// 创建使用量收集器
|
// 创建使用量收集器
|
||||||
let usage_collector = create_usage_collector(ctx, state, status.as_u16(), parser_config);
|
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(
|
pub async fn handle_non_streaming(
|
||||||
response: reqwest::Response,
|
response: ProxyResponse,
|
||||||
ctx: &RequestContext,
|
ctx: &RequestContext,
|
||||||
state: &ProxyState,
|
state: &ProxyState,
|
||||||
parser_config: &UsageParserConfig,
|
parser_config: &UsageParserConfig,
|
||||||
) -> Result<Response, ProxyError> {
|
) -> Result<Response, ProxyError> {
|
||||||
let response_headers = response.headers().clone();
|
// 整包超时:仅在故障转移开启且配置值非零时生效
|
||||||
let status = response.status();
|
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)
|
||||||
let body_bytes = response.bytes().await.map_err(|e| {
|
} else {
|
||||||
log::error!("[{}] 读取响应失败: {e}", ctx.tag);
|
Duration::ZERO
|
||||||
ProxyError::ForwardFailed(format!("Failed to read response body: {e}"))
|
};
|
||||||
})?;
|
let (response_headers, status, body_bytes) =
|
||||||
log::debug!(
|
read_decoded_body(response, ctx.tag, body_timeout).await?;
|
||||||
"[{}] 已接收上游响应体: status={}, bytes={}, headers={}",
|
|
||||||
ctx.tag,
|
|
||||||
status.as_u16(),
|
|
||||||
body_bytes.len(),
|
|
||||||
format_headers(&response_headers)
|
|
||||||
);
|
|
||||||
|
|
||||||
log::debug!(
|
log::debug!(
|
||||||
"[{}] 上游响应体内容: {}",
|
"[{}] 上游响应体内容: {}",
|
||||||
@@ -190,7 +292,7 @@ pub async fn handle_non_streaming(
|
|||||||
///
|
///
|
||||||
/// 根据响应类型自动选择流式或非流式处理
|
/// 根据响应类型自动选择流式或非流式处理
|
||||||
pub async fn process_response(
|
pub async fn process_response(
|
||||||
response: reqwest::Response,
|
response: ProxyResponse,
|
||||||
ctx: &RequestContext,
|
ctx: &RequestContext,
|
||||||
state: &ProxyState,
|
state: &ProxyState,
|
||||||
parser_config: &UsageParserConfig,
|
parser_config: &UsageParserConfig,
|
||||||
|
|||||||
@@ -1,6 +1,12 @@
|
|||||||
//! HTTP代理服务器
|
//! HTTP代理服务器
|
||||||
//!
|
//!
|
||||||
//! 基于Axum的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::{
|
use super::{
|
||||||
failover_switch::FailoverSwitchManager, handlers, log_codes::srv as log_srv,
|
failover_switch::FailoverSwitchManager, handlers, log_codes::srv as log_srv,
|
||||||
@@ -12,6 +18,7 @@ use axum::{
|
|||||||
routing::{get, post},
|
routing::{get, post},
|
||||||
Router,
|
Router,
|
||||||
};
|
};
|
||||||
|
use hyper_util::rt::TokioIo;
|
||||||
use std::net::SocketAddr;
|
use std::net::SocketAddr;
|
||||||
use std::sync::Arc;
|
use std::sync::Arc;
|
||||||
use tokio::sync::{oneshot, RwLock};
|
use tokio::sync::{oneshot, RwLock};
|
||||||
@@ -114,15 +121,77 @@ impl ProxyServer {
|
|||||||
// 记录启动时间
|
// 记录启动时间
|
||||||
*self.state.start_time.write().await = Some(std::time::Instant::now());
|
*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 state = self.state.clone();
|
||||||
let handle = tokio::spawn(async move {
|
let handle = tokio::spawn(async move {
|
||||||
axum::serve(listener, app)
|
let mut shutdown_rx = shutdown_rx;
|
||||||
.with_graceful_shutdown(async {
|
loop {
|
||||||
shutdown_rx.await.ok();
|
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<hyper::body::Incoming>| {
|
||||||
|
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);
|
||||||
|
<Router as tower::Service<http::Request<axum::body::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
|
.await
|
||||||
.ok();
|
{
|
||||||
|
// 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;
|
state.status.write().await.running = false;
|
||||||
|
|||||||
@@ -25,7 +25,7 @@ pub fn optimize(body: &mut Value, config: &OptimizerConfig) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
if model.contains("opus-4-6") || model.contains("sonnet-4-6") {
|
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["thinking"] = json!({"type": "adaptive"});
|
||||||
body["output_config"] = json!({"effort": "max"});
|
body["output_config"] = json!({"effort": "max"});
|
||||||
append_beta(body, "context-1m-2025-08-07");
|
append_beta(body, "context-1m-2025-08-07");
|
||||||
@@ -33,7 +33,7 @@ pub fn optimize(body: &mut Value, config: &OptimizerConfig) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
// legacy path
|
// legacy path
|
||||||
log::info!("[OPT] thinking: legacy({})", model);
|
log::info!("[OPT] thinking: legacy({model})");
|
||||||
|
|
||||||
let max_tokens = body
|
let max_tokens = body
|
||||||
.get("max_tokens")
|
.get("max_tokens")
|
||||||
|
|||||||
@@ -957,7 +957,7 @@ impl SkillService {
|
|||||||
if source_path.is_none() {
|
if source_path.is_none() {
|
||||||
source_path = Some(skill_path);
|
source_path = Some(skill_path);
|
||||||
}
|
}
|
||||||
log::debug!("Skill '{}' found in source '{}'", dir_name, label);
|
log::debug!("Skill '{dir_name}' found in source '{label}'");
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
@@ -12,8 +12,8 @@ use std::time::Instant;
|
|||||||
use crate::app_config::AppType;
|
use crate::app_config::AppType;
|
||||||
use crate::error::AppError;
|
use crate::error::AppError;
|
||||||
use crate::provider::Provider;
|
use crate::provider::Provider;
|
||||||
use crate::proxy::providers::transform::anthropic_to_openai;
|
|
||||||
use crate::proxy::providers::copilot_auth;
|
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::transform_responses::anthropic_to_responses;
|
||||||
use crate::proxy::providers::{get_adapter, AuthInfo, AuthStrategy};
|
use crate::proxy::providers::{get_adapter, AuthInfo, AuthStrategy};
|
||||||
|
|
||||||
@@ -95,8 +95,7 @@ impl StreamCheckService {
|
|||||||
let mut last_result = None;
|
let mut last_result = None;
|
||||||
|
|
||||||
for attempt in 0..=effective_config.max_retries {
|
for attempt in 0..=effective_config.max_retries {
|
||||||
let result =
|
let result = Self::check_once(
|
||||||
Self::check_once(
|
|
||||||
app_type,
|
app_type,
|
||||||
provider,
|
provider,
|
||||||
&effective_config,
|
&effective_config,
|
||||||
@@ -304,6 +303,7 @@ impl StreamCheckService {
|
|||||||
/// 根据供应商的 api_format 选择请求格式:
|
/// 根据供应商的 api_format 选择请求格式:
|
||||||
/// - "anthropic" (默认): Anthropic Messages API (/v1/messages)
|
/// - "anthropic" (默认): Anthropic Messages API (/v1/messages)
|
||||||
/// - "openai_chat": OpenAI Chat Completions API (/v1/chat/completions)
|
/// - "openai_chat": OpenAI Chat Completions API (/v1/chat/completions)
|
||||||
|
#[allow(clippy::too_many_arguments)]
|
||||||
async fn check_claude_stream(
|
async fn check_claude_stream(
|
||||||
client: &Client,
|
client: &Client,
|
||||||
base_url: &str,
|
base_url: &str,
|
||||||
@@ -339,12 +339,8 @@ impl StreamCheckService {
|
|||||||
.unwrap_or(false);
|
.unwrap_or(false);
|
||||||
let is_openai_chat = effective_api_format == "openai_chat";
|
let is_openai_chat = effective_api_format == "openai_chat";
|
||||||
let is_openai_responses = effective_api_format == "openai_responses";
|
let is_openai_responses = effective_api_format == "openai_responses";
|
||||||
let url = Self::resolve_claude_stream_url(
|
let url =
|
||||||
base,
|
Self::resolve_claude_stream_url(base, auth.strategy, effective_api_format, is_full_url);
|
||||||
auth.strategy,
|
|
||||||
effective_api_format,
|
|
||||||
is_full_url,
|
|
||||||
);
|
|
||||||
|
|
||||||
let max_tokens = if is_openai_responses { 16 } else { 1 };
|
let max_tokens = if is_openai_responses { 16 } else { 1 };
|
||||||
|
|
||||||
@@ -375,8 +371,14 @@ impl StreamCheckService {
|
|||||||
.header("accept-encoding", "identity")
|
.header("accept-encoding", "identity")
|
||||||
.header("user-agent", copilot_auth::COPILOT_USER_AGENT)
|
.header("user-agent", copilot_auth::COPILOT_USER_AGENT)
|
||||||
.header("editor-version", copilot_auth::COPILOT_EDITOR_VERSION)
|
.header("editor-version", copilot_auth::COPILOT_EDITOR_VERSION)
|
||||||
.header("editor-plugin-version", copilot_auth::COPILOT_PLUGIN_VERSION)
|
.header(
|
||||||
.header("copilot-integration-id", copilot_auth::COPILOT_INTEGRATION_ID)
|
"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("x-github-api-version", copilot_auth::COPILOT_API_VERSION)
|
||||||
.header("openai-intent", "conversation-panel");
|
.header("openai-intent", "conversation-panel");
|
||||||
} else if is_openai_chat || is_openai_responses {
|
} else if is_openai_chat || is_openai_responses {
|
||||||
|
|||||||
@@ -463,12 +463,10 @@ fn validate_manifest_compat(manifest: &SyncManifest, layout: RemoteLayout) -> Re
|
|||||||
return Err(localized(
|
return Err(localized(
|
||||||
"webdav.sync.manifest_db_version_incompatible",
|
"webdav.sync.manifest_db_version_incompatible",
|
||||||
format!(
|
format!(
|
||||||
"远端数据库快照版本不兼容: db-v{} (本地 db-v{DB_COMPAT_VERSION})",
|
"远端数据库快照版本不兼容: db-v{db_compat_version} (本地 db-v{DB_COMPAT_VERSION})"
|
||||||
db_compat_version
|
|
||||||
),
|
),
|
||||||
format!(
|
format!(
|
||||||
"Remote database snapshot version is incompatible: db-v{} (local db-v{DB_COMPAT_VERSION})",
|
"Remote database snapshot version is incompatible: db-v{db_compat_version} (local db-v{DB_COMPAT_VERSION})"
|
||||||
db_compat_version
|
|
||||||
),
|
),
|
||||||
));
|
));
|
||||||
}
|
}
|
||||||
@@ -476,12 +474,10 @@ fn validate_manifest_compat(manifest: &SyncManifest, layout: RemoteLayout) -> Re
|
|||||||
return Err(localized(
|
return Err(localized(
|
||||||
"webdav.sync.manifest_db_version_incompatible",
|
"webdav.sync.manifest_db_version_incompatible",
|
||||||
format!(
|
format!(
|
||||||
"远端数据库快照版本不兼容: db-v{} (本地最高支持 db-v{DB_COMPAT_VERSION})",
|
"远端数据库快照版本不兼容: db-v{db_compat_version} (本地最高支持 db-v{DB_COMPAT_VERSION})"
|
||||||
db_compat_version
|
|
||||||
),
|
),
|
||||||
format!(
|
format!(
|
||||||
"Remote database snapshot version is incompatible: db-v{} (local supports up to db-v{DB_COMPAT_VERSION})",
|
"Remote database snapshot version is incompatible: db-v{db_compat_version} (local supports up to db-v{DB_COMPAT_VERSION})"
|
||||||
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;
|
let max_mb = MAX_SYNC_ARTIFACT_BYTES / 1024 / 1024;
|
||||||
return Err(localized(
|
return Err(localized(
|
||||||
"webdav.sync.artifact_too_large",
|
"webdav.sync.artifact_too_large",
|
||||||
format!("artifact {artifact_name} 超过下载上限({} MB)", max_mb),
|
format!("artifact {artifact_name} 超过下载上限({max_mb} MB)"),
|
||||||
format!(
|
format!("Artifact {artifact_name} exceeds download limit ({max_mb} MB)"),
|
||||||
"Artifact {artifact_name} exceeds download limit ({} MB)",
|
|
||||||
max_mb
|
|
||||||
),
|
|
||||||
));
|
));
|
||||||
}
|
}
|
||||||
Ok(())
|
Ok(())
|
||||||
|
|||||||
@@ -350,8 +350,8 @@ fn copy_entry_with_total_limit<R: Read, W: Write>(
|
|||||||
let max_mb = max_total_bytes / 1024 / 1024;
|
let max_mb = max_total_bytes / 1024 / 1024;
|
||||||
return Err(localized(
|
return Err(localized(
|
||||||
"webdav.sync.skills_zip_too_large",
|
"webdav.sync.skills_zip_too_large",
|
||||||
format!("skills.zip 解压后体积超过上限({} MB)", max_mb),
|
format!("skills.zip 解压后体积超过上限({max_mb} MB)"),
|
||||||
format!("skills.zip extracted size exceeds limit ({} MB)", max_mb),
|
format!("skills.zip extracted size exceeds limit ({max_mb} MB)"),
|
||||||
));
|
));
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
@@ -148,7 +148,7 @@ fn scan_sessions_sqlite() -> Vec<SessionMeta> {
|
|||||||
},
|
},
|
||||||
created_at: Some(created),
|
created_at: Some(created),
|
||||||
last_active_at: Some(updated),
|
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}")),
|
resume_command: Some(format!("opencode session resume {session_id}")),
|
||||||
});
|
});
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -42,9 +42,9 @@ export function AddProviderDialog({
|
|||||||
const { t } = useTranslation();
|
const { t } = useTranslation();
|
||||||
// OpenCode and OpenClaw don't support universal providers
|
// OpenCode and OpenClaw don't support universal providers
|
||||||
const showUniversalTab = appId !== "opencode" && appId !== "openclaw";
|
const showUniversalTab = appId !== "opencode" && appId !== "openclaw";
|
||||||
const [activeTab, setActiveTab] = useState<
|
const [activeTab, setActiveTab] = useState<"app-specific" | "universal">(
|
||||||
"app-specific" | "universal"
|
"app-specific",
|
||||||
>("app-specific");
|
);
|
||||||
const [universalFormOpen, setUniversalFormOpen] = useState(false);
|
const [universalFormOpen, setUniversalFormOpen] = useState(false);
|
||||||
const [selectedUniversalPreset, setSelectedUniversalPreset] =
|
const [selectedUniversalPreset, setSelectedUniversalPreset] =
|
||||||
useState<UniversalProviderPreset | null>(null);
|
useState<UniversalProviderPreset | null>(null);
|
||||||
@@ -284,9 +284,7 @@ export function AddProviderDialog({
|
|||||||
{showUniversalTab ? (
|
{showUniversalTab ? (
|
||||||
<Tabs
|
<Tabs
|
||||||
value={activeTab}
|
value={activeTab}
|
||||||
onValueChange={(v) =>
|
onValueChange={(v) => setActiveTab(v as "app-specific" | "universal")}
|
||||||
setActiveTab(v as "app-specific" | "universal")
|
|
||||||
}
|
|
||||||
>
|
>
|
||||||
<TabsList className="grid w-full grid-cols-2 mb-6">
|
<TabsList className="grid w-full grid-cols-2 mb-6">
|
||||||
<TabsTrigger value="app-specific">
|
<TabsTrigger value="app-specific">
|
||||||
|
|||||||
@@ -166,8 +166,7 @@ export function ClaudeFormFields({
|
|||||||
apiFormat !== "anthropic" ||
|
apiFormat !== "anthropic" ||
|
||||||
apiKeyField !== "ANTHROPIC_AUTH_TOKEN"
|
apiKeyField !== "ANTHROPIC_AUTH_TOKEN"
|
||||||
);
|
);
|
||||||
const [advancedExpanded, setAdvancedExpanded] =
|
const [advancedExpanded, setAdvancedExpanded] = useState(hasAnyAdvancedValue);
|
||||||
useState(hasAnyAdvancedValue);
|
|
||||||
|
|
||||||
// 预设填充高级值后自动展开(仅从折叠→展开,不会自动折叠)
|
// 预设填充高级值后自动展开(仅从折叠→展开,不会自动折叠)
|
||||||
useEffect(() => {
|
useEffect(() => {
|
||||||
@@ -409,10 +408,7 @@ export function ClaudeFormFields({
|
|||||||
|
|
||||||
{/* 高级选项(API 格式 + 认证字段 + 模型映射) */}
|
{/* 高级选项(API 格式 + 认证字段 + 模型映射) */}
|
||||||
{shouldShowModelSelector && (
|
{shouldShowModelSelector && (
|
||||||
<Collapsible
|
<Collapsible open={advancedExpanded} onOpenChange={setAdvancedExpanded}>
|
||||||
open={advancedExpanded}
|
|
||||||
onOpenChange={setAdvancedExpanded}
|
|
||||||
>
|
|
||||||
<CollapsibleTrigger asChild>
|
<CollapsibleTrigger asChild>
|
||||||
<Button
|
<Button
|
||||||
type="button"
|
type="button"
|
||||||
|
|||||||
@@ -1,4 +1,10 @@
|
|||||||
import React, { useCallback, useEffect, useMemo, useRef, useState } from "react";
|
import React, {
|
||||||
|
useCallback,
|
||||||
|
useEffect,
|
||||||
|
useMemo,
|
||||||
|
useRef,
|
||||||
|
useState,
|
||||||
|
} from "react";
|
||||||
import { useTranslation } from "react-i18next";
|
import { useTranslation } from "react-i18next";
|
||||||
import JsonEditor from "@/components/JsonEditor";
|
import JsonEditor from "@/components/JsonEditor";
|
||||||
import {
|
import {
|
||||||
@@ -175,10 +181,7 @@ export const CodexConfigSection: React.FC<CodexConfigSectionProps> = ({
|
|||||||
}
|
}
|
||||||
} else {
|
} else {
|
||||||
toml = removeCodexTopLevelField(toml, "model_context_window");
|
toml = removeCodexTopLevelField(toml, "model_context_window");
|
||||||
toml = removeCodexTopLevelField(
|
toml = removeCodexTopLevelField(toml, "model_auto_compact_token_limit");
|
||||||
toml,
|
|
||||||
"model_auto_compact_token_limit",
|
|
||||||
);
|
|
||||||
}
|
}
|
||||||
handleLocalChange(toml);
|
handleLocalChange(toml);
|
||||||
},
|
},
|
||||||
|
|||||||
@@ -11,5 +11,4 @@ export const TEMPLATE_TYPES = {
|
|||||||
GITHUB_COPILOT: "github_copilot",
|
GITHUB_COPILOT: "github_copilot",
|
||||||
} as const;
|
} as const;
|
||||||
|
|
||||||
export type TemplateType =
|
export type TemplateType = (typeof TEMPLATE_TYPES)[keyof typeof TEMPLATE_TYPES];
|
||||||
(typeof TEMPLATE_TYPES)[keyof typeof TEMPLATE_TYPES];
|
|
||||||
|
|||||||
@@ -142,20 +142,44 @@ export function useProviderActions(activeApp: AppId, isProxyRunning?: boolean) {
|
|||||||
const isCopilotProvider =
|
const isCopilotProvider =
|
||||||
activeApp === "claude" &&
|
activeApp === "claude" &&
|
||||||
provider.meta?.providerType === "github_copilot";
|
provider.meta?.providerType === "github_copilot";
|
||||||
const requiresProxyForSwitch =
|
|
||||||
!isProxyRunning &&
|
|
||||||
provider.category !== "official" &&
|
|
||||||
((activeApp === "claude" &&
|
|
||||||
(isCopilotProvider ||
|
|
||||||
provider.meta?.isFullUrl ||
|
|
||||||
provider.meta?.apiFormat === "openai_chat" ||
|
|
||||||
provider.meta?.apiFormat === "openai_responses")) ||
|
|
||||||
(activeApp === "codex" && provider.meta?.isFullUrl));
|
|
||||||
|
|
||||||
if (requiresProxyForSwitch) {
|
// Determine why this provider requires the proxy
|
||||||
|
let proxyRequiredReason: string | null = null;
|
||||||
|
if (!isProxyRunning && provider.category !== "official") {
|
||||||
|
if (isCopilotProvider) {
|
||||||
|
proxyRequiredReason = t("notifications.proxyReasonCopilot", {
|
||||||
|
defaultValue: "使用 GitHub Copilot 作为 Claude 供应商",
|
||||||
|
});
|
||||||
|
} else if (
|
||||||
|
provider.meta?.apiFormat === "openai_chat" &&
|
||||||
|
activeApp === "claude"
|
||||||
|
) {
|
||||||
|
proxyRequiredReason = t("notifications.proxyReasonOpenAIChat", {
|
||||||
|
defaultValue: "使用 OpenAI Chat 接口格式",
|
||||||
|
});
|
||||||
|
} else if (
|
||||||
|
provider.meta?.apiFormat === "openai_responses" &&
|
||||||
|
activeApp === "claude"
|
||||||
|
) {
|
||||||
|
proxyRequiredReason = t("notifications.proxyReasonOpenAIResponses", {
|
||||||
|
defaultValue: "使用 OpenAI Responses 接口格式",
|
||||||
|
});
|
||||||
|
} else if (
|
||||||
|
provider.meta?.isFullUrl &&
|
||||||
|
(activeApp === "claude" || activeApp === "codex")
|
||||||
|
) {
|
||||||
|
proxyRequiredReason = t("notifications.proxyReasonFullUrl", {
|
||||||
|
defaultValue: "开启了完整 URL 连接模式",
|
||||||
|
});
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
if (proxyRequiredReason) {
|
||||||
toast.warning(
|
toast.warning(
|
||||||
t("notifications.proxyRequiredForSwitch", {
|
t("notifications.proxyRequiredForSwitch", {
|
||||||
defaultValue: "此供应商需要代理服务,请先启动代理",
|
reason: proxyRequiredReason,
|
||||||
|
defaultValue:
|
||||||
|
"此供应商{{reason}},需要代理服务才能正常使用,请先启动代理",
|
||||||
}),
|
}),
|
||||||
);
|
);
|
||||||
return;
|
return;
|
||||||
|
|||||||
+10
-8
@@ -1,4 +1,9 @@
|
|||||||
import { useMutation, useQuery, useQueryClient, keepPreviousData } from "@tanstack/react-query";
|
import {
|
||||||
|
useMutation,
|
||||||
|
useQuery,
|
||||||
|
useQueryClient,
|
||||||
|
keepPreviousData,
|
||||||
|
} from "@tanstack/react-query";
|
||||||
import {
|
import {
|
||||||
skillsApi,
|
skillsApi,
|
||||||
type SkillBackupEntry,
|
type SkillBackupEntry,
|
||||||
@@ -108,13 +113,10 @@ export function useInstallSkill() {
|
|||||||
export function useUninstallSkill() {
|
export function useUninstallSkill() {
|
||||||
const queryClient = useQueryClient();
|
const queryClient = useQueryClient();
|
||||||
return useMutation({
|
return useMutation({
|
||||||
mutationFn: ({
|
mutationFn: ({ id, skillKey }: { id: string; skillKey: string }) =>
|
||||||
id,
|
skillsApi
|
||||||
skillKey,
|
.uninstallUnified(id)
|
||||||
}: {
|
.then((result) => ({ ...result, skillKey })),
|
||||||
id: string;
|
|
||||||
skillKey: string;
|
|
||||||
}) => skillsApi.uninstallUnified(id).then((result) => ({ ...result, skillKey })),
|
|
||||||
onSuccess: ({ skillKey }, _vars) => {
|
onSuccess: ({ skillKey }, _vars) => {
|
||||||
// 直接更新 installed 缓存,移除该 skill
|
// 直接更新 installed 缓存,移除该 skill
|
||||||
queryClient.setQueryData<InstalledSkill[]>(
|
queryClient.setQueryData<InstalledSkill[]>(
|
||||||
|
|||||||
@@ -174,10 +174,13 @@
|
|||||||
"deleteFailed": "Failed to delete provider: {{error}}",
|
"deleteFailed": "Failed to delete provider: {{error}}",
|
||||||
"settingsSaved": "Settings saved",
|
"settingsSaved": "Settings saved",
|
||||||
"settingsSaveFailed": "Failed to save settings: {{error}}",
|
"settingsSaveFailed": "Failed to save settings: {{error}}",
|
||||||
"openAIChatFormatHint": "This provider uses OpenAI Chat format and requires the proxy service to be enabled",
|
"proxyRequiredForSwitch": "This provider {{reason}}, requires the proxy service to work properly. Start the proxy first.",
|
||||||
|
"proxyReasonCopilot": "uses GitHub Copilot as a Claude provider",
|
||||||
|
"proxyReasonOpenAIChat": "uses OpenAI Chat API format",
|
||||||
|
"proxyReasonOpenAIResponses": "uses OpenAI Responses API format",
|
||||||
|
"proxyReasonFullUrl": "has full URL connection mode enabled",
|
||||||
"openAIFormatHint": "This provider uses OpenAI-compatible format and requires the proxy service to be enabled",
|
"openAIFormatHint": "This provider uses OpenAI-compatible format and requires the proxy service to be enabled",
|
||||||
"copilotProxyHint": "GitHub Copilot as a Claude provider always requires the local proxy; the proxy automatically selects Chat Completions or Responses based on the current model.",
|
"copilotProxyHint": "GitHub Copilot as a Claude provider always requires the local proxy; the proxy automatically selects Chat Completions or Responses based on the current model.",
|
||||||
"proxyRequiredForSwitch": "This provider requires the proxy service. Start the proxy first.",
|
|
||||||
"openLinkFailed": "Failed to open link",
|
"openLinkFailed": "Failed to open link",
|
||||||
"openclawModelsRegistered": "Models have been registered to /model list",
|
"openclawModelsRegistered": "Models have been registered to /model list",
|
||||||
"openclawDefaultModelSet": "Set as default model",
|
"openclawDefaultModelSet": "Set as default model",
|
||||||
|
|||||||
@@ -174,10 +174,13 @@
|
|||||||
"deleteFailed": "プロバイダーの削除に失敗しました: {{error}}",
|
"deleteFailed": "プロバイダーの削除に失敗しました: {{error}}",
|
||||||
"settingsSaved": "設定を保存しました",
|
"settingsSaved": "設定を保存しました",
|
||||||
"settingsSaveFailed": "設定の保存に失敗しました: {{error}}",
|
"settingsSaveFailed": "設定の保存に失敗しました: {{error}}",
|
||||||
"openAIChatFormatHint": "このプロバイダーは OpenAI Chat フォーマットを使用しており、プロキシサービスの有効化が必要です",
|
"proxyRequiredForSwitch": "このプロバイダーは{{reason}}、プロキシサービスが必要です。先にプロキシを起動してください",
|
||||||
|
"proxyReasonCopilot": "GitHub Copilot を Claude プロバイダーとして使用しており",
|
||||||
|
"proxyReasonOpenAIChat": "OpenAI Chat API フォーマットを使用しており",
|
||||||
|
"proxyReasonOpenAIResponses": "OpenAI Responses API フォーマットを使用しており",
|
||||||
|
"proxyReasonFullUrl": "完全 URL 接続モードが有効になっており",
|
||||||
"openAIFormatHint": "このプロバイダーは OpenAI 互換フォーマットを使用しており、プロキシサービスの有効化が必要です",
|
"openAIFormatHint": "このプロバイダーは OpenAI 互換フォーマットを使用しており、プロキシサービスの有効化が必要です",
|
||||||
"copilotProxyHint": "GitHub Copilot を Claude プロバイダーとして使用する場合、ローカルプロキシが常に必要です。プロキシは現在のモデルに応じて Chat Completions または Responses を自動的に選択します。",
|
"copilotProxyHint": "GitHub Copilot を Claude プロバイダーとして使用する場合、ローカルプロキシが常に必要です。プロキシは現在のモデルに応じて Chat Completions または Responses を自動的に選択します。",
|
||||||
"proxyRequiredForSwitch": "このプロバイダーにはプロキシサービスが必要です。先にプロキシを起動してください",
|
|
||||||
"openLinkFailed": "リンクを開けませんでした",
|
"openLinkFailed": "リンクを開けませんでした",
|
||||||
"openclawModelsRegistered": "モデルが /model リストに登録されました",
|
"openclawModelsRegistered": "モデルが /model リストに登録されました",
|
||||||
"openclawDefaultModelSet": "デフォルトモデルに設定しました",
|
"openclawDefaultModelSet": "デフォルトモデルに設定しました",
|
||||||
|
|||||||
@@ -174,10 +174,13 @@
|
|||||||
"deleteFailed": "删除供应商失败:{{error}}",
|
"deleteFailed": "删除供应商失败:{{error}}",
|
||||||
"settingsSaved": "设置已保存",
|
"settingsSaved": "设置已保存",
|
||||||
"settingsSaveFailed": "保存设置失败:{{error}}",
|
"settingsSaveFailed": "保存设置失败:{{error}}",
|
||||||
"openAIChatFormatHint": "此供应商使用 OpenAI Chat 格式,需要开启代理服务才能正常使用",
|
"proxyRequiredForSwitch": "此供应商{{reason}},需要代理服务才能正常使用,请先启动代理",
|
||||||
|
"proxyReasonCopilot": "使用 GitHub Copilot 作为 Claude 供应商",
|
||||||
|
"proxyReasonOpenAIChat": "使用 OpenAI Chat 接口格式",
|
||||||
|
"proxyReasonOpenAIResponses": "使用 OpenAI Responses 接口格式",
|
||||||
|
"proxyReasonFullUrl": "开启了完整 URL 连接模式",
|
||||||
"openAIFormatHint": "此供应商使用 OpenAI 兼容格式,需要开启代理服务才能正常使用",
|
"openAIFormatHint": "此供应商使用 OpenAI 兼容格式,需要开启代理服务才能正常使用",
|
||||||
"copilotProxyHint": "GitHub Copilot 作为 Claude 供应商时始终需要本地代理;代理会根据当前模型自动选择 Chat Completions 或 Responses。",
|
"copilotProxyHint": "GitHub Copilot 作为 Claude 供应商时始终需要本地代理;代理会根据当前模型自动选择 Chat Completions 或 Responses。",
|
||||||
"proxyRequiredForSwitch": "此供应商需要代理服务,请先启动代理",
|
|
||||||
"openLinkFailed": "链接打开失败",
|
"openLinkFailed": "链接打开失败",
|
||||||
"openclawModelsRegistered": "模型已注册到 /model 列表",
|
"openclawModelsRegistered": "模型已注册到 /model 列表",
|
||||||
"openclawDefaultModelSet": "已设为默认模型",
|
"openclawDefaultModelSet": "已设为默认模型",
|
||||||
|
|||||||
Reference in New Issue
Block a user