From 146db0d11014a4a455501e0c1f4b26d54caa89e9 Mon Sep 17 00:00:00 2001 From: coso Date: Sun, 8 Feb 2026 22:34:56 +0800 Subject: [PATCH] =?UTF-8?q?refactor:=20=E8=BF=81=E7=A7=BB=20server=20?= =?UTF-8?q?=E6=A8=A1=E5=9D=97=E5=88=B0=20proxycast-server=20crate?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit - 创建 proxycast-server crate,包含完整的 HTTP 服务器逻辑 - 迁移 server/handlers/ 下所有 8 个处理器模块 - 迁移 client_detector 模块 - 主 crate server/mod.rs 替换为 re-export 层 - 升级 proxycast-core logger:完整的日志轮转、压缩归档、正则脱敏 - 提取 network 模块到 proxycast-core(从 network_cmd 中分离纯逻辑) - 修复 bootstrap.rs/state.rs 中 LogStore::with_config 调用 - 修复 provider_calls.rs doctest 代码块标记 --- src-tauri/Cargo.lock | 87 + src-tauri/Cargo.toml | 8 + src-tauri/crates/core/Cargo.toml | 6 + src-tauri/crates/core/src/lib.rs | 7 + src-tauri/crates/core/src/logger.rs | 172 +- src-tauri/crates/core/src/network.rs | 158 ++ src-tauri/crates/server/Cargo.toml | 40 + .../server/src}/client_detector.rs | 2 +- .../server/src}/handlers/api.rs | 96 +- .../server/src}/handlers/credentials_api.rs | 16 +- .../server/src}/handlers/image_handler.rs | 12 +- .../server/src}/handlers/kiro_credential.rs | 8 +- .../server/src}/handlers/management.rs | 8 +- .../server/src}/handlers/mod.rs | 0 .../server/src}/handlers/provider_calls.rs | 61 +- .../server/src}/handlers/websocket.rs | 24 +- src-tauri/crates/server/src/lib.rs | 1830 ++++++++++++++++ src-tauri/src/app/bootstrap.rs | 4 +- src-tauri/src/app/state.rs | 4 +- src-tauri/src/commands/network_cmd.rs | 208 +- src-tauri/src/logger.rs | 377 +--- src-tauri/src/server/mod.rs | 1839 +---------------- 22 files changed, 2417 insertions(+), 2550 deletions(-) create mode 100644 src-tauri/crates/core/src/network.rs create mode 100644 src-tauri/crates/server/Cargo.toml rename src-tauri/{src/server => crates/server/src}/client_detector.rs (97%) rename src-tauri/{src/server => crates/server/src}/handlers/api.rs (95%) rename src-tauri/{src/server => crates/server/src}/handlers/credentials_api.rs (97%) rename src-tauri/{src/server => crates/server/src}/handlers/image_handler.rs (97%) rename src-tauri/{src/server => crates/server/src}/handlers/kiro_credential.rs (99%) rename src-tauri/{src/server => crates/server/src}/handlers/management.rs (98%) rename src-tauri/{src/server => crates/server/src}/handlers/mod.rs (100%) rename src-tauri/{src/server => crates/server/src}/handlers/provider_calls.rs (98%) rename src-tauri/{src/server => crates/server/src}/handlers/websocket.rs (97%) create mode 100644 src-tauri/crates/server/src/lib.rs diff --git a/src-tauri/Cargo.lock b/src-tauri/Cargo.lock index 0f05f8845..f1b21f172 100644 --- a/src-tauri/Cargo.lock +++ b/src-tauri/Cargo.lock @@ -6669,14 +6669,18 @@ dependencies = [ "parking_lot", "portable-pty", "proptest", + "proxycast-agent", "proxycast-config", "proxycast-core", "proxycast-credential", "proxycast-infra", + "proxycast-mcp", "proxycast-processor", "proxycast-providers", + "proxycast-server", "proxycast-server-utils", "proxycast-services", + "proxycast-skills", "proxycast-terminal", "proxycast-websocket", "rand 0.8.5", @@ -6724,6 +6728,23 @@ dependencies = [ "zip", ] +[[package]] +name = "proxycast-agent" +version = "0.60.0" +dependencies = [ + "aster", + "async-trait", + "chrono", + "proxycast-core", + "proxycast-mcp", + "rmcp 0.6.4", + "serde", + "serde_json", + "tokio", + "tokio-util", + "tracing", +] + [[package]] name = "proxycast-config" version = "0.60.0" @@ -6752,11 +6773,13 @@ dependencies = [ "dirs 5.0.1", "flate2", "futures", + "if-addrs", "indexmap 2.13.0", "notify 6.1.1", "parking_lot", "proptest", "rand 0.8.5", + "regex", "reqwest 0.12.28", "rusqlite", "serde", @@ -6815,6 +6838,21 @@ dependencies = [ "uuid", ] +[[package]] +name = "proxycast-mcp" +version = "0.60.0" +dependencies = [ + "async-trait", + "glob", + "proxycast-core", + "rmcp 0.6.4", + "serde", + "serde_json", + "thiserror 1.0.69", + "tokio", + "tracing", +] + [[package]] name = "proxycast-processor" version = "0.60.0" @@ -6868,6 +6906,43 @@ dependencies = [ "uuid", ] +[[package]] +name = "proxycast-server" +version = "0.60.0" +dependencies = [ + "async-stream", + "axum 0.7.9", + "base64 0.22.1", + "bytes", + "chrono", + "dirs 5.0.1", + "futures", + "parking_lot", + "proptest", + "proxycast-config", + "proxycast-core", + "proxycast-credential", + "proxycast-infra", + "proxycast-processor", + "proxycast-providers", + "proxycast-server-utils", + "proxycast-services", + "proxycast-websocket", + "regex", + "reqwest 0.12.28", + "rusqlite", + "serde", + "serde_json", + "subtle", + "tokio", + "tokio-util", + "tower 0.5.3", + "tower-http", + "tracing", + "urlencoding", + "uuid", +] + [[package]] name = "proxycast-server-utils" version = "0.60.0" @@ -6921,6 +6996,18 @@ dependencies = [ "zip", ] +[[package]] +name = "proxycast-skills" +version = "0.60.0" +dependencies = [ + "async-trait", + "dirs 5.0.1", + "regex", + "serde", + "serde_json", + "tracing", +] + [[package]] name = "proxycast-terminal" version = "0.60.0" diff --git a/src-tauri/Cargo.toml b/src-tauri/Cargo.toml index 7ea4fa8ec..e43808baa 100644 --- a/src-tauri/Cargo.toml +++ b/src-tauri/Cargo.toml @@ -21,6 +21,10 @@ proxycast-credential = { path = "crates/credential" } proxycast-websocket = { path = "crates/websocket" } proxycast-processor = { path = "crates/processor" } proxycast-server-utils = { path = "crates/server-utils" } +proxycast-server = { path = "crates/server" } +proxycast-skills = { path = "crates/skills" } +proxycast-mcp = { path = "crates/mcp" } +proxycast-agent = { path = "crates/agent" } voice-core = { path = "crates/voice-core" } # 序列化 @@ -201,6 +205,10 @@ proxycast-credential.workspace = true proxycast-websocket.workspace = true proxycast-processor.workspace = true proxycast-server-utils.workspace = true +proxycast-server.workspace = true +proxycast-skills.workspace = true +proxycast-mcp.workspace = true +proxycast-agent.workspace = true voice-core.workspace = true # Tauri diff --git a/src-tauri/crates/core/Cargo.toml b/src-tauri/crates/core/Cargo.toml index 0526eba8d..6c3e3bae3 100644 --- a/src-tauri/crates/core/Cargo.toml +++ b/src-tauri/crates/core/Cargo.toml @@ -51,12 +51,18 @@ flate2.workspace = true tar.workspace = true zip.workspace = true +# 正则表达式(logger 脱敏需要) +regex.workspace = true + # YAML 配置 serde_yaml.workspace = true # 数据库(errors 模块需要 rusqlite::Error) rusqlite.workspace = true +# 网络接口(network 模块需要) +if-addrs.workspace = true + [dev-dependencies] proptest.workspace = true tempfile.workspace = true \ No newline at end of file diff --git a/src-tauri/crates/core/src/lib.rs b/src-tauri/crates/core/src/lib.rs index b59aad146..f81d7a2ff 100644 --- a/src-tauri/crates/core/src/lib.rs +++ b/src-tauri/crates/core/src/lib.rs @@ -47,6 +47,12 @@ pub mod processor; // WebSocket 核心类型 pub mod websocket; +// 事件发射抽象(供独立 crate 解耦 Tauri 依赖) +pub mod event_emit; + +// 网络工具 +pub mod network; + // 数据层 pub mod content; pub mod database; @@ -54,6 +60,7 @@ pub mod memory; pub mod workspace; // 重新导出常用类型 +pub use event_emit::{DynEmitter, EventEmit, NoOpEmitter}; pub use logger::{LogEntry, LogStore, LogStoreConfig, SharedLogStore}; pub use models::provider_type::ProviderType; pub use models::*; diff --git a/src-tauri/crates/core/src/logger.rs b/src-tauri/crates/core/src/logger.rs index 69af7b85c..1c56b9917 100644 --- a/src-tauri/crates/core/src/logger.rs +++ b/src-tauri/crates/core/src/logger.rs @@ -1,9 +1,10 @@ //! 日志管理模块 use chrono::{Duration, Local, Utc}; +use regex::Regex; use serde::{Deserialize, Serialize}; use std::collections::VecDeque; use std::fs::{self, OpenOptions}; -use std::io::Write; +use std::io::{Read, Write}; use std::path::PathBuf; use std::sync::Arc; @@ -42,19 +43,13 @@ pub struct LogStore { impl Default for LogStore { fn default() -> Self { - // 默认日志文件路径: ~/.proxycast/logs/proxycast.log let log_dir = dirs::home_dir() .unwrap_or_else(|| PathBuf::from(".")) .join(".proxycast") .join("logs"); - - // 创建日志目录 let _ = fs::create_dir_all(&log_dir); - let log_file = log_dir.join("proxycast.log"); - let config = LogStoreConfig::default(); - Self { logs: VecDeque::new(), max_logs: config.max_logs, @@ -86,24 +81,18 @@ impl LogStore { level: level.to_string(), message: sanitized.clone(), }; - self.logs.push_back(entry.clone()); - - // 写入日志文件 if self.config.enable_file_logging { if let Some(ref path) = self.log_file_path { self.rotate_log_file_if_needed(path); let local_time = Local::now().format("%Y-%m-%d %H:%M:%S%.3f"); let log_line = format!("{} [{}] {}\n", local_time, level.to_uppercase(), sanitized); - if let Ok(mut file) = OpenOptions::new().create(true).append(true).open(path) { let _ = file.write_all(log_line.as_bytes()); } self.prune_old_logs(path); } } - - // 保持日志数量在限制内 if self.logs.len() > self.max_logs { self.logs.pop_front(); } @@ -115,7 +104,6 @@ impl LogStore { let log_dir = log_path.parent().unwrap_or(std::path::Path::new(".")); let raw_file = log_dir.join(format!("raw_response_{request_id}.txt")); let sanitized = sanitize_log_message(body); - if let Ok(mut file) = OpenOptions::new() .create(true) .truncate(true) @@ -145,26 +133,22 @@ impl LogStore { let Ok(metadata) = fs::metadata(path) else { return; }; - if metadata.len() <= self.config.max_file_size { return; } - let suffix = Local::now().format("%Y%m%d-%H%M%S"); let rotated = path.with_file_name(format!( "{}.{}", path.file_name().unwrap_or_default().to_string_lossy(), suffix )); - let _ = fs::rename(path, &rotated); self.prune_old_logs(path); } - fn prune_old_logs(&self, path: &std::path::Path) { - let Some(dir) = path.parent() else { - return; - }; + fn prune_old_logs(&self, path: &PathBuf) { + let Some(dir) = path.parent() else { return }; + self.archive_old_logs(path); let Ok(entries) = fs::read_dir(dir) else { return; }; @@ -173,7 +157,6 @@ impl LogStore { "{}.", path.file_name().unwrap_or_default().to_string_lossy() ); - for entry in entries.flatten() { let file_name = entry.file_name(); let file_name = file_name.to_string_lossy(); @@ -192,26 +175,106 @@ impl LogStore { } } } + + fn archive_old_logs(&self, path: &PathBuf) { + let Some(dir) = path.parent() else { return }; + let Ok(entries) = fs::read_dir(dir) else { + return; + }; + let archive_cutoff = Utc::now() - Duration::days(7); + let delete_cutoff = Utc::now() - Duration::days(30); + let prefix = format!( + "{}.", + path.file_name().unwrap_or_default().to_string_lossy() + ); + for entry in entries.flatten() { + let file_name = entry.file_name(); + let file_name = file_name.to_string_lossy(); + if !file_name.starts_with(&prefix) { + continue; + } + let path = entry.path(); + let Ok(metadata) = entry.metadata() else { + continue; + }; + let Ok(modified) = metadata.modified() else { + continue; + }; + let modified = chrono::DateTime::::from(modified); + // 删除超过 30 天的 gz 文件 + if file_name.ends_with(".gz") { + if modified < delete_cutoff { + let _ = fs::remove_file(path); + } + continue; + } + // 跳过不到 7 天的文件 + if modified >= archive_cutoff { + continue; + } + // 压缩超过 7 天的日志文件 + let mut input = Vec::new(); + if let Ok(mut file) = fs::File::open(&path) { + if file.read_to_end(&mut input).is_err() { + continue; + } + } else { + continue; + } + let gz_path = path.with_extension(format!( + "{}.gz", + path.extension().unwrap_or_default().to_string_lossy() + )); + if let Ok(gz_file) = fs::File::create(&gz_path) { + let mut encoder = + flate2::write::GzEncoder::new(gz_file, flate2::Compression::default()); + if encoder.write_all(&input).is_ok() && encoder.finish().is_ok() { + let _ = fs::remove_file(&path); + } + } + } + } } -/// 简化的共享日志存储类型(使用 parking_lot) pub type SharedLogStore = Arc>; /// P2 安全修复:扩展日志脱敏规则,覆盖更多敏感字段 pub fn sanitize_log_message(message: &str) -> String { - // 简化版本:使用字符串替换而不是正则表达式 + let patterns = [ + (r"Bearer\s+[A-Za-z0-9._-]+", "Bearer ***"), + ( + r#"api[_-]?key["']?\s*[:=]\s*["']?[A-Za-z0-9._-]+"#, + "api_key: ***", + ), + (r#"token["']?\s*[:=]\s*["']?[A-Za-z0-9._-]+"#, "token: ***"), + ( + r#"access[_-]?token["']?\s*[:=]\s*["']?[A-Za-z0-9._-]+"#, + "access_token: ***", + ), + ( + r#"refresh[_-]?token["']?\s*[:=]\s*["']?[A-Za-z0-9._-]+"#, + "refresh_token: ***", + ), + ( + r#"client[_-]?secret["']?\s*[:=]\s*["']?[A-Za-z0-9._-]+"#, + "client_secret: ***", + ), + ( + r#"[Aa]uthorization["']?\s*[:=]\s*["']?[A-Za-z0-9._\s-]+"#, + "authorization: ***", + ), + (r#"password["']?\s*[:=]\s*["']?[^\s"',}]+"#, "password: ***"), + ( + r#"secret["']?\s*[:=]\s*["']?[A-Za-z0-9._-]+"#, + "secret: ***", + ), + ]; let mut sanitized = message.to_string(); - - // Bearer token - if let Some(pos) = sanitized.find("Bearer ") { - let start = pos + 7; - if let Some(end) = - sanitized[start..].find(|c: char| c.is_whitespace() || c == '"' || c == '\'') - { - sanitized.replace_range(start..start + end, "***"); + for (pattern, replacement) in patterns { + if let Ok(re) = Regex::new(pattern) { + sanitized = re.replace_all(&sanitized, replacement).to_string(); } } - sanitized } @@ -221,11 +284,52 @@ mod tests { #[test] fn test_sanitize_bearer_token() { - let input = "Authorization: Bearer abcDEF123 end"; + let input = "Authorization: Bearer abcDEF123._-XYZ"; let output = sanitize_log_message(input); + assert!(!output.contains("abcDEF123")); assert!(output.contains("***")); } + #[test] + fn test_sanitize_api_key() { + let input = r#"request api_key="sk-test_123.456-ABC" end"#; + let output = sanitize_log_message(input); + assert!(output.contains("api_key: ***")); + assert!(!output.contains("sk-test_123")); + } + + #[test] + fn test_sanitize_access_token() { + let input = "access_token=atk_12345"; + let output = sanitize_log_message(input); + assert!(output.contains("access_token: ***")); + assert!(!output.contains("atk_12345")); + } + + #[test] + fn test_sanitize_refresh_token() { + let input = "refresh_token: rtk_ABCDE-123"; + let output = sanitize_log_message(input); + assert!(output.contains("refresh_token: ***")); + assert!(!output.contains("rtk_ABCDE")); + } + + #[test] + fn test_sanitize_client_secret() { + let input = "client_secret = \"cs_SeCreT-999\""; + let output = sanitize_log_message(input); + assert!(output.contains("client_secret: ***")); + assert!(!output.contains("cs_SeCreT")); + } + + #[test] + fn test_sanitize_password() { + let input = r#"{"password":"p@ssW0rd!"}"#; + let output = sanitize_log_message(input); + assert!(output.contains("password: ***")); + assert!(!output.contains("p@ssW0rd!")); + } + #[test] fn test_plain_text_unchanged() { let input = "这是一段普通日志,不包含任何敏感字段。"; diff --git a/src-tauri/crates/core/src/network.rs b/src-tauri/crates/core/src/network.rs new file mode 100644 index 000000000..42ebdc546 --- /dev/null +++ b/src-tauri/crates/core/src/network.rs @@ -0,0 +1,158 @@ +//! 网络工具模块 +//! +//! 提供获取本地网络接口信息的功能。 +//! 从主 crate 的 commands/network_cmd.rs 迁移而来。 + +use serde::Serialize; +use std::net::{IpAddr, UdpSocket}; + +/// 网络接口信息 +#[derive(Debug, Clone, Serialize)] +pub struct NetworkInfo { + /// 本地回环地址 + pub localhost: String, + /// 内网 IP 地址(局域网) + pub lan_ip: Option, + /// 所有可用的网络接口 IP 地址 + pub all_ips: Vec, +} + +/// 获取本地网络信息 +/// +/// 返回 localhost 和内网 IP 地址,用于客户端连接 +pub fn get_network_info() -> Result { + let lan_ip = get_local_ip(); + let all_ips = get_all_local_ips(); + + Ok(NetworkInfo { + localhost: "127.0.0.1".to_string(), + lan_ip, + all_ips, + }) +} + +/// 获取本机内网 IP 地址 +/// +/// 通过创建 UDP socket 连接外部地址来获取本机的内网 IP +fn get_local_ip() -> Option { + let socket = UdpSocket::bind("0.0.0.0:0").ok()?; + socket.connect("8.8.8.8:80").ok()?; + let local_addr = socket.local_addr().ok()?; + let ip_str = local_addr.ip().to_string(); + + if let IpAddr::V4(ipv4) = local_addr.ip() { + if ipv4.octets()[0] == 198 && (ipv4.octets()[1] == 18 || ipv4.octets()[1] == 19) { + let all_ips = get_all_local_ips(); + if let Some(ip) = all_ips.iter().find(|ip| ip.starts_with("192.168.")) { + return Some(ip.clone()); + } + if let Some(ip) = all_ips.first() { + return Some(ip.clone()); + } + return Some("127.0.0.1".to_string()); + } + } + + Some(ip_str) +} + +/// 获取所有本地网络接口的 IP 地址 +/// +/// 返回所有非回环的 IPv4 私有地址 +fn get_all_local_ips() -> Vec { + let mut ips = Vec::new(); + + if let Ok(interfaces) = if_addrs::get_if_addrs() { + for iface in interfaces { + if let IpAddr::V4(ipv4) = iface.ip() { + if ipv4.is_loopback() { + continue; + } + if ipv4.octets()[0] == 169 && ipv4.octets()[1] == 254 { + continue; + } + if ipv4.octets()[0] == 198 && (ipv4.octets()[1] == 18 || ipv4.octets()[1] == 19) { + continue; + } + let is_private = ipv4.octets()[0] == 10 + || (ipv4.octets()[0] == 172 + && (ipv4.octets()[1] >= 16 && ipv4.octets()[1] <= 31)) + || (ipv4.octets()[0] == 192 && ipv4.octets()[1] == 168); + + if is_private { + ips.push(ipv4.to_string()); + } + } + } + } + + ips +} + +/// 根据监听地址生成可访问的 host +pub fn get_accessible_host(listen_host: &str) -> String { + match listen_host { + "0.0.0.0" => get_network_info() + .ok() + .and_then(|info| info.lan_ip) + .unwrap_or_else(|| "127.0.0.1".to_string()), + "localhost" => "127.0.0.1".to_string(), + _ => listen_host.to_string(), + } +} + +/// 根据监听地址生成可访问的 URL +pub fn get_accessible_url(listen_host: &str, port: u16) -> String { + let host = get_accessible_host(listen_host); + format!("http://{host}:{port}") +} + +/// 根据监听地址生成本地访问的 URL +#[allow(dead_code)] +pub fn get_local_url(listen_host: &str, port: u16) -> String { + let host = match listen_host { + "0.0.0.0" | "localhost" => "127.0.0.1".to_string(), + _ => listen_host.to_string(), + }; + format!("http://{host}:{port}") +} + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn test_get_accessible_host_localhost() { + assert_eq!(get_accessible_host("127.0.0.1"), "127.0.0.1"); + assert_eq!(get_accessible_host("localhost"), "127.0.0.1"); + } + + #[test] + fn test_get_accessible_host_specific_ip() { + assert_eq!(get_accessible_host("192.168.1.100"), "192.168.1.100"); + assert_eq!(get_accessible_host("10.0.0.1"), "10.0.0.1"); + } + + #[test] + fn test_get_local_url() { + assert_eq!(get_local_url("0.0.0.0", 8999), "http://127.0.0.1:8999"); + assert_eq!(get_local_url("127.0.0.1", 8999), "http://127.0.0.1:8999"); + assert_eq!(get_local_url("localhost", 8999), "http://127.0.0.1:8999"); + assert_eq!( + get_local_url("192.168.1.100", 8999), + "http://192.168.1.100:8999" + ); + } + + #[test] + fn test_get_accessible_url_specific_ip() { + assert_eq!( + get_accessible_url("192.168.1.100", 8999), + "http://192.168.1.100:8999" + ); + assert_eq!( + get_accessible_url("127.0.0.1", 8999), + "http://127.0.0.1:8999" + ); + } +} diff --git a/src-tauri/crates/server/Cargo.toml b/src-tauri/crates/server/Cargo.toml new file mode 100644 index 000000000..a10a3c23c --- /dev/null +++ b/src-tauri/crates/server/Cargo.toml @@ -0,0 +1,40 @@ +[package] +name = "proxycast-server" +version.workspace = true +edition.workspace = true + +[dependencies] +proxycast-core.workspace = true +proxycast-config.workspace = true +proxycast-infra.workspace = true +proxycast-providers.workspace = true +proxycast-services.workspace = true +proxycast-credential.workspace = true +proxycast-websocket.workspace = true +proxycast-processor.workspace = true +proxycast-server-utils.workspace = true + +serde.workspace = true +serde_json.workspace = true +tokio.workspace = true +futures.workspace = true +axum.workspace = true +tower.workspace = true +tower-http.workspace = true +tracing.workspace = true +reqwest.workspace = true +rusqlite.workspace = true +chrono.workspace = true +uuid.workspace = true +base64.workspace = true +bytes.workspace = true +regex.workspace = true +subtle.workspace = true +async-stream.workspace = true +urlencoding.workspace = true +parking_lot.workspace = true +tokio-util.workspace = true +dirs.workspace = true + +[dev-dependencies] +proptest.workspace = true diff --git a/src-tauri/src/server/client_detector.rs b/src-tauri/crates/server/src/client_detector.rs similarity index 97% rename from src-tauri/src/server/client_detector.rs rename to src-tauri/crates/server/src/client_detector.rs index 768872f8b..20c2fc236 100644 --- a/src-tauri/src/server/client_detector.rs +++ b/src-tauri/crates/server/src/client_detector.rs @@ -7,8 +7,8 @@ pub use proxycast_core::models::client_type::*; #[cfg(test)] mod property_tests { use super::*; - use crate::config::EndpointProvidersConfig; use proptest::prelude::*; + use proxycast_core::config::EndpointProvidersConfig; fn arb_client_type() -> impl Strategy { prop_oneof![ diff --git a/src-tauri/src/server/handlers/api.rs b/src-tauri/crates/server/src/handlers/api.rs similarity index 95% rename from src-tauri/src/server/handlers/api.rs rename to src-tauri/crates/server/src/handlers/api.rs index 93ef7ec8b..85c6c5cf4 100644 --- a/src-tauri/src/server/handlers/api.rs +++ b/src-tauri/crates/server/src/handlers/api.rs @@ -25,18 +25,18 @@ use axum::{ use serde_json::json; use std::collections::HashMap; -use crate::converter::anthropic_to_openai::convert_anthropic_to_openai; -use crate::models::anthropic::AnthropicMessagesRequest; -use crate::models::openai::ChatCompletionRequest; -use crate::processor::RequestContext; -use crate::server::client_detector::ClientType; -use crate::server::{record_request_telemetry, record_token_usage, AppState}; -use crate::server_utils::{ +use crate::client_detector::ClientType; +use crate::{record_request_telemetry, record_token_usage, AppState}; +use proxycast_core::models::anthropic::AnthropicMessagesRequest; +use proxycast_core::models::openai::ChatCompletionRequest; +use proxycast_core::ProviderType; +use proxycast_processor::RequestContext; +use proxycast_providers::converter::anthropic_to_openai::convert_anthropic_to_openai; +use proxycast_providers::streaming::StreamFormat as StreamingFormat; +use proxycast_server_utils::{ build_anthropic_response, build_anthropic_stream_response, message_content_len, parse_cw_response, safe_truncate, }; -use crate::streaming::StreamFormat as StreamingFormat; -use crate::ProviderType; use super::{call_provider_anthropic, call_provider_openai}; @@ -331,13 +331,14 @@ pub async fn chat_completions( let credential = if credential.is_none() { eprintln!("[CHAT_COMPLETIONS] Provider Pool 中未找到凭证,尝试 API Key Provider..."); - use crate::database::dao::api_key_provider::ApiProviderType; + use proxycast_core::database::dao::api_key_provider::ApiProviderType; let provider_id_lower = selected_provider.to_lowercase(); // 策略 1: 优先按 provider_id 直接查找(支持 deepseek, moonshot 等 60+ Provider) // 这些 Provider 在 API Key Provider 中有独立配置 - let mut found_credential: Option = - None; + let mut found_credential: Option< + proxycast_core::models::provider_pool_model::ProviderCredential, + > = None; if let Some(db) = &state.db { // 先尝试按 provider_id 直接查找 @@ -347,7 +348,7 @@ pub async fn chat_completions( .api_key_service .get_fallback_credential( db, - &crate::models::provider_pool_model::PoolProviderType::OpenAI, + &proxycast_core::models::provider_pool_model::PoolProviderType::OpenAI, Some(&provider_id_lower), Some(&client_type), ) @@ -408,36 +409,40 @@ pub async fn chat_completions( }; let provider_type = match provider_info.provider_type { - ApiProviderType::Anthropic => crate::ProviderType::Anthropic, - ApiProviderType::Openai | ApiProviderType::OpenaiResponse => { - crate::ProviderType::OpenAI + ApiProviderType::Anthropic => { + proxycast_core::ProviderType::Anthropic } - ApiProviderType::Gemini => crate::ProviderType::GeminiApiKey, - _ => crate::ProviderType::OpenAI, + ApiProviderType::Openai | ApiProviderType::OpenaiResponse => { + proxycast_core::ProviderType::OpenAI + } + ApiProviderType::Gemini => { + proxycast_core::ProviderType::GeminiApiKey + } + _ => proxycast_core::ProviderType::OpenAI, }; let credential_data = match provider_type { - crate::ProviderType::Anthropic => { - crate::models::provider_pool_model::CredentialData::AnthropicKey { + proxycast_core::ProviderType::Anthropic => { + proxycast_core::models::provider_pool_model::CredentialData::AnthropicKey { api_key: api_key.clone(), base_url, } } - crate::ProviderType::GeminiApiKey => { - crate::models::provider_pool_model::CredentialData::GeminiApiKey { + proxycast_core::ProviderType::GeminiApiKey => { + proxycast_core::models::provider_pool_model::CredentialData::GeminiApiKey { api_key: api_key.clone(), base_url, excluded_models: vec![], } } - _ => crate::models::provider_pool_model::CredentialData::OpenAIKey { + _ => proxycast_core::models::provider_pool_model::CredentialData::OpenAIKey { api_key: api_key.clone(), base_url, }, }; let mut cred = - crate::models::provider_pool_model::ProviderCredential::new( + proxycast_core::models::provider_pool_model::ProviderCredential::new( provider_type, credential_data, ); @@ -522,9 +527,9 @@ pub async fn chat_completions( let is_success = response.status().is_success(); let _status_code = response.status().as_u16(); let status = if is_success { - crate::telemetry::RequestStatus::Success + proxycast_infra::telemetry::RequestStatus::Success } else { - crate::telemetry::RequestStatus::Failed + proxycast_infra::telemetry::RequestStatus::Failed }; record_request_telemetry(&state, &ctx, status, None); @@ -674,7 +679,7 @@ pub async fn chat_completions( record_request_telemetry( &state, &ctx, - crate::telemetry::RequestStatus::Success, + proxycast_infra::telemetry::RequestStatus::Success, None, ); // 记录 Token 使用量 @@ -693,7 +698,7 @@ pub async fn chat_completions( record_request_telemetry( &state, &ctx, - crate::telemetry::RequestStatus::Failed, + proxycast_infra::telemetry::RequestStatus::Failed, Some(e.to_string()), ); // 标记 Flow 失败 @@ -1059,8 +1064,9 @@ pub async fn anthropic_messages( eprintln!("[ANTHROPIC_MESSAGES] Provider Pool 中未找到凭证,尝试 API Key Provider..."); // 策略 1: 优先按 provider_id 直接查找(支持自定义 Provider) - let mut found_credential: Option = - None; + let mut found_credential: Option< + proxycast_core::models::provider_pool_model::ProviderCredential, + > = None; if let Some(db) = &state.db { eprintln!("[ANTHROPIC_MESSAGES] 尝试按 provider_id '{selected_provider}' 直接查找凭证"); @@ -1069,7 +1075,7 @@ pub async fn anthropic_messages( .api_key_service .get_fallback_credential( db, - &crate::models::provider_pool_model::PoolProviderType::Anthropic, + &proxycast_core::models::provider_pool_model::PoolProviderType::Anthropic, Some(&selected_provider), Some(&client_type), ) @@ -1133,9 +1139,9 @@ pub async fn anthropic_messages( // 记录请求统计 let is_success = response.status().is_success(); let status = if is_success { - crate::telemetry::RequestStatus::Success + proxycast_infra::telemetry::RequestStatus::Success } else { - crate::telemetry::RequestStatus::Failed + proxycast_infra::telemetry::RequestStatus::Failed }; record_request_telemetry(&state, &ctx, status, None); @@ -1552,9 +1558,9 @@ fn get_target_stream_format(path: &str) -> StreamingFormat { /// 当前所有 Provider 都返回 false,因为 StreamingProvider trait 尚未实现。 /// 一旦任务 6 完成,此函数将根据凭证类型返回适当的值。 fn should_use_true_streaming( - credential: &crate::models::provider_pool_model::ProviderCredential, + credential: &proxycast_core::models::provider_pool_model::ProviderCredential, ) -> bool { - use crate::models::provider_pool_model::CredentialData; + use proxycast_core::models::provider_pool_model::CredentialData; // TODO: 当 StreamingProvider trait 实现后,根据凭证类型返回 true // 目前所有 Provider 都使用伪流式模式 @@ -1665,10 +1671,10 @@ fn map_to_api_key_provider_id(provider_type: &str) -> String { /// 根据 API Provider 类型构建额外的请求头 fn build_api_key_headers( - provider_type: &crate::database::dao::api_key_provider::ApiProviderType, + provider_type: &proxycast_core::database::dao::api_key_provider::ApiProviderType, api_key: &str, ) -> HashMap { - use crate::database::dao::api_key_provider::ApiProviderType; + use proxycast_core::database::dao::api_key_provider::ApiProviderType; let mut headers = HashMap::new(); @@ -1693,9 +1699,9 @@ fn build_api_key_headers( /// 获取默认的 API Host fn get_default_api_host( - provider_type: &crate::database::dao::api_key_provider::ApiProviderType, + provider_type: &proxycast_core::database::dao::api_key_provider::ApiProviderType, ) -> String { - use crate::database::dao::api_key_provider::ApiProviderType; + use proxycast_core::database::dao::api_key_provider::ApiProviderType; match provider_type { ApiProviderType::Openai | ApiProviderType::OpenaiResponse => { @@ -1718,11 +1724,11 @@ fn convert_openai_to_anthropic(request: &ChatCompletionRequest) -> serde_json::V // 提取 system prompt if let Some(content) = &msg.content { system_prompt = Some(match content { - crate::models::openai::MessageContent::Text(s) => s.clone(), - crate::models::openai::MessageContent::Parts(parts) => parts + proxycast_core::models::openai::MessageContent::Text(s) => s.clone(), + proxycast_core::models::openai::MessageContent::Parts(parts) => parts .iter() .filter_map(|p| { - if let crate::models::openai::ContentPart::Text { text } = p { + if let proxycast_core::models::openai::ContentPart::Text { text } = p { Some(text.clone()) } else { None @@ -1736,11 +1742,11 @@ fn convert_openai_to_anthropic(request: &ChatCompletionRequest) -> serde_json::V // 转换其他消息 let content = match &msg.content { Some(c) => match c { - crate::models::openai::MessageContent::Text(s) => s.clone(), - crate::models::openai::MessageContent::Parts(parts) => parts + proxycast_core::models::openai::MessageContent::Text(s) => s.clone(), + proxycast_core::models::openai::MessageContent::Parts(parts) => parts .iter() .filter_map(|p| { - if let crate::models::openai::ContentPart::Text { text } = p { + if let proxycast_core::models::openai::ContentPart::Text { text } = p { Some(text.clone()) } else { None diff --git a/src-tauri/src/server/handlers/credentials_api.rs b/src-tauri/crates/server/src/handlers/credentials_api.rs similarity index 97% rename from src-tauri/src/server/handlers/credentials_api.rs rename to src-tauri/crates/server/src/handlers/credentials_api.rs index 1266efd06..adb7195d1 100644 --- a/src-tauri/src/server/handlers/credentials_api.rs +++ b/src-tauri/crates/server/src/handlers/credentials_api.rs @@ -16,10 +16,10 @@ use axum::{ use chrono::{DateTime, Utc}; use serde::{Deserialize, Serialize}; -use crate::database::dao::api_key_provider::{ApiKeyProviderDao, ApiProviderType}; -use crate::database::dao::provider_pool::ProviderPoolDao; -use crate::models::provider_pool_model::PoolProviderType; -use crate::server::AppState; +use crate::AppState; +use proxycast_core::database::dao::api_key_provider::{ApiKeyProviderDao, ApiProviderType}; +use proxycast_core::database::dao::provider_pool::ProviderPoolDao; +use proxycast_core::models::provider_pool_model::PoolProviderType; /// 选择凭证请求参数 #[derive(Debug, Deserialize)] @@ -155,7 +155,7 @@ pub async fn credentials_select( /// 尝试从 OAuth 凭证池选择凭证 async fn try_select_oauth_credential( state: &AppState, - db: &crate::database::DbConnection, + db: &proxycast_core::database::DbConnection, request: &SelectCredentialRequest, ) -> Result, CredentialApiError> { // 使用 ProviderPoolService 智能选择凭证 @@ -208,7 +208,7 @@ async fn try_select_oauth_credential( /// 尝试从 API Key Provider 选择凭证 async fn try_select_api_key_credential( state: &AppState, - db: &crate::database::DbConnection, + db: &proxycast_core::database::DbConnection, request: &SelectCredentialRequest, ) -> Result, CredentialApiError> { // 将 provider_type 映射到 API Key Provider ID @@ -363,7 +363,7 @@ pub async fn credentials_get_token( /// 尝试从 OAuth 凭证池获取 Token async fn try_get_oauth_token( state: &AppState, - db: &crate::database::DbConnection, + db: &proxycast_core::database::DbConnection, uuid: &str, ) -> Result, CredentialApiError> { // 查询凭证 @@ -468,7 +468,7 @@ async fn try_get_oauth_token( /// 尝试从 API Key Provider 获取 Token async fn try_get_api_key_token( state: &AppState, - db: &crate::database::DbConnection, + db: &proxycast_core::database::DbConnection, uuid: &str, ) -> Result, CredentialApiError> { let conn = db.lock().map_err(|e| CredentialApiError { diff --git a/src-tauri/src/server/handlers/image_handler.rs b/src-tauri/crates/server/src/handlers/image_handler.rs similarity index 97% rename from src-tauri/src/server/handlers/image_handler.rs rename to src-tauri/crates/server/src/handlers/image_handler.rs index a2c673129..1cd537662 100644 --- a/src-tauri/src/server/handlers/image_handler.rs +++ b/src-tauri/crates/server/src/handlers/image_handler.rs @@ -23,14 +23,14 @@ use axum::{ Json, }; -use crate::converter::openai_to_antigravity::{ +use crate::handlers::verify_api_key; +use crate::AppState; +use proxycast_core::models::openai::ImageGenerationRequest; +use proxycast_core::models::provider_pool_model::CredentialData; +use proxycast_providers::converter::openai_to_antigravity::{ convert_antigravity_image_response, convert_image_request_to_antigravity, }; -use crate::models::openai::ImageGenerationRequest; -use crate::models::provider_pool_model::CredentialData; -use crate::providers::AntigravityProvider; -use crate::server::handlers::verify_api_key; -use crate::server::AppState; +use proxycast_providers::providers::AntigravityProvider; /// 处理图像生成请求 /// diff --git a/src-tauri/src/server/handlers/kiro_credential.rs b/src-tauri/crates/server/src/handlers/kiro_credential.rs similarity index 99% rename from src-tauri/src/server/handlers/kiro_credential.rs rename to src-tauri/crates/server/src/handlers/kiro_credential.rs index ace9d57cc..64250fe2b 100644 --- a/src-tauri/src/server/handlers/kiro_credential.rs +++ b/src-tauri/crates/server/src/handlers/kiro_credential.rs @@ -15,9 +15,11 @@ use axum::{ use chrono::{DateTime, Utc}; use serde::{Deserialize, Serialize}; -use crate::database::dao::provider_pool::ProviderPoolDao; -use crate::models::provider_pool_model::{CachedTokenInfo, PoolProviderType, ProviderCredential}; -use crate::server::AppState; +use crate::AppState; +use proxycast_core::database::dao::provider_pool::ProviderPoolDao; +use proxycast_core::models::provider_pool_model::{ + CachedTokenInfo, PoolProviderType, ProviderCredential, +}; /// 可用凭证信息 #[derive(Debug, Clone, Serialize)] diff --git a/src-tauri/src/server/handlers/management.rs b/src-tauri/crates/server/src/handlers/management.rs similarity index 98% rename from src-tauri/src/server/handlers/management.rs rename to src-tauri/crates/server/src/handlers/management.rs index 7ca9f66da..347d039a4 100644 --- a/src-tauri/src/server/handlers/management.rs +++ b/src-tauri/crates/server/src/handlers/management.rs @@ -7,8 +7,8 @@ use axum::{extract::State, http::StatusCode, response::IntoResponse, Json}; use serde::{Deserialize, Serialize}; -use crate::database::dao::provider_pool::ProviderPoolDao; -use crate::server::AppState; +use crate::AppState; +use proxycast_core::database::dao::provider_pool::ProviderPoolDao; // ============ Types ============ @@ -201,7 +201,7 @@ pub async fn management_add_credential( State(state): State, Json(request): Json, ) -> impl IntoResponse { - use crate::models::provider_pool_model::{ + use proxycast_core::models::provider_pool_model::{ CredentialData, PoolProviderType, ProviderCredential, }; @@ -538,7 +538,7 @@ pub async fn management_update_config( // 更新默认 Provider if let Some(provider) = request.default_provider { // 验证 provider 类型 - if provider.parse::().is_ok() { + if provider.parse::().is_ok() { let mut dp = state.default_provider.write().await; *dp = provider.clone(); tracing::info!("[MANAGEMENT] Updated default_provider to: {}", provider); diff --git a/src-tauri/src/server/handlers/mod.rs b/src-tauri/crates/server/src/handlers/mod.rs similarity index 100% rename from src-tauri/src/server/handlers/mod.rs rename to src-tauri/crates/server/src/handlers/mod.rs diff --git a/src-tauri/src/server/handlers/provider_calls.rs b/src-tauri/crates/server/src/handlers/provider_calls.rs similarity index 98% rename from src-tauri/src/server/handlers/provider_calls.rs rename to src-tauri/crates/server/src/handlers/provider_calls.rs index ef1504da5..104efa84f 100644 --- a/src-tauri/src/server/handlers/provider_calls.rs +++ b/src-tauri/crates/server/src/handlers/provider_calls.rs @@ -49,29 +49,29 @@ use axum::{ }; use futures::StreamExt; -use crate::converter::anthropic_to_openai::convert_anthropic_to_openai; -use crate::converter::openai_to_antigravity::{ +use crate::AppState; +use proxycast_core::models::anthropic::AnthropicMessagesRequest; +use proxycast_core::models::openai::ChatCompletionRequest; +use proxycast_core::models::provider_pool_model::{CredentialData, ProviderCredential}; +use proxycast_providers::converter::anthropic_to_openai::convert_anthropic_to_openai; +use proxycast_providers::converter::openai_to_antigravity::{ convert_antigravity_to_openai_response, convert_openai_to_antigravity_with_context, }; -use crate::models::anthropic::AnthropicMessagesRequest; -use crate::models::openai::ChatCompletionRequest; -use crate::models::provider_pool_model::{CredentialData, ProviderCredential}; -use crate::providers::{ +use proxycast_providers::providers::{ AntigravityProvider, ClaudeCustomProvider, CodexProvider, KiroProvider, OpenAICustomProvider, VertexProvider, }; -use crate::server::AppState; -use crate::server_utils::{ - build_anthropic_response, build_anthropic_stream_response, build_error_response, - build_error_response_with_status, parse_cw_response, safe_truncate, CWParsedResponse, -}; -use crate::session::store_thought_signature; -use crate::stream::{PipelineConfig, StreamPipeline}; -use crate::streaming::traits::StreamingProvider; -use crate::streaming::{ +use proxycast_providers::session::store_thought_signature; +use proxycast_providers::stream::{PipelineConfig, StreamPipeline}; +use proxycast_providers::streaming::traits::StreamingProvider; +use proxycast_providers::streaming::{ StreamConfig, StreamContext, StreamError, StreamFormat as StreamingFormat, StreamManager, StreamResponse, }; +use proxycast_server_utils::{ + build_anthropic_response, build_anthropic_stream_response, build_error_response, + build_error_response_with_status, parse_cw_response, safe_truncate, CWParsedResponse, +}; /// 根据凭证调用 Provider (Anthropic 格式) /// @@ -1658,9 +1658,9 @@ pub async fn call_provider_openai( // 创建 StreamConverter 将 Anthropic SSE 转换为 OpenAI SSE let converter = std::sync::Arc::new(tokio::sync::Mutex::new( - crate::streaming::converter::StreamConverter::with_model( - crate::streaming::converter::StreamFormat::AnthropicSse, - crate::streaming::converter::StreamFormat::OpenAiSse, + proxycast_providers::streaming::converter::StreamConverter::with_model( + proxycast_providers::streaming::converter::StreamFormat::AnthropicSse, + proxycast_providers::streaming::converter::StreamFormat::OpenAiSse, &request.model, ), )); @@ -1681,7 +1681,7 @@ pub async fn call_provider_openai( }; for sse_str in sse_events { - yield Ok::(sse_str); + yield Ok::(sse_str); } } Err(e) => { @@ -1699,7 +1699,7 @@ pub async fn call_provider_openai( }; for sse_str in final_events { - yield Ok::(sse_str); + yield Ok::(sse_str); } }; @@ -2257,9 +2257,14 @@ pub async fn handle_streaming_response_with_timeout( // 获取 flow_id 的克隆用于回调 // 创建带超时的流式处理,使用 BoxStream 统一类型 - let timeout_stream: BoxStream<'static, Result> = { + let timeout_stream: BoxStream< + 'static, + Result, + > = { let stream = manager.handle_stream(context, source_stream); - Box::pin(crate::streaming::with_timeout(stream, &config)) + Box::pin(proxycast_providers::streaming::with_timeout( + stream, &config, + )) }; // 转换为 Body 流 @@ -2299,7 +2304,7 @@ pub async fn handle_streaming_response_with_timeout( /// # 返回 /// 统一的流式响应类型 pub fn response_to_stream(response: reqwest::Response) -> StreamResponse { - crate::streaming::reqwest_stream_to_stream_response(response) + proxycast_providers::streaming::reqwest_stream_to_stream_response(response) } // ============================================================================ @@ -2355,7 +2360,7 @@ pub async fn handle_streaming_with_disconnect_detection( // 创建流式处理 let managed_stream: futures::stream::BoxStream< 'static, - Result, + Result, > = Box::pin(manager.handle_stream(context, source_stream)); // 如果有取消令牌,创建一个可取消的流 @@ -2624,8 +2629,8 @@ pub async fn handle_kiro_stream( // 检查是否是 401/403 错误或 Token 过期,需要刷新 token 重试(需求 4.1) let needs_token_refresh = matches!( &e, - crate::providers::ProviderError::AuthenticationError(_) - | crate::providers::ProviderError::TokenExpired(_) + proxycast_providers::providers::ProviderError::AuthenticationError(_) + | proxycast_providers::providers::ProviderError::TokenExpired(_) ); if needs_token_refresh { @@ -3109,7 +3114,7 @@ fn build_sse_response( /// ``` /// /// OpenAI SSE 格式: -/// ``` +/// ```text /// data: {"id":"chatcmpl-xxx","object":"chat.completion.chunk","choices":[{"index":0,"delta":{"content":"Hello"},"finish_reason":null}]} /// ``` fn convert_gemini_chunk_to_openai_sse(json: &serde_json::Value, model: &str) -> Option { @@ -3207,7 +3212,7 @@ fn convert_gemini_chunk_to_openai_sse(json: &serde_json::Value, model: &str) -> /// 将 OpenAI ChatCompletionResponse 转换为 Anthropic MessagesResponse 格式 fn convert_openai_response_to_anthropic( - openai_resp: &crate::models::openai::ChatCompletionResponse, + openai_resp: &proxycast_core::models::openai::ChatCompletionResponse, model: &str, ) -> serde_json::Value { // 提取第一个 choice 的内容 diff --git a/src-tauri/src/server/handlers/websocket.rs b/src-tauri/crates/server/src/handlers/websocket.rs similarity index 97% rename from src-tauri/src/server/handlers/websocket.rs rename to src-tauri/crates/server/src/handlers/websocket.rs index 2fe842caf..42b6da4a4 100644 --- a/src-tauri/src/server/handlers/websocket.rs +++ b/src-tauri/crates/server/src/handlers/websocket.rs @@ -16,20 +16,20 @@ use serde::Deserialize; use std::sync::Arc; use tokio::sync::Mutex; -use crate::converter::anthropic_to_openai::convert_anthropic_to_openai; -use crate::converter::openai_to_antigravity::{ +use crate::AppState; +use proxycast_core::models::anthropic::AnthropicMessagesRequest; +use proxycast_core::models::openai::ChatCompletionRequest; +use proxycast_core::models::provider_pool_model::ProviderCredential; +use proxycast_processor::RequestContext; +use proxycast_providers::converter::anthropic_to_openai::convert_anthropic_to_openai; +use proxycast_providers::converter::openai_to_antigravity::{ convert_antigravity_to_openai_response, convert_openai_to_antigravity_with_context, }; -use crate::models::anthropic::AnthropicMessagesRequest; -use crate::models::openai::ChatCompletionRequest; -use crate::models::provider_pool_model::ProviderCredential; -use crate::processor::RequestContext; -use crate::providers::{ +use proxycast_providers::providers::{ AntigravityProvider, ClaudeCustomProvider, KiroProvider, OpenAICustomProvider, }; -use crate::server::AppState; -use crate::server_utils::parse_cw_response; -use crate::websocket::{ +use proxycast_server_utils::parse_cw_response; +use proxycast_websocket::{ WsApiRequest, WsApiResponse, WsEndpoint, WsError, WsMessage as WsProtoMessage, }; @@ -447,7 +447,7 @@ pub async fn call_provider_openai_for_ws( credential: &ProviderCredential, request: &ChatCompletionRequest, ) -> Result { - use crate::models::provider_pool_model::CredentialData; + use proxycast_core::models::provider_pool_model::CredentialData; match &credential.credential { CredentialData::KiroOAuth { creds_file_path } => { @@ -726,7 +726,7 @@ pub async fn call_provider_anthropic_for_ws( credential: &ProviderCredential, request: &AnthropicMessagesRequest, ) -> Result { - use crate::models::provider_pool_model::CredentialData; + use proxycast_core::models::provider_pool_model::CredentialData; match &credential.credential { CredentialData::ClaudeKey { api_key, base_url } => { diff --git a/src-tauri/crates/server/src/lib.rs b/src-tauri/crates/server/src/lib.rs new file mode 100644 index 000000000..821603522 --- /dev/null +++ b/src-tauri/crates/server/src/lib.rs @@ -0,0 +1,1830 @@ +//! HTTP API 服务器 + +pub mod client_detector; + +use axum::{ + extract::{DefaultBodyLimit, Path, State}, + http::{HeaderMap, StatusCode}, + response::{IntoResponse, Response}, + routing::{get, post}, + Json, Router, +}; +use proxycast_core::config::{ + Config, ConfigChangeKind, ConfigManager, EndpointProvidersConfig, FileChangeEvent, FileWatcher, + HotReloadManager, ReloadResult, +}; +use proxycast_core::database::dao::provider_pool::ProviderPoolDao; +use proxycast_core::database::DbConnection; +use proxycast_core::logger::LogStore; +use proxycast_core::models::anthropic::*; +use proxycast_core::models::openai::*; +use proxycast_core::models::provider_pool_model::CredentialData; +use proxycast_core::models::route_model::{RouteInfo, RouteListResponse}; +use proxycast_credential::CredentialSyncService; +use proxycast_infra::injection::Injector; +use proxycast_processor::{RequestContext, RequestProcessor}; +use proxycast_providers::converter::anthropic_to_openai::convert_anthropic_to_openai; +use proxycast_providers::providers::antigravity::AntigravityProvider; +use proxycast_providers::providers::claude_custom::ClaudeCustomProvider; +use proxycast_providers::providers::gemini::GeminiProvider; +use proxycast_providers::providers::kiro::KiroProvider; +use proxycast_providers::providers::openai_custom::OpenAICustomProvider; +use proxycast_server_utils::{ + build_anthropic_response, build_anthropic_stream_response, build_error_response, + build_error_response_with_status, build_gemini_cli_request, build_gemini_native_request, + health, models, parse_cw_response, +}; +use proxycast_services::kiro_event_service::KiroEventService; +use proxycast_services::provider_pool_service::ProviderPoolService; +use proxycast_services::token_cache_service::TokenCacheService; +use proxycast_websocket::{WsConfig, WsConnectionManager, WsStats}; +use serde::{Deserialize, Serialize}; +use std::path::PathBuf; +use std::sync::Arc; +use tokio::sync::{oneshot, RwLock}; + +/// 记录请求统计到遥测系统 +pub fn record_request_telemetry( + state: &AppState, + ctx: &RequestContext, + status: proxycast_infra::telemetry::RequestStatus, + error_message: Option, +) { + use proxycast_infra::telemetry::RequestLog; + + let provider = ctx.provider.unwrap_or(proxycast_core::ProviderType::Kiro); + let mut log = RequestLog::new( + ctx.request_id.clone(), + provider, + ctx.resolved_model.clone(), + ctx.is_stream, + ); + + // 设置状态和持续时间 + match status { + proxycast_infra::telemetry::RequestStatus::Success => { + log.mark_success(ctx.elapsed_ms(), 200) + } + proxycast_infra::telemetry::RequestStatus::Failed => log.mark_failed( + ctx.elapsed_ms(), + None, + error_message.clone().unwrap_or_default(), + ), + proxycast_infra::telemetry::RequestStatus::Timeout => log.mark_timeout(ctx.elapsed_ms()), + proxycast_infra::telemetry::RequestStatus::Cancelled => { + log.mark_cancelled(ctx.elapsed_ms()) + } + proxycast_infra::telemetry::RequestStatus::Retrying => { + log.duration_ms = ctx.elapsed_ms(); + } + } + + // 设置凭证 ID + if let Some(cred_id) = &ctx.credential_id { + log.set_credential_id(cred_id.clone()); + } + + // 设置重试次数 + log.retry_count = ctx.retry_count; + + // 记录到统计聚合器 + { + let stats = state.processor.stats.write(); + stats.record(log.clone()); + } + + // 记录到请求日志记录器(用于前端日志列表显示) + if let Some(logger) = &state.request_logger { + let _ = logger.record(log.clone()); + } + + tracing::info!( + "[TELEMETRY] request_id={} provider={:?} model={} status={:?} duration_ms={}", + ctx.request_id, + provider, + ctx.resolved_model, + status, + ctx.elapsed_ms() + ); +} + +/// 记录 Token 使用量到遥测系统 +pub fn record_token_usage( + state: &AppState, + ctx: &RequestContext, + input_tokens: Option, + output_tokens: Option, +) { + use proxycast_infra::telemetry::{TokenSource, TokenUsageRecord}; + + // 只有当至少有一个 Token 值时才记录 + if input_tokens.is_none() && output_tokens.is_none() { + return; + } + + let provider = ctx.provider.unwrap_or(proxycast_core::ProviderType::Kiro); + let record = TokenUsageRecord::new( + uuid::Uuid::new_v4().to_string(), + provider, + ctx.resolved_model.clone(), + input_tokens.unwrap_or(0), + output_tokens.unwrap_or(0), + TokenSource::Actual, + ) + .with_request_id(ctx.request_id.clone()); + + // 记录到 Token 追踪器 + { + let tokens = state.processor.tokens.write(); + tokens.record(record); + } + + tracing::debug!( + "[TOKEN] request_id={} input={} output={}", + ctx.request_id, + input_tokens.unwrap_or(0), + output_tokens.unwrap_or(0) + ); +} + +#[derive(Debug, Clone, Serialize, Deserialize)] +pub struct ServerStatus { + pub running: bool, + pub host: String, + pub port: u16, + pub requests: u64, + pub uptime_secs: u64, +} + +pub struct ServerState { + pub config: Config, + pub running: bool, + pub requests: u64, + pub start_time: Option, + pub kiro_provider: KiroProvider, + pub gemini_provider: GeminiProvider, + pub openai_custom_provider: OpenAICustomProvider, + pub claude_custom_provider: ClaudeCustomProvider, + pub default_provider_ref: Arc>, + /// 路由器引用(用于动态更新默认 Provider) + pub router_ref: Option>>, + shutdown_tx: Option>, + /// 服务器运行时使用的 API key(启动时从配置复制) + /// 用于 test_api 命令,确保测试使用的 API key 和服务器一致 + pub running_api_key: Option, + /// 服务器实际监听的 host(可能与配置不同,因为会自动切换到有效的 IP) + pub running_host: Option, +} + +impl ServerState { + pub fn new(config: Config) -> Self { + let kiro = KiroProvider::new(); + let gemini = GeminiProvider::new(); + let openai_custom = OpenAICustomProvider::new(); + let claude_custom = ClaudeCustomProvider::new(); + let default_provider_ref = Arc::new(RwLock::new(config.default_provider.clone())); + + Self { + config, + running: false, + requests: 0, + start_time: None, + kiro_provider: kiro, + gemini_provider: gemini, + openai_custom_provider: openai_custom, + claude_custom_provider: claude_custom, + default_provider_ref, + router_ref: None, + shutdown_tx: None, + running_api_key: None, + running_host: None, + } + } + + pub fn status(&self) -> ServerStatus { + ServerStatus { + running: self.running, + // 使用实际运行的 host,如果没有则使用配置的 host + host: self + .running_host + .clone() + .unwrap_or_else(|| self.config.server.host.clone()), + port: self.config.server.port, + requests: self.requests, + uptime_secs: self.start_time.map(|t| t.elapsed().as_secs()).unwrap_or(0), + } + } + + /// 增加请求计数 + pub fn increment_request_count(&mut self) { + self.requests = self.requests.saturating_add(1); + } + + /// 解析绑定地址 + /// + /// 直接返回用户配置的地址,不做任何自动替换。 + /// 如果地址无效,绑定时会失败并返回错误。 + fn resolve_bind_host(&self, configured_host: &str) -> String { + tracing::info!("[SERVER] 使用配置的监听地址: {}", configured_host); + configured_host.to_string() + } + + pub async fn start( + &mut self, + logs: Arc>, + pool_service: Arc, + token_cache: Arc, + db: Option, + ) -> Result<(), Box> { + self.start_with_telemetry(logs, pool_service, token_cache, db, None, None, None) + .await + } + + /// 启动服务器(使用共享的遥测实例) + /// + /// 这允许服务器与 TelemetryState 共享同一个 StatsAggregator、TokenTracker 和 RequestLogger, + /// 使得请求处理过程中记录的统计数据能够在前端监控页面中显示。 + pub async fn start_with_telemetry( + &mut self, + logs: Arc>, + pool_service: Arc, + token_cache: Arc, + db: Option, + shared_stats: Option>>, + shared_tokens: Option>>, + shared_logger: Option>, + ) -> Result<(), Box> { + self.start_with_telemetry_and_flow_monitor( + logs, + pool_service, + token_cache, + db, + shared_stats, + shared_tokens, + shared_logger, + ) + .await + } + + /// 启动服务器(使用共享的遥测实例) + /// + /// 这允许服务器与 TelemetryState 共享同一个 StatsAggregator、TokenTracker 和 RequestLogger, + /// 使得请求处理过程中记录的统计数据能够在前端监控页面中显示。 + pub async fn start_with_telemetry_and_flow_monitor( + &mut self, + logs: Arc>, + pool_service: Arc, + token_cache: Arc, + db: Option, + shared_stats: Option>>, + shared_tokens: Option>>, + shared_logger: Option>, + ) -> Result<(), Box> { + if self.running { + return Ok(()); + } + + let (tx, rx) = oneshot::channel(); + self.shutdown_tx = Some(tx); + + // 智能选择监听地址 + // - 127.0.0.1, localhost, 0.0.0.0, :: 直接使用 + // - 局域网 IP:检查是否在当前网卡列表中,如果不在则自动切换到当前局域网 IP + let configured_host = self.config.server.host.clone(); + let host = self.resolve_bind_host(&configured_host); + + // 如果地址发生了变化,记录日志 + if host != configured_host { + tracing::warn!( + "[SERVER] 配置的监听地址 {} 不可用,自动切换到 {}", + configured_host, + host + ); + } + + let port = self.config.server.port; + let api_key = self.config.server.api_key.clone(); + let api_key_for_state = api_key.clone(); // 用于保存到 running_api_key + let default_provider_ref = self.default_provider_ref.clone(); + + // 重新加载凭证 + let _ = self.kiro_provider.load_credentials().await; + let kiro = self.kiro_provider.clone(); + + // 创建参数注入器 + let injection_enabled = self.config.injection.enabled; + let injector = Injector::with_rules( + self.config + .injection + .rules + .iter() + .map(|r| r.clone().into()) + .collect(), + ); + + // 获取配置和配置路径用于热重载 + let config = self.config.clone(); + let config_path = proxycast_core::config::ConfigManager::default_config_path(); + + // 创建请求处理器(在 spawn 之前创建,以便保存 router_ref) + let processor = match (&shared_stats, &shared_tokens) { + (Some(stats), Some(tokens)) => Arc::new(RequestProcessor::with_shared_telemetry( + pool_service.clone(), + stats.clone(), + tokens.clone(), + )), + _ => Arc::new(RequestProcessor::with_defaults(pool_service.clone())), + }; + + // 从配置初始化 Router 的默认 Provider + { + let default_provider_str = &config.routing.default_provider; + + // 尝试解析为 ProviderType 枚举 + match default_provider_str.parse::() { + Ok(provider_type) => { + let mut router = processor.router.write().await; + router.set_default_provider(provider_type); + tracing::info!( + "[SERVER] 从配置初始化 Router 默认 Provider: {} (ProviderType)", + default_provider_str + ); + } + Err(_) => { + // 如果解析失败,可能是自定义 provider ID + // 这种情况下,路由器保持空状态,请求会直接使用 provider_id 进行凭证查找 + tracing::warn!( + "[SERVER] 配置的默认 Provider '{}' 不是有效的 ProviderType 枚举值,可能是自定义 Provider ID。\ + 路由器将保持空状态,请求将直接使用 provider_id 进行凭证查找。", + default_provider_str + ); + eprintln!( + "[SERVER] 警告:默认 Provider '{default_provider_str}' 不是标准 Provider 类型(kiro/openai/claude等),\ + 可能是自定义 Provider ID。如果这是预期行为,请忽略此警告。" + ); + } + } + } + + // 保存 router_ref 以便后续动态更新 + self.router_ref = Some(processor.router.clone()); + + // 保存实际使用的 host(在移动到 spawn 之前克隆) + let running_host = host.clone(); + + tokio::spawn(async move { + if let Err(e) = run_server( + &host, + port, + &api_key, + default_provider_ref, + kiro, + logs, + rx, + pool_service, + token_cache, + db, + injector, + injection_enabled, + shared_stats, + shared_tokens, + shared_logger, + Some(config), + Some(config_path), + Some(processor), + None, // dev_bridge_callback: 由主 crate 在重新导出层注入 + ) + .await + { + tracing::error!("Server error: {}", e); + } + }); + + self.running = true; + self.start_time = Some(std::time::Instant::now()); + // 保存服务器运行时使用的 API key,用于 test_api 命令 + self.running_api_key = Some(api_key_for_state); + // 保存服务器实际监听的 host(可能与配置不同) + self.running_host = Some(running_host); + Ok(()) + } + + pub async fn stop(&mut self) { + if let Some(tx) = self.shutdown_tx.take() { + let _ = tx.send(()); + } + self.running = false; + self.start_time = None; + self.running_api_key = None; + self.running_host = None; + self.router_ref = None; + } +} + +pub mod handlers; + +#[derive(Clone)] +#[allow(dead_code)] +pub struct AppState { + pub api_key: String, + pub base_url: String, + pub default_provider: Arc>, + pub kiro: Arc>, + pub logs: Arc>, + pub kiro_refresh_lock: Arc>, + pub gemini_refresh_lock: Arc>, + pub pool_service: Arc, + pub token_cache: Arc, + pub db: Option, + /// 参数注入器 + pub injector: Arc>, + /// 是否启用参数注入 + pub injection_enabled: Arc>, + /// 请求处理器 + pub processor: Arc, + /// WebSocket 连接管理器 + pub ws_manager: Arc, + /// WebSocket 统计信息 + pub ws_stats: Arc, + /// 热重载管理器 + pub hot_reload_manager: Option>, + /// 请求日志记录器(与 TelemetryState 共享) + pub request_logger: Option>, + /// Amp CLI 路由器 + pub amp_router: Arc, + /// 端点 Provider 配置 + pub endpoint_providers: Arc>, + /// Kiro 事件服务 + pub kiro_event_service: Arc, + /// API Key Provider 服务(用于智能降级) + pub api_key_service: Arc, +} + +/// 启动配置文件监控 +/// +/// 监控配置文件变化并触发热重载。 +/// +/// # 连接保持 +/// +/// 热重载过程不会中断现有连接: +/// - 配置更新在独立的 tokio 任务中异步执行 +/// - 使用 RwLock 进行原子性更新,不会阻塞正在处理的请求 +/// - 服务器继续运行,不需要重启 +/// - HTTP 和 WebSocket 连接保持活跃 +async fn start_config_watcher( + config_path: PathBuf, + hot_reload_manager: Option>, + processor: Arc, + logs: Arc>, + db: Option, + config_manager: Option>>, +) -> Option { + let (tx, mut rx) = tokio::sync::mpsc::unbounded_channel::(); + + // 创建文件监控器 + let mut watcher = match FileWatcher::new(&config_path, tx) { + Ok(w) => w, + Err(e) => { + tracing::error!("[HOT_RELOAD] 创建文件监控器失败: {}", e); + return None; + } + }; + + // 启动监控 + if let Err(e) = watcher.start() { + tracing::error!("[HOT_RELOAD] 启动文件监控失败: {}", e); + return None; + } + + tracing::info!("[HOT_RELOAD] 配置文件监控已启动: {:?}", config_path); + + // 启动事件处理任务 + let hot_reload_manager_clone = hot_reload_manager.clone(); + let processor_clone = processor.clone(); + let logs_clone = logs.clone(); + let db_clone = db.clone(); + let config_manager_clone = config_manager.clone(); + + tokio::spawn(async move { + while let Some(event) = rx.recv().await { + // 只处理修改事件 + if event.kind != ConfigChangeKind::Modified { + continue; + } + + tracing::info!("[HOT_RELOAD] 检测到配置文件变更: {:?}", event.path); + logs_clone.write().await.add( + "info", + &format!("[HOT_RELOAD] 检测到配置文件变更: {:?}", event.path), + ); + + // 执行热重载 + if let Some(ref manager) = hot_reload_manager_clone { + let result = manager.reload(); + match &result { + ReloadResult::Success { .. } => { + tracing::info!("[HOT_RELOAD] 配置热重载成功"); + logs_clone + .write() + .await + .add("info", "[HOT_RELOAD] 配置热重载成功"); + + // 更新处理器中的组件 + let new_config = manager.config(); + update_processor_config(&processor_clone, &new_config).await; + + // 同步凭证池 + if let (Some(ref db), Some(ref cfg_manager)) = + (&db_clone, &config_manager_clone) + { + match sync_credential_pool_from_config(db, cfg_manager, &logs_clone) + .await + { + Ok(count) => { + tracing::info!( + "[HOT_RELOAD] 凭证池同步完成,共 {} 个凭证", + count + ); + logs_clone.write().await.add( + "info", + &format!("[HOT_RELOAD] 凭证池同步完成,共 {count} 个凭证"), + ); + } + Err(e) => { + tracing::warn!("[HOT_RELOAD] 凭证池同步失败: {}", e); + logs_clone + .write() + .await + .add("warn", &format!("[HOT_RELOAD] 凭证池同步失败: {e}")); + } + } + } + } + ReloadResult::RolledBack { error, .. } => { + tracing::warn!("[HOT_RELOAD] 配置热重载失败,已回滚: {}", error); + logs_clone.write().await.add( + "warn", + &format!("[HOT_RELOAD] 配置热重载失败,已回滚: {error}"), + ); + } + ReloadResult::Failed { + error, + rollback_error, + .. + } => { + tracing::error!( + "[HOT_RELOAD] 配置热重载失败: {}, 回滚错误: {:?}", + error, + rollback_error + ); + logs_clone.write().await.add( + "error", + &format!( + "[HOT_RELOAD] 配置热重载失败: {error}, 回滚错误: {rollback_error:?}" + ), + ); + } + } + } + } + }); + + Some(watcher) +} + +/// 更新处理器配置 +/// +/// 当配置热重载成功后,更新 RequestProcessor 中的各个组件。 +/// +/// # 原子性更新 +/// +/// 每个组件的更新都是原子性的,使用 RwLock 确保: +/// - 正在处理的请求不会看到部分更新的状态 +/// - 更新过程不会阻塞新请求的处理 +/// - 现有连接不受影响 +async fn update_processor_config(processor: &RequestProcessor, config: &Config) { + // 更新注入器规则 + { + let mut injector = processor.injector.write().await; + injector.clear(); + for rule in &config.injection.rules { + injector.add_rule(rule.clone().into()); + } + tracing::debug!( + "[HOT_RELOAD] 注入器规则已更新: {} 条规则", + config.injection.rules.len() + ); + } + + // 更新路由器默认 Provider + { + let mut router = processor.router.write().await; + + // 尝试解析为 ProviderType 枚举 + match config + .routing + .default_provider + .parse::() + { + Ok(provider_type) => { + router.set_default_provider(provider_type); + tracing::debug!( + "[HOT_RELOAD] 路由器默认 Provider 已更新: {} (ProviderType)", + config.routing.default_provider + ); + } + Err(_) => { + // 如果解析失败,可能是自定义 provider ID + // 清空路由器的默认 provider,让请求直接使用 provider_id + tracing::warn!( + "[HOT_RELOAD] 配置的默认 Provider '{}' 不是有效的 ProviderType 枚举值,可能是自定义 Provider ID。\ + 路由器默认 Provider 将被清空。", + config.routing.default_provider + ); + } + } + } + + // 更新模型映射器 + { + let mut mapper = processor.mapper.write().await; + mapper.clear(); + for (alias, model) in &config.routing.model_aliases { + mapper.add_alias(alias, model); + } + tracing::debug!( + "[HOT_RELOAD] 模型别名已更新: {} 个别名", + config.routing.model_aliases.len() + ); + } + + // 注意:重试配置目前不支持热更新,因为 Retrier 是不可变的 + // 如果需要更新重试配置,需要重启服务器 + tracing::debug!( + "[HOT_RELOAD] 重试配置: max_retries={}, base_delay={}ms (需重启生效)", + config.retry.max_retries, + config.retry.base_delay_ms + ); + + tracing::info!("[HOT_RELOAD] 处理器配置更新完成"); +} + +/// 从配置同步凭证池 +/// +/// 当配置热重载成功后,从 YAML 配置中加载凭证并同步到数据库。 +/// +/// # 同步策略 +/// +/// - 从配置中加载所有凭证 +/// - 对于配置中存在但数据库中不存在的凭证,添加到数据库 +/// - 对于配置中存在且数据库中也存在的凭证,更新数据库中的记录 +/// - 对于数据库中存在但配置中不存在的凭证,保留(不删除,避免丢失运行时状态) +async fn sync_credential_pool_from_config( + db: &DbConnection, + config_manager: &Arc>, + _logs: &Arc>, +) -> Result { + // 创建凭证同步服务 + let sync_service = CredentialSyncService::new(config_manager.clone()); + + // 从配置加载凭证 + let credentials = sync_service.load_from_config().map_err(|e| e.to_string())?; + + let conn = proxycast_core::database::lock_db(db)?; + let mut synced_count = 0; + + for cred in &credentials { + // 检查凭证是否已存在 + let existing = + ProviderPoolDao::get_by_uuid(&conn, &cred.uuid).map_err(|e| e.to_string())?; + + if existing.is_some() { + // 更新现有凭证 + ProviderPoolDao::update(&conn, cred).map_err(|e| e.to_string())?; + tracing::debug!( + "[HOT_RELOAD] 更新凭证: {} ({})", + cred.uuid, + cred.provider_type + ); + } else { + // 添加新凭证 + ProviderPoolDao::insert(&conn, cred).map_err(|e| e.to_string())?; + tracing::debug!( + "[HOT_RELOAD] 添加凭证: {} ({})", + cred.uuid, + cred.provider_type + ); + } + synced_count += 1; + } + + Ok(synced_count) +} + +/// 开发桥接启动回调类型 +pub type DevBridgeCallback = Box; + +async fn run_server( + host: &str, + port: u16, + api_key: &str, + default_provider: Arc>, + kiro: KiroProvider, + logs: Arc>, + shutdown: oneshot::Receiver<()>, + pool_service: Arc, + token_cache: Arc, + db: Option, + injector: Injector, + injection_enabled: bool, + shared_stats: Option>>, + shared_tokens: Option>>, + shared_logger: Option>, + config: Option, + config_path: Option, + processor: Option>, + dev_bridge_callback: Option, +) -> Result<(), Box> { + let base_url = format!("http://{host}:{port}"); + + // 使用传入的 processor 或创建新的 + let processor = match processor { + Some(p) => p, + None => match (&shared_stats, &shared_tokens) { + (Some(stats), Some(tokens)) => Arc::new(RequestProcessor::with_shared_telemetry( + pool_service.clone(), + stats.clone(), + tokens.clone(), + )), + _ => Arc::new(RequestProcessor::with_defaults(pool_service.clone())), + }, + }; + + // 将注入器规则同步到处理器 + { + let mut proc_injector = processor.injector.write().await; + for rule in injector.rules() { + proc_injector.add_rule(rule.clone()); + } + } + + // 从配置初始化 Router 的默认 Provider + if let Some(cfg) = &config { + let default_provider_str = &cfg.routing.default_provider; + + // 尝试解析为 ProviderType 枚举 + match default_provider_str.parse::() { + Ok(provider_type) => { + let mut router = processor.router.write().await; + router.set_default_provider(provider_type); + tracing::info!( + "[SERVER] 从配置初始化 Router 默认 Provider: {} (ProviderType)", + default_provider_str + ); + } + Err(_) => { + // 如果解析失败,可能是自定义 provider ID + tracing::warn!( + "[SERVER] 配置的默认 Provider '{}' 不是有效的 ProviderType 枚举值,可能是自定义 Provider ID。\ + 路由器将保持空状态,请求将直接使用 provider_id 进行凭证查找。", + default_provider_str + ); + eprintln!( + "[SERVER] 警告:默认 Provider '{default_provider_str}' 不是标准 Provider 类型,可能是自定义 Provider ID" + ); + } + } + } + + // 初始化 WebSocket 管理器 + let ws_manager = Arc::new(WsConnectionManager::new(WsConfig::default())); + let ws_stats = ws_manager.stats().clone(); + + // 初始化热重载管理器 + let hot_reload_manager = match (&config, &config_path) { + (Some(cfg), Some(path)) => Some(Arc::new(HotReloadManager::new(cfg.clone(), path.clone()))), + _ => None, + }; + + // 初始化配置管理器(用于凭证池同步) + let config_manager: Option>> = + match (&config, &config_path) { + (Some(cfg), Some(path)) => Some(Arc::new(std::sync::RwLock::new( + ConfigManager::with_config(cfg.clone(), path.clone()), + ))), + _ => None, + }; + + let logs_clone = logs.clone(); + let db_clone = db.clone(); + + // 初始化 Amp CLI 路由器 + let amp_router = Arc::new(proxycast_core::router::AmpRouter::new( + config + .as_ref() + .map(|c| c.ampcode.clone()) + .unwrap_or_default(), + )); + + // 初始化端点 Provider 配置 + let endpoint_providers = Arc::new(RwLock::new( + config + .as_ref() + .map(|c| c.endpoint_providers.clone()) + .unwrap_or_default(), + )); + + // 创建 Kiro 事件服务 + let kiro_event_service = Arc::new(KiroEventService::new()); + + // 创建 API Key Provider 服务 + let api_key_service = + Arc::new(proxycast_services::api_key_provider_service::ApiKeyProviderService::new()); + + let state = AppState { + api_key: api_key.to_string(), + base_url, + default_provider, + kiro: Arc::new(RwLock::new(kiro)), + logs, + kiro_refresh_lock: Arc::new(tokio::sync::Mutex::new(())), + gemini_refresh_lock: Arc::new(tokio::sync::Mutex::new(())), + pool_service, + token_cache, + db, + injector: Arc::new(RwLock::new(injector)), + injection_enabled: Arc::new(RwLock::new(injection_enabled)), + processor: processor.clone(), + ws_manager, + ws_stats, + hot_reload_manager: hot_reload_manager.clone(), + request_logger: shared_logger, + amp_router, + endpoint_providers, + kiro_event_service, + api_key_service, + }; + + // ========== 开发模式:通过回调启动桥接服务器 ========== + if let Some(callback) = dev_bridge_callback { + callback(state.clone()); + } + + // 启动配置文件监控 + let _file_watcher = if let Some(path) = config_path { + start_config_watcher( + path, + hot_reload_manager, + processor, + logs_clone, + db_clone, + config_manager, + ) + .await + } else { + None + }; + + // 设置请求体大小限制为 100MB,支持大型上下文请求(如 Claude Code 的 /compact 命令) + let body_limit = 100 * 1024 * 1024; // 100MB + + // 创建管理 API 路由(带认证中间件) + let management_config = config + .as_ref() + .map(|c| c.remote_management.clone()) + .unwrap_or_default(); + + let management_routes = Router::new() + .route("/v0/management/status", get(handlers::management_status)) + .route( + "/v0/management/credentials", + get(handlers::management_list_credentials), + ) + .route( + "/v0/management/credentials", + post(handlers::management_add_credential), + ) + .route( + "/v0/management/config", + get(handlers::management_get_config), + ) + .route( + "/v0/management/config", + axum::routing::put(handlers::management_update_config), + ) + .layer(proxycast_core::middleware::ManagementAuthLayer::new( + management_config, + )); + + // Kiro凭证管理API路由 + let kiro_api_routes = Router::new() + .route( + "/api/kiro/credentials/available", + get(handlers::get_available_credentials), + ) + .route( + "/api/kiro/credentials/select", + post(handlers::select_credential), + ) + .route( + "/api/kiro/credentials/{uuid}/refresh", + axum::routing::put(handlers::refresh_credential), + ) + .route( + "/api/kiro/credentials/{uuid}/status", + get(handlers::get_credential_status), + ); + + // 凭证 API 路由(用于 aster Agent 集成) + let credentials_api_routes = Router::new() + .route("/v1/credentials/select", post(handlers::credentials_select)) + .route( + "/v1/credentials/{uuid}/token", + get(handlers::credentials_get_token), + ); + + let app = Router::new() + .route("/health", get(health)) + .route("/v1/models", get(models)) + .route("/v1/routes", get(list_routes)) + .route("/v1/chat/completions", post( + |State(state): State, + headers: HeaderMap, + Json(request): Json| async { + handlers::chat_completions(State(state), headers, Json(request)).await + } + )) + .route("/v1/messages", post( + |State(state): State, + headers: HeaderMap, + Json(request): Json| async { + handlers::anthropic_messages(State(state), headers, Json(request)).await + } + )) + .route("/v1/messages/count_tokens", post(count_tokens)) + // 图像生成 API 路由 + .route( + "/v1/images/generations", + post(handlers::handle_image_generation), + ) + // WebSocket 路由 + .route("/v1/ws", get(handlers::ws_upgrade_handler)) + .route("/ws", get(handlers::ws_upgrade_handler)) + // 多供应商路由 + .route( + "/{selector}/v1/messages", + post(anthropic_messages_with_selector), + ) + .route( + "/{selector}/v1/chat/completions", + post(chat_completions_with_selector), + ) + // 管理 API 路由 + .merge(management_routes) + // Kiro凭证管理API路由 + .merge(kiro_api_routes) + // 凭证 API 路由(用于 aster Agent 集成) + .merge(credentials_api_routes) + .layer(DefaultBodyLimit::max(body_limit)) + .with_state(state); + + let addr: std::net::SocketAddr = format!("{host}:{port}") + .parse() + .map_err(|e| format!("无效的监听地址 {host}:{port} - {e}"))?; + + let listener = tokio::net::TcpListener::bind(addr).await.map_err(|e| { + format!("无法绑定到 {host}:{port},错误: {e}。请检查地址是否有效或端口是否被占用。") + })?; + + tracing::info!("Server listening on {}", addr); + + axum::serve(listener, app) + .with_graceful_shutdown(async move { + let _ = shutdown.await; + }) + .await?; + + Ok(()) +} + +async fn count_tokens( + State(state): State, + headers: HeaderMap, + Json(_request): Json, +) -> Response { + if let Err(e) = handlers::verify_api_key(&headers, &state.api_key).await { + return e.into_response(); + } + + // Claude Code 需要这个端点,返回估算值 + Json(serde_json::json!({ + "input_tokens": 100 + })) + .into_response() +} + +/// Gemini 原生协议处理 +/// 路由: POST /v1/gemini/{model}:{method} +/// 例如: /v1/gemini/gemini-3-pro-preview:generateContent +#[allow(dead_code)] +async fn gemini_generate_content( + State(state): State, + headers: HeaderMap, + Path(path): Path, + Json(request): Json, +) -> Response { + if let Err(e) = handlers::verify_api_key(&headers, &state.api_key).await { + return e.into_response(); + } + + // 解析路径: {model}:{method} + // 例如: gemini-3-pro-preview:generateContent + let parts: Vec<&str> = path.splitn(2, ':').collect(); + if parts.len() != 2 { + return ( + StatusCode::BAD_REQUEST, + Json(serde_json::json!({ + "error": { + "message": format!("无效的路径格式: {},期望格式: model:method", path) + } + })), + ) + .into_response(); + } + + let model = parts[0]; + let method = parts[1]; + + state.logs.write().await.add( + "info", + &format!("[GEMINI] POST /v1/gemini/{path} model={model} method={method}"), + ); + + // 目前只支持 generateContent 方法 + if method != "generateContent" && method != "streamGenerateContent" { + return ( + StatusCode::BAD_REQUEST, + Json(serde_json::json!({ + "error": { + "message": format!("不支持的方法: {},目前只支持 generateContent", method) + } + })), + ) + .into_response(); + } + + let is_stream = method == "streamGenerateContent"; + + // 获取默认 provider + let default_provider = state.default_provider.read().await.clone(); + + // 尝试从凭证池中选择凭证(不降级,指定什么就用什么) + let credential = match &state.db { + Some(db) => state + .pool_service + .select_credential(db, &default_provider, Some(model)) + .ok() + .flatten(), + None => None, + }; + + let cred = match credential { + Some(c) => c, + None => { + return ( + StatusCode::SERVICE_UNAVAILABLE, + Json(serde_json::json!({ + "error": { + "message": format!("No available credentials for provider '{}'. Please add credentials in the Provider Pool.", default_provider) + } + })), + ) + .into_response(); + } + }; + + state.logs.write().await.add( + "info", + &format!( + "[GEMINI] 使用凭证: type={} name={:?} uuid={}", + cred.provider_type, + cred.name, + &cred.uuid[..8] + ), + ); + + // 调用 Antigravity Provider + match &cred.credential { + CredentialData::AntigravityOAuth { + creds_file_path, + project_id, + } => { + let mut antigravity = AntigravityProvider::new(); + if let Err(e) = antigravity + .load_credentials_from_path(creds_file_path) + .await + { + return ( + StatusCode::INTERNAL_SERVER_ERROR, + Json(serde_json::json!({ + "error": { + "message": format!("加载 Antigravity 凭证失败: {}", e) + } + })), + ) + .into_response(); + } + + // 使用新的 validate_token() 方法检查 Token 状态 + let validation_result = antigravity.validate_token(); + tracing::info!( + "[Antigravity Gemini] Token 验证结果: {:?}", + validation_result + ); + + // 根据验证结果决定是否刷新 + if validation_result.needs_refresh() { + tracing::info!("[Antigravity Gemini] Token 需要刷新,开始刷新..."); + match antigravity.refresh_token_with_retry(3).await { + Ok(new_token) => { + tracing::info!( + "[Antigravity Gemini] Token 刷新成功,新 token 长度: {}", + new_token.len() + ); + } + Err(refresh_error) => { + tracing::error!("[Antigravity Gemini] Token 刷新失败: {:?}", refresh_error); + + // 根据错误类型返回不同的状态码和消息 + let (status, message) = if refresh_error.requires_reauth() { + (StatusCode::UNAUTHORIZED, refresh_error.user_message()) + } else { + ( + StatusCode::INTERNAL_SERVER_ERROR, + refresh_error.user_message(), + ) + }; + + return ( + status, + Json(serde_json::json!({ + "error": { + "message": message + } + })), + ) + .into_response(); + } + } + } + + // 设置项目 ID + if let Some(pid) = project_id { + antigravity.project_id = Some(pid.clone()); + } else if antigravity.project_id.is_none() { + // 如果凭证中没有 project_id,尝试从 API 获取或生成随机 ID + if let Err(e) = antigravity.discover_project().await { + tracing::warn!("[Antigravity] 获取项目 ID 失败: {},使用随机生成的 ID", e); + // 生成随机项目 ID + let uuid = uuid::Uuid::new_v4(); + let bytes = uuid.as_bytes(); + let adjectives = ["useful", "bright", "swift", "calm", "bold"]; + let nouns = ["fuze", "wave", "spark", "flow", "core"]; + let adj = adjectives[(bytes[0] as usize) % adjectives.len()]; + let noun = nouns[(bytes[1] as usize) % nouns.len()]; + let random_part: String = uuid.to_string()[..5].to_lowercase(); + antigravity.project_id = Some(format!("{adj}-{noun}-{random_part}")); + } + } + + let proj_id = antigravity.project_id.clone().unwrap_or_else(|| { + // 最后的后备:生成随机 ID + let uuid = uuid::Uuid::new_v4(); + format!("proxycast-{}", &uuid.to_string()[..8]) + }); + + state + .logs + .write() + .await + .add("debug", &format!("[GEMINI] 使用 project_id: {proj_id}")); + + // 构建 Antigravity 请求体 + // 直接使用用户传入的 Gemini 格式请求,只添加必要的字段 + let antigravity_request = build_gemini_native_request(&request, model, &proj_id); + + state.logs.write().await.add( + "debug", + &format!( + "[GEMINI] 请求体: {}", + serde_json::to_string(&antigravity_request).unwrap_or_default() + ), + ); + + if is_stream { + // 流式响应 - 暂不支持,返回错误 + return ( + StatusCode::NOT_IMPLEMENTED, + Json(serde_json::json!({ + "error": { + "message": "流式响应暂不支持,请使用 generateContent" + } + })), + ) + .into_response(); + } + + // 非流式响应 + match antigravity + .call_api("generateContent", &antigravity_request) + .await + { + Ok(resp) => { + state.logs.write().await.add( + "info", + &format!( + "[GEMINI] 响应成功: {}", + serde_json::to_string(&resp) + .unwrap_or_default() + .chars() + .take(200) + .collect::() + ), + ); + + // 直接返回 Gemini 格式响应 + Json(resp).into_response() + } + Err(api_err) => { + state.logs.write().await.add( + "error", + &format!( + "[GEMINI] 请求失败 (HTTP {}): {}", + api_err.status_code, api_err.message + ), + ); + + // 直接使用 AntigravityApiError 的状态码构建响应 + build_error_response_with_status(api_err.status_code, &api_err.to_string()) + } + } + } + CredentialData::GeminiOAuth { + creds_file_path, + project_id, + } => { + // 使用 GeminiProvider 处理 Gemini CLI OAuth 凭证 + let mut gemini = GeminiProvider::new(); + if let Err(e) = gemini.load_credentials_from_path(creds_file_path).await { + return ( + StatusCode::INTERNAL_SERVER_ERROR, + Json(serde_json::json!({ + "error": { + "message": format!("加载 Gemini 凭证失败: {}", e) + } + })), + ) + .into_response(); + } + + // 检查并刷新 Token + if !gemini.is_token_valid() { + tracing::info!("[Gemini CLI] Token 需要刷新,开始刷新..."); + match gemini.refresh_token_with_retry(3).await { + Ok(new_token) => { + tracing::info!( + "[Gemini CLI] Token 刷新成功,新 token 长度: {}", + new_token.len() + ); + } + Err(refresh_error) => { + tracing::error!("[Gemini CLI] Token 刷新失败: {:?}", refresh_error); + return ( + StatusCode::UNAUTHORIZED, + Json(serde_json::json!({ + "error": { + "message": format!("Token 刷新失败: {}", refresh_error) + } + })), + ) + .into_response(); + } + } + } + + // 设置项目 ID + if let Some(pid) = project_id { + gemini.project_id = Some(pid.clone()); + } else if gemini.project_id.is_none() { + // 尝试从 API 获取项目 ID + if let Err(e) = gemini.discover_project().await { + tracing::warn!("[Gemini CLI] 获取项目 ID 失败: {},使用随机生成的 ID", e); + let uuid = uuid::Uuid::new_v4(); + let bytes = uuid.as_bytes(); + let adjectives = ["useful", "bright", "swift", "calm", "bold"]; + let nouns = ["fuze", "wave", "spark", "flow", "core"]; + let adj = adjectives[(bytes[0] as usize) % adjectives.len()]; + let noun = nouns[(bytes[1] as usize) % nouns.len()]; + let random_part: String = uuid.to_string()[..5].to_lowercase(); + gemini.project_id = Some(format!("{adj}-{noun}-{random_part}")); + } + } + + let proj_id = gemini.project_id.clone().unwrap_or_else(|| { + let uuid = uuid::Uuid::new_v4(); + format!("proxycast-{}", &uuid.to_string()[..8]) + }); + + state + .logs + .write() + .await + .add("debug", &format!("[GEMINI CLI] 使用 project_id: {proj_id}")); + + // 构建 Gemini CLI 请求体 + // Gemini CLI 使用 Cloud Code Assist 端点,不做模型名称映射 + let gemini_request = build_gemini_cli_request(&request, model, &proj_id); + + state.logs.write().await.add( + "debug", + &format!( + "[GEMINI CLI] 请求体: {}", + serde_json::to_string(&gemini_request).unwrap_or_default() + ), + ); + + if is_stream { + // 流式响应 - 暂不支持 + return ( + StatusCode::NOT_IMPLEMENTED, + Json(serde_json::json!({ + "error": { + "message": "Gemini CLI 流式响应暂不支持,请使用 generateContent" + } + })), + ) + .into_response(); + } + + // 非流式响应 + match gemini.call_api("generateContent", &gemini_request).await { + Ok(resp) => { + state.logs.write().await.add( + "info", + &format!( + "[GEMINI CLI] 响应成功: {}", + serde_json::to_string(&resp) + .unwrap_or_default() + .chars() + .take(200) + .collect::() + ), + ); + + // 直接返回 Gemini 格式响应 + Json(resp).into_response() + } + Err(api_err) => { + state + .logs + .write() + .await + .add("error", &format!("[GEMINI CLI] 请求失败: {api_err}")); + + build_error_response(&api_err.to_string()) + } + } + } + _ => ( + StatusCode::BAD_REQUEST, + Json(serde_json::json!({ + "error": { + "message": "Gemini 原生协议只支持 Antigravity 或 Gemini CLI OAuth 凭证" + } + })), + ) + .into_response(), + } +} + +/// 列出所有可用路由 +async fn list_routes(State(state): State) -> impl IntoResponse { + // 处理 base_url:检查 IP 是否有效(在当前网卡列表中或是特殊地址) + let display_base_url = { + // 从 base_url 中提取 host 部分 + let url_parts: Vec<&str> = state.base_url.split("://").collect(); + let host_port = if url_parts.len() > 1 { + url_parts[1] + } else { + &state.base_url + }; + let host = host_port.split(':').next().unwrap_or("localhost"); + + // 检查是否需要替换 IP + let should_replace = if host == "0.0.0.0" || host == "127.0.0.1" || host == "localhost" { + // 0.0.0.0 需要替换为局域网 IP,127.0.0.1 和 localhost 保持不变 + host == "0.0.0.0" + } else { + // 检查 IP 是否在当前网卡列表中 + if let Ok(network_info) = proxycast_core::network::get_network_info() { + !network_info.all_ips.contains(&host.to_string()) + } else { + false + } + }; + + if should_replace { + // 获取局域网 IP 进行替换 + // 优先选择 192.168.x.x 或 10.x.x.x 开头的 IP(真正的局域网 IP) + if let Ok(network_info) = proxycast_core::network::get_network_info() { + let new_ip = network_info + .all_ips + .iter() + .find(|ip| ip.starts_with("192.168.") || ip.starts_with("10.")) + .or(network_info.lan_ip.as_ref()) + .or_else(|| network_info.all_ips.first()) + .cloned() + .unwrap_or_else(|| "localhost".to_string()); + state.base_url.replace(host, &new_ip) + } else { + state.base_url.replace(host, "localhost") + } + } else { + state.base_url.clone() + } + }; + + let routes = match &state.db { + Some(db) => state + .pool_service + .get_available_routes(db, &display_base_url) + .unwrap_or_default(), + None => Vec::new(), + }; + + // 获取默认 Provider + let default_provider = state.default_provider.read().await.clone(); + + // 添加默认路由 + let mut all_routes = vec![RouteInfo { + selector: "default".to_string(), + provider_type: default_provider.clone(), + credential_count: 1, + endpoints: vec![ + proxycast_core::models::route_model::RouteEndpoint { + path: "/v1/messages".to_string(), + protocol: "claude".to_string(), + url: format!("{display_base_url}/v1/messages"), + }, + proxycast_core::models::route_model::RouteEndpoint { + path: "/v1/chat/completions".to_string(), + protocol: "openai".to_string(), + url: format!("{display_base_url}/v1/chat/completions"), + }, + ], + tags: vec!["默认".to_string()], + enabled: true, + }]; + all_routes.extend(routes); + + let response = RouteListResponse { + base_url: display_base_url, + default_provider, + routes: all_routes, + }; + + Json(response) +} + +/// 带选择器的 Anthropic messages 处理 +async fn anthropic_messages_with_selector( + State(state): State, + Path(selector): Path, + headers: HeaderMap, + Json(request): Json, +) -> Response { + // 使用 Anthropic 格式的认证验证 + if let Err(e) = handlers::verify_api_key_anthropic(&headers, &state.api_key).await { + state.logs.write().await.add( + "warn", + &format!("Unauthorized request to /{selector}/v1/messages"), + ); + return e.into_response(); + } + + state.logs.write().await.add( + "info", + &format!( + "[REQ] POST /{}/v1/messages model={} stream={}", + selector, request.model, request.stream + ), + ); + + // 尝试解析凭证(不降级,指定什么就用什么) + let credential = match &state.db { + Some(db) => { + // 首先尝试按名称查找 + if let Ok(Some(cred)) = state.pool_service.get_by_name(db, &selector) { + Some(cred) + } + // 然后尝试按 UUID 查找 + else if let Ok(Some(cred)) = state.pool_service.get_by_uuid(db, &selector) { + Some(cred) + } + // 最后尝试按 provider 类型选择(不降级) + else if let Ok(Some(cred)) = + state + .pool_service + .select_credential(db, &selector, Some(&request.model)) + { + Some(cred) + } else { + None + } + } + None => None, + }; + + match credential { + Some(cred) => { + state.logs.write().await.add( + "info", + &format!( + "[ROUTE] Using credential: type={} name={:?} uuid={}", + cred.provider_type, + cred.name, + &cred.uuid[..8] + ), + ); + + // 根据凭证类型调用相应的 Provider + // 注意:这里没有 Flow 捕获,因为是通过 selector 路由的请求 + handlers::call_provider_anthropic(&state, &cred, &request, None).await + } + None => { + // 不再回退到默认 provider,直接返回错误 + state.logs.write().await.add( + "error", + &format!( + "[ROUTE] No available credentials for selector '{selector}', refusing to fallback" + ), + ); + ( + StatusCode::SERVICE_UNAVAILABLE, + Json(serde_json::json!({ + "error": { + "type": "provider_unavailable", + "message": format!("No available credentials for selector '{}'", selector) + } + })), + ) + .into_response() + } + } +} + +/// 带选择器的 OpenAI chat completions 处理 +async fn chat_completions_with_selector( + State(state): State, + Path(selector): Path, + headers: HeaderMap, + Json(request): Json, +) -> Response { + if let Err(e) = handlers::verify_api_key(&headers, &state.api_key).await { + state.logs.write().await.add( + "warn", + &format!("Unauthorized request to /{selector}/v1/chat/completions"), + ); + return e.into_response(); + } + + state.logs.write().await.add( + "info", + &format!( + "[REQ] POST /{}/v1/chat/completions model={} stream={}", + selector, request.model, request.stream + ), + ); + + // 尝试解析凭证(不降级,指定什么就用什么) + let credential = match &state.db { + Some(db) => { + if let Ok(Some(cred)) = state.pool_service.get_by_name(db, &selector) { + Some(cred) + } else if let Ok(Some(cred)) = state.pool_service.get_by_uuid(db, &selector) { + Some(cred) + } else if let Ok(Some(cred)) = + state + .pool_service + .select_credential(db, &selector, Some(&request.model)) + { + Some(cred) + } else { + None + } + } + None => None, + }; + + match credential { + Some(cred) => { + state.logs.write().await.add( + "info", + &format!( + "[ROUTE] Using credential: type={} name={:?} uuid={}", + cred.provider_type, + cred.name, + &cred.uuid[..8] + ), + ); + + // 注意:这里没有 Flow 捕获,因为是通过 selector 路由的请求 + handlers::call_provider_openai(&state, &cred, &request, None).await + } + None => { + // 不再回退到默认 provider,直接返回错误 + state.logs.write().await.add( + "error", + &format!( + "[ROUTE] No available credentials for selector '{selector}', refusing to fallback" + ), + ); + ( + StatusCode::SERVICE_UNAVAILABLE, + Json(serde_json::json!({ + "error": { + "message": format!("No available credentials for selector '{}'", selector), + "type": "provider_unavailable", + "code": "no_credentials" + } + })), + ) + .into_response() + } + } +} + +/// 内部 Anthropic messages 处理 (使用默认 Kiro) +/// 预留:用于内部直接调用 Kiro API +#[allow(dead_code)] +async fn anthropic_messages_internal( + state: &AppState, + request: &AnthropicMessagesRequest, +) -> Response { + // 检查 token + { + let _guard = state.kiro_refresh_lock.lock().await; + let mut kiro = state.kiro.write().await; + let needs_refresh = + kiro.credentials.access_token.is_none() || kiro.is_token_expiring_soon(); + if needs_refresh { + if let Err(e) = kiro.refresh_token().await { + state + .logs + .write() + .await + .add("error", &format!("[AUTH] Token refresh failed: {e}")); + return ( + StatusCode::UNAUTHORIZED, + Json(serde_json::json!({"error": {"message": format!("Token refresh failed: {e}")}})), + ) + .into_response(); + } + } + } + + let openai_request = convert_anthropic_to_openai(request); + let kiro = state.kiro.read().await; + + match kiro.call_api(&openai_request).await { + Ok(resp) => { + let status = resp.status(); + if status.is_success() { + match resp.bytes().await { + Ok(bytes) => { + let body = String::from_utf8_lossy(&bytes).to_string(); + let parsed = parse_cw_response(&body); + if request.stream { + build_anthropic_stream_response(&request.model, &parsed) + } else { + build_anthropic_response(&request.model, &parsed) + } + } + Err(e) => ( + StatusCode::INTERNAL_SERVER_ERROR, + Json(serde_json::json!({"error": {"message": e.to_string()}})), + ) + .into_response(), + } + } else { + let body = resp.text().await.unwrap_or_default(); + ( + StatusCode::from_u16(status.as_u16()).unwrap_or(StatusCode::INTERNAL_SERVER_ERROR), + Json(serde_json::json!({"error": {"message": format!("Upstream error: {}", body)}})), + ) + .into_response() + } + } + Err(e) => ( + StatusCode::INTERNAL_SERVER_ERROR, + Json(serde_json::json!({"error": {"message": e.to_string()}})), + ) + .into_response(), + } +} + +/// 内部 OpenAI chat completions 处理 (使用默认 Kiro) +/// 预留:用于内部直接调用 Kiro API +#[allow(dead_code)] +async fn chat_completions_internal(state: &AppState, request: &ChatCompletionRequest) -> Response { + { + let _guard = state.kiro_refresh_lock.lock().await; + let mut kiro = state.kiro.write().await; + let needs_refresh = + kiro.credentials.access_token.is_none() || kiro.is_token_expiring_soon(); + if needs_refresh { + if let Err(e) = kiro.refresh_token().await { + return ( + StatusCode::UNAUTHORIZED, + Json(serde_json::json!({"error": {"message": format!("Token refresh failed: {e}")}})), + ) + .into_response(); + } + } + } + + let kiro = state.kiro.read().await; + match kiro.call_api(request).await { + Ok(resp) => { + let status = resp.status(); + if status.is_success() { + match resp.text().await { + Ok(body) => { + let parsed = parse_cw_response(&body); + let has_tool_calls = !parsed.tool_calls.is_empty(); + + let message = if has_tool_calls { + serde_json::json!({ + "role": "assistant", + "content": if parsed.content.is_empty() { serde_json::Value::Null } else { serde_json::json!(parsed.content) }, + "tool_calls": parsed.tool_calls.iter().map(|tc| { + serde_json::json!({ + "id": tc.id, + "type": "function", + "function": { + "name": tc.function.name, + "arguments": tc.function.arguments + } + }) + }).collect::>() + }) + } else { + serde_json::json!({ + "role": "assistant", + "content": parsed.content + }) + }; + + let response = serde_json::json!({ + "id": format!("chatcmpl-{}", uuid::Uuid::new_v4()), + "object": "chat.completion", + "created": std::time::SystemTime::now() + .duration_since(std::time::UNIX_EPOCH) + .unwrap_or_default() + .as_secs(), + "model": request.model, + "choices": [{ + "index": 0, + "message": message, + "finish_reason": if has_tool_calls { "tool_calls" } else { "stop" } + }], + "usage": { + "prompt_tokens": 0, + "completion_tokens": 0, + "total_tokens": 0 + } + }); + Json(response).into_response() + } + Err(e) => ( + StatusCode::INTERNAL_SERVER_ERROR, + Json(serde_json::json!({"error": {"message": e.to_string()}})), + ) + .into_response(), + } + } else { + let body = resp.text().await.unwrap_or_default(); + ( + StatusCode::from_u16(status.as_u16()).unwrap_or(StatusCode::INTERNAL_SERVER_ERROR), + Json(serde_json::json!({"error": {"message": format!("Upstream error: {}", body)}})), + ) + .into_response() + } + } + Err(e) => ( + StatusCode::INTERNAL_SERVER_ERROR, + Json(serde_json::json!({"error": {"message": e.to_string()}})), + ) + .into_response(), + } +} diff --git a/src-tauri/src/app/bootstrap.rs b/src-tauri/src/app/bootstrap.rs index f99062c9c..fb91968f5 100644 --- a/src-tauri/src/app/bootstrap.rs +++ b/src-tauri/src/app/bootstrap.rs @@ -144,7 +144,9 @@ pub struct AppStates { pub fn init_states(config: &Config) -> Result { // 核心状态 let state: AppState = Arc::new(RwLock::new(server::ServerState::new(config.clone()))); - let logs: LogState = Arc::new(RwLock::new(logger::LogStore::with_config(&config.logging))); + let logs: LogState = Arc::new(RwLock::new(logger::create_log_store_from_config( + &config.logging, + ))); // 数据库 let db = database::init_database().map_err(|e| format!("数据库初始化失败: {e}"))?; diff --git a/src-tauri/src/app/state.rs b/src-tauri/src/app/state.rs index edeb93ad6..2104e2364 100644 --- a/src-tauri/src/app/state.rs +++ b/src-tauri/src/app/state.rs @@ -33,7 +33,9 @@ use crate::server; /// 初始化核心应用状态 pub fn init_core_state(config: Config) -> (AppState, LogState) { let state: AppState = Arc::new(RwLock::new(server::ServerState::new(config.clone()))); - let logs: LogState = Arc::new(RwLock::new(logger::LogStore::with_config(&config.logging))); + let logs: LogState = Arc::new(RwLock::new(logger::create_log_store_from_config( + &config.logging, + ))); (state, logs) } diff --git a/src-tauri/src/commands/network_cmd.rs b/src-tauri/src/commands/network_cmd.rs index f0a5a0d3c..7f56ee1c2 100644 --- a/src-tauri/src/commands/network_cmd.rs +++ b/src-tauri/src/commands/network_cmd.rs @@ -1,208 +1,14 @@ //! 网络相关命令 //! -//! 提供获取本地网络接口信息的功能 +//! 核心逻辑已迁移到 proxycast-core::network,本文件保留 Tauri 命令包装。 -use serde::Serialize; -use std::net::{IpAddr, UdpSocket}; +// 重新导出核心类型 +pub use proxycast_core::network::{ + get_accessible_host, get_accessible_url, get_local_url, NetworkInfo, +}; -/// 网络接口信息 -#[derive(Debug, Clone, Serialize)] -pub struct NetworkInfo { - /// 本地回环地址 - pub localhost: String, - /// 内网 IP 地址(局域网) - pub lan_ip: Option, - /// 所有可用的网络接口 IP 地址 - pub all_ips: Vec, -} - -/// 获取本地网络信息 -/// -/// 返回 localhost 和内网 IP 地址,用于客户端连接 +/// 获取本地网络信息(Tauri 命令包装) #[tauri::command] pub fn get_network_info() -> Result { - let lan_ip = get_local_ip(); - let all_ips = get_all_local_ips(); - - Ok(NetworkInfo { - localhost: "127.0.0.1".to_string(), - lan_ip, - all_ips, - }) -} - -/// 获取本机内网 IP 地址 -/// -/// 通过创建 UDP socket 连接外部地址来获取本机的内网 IP -/// 如果获取到的是 VPN 地址,则从 all_ips 中选择一个合适的 -fn get_local_ip() -> Option { - // 创建一个 UDP socket 并连接到外部地址(不会真正发送数据) - // 这样可以获取到本机用于出站连接的 IP 地址 - let socket = UdpSocket::bind("0.0.0.0:0").ok()?; - socket.connect("8.8.8.8:80").ok()?; - let local_addr = socket.local_addr().ok()?; - let ip_str = local_addr.ip().to_string(); - - // 检查是否是 VPN 地址 (198.18.x.x) - if let IpAddr::V4(ipv4) = local_addr.ip() { - if ipv4.octets()[0] == 198 && (ipv4.octets()[1] == 18 || ipv4.octets()[1] == 19) { - // 是 VPN 地址,尝试从 all_ips 中获取真实的局域网 IP - let all_ips = get_all_local_ips(); - // 优先选择 192.168.x.x - if let Some(ip) = all_ips.iter().find(|ip| ip.starts_with("192.168.")) { - return Some(ip.clone()); - } - // 其次选择任意私有 IP - if let Some(ip) = all_ips.first() { - return Some(ip.clone()); - } - // 如果没有私有 IP,返回 127.0.0.1 - return Some("127.0.0.1".to_string()); - } - } - - Some(ip_str) -} - -/// 获取所有本地网络接口的 IP 地址 -/// -/// 返回所有非回环的 IPv4 地址,过滤掉 VPN 和虚拟网卡 -fn get_all_local_ips() -> Vec { - let mut ips = Vec::new(); - - // 使用 if-addrs crate 获取所有网络接口 - if let Ok(interfaces) = if_addrs::get_if_addrs() { - for iface in interfaces { - // 只处理 IPv4 地址 - if let IpAddr::V4(ipv4) = iface.ip() { - // 过滤掉回环地址 - if ipv4.is_loopback() { - continue; - } - - // 过滤掉链路本地地址 (169.254.x.x) - if ipv4.octets()[0] == 169 && ipv4.octets()[1] == 254 { - continue; - } - - // 过滤掉常见的 VPN 地址段 - // 198.18.0.0/15 (用于基准测试) - if ipv4.octets()[0] == 198 && (ipv4.octets()[1] == 18 || ipv4.octets()[1] == 19) { - continue; - } - - // 只保留私有网络地址 - // 10.0.0.0/8, 172.16.0.0/12, 192.168.0.0/16 - let is_private = ipv4.octets()[0] == 10 - || (ipv4.octets()[0] == 172 - && (ipv4.octets()[1] >= 16 && ipv4.octets()[1] <= 31)) - || (ipv4.octets()[0] == 192 && ipv4.octets()[1] == 168); - - if is_private { - ips.push(ipv4.to_string()); - } - } - } - } - - ips -} - -/// 根据监听地址生成可访问的 URL -/// -/// 用于生成客户端配置中的 API URL。 -/// -/// # 参数 -/// - `listen_host`: 服务器监听地址 -/// - `port`: 服务器端口 -/// -/// # 返回 -/// - 如果监听地址为 `0.0.0.0`,返回局域网 IP 或 `127.0.0.1` -/// - 如果监听地址为 `127.0.0.1` 或 `localhost`,返回 `127.0.0.1` -/// - 其他情况返回原始地址 -pub fn get_accessible_host(listen_host: &str) -> String { - match listen_host { - "0.0.0.0" => { - // 获取局域网 IP,如果没有则使用 127.0.0.1 - get_network_info() - .ok() - .and_then(|info| info.lan_ip) - .unwrap_or_else(|| "127.0.0.1".to_string()) - } - "localhost" => "127.0.0.1".to_string(), - _ => listen_host.to_string(), - } -} - -/// 根据监听地址生成可访问的 URL -/// -/// # 参数 -/// - `listen_host`: 服务器监听地址 -/// - `port`: 服务器端口 -/// -/// # 返回 -/// 格式为 `http://{host}:{port}` 的 URL -pub fn get_accessible_url(listen_host: &str, port: u16) -> String { - let host = get_accessible_host(listen_host); - format!("http://{host}:{port}") -} - -/// 根据监听地址生成本地访问的 URL -/// -/// 用于 Agent 等本地组件访问服务器。 -/// 对于 `0.0.0.0`,返回 `127.0.0.1`(本地访问)。 -/// -/// # 参数 -/// - `listen_host`: 服务器监听地址 -/// - `port`: 服务器端口 -/// -/// # 返回 -/// 格式为 `http://{host}:{port}` 的 URL -#[allow(dead_code)] -pub fn get_local_url(listen_host: &str, port: u16) -> String { - let host = match listen_host { - "0.0.0.0" | "localhost" => "127.0.0.1".to_string(), - _ => listen_host.to_string(), - }; - format!("http://{host}:{port}") -} - -#[cfg(test)] -mod tests { - use super::*; - - #[test] - fn test_get_accessible_host_localhost() { - assert_eq!(get_accessible_host("127.0.0.1"), "127.0.0.1"); - assert_eq!(get_accessible_host("localhost"), "127.0.0.1"); - } - - #[test] - fn test_get_accessible_host_specific_ip() { - assert_eq!(get_accessible_host("192.168.1.100"), "192.168.1.100"); - assert_eq!(get_accessible_host("10.0.0.1"), "10.0.0.1"); - } - - #[test] - fn test_get_local_url() { - assert_eq!(get_local_url("0.0.0.0", 8999), "http://127.0.0.1:8999"); - assert_eq!(get_local_url("127.0.0.1", 8999), "http://127.0.0.1:8999"); - assert_eq!(get_local_url("localhost", 8999), "http://127.0.0.1:8999"); - assert_eq!( - get_local_url("192.168.1.100", 8999), - "http://192.168.1.100:8999" - ); - } - - #[test] - fn test_get_accessible_url_specific_ip() { - assert_eq!( - get_accessible_url("192.168.1.100", 8999), - "http://192.168.1.100:8999" - ); - assert_eq!( - get_accessible_url("127.0.0.1", 8999), - "http://127.0.0.1:8999" - ); - } + proxycast_core::network::get_network_info() } diff --git a/src-tauri/src/logger.rs b/src-tauri/src/logger.rs index cf2871ca1..28d7e4f9f 100644 --- a/src-tauri/src/logger.rs +++ b/src-tauri/src/logger.rs @@ -1,375 +1,12 @@ //! 日志管理模块 -use chrono::{Duration, Local, Utc}; -use flate2::write::GzEncoder; -use flate2::Compression; -use regex::Regex; -use serde::{Deserialize, Serialize}; -use std::collections::VecDeque; -use std::fs::{self, OpenOptions}; -use std::io::{Read, Write}; -use std::path::PathBuf; -use std::sync::Arc; -use tokio::sync::RwLock; +//! +//! 核心逻辑已迁移到 proxycast-core crate,本文件保留扩展函数。 -#[derive(Debug, Clone)] -pub struct LogStoreConfig { - pub max_logs: usize, - pub retention_days: u32, - pub max_file_size: u64, - pub enable_file_logging: bool, -} +pub use proxycast_core::logger::*; -impl Default for LogStoreConfig { - fn default() -> Self { - Self { - max_logs: 1000, - retention_days: 7, - max_file_size: 10 * 1024 * 1024, - enable_file_logging: true, - } - } -} +use crate::config::LoggingConfig; -#[derive(Debug, Clone, Serialize, Deserialize)] -pub struct LogEntry { - pub timestamp: String, - pub level: String, - pub message: String, -} - -pub struct LogStore { - logs: VecDeque, - max_logs: usize, - config: LogStoreConfig, - log_file_path: Option, -} - -impl Default for LogStore { - fn default() -> Self { - // 默认日志文件路径: ~/.proxycast/logs/proxycast.log - let log_dir = dirs::home_dir() - .unwrap_or_else(|| PathBuf::from(".")) - .join(".proxycast") - .join("logs"); - - // 创建日志目录 - let _ = fs::create_dir_all(&log_dir); - - let log_file = log_dir.join("proxycast.log"); - - let config = LogStoreConfig::default(); - - Self { - logs: VecDeque::new(), - max_logs: config.max_logs, - config, - log_file_path: Some(log_file), - } - } -} - -impl LogStore { - pub fn new() -> Self { - Self::default() - } - - pub fn with_config(logging: &crate::config::LoggingConfig) -> Self { - let mut store = Self::default(); - store.config.retention_days = logging.retention_days; - store.config.enable_file_logging = logging.enabled; - store.max_logs = store.config.max_logs; - store - } - - pub fn add(&mut self, level: &str, message: &str) { - let sanitized = sanitize_log_message(message); - let now = Utc::now(); - let entry = LogEntry { - timestamp: now.to_rfc3339(), - level: level.to_string(), - message: sanitized.clone(), - }; - - self.logs.push_back(entry.clone()); - - // 写入日志文件 - if self.config.enable_file_logging { - if let Some(ref path) = self.log_file_path { - self.rotate_log_file_if_needed(path); - let local_time = Local::now().format("%Y-%m-%d %H:%M:%S%.3f"); - let log_line = format!("{} [{}] {}\n", local_time, level.to_uppercase(), sanitized); - - if let Ok(mut file) = OpenOptions::new().create(true).append(true).open(path) { - let _ = file.write_all(log_line.as_bytes()); - } - self.prune_old_logs(path); - } - } - - // 保持日志数量在限制内 - if self.logs.len() > self.max_logs { - self.logs.pop_front(); - } - } - - /// 记录原始响应到单独的文件(用于调试) - pub fn log_raw_response(&self, request_id: &str, body: &str) { - if let Some(ref log_path) = self.log_file_path { - let log_dir = log_path.parent().unwrap_or(std::path::Path::new(".")); - let raw_file = log_dir.join(format!("raw_response_{request_id}.txt")); - let sanitized = sanitize_log_message(body); - - if let Ok(mut file) = OpenOptions::new() - .create(true) - .truncate(true) - .write(true) - .open(&raw_file) - { - let _ = file.write_all(sanitized.as_bytes()); - } - } - } - - pub fn get_logs(&self) -> Vec { - self.logs.iter().cloned().collect() - } - - pub fn clear(&mut self) { - self.logs.clear(); - } - - pub fn get_log_file_path(&self) -> Option { - self.log_file_path - .as_ref() - .map(|p| p.to_string_lossy().to_string()) - } - - fn rotate_log_file_if_needed(&self, path: &PathBuf) { - let Ok(metadata) = fs::metadata(path) else { - return; - }; - - if metadata.len() <= self.config.max_file_size { - return; - } - - let suffix = Local::now().format("%Y%m%d-%H%M%S"); - let rotated = path.with_file_name(format!( - "{}.{}", - path.file_name().unwrap_or_default().to_string_lossy(), - suffix - )); - - let _ = fs::rename(path, &rotated); - self.prune_old_logs(path); - } - - fn prune_old_logs(&self, path: &PathBuf) { - let Some(dir) = path.parent() else { - return; - }; - self.archive_old_logs(path); - let Ok(entries) = fs::read_dir(dir) else { - return; - }; - let cutoff = Utc::now() - Duration::days(self.config.retention_days as i64); - let prefix = format!( - "{}.", - path.file_name().unwrap_or_default().to_string_lossy() - ); - - for entry in entries.flatten() { - let file_name = entry.file_name(); - let file_name = file_name.to_string_lossy(); - if !file_name.starts_with(&prefix) { - continue; - } - let Ok(metadata) = entry.metadata() else { - continue; - }; - let Ok(modified) = metadata.modified() else { - continue; - }; - let modified = chrono::DateTime::::from(modified); - if modified < cutoff { - let _ = fs::remove_file(entry.path()); - } - } - } - - fn archive_old_logs(&self, path: &PathBuf) { - let Some(dir) = path.parent() else { - return; - }; - let Ok(entries) = fs::read_dir(dir) else { - return; - }; - let archive_cutoff = Utc::now() - Duration::days(7); - let delete_cutoff = Utc::now() - Duration::days(30); - let prefix = format!( - "{}.", - path.file_name().unwrap_or_default().to_string_lossy() - ); - - for entry in entries.flatten() { - let file_name = entry.file_name(); - let file_name = file_name.to_string_lossy(); - if !file_name.starts_with(&prefix) { - continue; - } - let path = entry.path(); - let Ok(metadata) = entry.metadata() else { - continue; - }; - let Ok(modified) = metadata.modified() else { - continue; - }; - let modified = chrono::DateTime::::from(modified); - - if file_name.ends_with(".gz") { - if modified < delete_cutoff { - let _ = fs::remove_file(path); - } - continue; - } - - if modified >= archive_cutoff { - continue; - } - - let mut input = Vec::new(); - if let Ok(mut file) = fs::File::open(&path) { - if file.read_to_end(&mut input).is_err() { - continue; - } - } else { - continue; - } - - let gz_path = path.with_extension(format!( - "{}.gz", - path.extension().unwrap_or_default().to_string_lossy() - )); - if let Ok(gz_file) = fs::File::create(&gz_path) { - let mut encoder = GzEncoder::new(gz_file, Compression::default()); - if encoder.write_all(&input).is_ok() && encoder.finish().is_ok() { - let _ = fs::remove_file(&path); - } - } - } - } -} - -#[allow(dead_code)] -pub type SharedLogStore = Arc>; - -/// P2 安全修复:扩展日志脱敏规则,覆盖更多敏感字段 -pub fn sanitize_log_message(message: &str) -> String { - let patterns = [ - // Bearer token - (r"Bearer\s+[A-Za-z0-9._-]+", "Bearer ***"), - // API key 各种格式 - ( - r#"api[_-]?key["']?\s*[:=]\s*["']?[A-Za-z0-9._-]+"#, - "api_key: ***", - ), - // 通用 token - (r#"token["']?\s*[:=]\s*["']?[A-Za-z0-9._-]+"#, "token: ***"), - // P2 新增:access_token - ( - r#"access[_-]?token["']?\s*[:=]\s*["']?[A-Za-z0-9._-]+"#, - "access_token: ***", - ), - // P2 新增:refresh_token - ( - r#"refresh[_-]?token["']?\s*[:=]\s*["']?[A-Za-z0-9._-]+"#, - "refresh_token: ***", - ), - // P2 新增:client_secret - ( - r#"client[_-]?secret["']?\s*[:=]\s*["']?[A-Za-z0-9._-]+"#, - "client_secret: ***", - ), - // P2 新增:authorization header - ( - r#"[Aa]uthorization["']?\s*[:=]\s*["']?[A-Za-z0-9._\s-]+"#, - "authorization: ***", - ), - // P2 新增:password - (r#"password["']?\s*[:=]\s*["']?[^\s"',}]+"#, "password: ***"), - // P2 新增:secret - ( - r#"secret["']?\s*[:=]\s*["']?[A-Za-z0-9._-]+"#, - "secret: ***", - ), - ]; - - let mut sanitized = message.to_string(); - for (pattern, replacement) in patterns { - if let Ok(re) = Regex::new(pattern) { - sanitized = re.replace_all(&sanitized, replacement).to_string(); - } - } - sanitized -} - -#[cfg(test)] -mod tests { - use super::sanitize_log_message; - - #[test] - fn test_sanitize_bearer_token() { - let input = "Authorization: Bearer abcDEF123._-XYZ"; - let output = sanitize_log_message(input); - // 验证敏感 token 被脱敏 - assert!(!output.contains("abcDEF123")); - assert!(output.contains("***")); - } - - #[test] - fn test_sanitize_api_key() { - let input = r#"request api_key="sk-test_123.456-ABC" end"#; - let output = sanitize_log_message(input); - assert!(output.contains("api_key: ***")); - assert!(!output.contains("sk-test_123")); - } - - #[test] - fn test_sanitize_access_token() { - let input = "access_token=atk_12345"; - let output = sanitize_log_message(input); - assert!(output.contains("access_token: ***")); - assert!(!output.contains("atk_12345")); - } - - #[test] - fn test_sanitize_refresh_token() { - let input = "refresh_token: rtk_ABCDE-123"; - let output = sanitize_log_message(input); - assert!(output.contains("refresh_token: ***")); - assert!(!output.contains("rtk_ABCDE")); - } - - #[test] - fn test_sanitize_client_secret() { - let input = "client_secret = \"cs_SeCreT-999\""; - let output = sanitize_log_message(input); - assert!(output.contains("client_secret: ***")); - assert!(!output.contains("cs_SeCreT")); - } - - #[test] - fn test_sanitize_password() { - let input = r#"{"password":"p@ssW0rd!"}"#; - let output = sanitize_log_message(input); - assert!(output.contains("password: ***")); - assert!(!output.contains("p@ssW0rd!")); - } - - #[test] - fn test_plain_text_unchanged() { - let input = "这是一段普通日志,不包含任何敏感字段。"; - let output = sanitize_log_message(input); - assert_eq!(output, input); - } +/// 使用 LoggingConfig 创建 LogStore +pub fn create_log_store_from_config(logging: &LoggingConfig) -> LogStore { + LogStore::with_custom_config(logging.retention_days, logging.enabled) } diff --git a/src-tauri/src/server/mod.rs b/src-tauri/src/server/mod.rs index 97d95dfbf..11e3b3b50 100644 --- a/src-tauri/src/server/mod.rs +++ b/src-tauri/src/server/mod.rs @@ -1,1838 +1,5 @@ //! HTTP API 服务器 +//! +//! 核心逻辑已迁移到 proxycast-server crate,本模块仅做重新导出。 -pub mod client_detector; - -use crate::config::{ - Config, ConfigChangeKind, ConfigManager, EndpointProvidersConfig, FileChangeEvent, FileWatcher, - HotReloadManager, ReloadResult, -}; -use crate::converter::anthropic_to_openai::convert_anthropic_to_openai; -use crate::credential::CredentialSyncService; -use crate::database::dao::provider_pool::ProviderPoolDao; -use crate::database::DbConnection; -use crate::injection::Injector; -use crate::logger::LogStore; -use crate::models::anthropic::*; -use crate::models::openai::*; -use crate::models::provider_pool_model::CredentialData; -use crate::models::route_model::{RouteInfo, RouteListResponse}; -use crate::processor::{RequestContext, RequestProcessor}; -use crate::providers::antigravity::AntigravityProvider; -use crate::providers::claude_custom::ClaudeCustomProvider; -use crate::providers::gemini::GeminiProvider; -use crate::providers::kiro::KiroProvider; -use crate::providers::openai_custom::OpenAICustomProvider; -use crate::server_utils::{ - build_anthropic_response, build_anthropic_stream_response, build_error_response, - build_error_response_with_status, build_gemini_cli_request, build_gemini_native_request, - health, models, parse_cw_response, -}; -use crate::services::kiro_event_service::KiroEventService; -use crate::services::provider_pool_service::ProviderPoolService; -use crate::services::token_cache_service::TokenCacheService; -use crate::websocket::{WsConfig, WsConnectionManager, WsStats}; -use axum::{ - extract::{DefaultBodyLimit, Path, State}, - http::{HeaderMap, StatusCode}, - response::{IntoResponse, Response}, - routing::{get, post}, - Json, Router, -}; -use serde::{Deserialize, Serialize}; -use std::path::PathBuf; -use std::sync::Arc; -use tokio::sync::{oneshot, RwLock}; - -/// 记录请求统计到遥测系统 -pub fn record_request_telemetry( - state: &AppState, - ctx: &RequestContext, - status: crate::telemetry::RequestStatus, - error_message: Option, -) { - use crate::telemetry::RequestLog; - - let provider = ctx.provider.unwrap_or(crate::ProviderType::Kiro); - let mut log = RequestLog::new( - ctx.request_id.clone(), - provider, - ctx.resolved_model.clone(), - ctx.is_stream, - ); - - // 设置状态和持续时间 - match status { - crate::telemetry::RequestStatus::Success => log.mark_success(ctx.elapsed_ms(), 200), - crate::telemetry::RequestStatus::Failed => log.mark_failed( - ctx.elapsed_ms(), - None, - error_message.clone().unwrap_or_default(), - ), - crate::telemetry::RequestStatus::Timeout => log.mark_timeout(ctx.elapsed_ms()), - crate::telemetry::RequestStatus::Cancelled => log.mark_cancelled(ctx.elapsed_ms()), - crate::telemetry::RequestStatus::Retrying => { - log.duration_ms = ctx.elapsed_ms(); - } - } - - // 设置凭证 ID - if let Some(cred_id) = &ctx.credential_id { - log.set_credential_id(cred_id.clone()); - } - - // 设置重试次数 - log.retry_count = ctx.retry_count; - - // 记录到统计聚合器 - { - let stats = state.processor.stats.write(); - stats.record(log.clone()); - } - - // 记录到请求日志记录器(用于前端日志列表显示) - if let Some(logger) = &state.request_logger { - let _ = logger.record(log.clone()); - } - - tracing::info!( - "[TELEMETRY] request_id={} provider={:?} model={} status={:?} duration_ms={}", - ctx.request_id, - provider, - ctx.resolved_model, - status, - ctx.elapsed_ms() - ); -} - -/// 记录 Token 使用量到遥测系统 -pub fn record_token_usage( - state: &AppState, - ctx: &RequestContext, - input_tokens: Option, - output_tokens: Option, -) { - use crate::telemetry::{TokenSource, TokenUsageRecord}; - - // 只有当至少有一个 Token 值时才记录 - if input_tokens.is_none() && output_tokens.is_none() { - return; - } - - let provider = ctx.provider.unwrap_or(crate::ProviderType::Kiro); - let record = TokenUsageRecord::new( - uuid::Uuid::new_v4().to_string(), - provider, - ctx.resolved_model.clone(), - input_tokens.unwrap_or(0), - output_tokens.unwrap_or(0), - TokenSource::Actual, - ) - .with_request_id(ctx.request_id.clone()); - - // 记录到 Token 追踪器 - { - let tokens = state.processor.tokens.write(); - tokens.record(record); - } - - tracing::debug!( - "[TOKEN] request_id={} input={} output={}", - ctx.request_id, - input_tokens.unwrap_or(0), - output_tokens.unwrap_or(0) - ); -} - -#[derive(Debug, Clone, Serialize, Deserialize)] -pub struct ServerStatus { - pub running: bool, - pub host: String, - pub port: u16, - pub requests: u64, - pub uptime_secs: u64, -} - -pub struct ServerState { - pub config: Config, - pub running: bool, - pub requests: u64, - pub start_time: Option, - pub kiro_provider: KiroProvider, - pub gemini_provider: GeminiProvider, - pub openai_custom_provider: OpenAICustomProvider, - pub claude_custom_provider: ClaudeCustomProvider, - pub default_provider_ref: Arc>, - /// 路由器引用(用于动态更新默认 Provider) - pub router_ref: Option>>, - shutdown_tx: Option>, - /// 服务器运行时使用的 API key(启动时从配置复制) - /// 用于 test_api 命令,确保测试使用的 API key 和服务器一致 - pub running_api_key: Option, - /// 服务器实际监听的 host(可能与配置不同,因为会自动切换到有效的 IP) - pub running_host: Option, -} - -impl ServerState { - pub fn new(config: Config) -> Self { - let kiro = KiroProvider::new(); - let gemini = GeminiProvider::new(); - let openai_custom = OpenAICustomProvider::new(); - let claude_custom = ClaudeCustomProvider::new(); - let default_provider_ref = Arc::new(RwLock::new(config.default_provider.clone())); - - Self { - config, - running: false, - requests: 0, - start_time: None, - kiro_provider: kiro, - gemini_provider: gemini, - openai_custom_provider: openai_custom, - claude_custom_provider: claude_custom, - default_provider_ref, - router_ref: None, - shutdown_tx: None, - running_api_key: None, - running_host: None, - } - } - - pub fn status(&self) -> ServerStatus { - ServerStatus { - running: self.running, - // 使用实际运行的 host,如果没有则使用配置的 host - host: self - .running_host - .clone() - .unwrap_or_else(|| self.config.server.host.clone()), - port: self.config.server.port, - requests: self.requests, - uptime_secs: self.start_time.map(|t| t.elapsed().as_secs()).unwrap_or(0), - } - } - - /// 增加请求计数 - pub fn increment_request_count(&mut self) { - self.requests = self.requests.saturating_add(1); - } - - /// 解析绑定地址 - /// - /// 直接返回用户配置的地址,不做任何自动替换。 - /// 如果地址无效,绑定时会失败并返回错误。 - fn resolve_bind_host(&self, configured_host: &str) -> String { - tracing::info!("[SERVER] 使用配置的监听地址: {}", configured_host); - configured_host.to_string() - } - - pub async fn start( - &mut self, - logs: Arc>, - pool_service: Arc, - token_cache: Arc, - db: Option, - ) -> Result<(), Box> { - self.start_with_telemetry(logs, pool_service, token_cache, db, None, None, None) - .await - } - - /// 启动服务器(使用共享的遥测实例) - /// - /// 这允许服务器与 TelemetryState 共享同一个 StatsAggregator、TokenTracker 和 RequestLogger, - /// 使得请求处理过程中记录的统计数据能够在前端监控页面中显示。 - pub async fn start_with_telemetry( - &mut self, - logs: Arc>, - pool_service: Arc, - token_cache: Arc, - db: Option, - shared_stats: Option>>, - shared_tokens: Option>>, - shared_logger: Option>, - ) -> Result<(), Box> { - self.start_with_telemetry_and_flow_monitor( - logs, - pool_service, - token_cache, - db, - shared_stats, - shared_tokens, - shared_logger, - ) - .await - } - - /// 启动服务器(使用共享的遥测实例) - /// - /// 这允许服务器与 TelemetryState 共享同一个 StatsAggregator、TokenTracker 和 RequestLogger, - /// 使得请求处理过程中记录的统计数据能够在前端监控页面中显示。 - pub async fn start_with_telemetry_and_flow_monitor( - &mut self, - logs: Arc>, - pool_service: Arc, - token_cache: Arc, - db: Option, - shared_stats: Option>>, - shared_tokens: Option>>, - shared_logger: Option>, - ) -> Result<(), Box> { - if self.running { - return Ok(()); - } - - let (tx, rx) = oneshot::channel(); - self.shutdown_tx = Some(tx); - - // 智能选择监听地址 - // - 127.0.0.1, localhost, 0.0.0.0, :: 直接使用 - // - 局域网 IP:检查是否在当前网卡列表中,如果不在则自动切换到当前局域网 IP - let configured_host = self.config.server.host.clone(); - let host = self.resolve_bind_host(&configured_host); - - // 如果地址发生了变化,记录日志 - if host != configured_host { - tracing::warn!( - "[SERVER] 配置的监听地址 {} 不可用,自动切换到 {}", - configured_host, - host - ); - } - - let port = self.config.server.port; - let api_key = self.config.server.api_key.clone(); - let api_key_for_state = api_key.clone(); // 用于保存到 running_api_key - let default_provider_ref = self.default_provider_ref.clone(); - - // 重新加载凭证 - let _ = self.kiro_provider.load_credentials().await; - let kiro = self.kiro_provider.clone(); - - // 创建参数注入器 - let injection_enabled = self.config.injection.enabled; - let injector = Injector::with_rules( - self.config - .injection - .rules - .iter() - .map(|r| r.clone().into()) - .collect(), - ); - - // 获取配置和配置路径用于热重载 - let config = self.config.clone(); - let config_path = crate::config::ConfigManager::default_config_path(); - - // 创建请求处理器(在 spawn 之前创建,以便保存 router_ref) - let processor = match (&shared_stats, &shared_tokens) { - (Some(stats), Some(tokens)) => Arc::new(RequestProcessor::with_shared_telemetry( - pool_service.clone(), - stats.clone(), - tokens.clone(), - )), - _ => Arc::new(RequestProcessor::with_defaults(pool_service.clone())), - }; - - // 从配置初始化 Router 的默认 Provider - { - let default_provider_str = &config.routing.default_provider; - - // 尝试解析为 ProviderType 枚举 - match default_provider_str.parse::() { - Ok(provider_type) => { - let mut router = processor.router.write().await; - router.set_default_provider(provider_type); - tracing::info!( - "[SERVER] 从配置初始化 Router 默认 Provider: {} (ProviderType)", - default_provider_str - ); - } - Err(_) => { - // 如果解析失败,可能是自定义 provider ID - // 这种情况下,路由器保持空状态,请求会直接使用 provider_id 进行凭证查找 - tracing::warn!( - "[SERVER] 配置的默认 Provider '{}' 不是有效的 ProviderType 枚举值,可能是自定义 Provider ID。\ - 路由器将保持空状态,请求将直接使用 provider_id 进行凭证查找。", - default_provider_str - ); - eprintln!( - "[SERVER] 警告:默认 Provider '{default_provider_str}' 不是标准 Provider 类型(kiro/openai/claude等),\ - 可能是自定义 Provider ID。如果这是预期行为,请忽略此警告。" - ); - } - } - } - - // 保存 router_ref 以便后续动态更新 - self.router_ref = Some(processor.router.clone()); - - // 保存实际使用的 host(在移动到 spawn 之前克隆) - let running_host = host.clone(); - - tokio::spawn(async move { - if let Err(e) = run_server( - &host, - port, - &api_key, - default_provider_ref, - kiro, - logs, - rx, - pool_service, - token_cache, - db, - injector, - injection_enabled, - shared_stats, - shared_tokens, - shared_logger, - Some(config), - Some(config_path), - Some(processor), - ) - .await - { - tracing::error!("Server error: {}", e); - } - }); - - self.running = true; - self.start_time = Some(std::time::Instant::now()); - // 保存服务器运行时使用的 API key,用于 test_api 命令 - self.running_api_key = Some(api_key_for_state); - // 保存服务器实际监听的 host(可能与配置不同) - self.running_host = Some(running_host); - Ok(()) - } - - pub async fn stop(&mut self) { - if let Some(tx) = self.shutdown_tx.take() { - let _ = tx.send(()); - } - self.running = false; - self.start_time = None; - self.running_api_key = None; - self.running_host = None; - self.router_ref = None; - } -} - -pub mod handlers; - -#[derive(Clone)] -#[allow(dead_code)] -pub struct AppState { - pub api_key: String, - pub base_url: String, - pub default_provider: Arc>, - pub kiro: Arc>, - pub logs: Arc>, - pub kiro_refresh_lock: Arc>, - pub gemini_refresh_lock: Arc>, - pub pool_service: Arc, - pub token_cache: Arc, - pub db: Option, - /// 参数注入器 - pub injector: Arc>, - /// 是否启用参数注入 - pub injection_enabled: Arc>, - /// 请求处理器 - pub processor: Arc, - /// WebSocket 连接管理器 - pub ws_manager: Arc, - /// WebSocket 统计信息 - pub ws_stats: Arc, - /// 热重载管理器 - pub hot_reload_manager: Option>, - /// 请求日志记录器(与 TelemetryState 共享) - pub request_logger: Option>, - /// Amp CLI 路由器 - pub amp_router: Arc, - /// 端点 Provider 配置 - pub endpoint_providers: Arc>, - /// Kiro 事件服务 - pub kiro_event_service: Arc, - /// API Key Provider 服务(用于智能降级) - pub api_key_service: Arc, -} - -/// 启动配置文件监控 -/// -/// 监控配置文件变化并触发热重载。 -/// -/// # 连接保持 -/// -/// 热重载过程不会中断现有连接: -/// - 配置更新在独立的 tokio 任务中异步执行 -/// - 使用 RwLock 进行原子性更新,不会阻塞正在处理的请求 -/// - 服务器继续运行,不需要重启 -/// - HTTP 和 WebSocket 连接保持活跃 -async fn start_config_watcher( - config_path: PathBuf, - hot_reload_manager: Option>, - processor: Arc, - logs: Arc>, - db: Option, - config_manager: Option>>, -) -> Option { - let (tx, mut rx) = tokio::sync::mpsc::unbounded_channel::(); - - // 创建文件监控器 - let mut watcher = match FileWatcher::new(&config_path, tx) { - Ok(w) => w, - Err(e) => { - tracing::error!("[HOT_RELOAD] 创建文件监控器失败: {}", e); - return None; - } - }; - - // 启动监控 - if let Err(e) = watcher.start() { - tracing::error!("[HOT_RELOAD] 启动文件监控失败: {}", e); - return None; - } - - tracing::info!("[HOT_RELOAD] 配置文件监控已启动: {:?}", config_path); - - // 启动事件处理任务 - let hot_reload_manager_clone = hot_reload_manager.clone(); - let processor_clone = processor.clone(); - let logs_clone = logs.clone(); - let db_clone = db.clone(); - let config_manager_clone = config_manager.clone(); - - tokio::spawn(async move { - while let Some(event) = rx.recv().await { - // 只处理修改事件 - if event.kind != ConfigChangeKind::Modified { - continue; - } - - tracing::info!("[HOT_RELOAD] 检测到配置文件变更: {:?}", event.path); - logs_clone.write().await.add( - "info", - &format!("[HOT_RELOAD] 检测到配置文件变更: {:?}", event.path), - ); - - // 执行热重载 - if let Some(ref manager) = hot_reload_manager_clone { - let result = manager.reload(); - match &result { - ReloadResult::Success { .. } => { - tracing::info!("[HOT_RELOAD] 配置热重载成功"); - logs_clone - .write() - .await - .add("info", "[HOT_RELOAD] 配置热重载成功"); - - // 更新处理器中的组件 - let new_config = manager.config(); - update_processor_config(&processor_clone, &new_config).await; - - // 同步凭证池 - if let (Some(ref db), Some(ref cfg_manager)) = - (&db_clone, &config_manager_clone) - { - match sync_credential_pool_from_config(db, cfg_manager, &logs_clone) - .await - { - Ok(count) => { - tracing::info!( - "[HOT_RELOAD] 凭证池同步完成,共 {} 个凭证", - count - ); - logs_clone.write().await.add( - "info", - &format!("[HOT_RELOAD] 凭证池同步完成,共 {count} 个凭证"), - ); - } - Err(e) => { - tracing::warn!("[HOT_RELOAD] 凭证池同步失败: {}", e); - logs_clone - .write() - .await - .add("warn", &format!("[HOT_RELOAD] 凭证池同步失败: {e}")); - } - } - } - } - ReloadResult::RolledBack { error, .. } => { - tracing::warn!("[HOT_RELOAD] 配置热重载失败,已回滚: {}", error); - logs_clone.write().await.add( - "warn", - &format!("[HOT_RELOAD] 配置热重载失败,已回滚: {error}"), - ); - } - ReloadResult::Failed { - error, - rollback_error, - .. - } => { - tracing::error!( - "[HOT_RELOAD] 配置热重载失败: {}, 回滚错误: {:?}", - error, - rollback_error - ); - logs_clone.write().await.add( - "error", - &format!( - "[HOT_RELOAD] 配置热重载失败: {error}, 回滚错误: {rollback_error:?}" - ), - ); - } - } - } - } - }); - - Some(watcher) -} - -/// 更新处理器配置 -/// -/// 当配置热重载成功后,更新 RequestProcessor 中的各个组件。 -/// -/// # 原子性更新 -/// -/// 每个组件的更新都是原子性的,使用 RwLock 确保: -/// - 正在处理的请求不会看到部分更新的状态 -/// - 更新过程不会阻塞新请求的处理 -/// - 现有连接不受影响 -async fn update_processor_config(processor: &RequestProcessor, config: &Config) { - // 更新注入器规则 - { - let mut injector = processor.injector.write().await; - injector.clear(); - for rule in &config.injection.rules { - injector.add_rule(rule.clone().into()); - } - tracing::debug!( - "[HOT_RELOAD] 注入器规则已更新: {} 条规则", - config.injection.rules.len() - ); - } - - // 更新路由器默认 Provider - { - let mut router = processor.router.write().await; - - // 尝试解析为 ProviderType 枚举 - match config - .routing - .default_provider - .parse::() - { - Ok(provider_type) => { - router.set_default_provider(provider_type); - tracing::debug!( - "[HOT_RELOAD] 路由器默认 Provider 已更新: {} (ProviderType)", - config.routing.default_provider - ); - } - Err(_) => { - // 如果解析失败,可能是自定义 provider ID - // 清空路由器的默认 provider,让请求直接使用 provider_id - tracing::warn!( - "[HOT_RELOAD] 配置的默认 Provider '{}' 不是有效的 ProviderType 枚举值,可能是自定义 Provider ID。\ - 路由器默认 Provider 将被清空。", - config.routing.default_provider - ); - } - } - } - - // 更新模型映射器 - { - let mut mapper = processor.mapper.write().await; - mapper.clear(); - for (alias, model) in &config.routing.model_aliases { - mapper.add_alias(alias, model); - } - tracing::debug!( - "[HOT_RELOAD] 模型别名已更新: {} 个别名", - config.routing.model_aliases.len() - ); - } - - // 注意:重试配置目前不支持热更新,因为 Retrier 是不可变的 - // 如果需要更新重试配置,需要重启服务器 - tracing::debug!( - "[HOT_RELOAD] 重试配置: max_retries={}, base_delay={}ms (需重启生效)", - config.retry.max_retries, - config.retry.base_delay_ms - ); - - tracing::info!("[HOT_RELOAD] 处理器配置更新完成"); -} - -/// 从配置同步凭证池 -/// -/// 当配置热重载成功后,从 YAML 配置中加载凭证并同步到数据库。 -/// -/// # 同步策略 -/// -/// - 从配置中加载所有凭证 -/// - 对于配置中存在但数据库中不存在的凭证,添加到数据库 -/// - 对于配置中存在且数据库中也存在的凭证,更新数据库中的记录 -/// - 对于数据库中存在但配置中不存在的凭证,保留(不删除,避免丢失运行时状态) -async fn sync_credential_pool_from_config( - db: &DbConnection, - config_manager: &Arc>, - _logs: &Arc>, -) -> Result { - // 创建凭证同步服务 - let sync_service = CredentialSyncService::new(config_manager.clone()); - - // 从配置加载凭证 - let credentials = sync_service.load_from_config().map_err(|e| e.to_string())?; - - let conn = crate::database::lock_db(db)?; - let mut synced_count = 0; - - for cred in &credentials { - // 检查凭证是否已存在 - let existing = - ProviderPoolDao::get_by_uuid(&conn, &cred.uuid).map_err(|e| e.to_string())?; - - if existing.is_some() { - // 更新现有凭证 - ProviderPoolDao::update(&conn, cred).map_err(|e| e.to_string())?; - tracing::debug!( - "[HOT_RELOAD] 更新凭证: {} ({})", - cred.uuid, - cred.provider_type - ); - } else { - // 添加新凭证 - ProviderPoolDao::insert(&conn, cred).map_err(|e| e.to_string())?; - tracing::debug!( - "[HOT_RELOAD] 添加凭证: {} ({})", - cred.uuid, - cred.provider_type - ); - } - synced_count += 1; - } - - Ok(synced_count) -} - -async fn run_server( - host: &str, - port: u16, - api_key: &str, - default_provider: Arc>, - kiro: KiroProvider, - logs: Arc>, - shutdown: oneshot::Receiver<()>, - pool_service: Arc, - token_cache: Arc, - db: Option, - injector: Injector, - injection_enabled: bool, - shared_stats: Option>>, - shared_tokens: Option>>, - shared_logger: Option>, - config: Option, - config_path: Option, - processor: Option>, -) -> Result<(), Box> { - let base_url = format!("http://{host}:{port}"); - - // 使用传入的 processor 或创建新的 - let processor = match processor { - Some(p) => p, - None => match (&shared_stats, &shared_tokens) { - (Some(stats), Some(tokens)) => Arc::new(RequestProcessor::with_shared_telemetry( - pool_service.clone(), - stats.clone(), - tokens.clone(), - )), - _ => Arc::new(RequestProcessor::with_defaults(pool_service.clone())), - }, - }; - - // 将注入器规则同步到处理器 - { - let mut proc_injector = processor.injector.write().await; - for rule in injector.rules() { - proc_injector.add_rule(rule.clone()); - } - } - - // 从配置初始化 Router 的默认 Provider - if let Some(cfg) = &config { - let default_provider_str = &cfg.routing.default_provider; - - // 尝试解析为 ProviderType 枚举 - match default_provider_str.parse::() { - Ok(provider_type) => { - let mut router = processor.router.write().await; - router.set_default_provider(provider_type); - tracing::info!( - "[SERVER] 从配置初始化 Router 默认 Provider: {} (ProviderType)", - default_provider_str - ); - } - Err(_) => { - // 如果解析失败,可能是自定义 provider ID - tracing::warn!( - "[SERVER] 配置的默认 Provider '{}' 不是有效的 ProviderType 枚举值,可能是自定义 Provider ID。\ - 路由器将保持空状态,请求将直接使用 provider_id 进行凭证查找。", - default_provider_str - ); - eprintln!( - "[SERVER] 警告:默认 Provider '{default_provider_str}' 不是标准 Provider 类型,可能是自定义 Provider ID" - ); - } - } - } - - // 初始化 WebSocket 管理器 - let ws_manager = Arc::new(WsConnectionManager::new(WsConfig::default())); - let ws_stats = ws_manager.stats().clone(); - - // 初始化热重载管理器 - let hot_reload_manager = match (&config, &config_path) { - (Some(cfg), Some(path)) => Some(Arc::new(HotReloadManager::new(cfg.clone(), path.clone()))), - _ => None, - }; - - // 初始化配置管理器(用于凭证池同步) - let config_manager: Option>> = - match (&config, &config_path) { - (Some(cfg), Some(path)) => Some(Arc::new(std::sync::RwLock::new( - ConfigManager::with_config(cfg.clone(), path.clone()), - ))), - _ => None, - }; - - let logs_clone = logs.clone(); - let db_clone = db.clone(); - - // 初始化 Amp CLI 路由器 - let amp_router = Arc::new(crate::router::AmpRouter::new( - config - .as_ref() - .map(|c| c.ampcode.clone()) - .unwrap_or_default(), - )); - - // 初始化端点 Provider 配置 - let endpoint_providers = Arc::new(RwLock::new( - config - .as_ref() - .map(|c| c.endpoint_providers.clone()) - .unwrap_or_default(), - )); - - // 创建 Kiro 事件服务 - let kiro_event_service = Arc::new(KiroEventService::new()); - - // 创建 API Key Provider 服务 - let api_key_service = - Arc::new(crate::services::api_key_provider_service::ApiKeyProviderService::new()); - - let state = AppState { - api_key: api_key.to_string(), - base_url, - default_provider, - kiro: Arc::new(RwLock::new(kiro)), - logs, - kiro_refresh_lock: Arc::new(tokio::sync::Mutex::new(())), - gemini_refresh_lock: Arc::new(tokio::sync::Mutex::new(())), - pool_service, - token_cache, - db, - injector: Arc::new(RwLock::new(injector)), - injection_enabled: Arc::new(RwLock::new(injection_enabled)), - processor: processor.clone(), - ws_manager, - ws_stats, - hot_reload_manager: hot_reload_manager.clone(), - request_logger: shared_logger, - amp_router, - endpoint_providers, - kiro_event_service, - api_key_service, - }; - - // ========== 开发模式:启动独立的 HTTP 桥接服务器 ========== - // 仅在 debug 模式下,启动一个独立的开发服务器在端口 3030 - // 允许浏览器 dev server 通过 HTTP 调用 Tauri 命令 - #[cfg(debug_assertions)] - { - eprintln!("[DevBridge] ===== 准备启动开发桥接服务器 ====="); - use tokio::sync::RwLock as TokioRwLock; - let dev_bridge_state = Arc::new(TokioRwLock::new(state.clone())); - eprintln!("[DevBridge] 状态已克隆,准备启动"); - tokio::spawn(async move { - eprintln!("[DevBridge] spawn 任务开始执行"); - match crate::dev_bridge::DevBridgeServer::start(dev_bridge_state, None).await { - Ok(_) => { - eprintln!("[DevBridge] 启动完成"); - } - Err(e) => { - eprintln!("[DevBridge] 启动失败: {e}"); - } - } - }); - } - - // 启动配置文件监控 - let _file_watcher = if let Some(path) = config_path { - start_config_watcher( - path, - hot_reload_manager, - processor, - logs_clone, - db_clone, - config_manager, - ) - .await - } else { - None - }; - - // 设置请求体大小限制为 100MB,支持大型上下文请求(如 Claude Code 的 /compact 命令) - let body_limit = 100 * 1024 * 1024; // 100MB - - // 创建管理 API 路由(带认证中间件) - let management_config = config - .as_ref() - .map(|c| c.remote_management.clone()) - .unwrap_or_default(); - - let management_routes = Router::new() - .route("/v0/management/status", get(handlers::management_status)) - .route( - "/v0/management/credentials", - get(handlers::management_list_credentials), - ) - .route( - "/v0/management/credentials", - post(handlers::management_add_credential), - ) - .route( - "/v0/management/config", - get(handlers::management_get_config), - ) - .route( - "/v0/management/config", - axum::routing::put(handlers::management_update_config), - ) - .layer(crate::middleware::ManagementAuthLayer::new( - management_config, - )); - - // Kiro凭证管理API路由 - let kiro_api_routes = Router::new() - .route( - "/api/kiro/credentials/available", - get(handlers::get_available_credentials), - ) - .route( - "/api/kiro/credentials/select", - post(handlers::select_credential), - ) - .route( - "/api/kiro/credentials/{uuid}/refresh", - axum::routing::put(handlers::refresh_credential), - ) - .route( - "/api/kiro/credentials/{uuid}/status", - get(handlers::get_credential_status), - ); - - // 凭证 API 路由(用于 aster Agent 集成) - let credentials_api_routes = Router::new() - .route("/v1/credentials/select", post(handlers::credentials_select)) - .route( - "/v1/credentials/{uuid}/token", - get(handlers::credentials_get_token), - ); - - let app = Router::new() - .route("/health", get(health)) - .route("/v1/models", get(models)) - .route("/v1/routes", get(list_routes)) - .route("/v1/chat/completions", post( - |State(state): State, - headers: HeaderMap, - Json(request): Json| async { - handlers::chat_completions(State(state), headers, Json(request)).await - } - )) - .route("/v1/messages", post( - |State(state): State, - headers: HeaderMap, - Json(request): Json| async { - handlers::anthropic_messages(State(state), headers, Json(request)).await - } - )) - .route("/v1/messages/count_tokens", post(count_tokens)) - // 图像生成 API 路由 - .route( - "/v1/images/generations", - post(handlers::handle_image_generation), - ) - // WebSocket 路由 - .route("/v1/ws", get(handlers::ws_upgrade_handler)) - .route("/ws", get(handlers::ws_upgrade_handler)) - // 多供应商路由 - .route( - "/{selector}/v1/messages", - post(anthropic_messages_with_selector), - ) - .route( - "/{selector}/v1/chat/completions", - post(chat_completions_with_selector), - ) - // 管理 API 路由 - .merge(management_routes) - // Kiro凭证管理API路由 - .merge(kiro_api_routes) - // 凭证 API 路由(用于 aster Agent 集成) - .merge(credentials_api_routes) - .layer(DefaultBodyLimit::max(body_limit)) - .with_state(state); - - let addr: std::net::SocketAddr = format!("{host}:{port}") - .parse() - .map_err(|e| format!("无效的监听地址 {host}:{port} - {e}"))?; - - let listener = tokio::net::TcpListener::bind(addr).await.map_err(|e| { - format!("无法绑定到 {host}:{port},错误: {e}。请检查地址是否有效或端口是否被占用。") - })?; - - tracing::info!("Server listening on {}", addr); - - axum::serve(listener, app) - .with_graceful_shutdown(async move { - let _ = shutdown.await; - }) - .await?; - - Ok(()) -} - -async fn count_tokens( - State(state): State, - headers: HeaderMap, - Json(_request): Json, -) -> Response { - if let Err(e) = handlers::verify_api_key(&headers, &state.api_key).await { - return e.into_response(); - } - - // Claude Code 需要这个端点,返回估算值 - Json(serde_json::json!({ - "input_tokens": 100 - })) - .into_response() -} - -/// Gemini 原生协议处理 -/// 路由: POST /v1/gemini/{model}:{method} -/// 例如: /v1/gemini/gemini-3-pro-preview:generateContent -#[allow(dead_code)] -async fn gemini_generate_content( - State(state): State, - headers: HeaderMap, - Path(path): Path, - Json(request): Json, -) -> Response { - if let Err(e) = handlers::verify_api_key(&headers, &state.api_key).await { - return e.into_response(); - } - - // 解析路径: {model}:{method} - // 例如: gemini-3-pro-preview:generateContent - let parts: Vec<&str> = path.splitn(2, ':').collect(); - if parts.len() != 2 { - return ( - StatusCode::BAD_REQUEST, - Json(serde_json::json!({ - "error": { - "message": format!("无效的路径格式: {},期望格式: model:method", path) - } - })), - ) - .into_response(); - } - - let model = parts[0]; - let method = parts[1]; - - state.logs.write().await.add( - "info", - &format!("[GEMINI] POST /v1/gemini/{path} model={model} method={method}"), - ); - - // 目前只支持 generateContent 方法 - if method != "generateContent" && method != "streamGenerateContent" { - return ( - StatusCode::BAD_REQUEST, - Json(serde_json::json!({ - "error": { - "message": format!("不支持的方法: {},目前只支持 generateContent", method) - } - })), - ) - .into_response(); - } - - let is_stream = method == "streamGenerateContent"; - - // 获取默认 provider - let default_provider = state.default_provider.read().await.clone(); - - // 尝试从凭证池中选择凭证(不降级,指定什么就用什么) - let credential = match &state.db { - Some(db) => state - .pool_service - .select_credential(db, &default_provider, Some(model)) - .ok() - .flatten(), - None => None, - }; - - let cred = match credential { - Some(c) => c, - None => { - return ( - StatusCode::SERVICE_UNAVAILABLE, - Json(serde_json::json!({ - "error": { - "message": format!("No available credentials for provider '{}'. Please add credentials in the Provider Pool.", default_provider) - } - })), - ) - .into_response(); - } - }; - - state.logs.write().await.add( - "info", - &format!( - "[GEMINI] 使用凭证: type={} name={:?} uuid={}", - cred.provider_type, - cred.name, - &cred.uuid[..8] - ), - ); - - // 调用 Antigravity Provider - match &cred.credential { - CredentialData::AntigravityOAuth { - creds_file_path, - project_id, - } => { - let mut antigravity = AntigravityProvider::new(); - if let Err(e) = antigravity - .load_credentials_from_path(creds_file_path) - .await - { - return ( - StatusCode::INTERNAL_SERVER_ERROR, - Json(serde_json::json!({ - "error": { - "message": format!("加载 Antigravity 凭证失败: {}", e) - } - })), - ) - .into_response(); - } - - // 使用新的 validate_token() 方法检查 Token 状态 - let validation_result = antigravity.validate_token(); - tracing::info!( - "[Antigravity Gemini] Token 验证结果: {:?}", - validation_result - ); - - // 根据验证结果决定是否刷新 - if validation_result.needs_refresh() { - tracing::info!("[Antigravity Gemini] Token 需要刷新,开始刷新..."); - match antigravity.refresh_token_with_retry(3).await { - Ok(new_token) => { - tracing::info!( - "[Antigravity Gemini] Token 刷新成功,新 token 长度: {}", - new_token.len() - ); - } - Err(refresh_error) => { - tracing::error!("[Antigravity Gemini] Token 刷新失败: {:?}", refresh_error); - - // 根据错误类型返回不同的状态码和消息 - let (status, message) = if refresh_error.requires_reauth() { - (StatusCode::UNAUTHORIZED, refresh_error.user_message()) - } else { - ( - StatusCode::INTERNAL_SERVER_ERROR, - refresh_error.user_message(), - ) - }; - - return ( - status, - Json(serde_json::json!({ - "error": { - "message": message - } - })), - ) - .into_response(); - } - } - } - - // 设置项目 ID - if let Some(pid) = project_id { - antigravity.project_id = Some(pid.clone()); - } else if antigravity.project_id.is_none() { - // 如果凭证中没有 project_id,尝试从 API 获取或生成随机 ID - if let Err(e) = antigravity.discover_project().await { - tracing::warn!("[Antigravity] 获取项目 ID 失败: {},使用随机生成的 ID", e); - // 生成随机项目 ID - let uuid = uuid::Uuid::new_v4(); - let bytes = uuid.as_bytes(); - let adjectives = ["useful", "bright", "swift", "calm", "bold"]; - let nouns = ["fuze", "wave", "spark", "flow", "core"]; - let adj = adjectives[(bytes[0] as usize) % adjectives.len()]; - let noun = nouns[(bytes[1] as usize) % nouns.len()]; - let random_part: String = uuid.to_string()[..5].to_lowercase(); - antigravity.project_id = Some(format!("{adj}-{noun}-{random_part}")); - } - } - - let proj_id = antigravity.project_id.clone().unwrap_or_else(|| { - // 最后的后备:生成随机 ID - let uuid = uuid::Uuid::new_v4(); - format!("proxycast-{}", &uuid.to_string()[..8]) - }); - - state - .logs - .write() - .await - .add("debug", &format!("[GEMINI] 使用 project_id: {proj_id}")); - - // 构建 Antigravity 请求体 - // 直接使用用户传入的 Gemini 格式请求,只添加必要的字段 - let antigravity_request = build_gemini_native_request(&request, model, &proj_id); - - state.logs.write().await.add( - "debug", - &format!( - "[GEMINI] 请求体: {}", - serde_json::to_string(&antigravity_request).unwrap_or_default() - ), - ); - - if is_stream { - // 流式响应 - 暂不支持,返回错误 - return ( - StatusCode::NOT_IMPLEMENTED, - Json(serde_json::json!({ - "error": { - "message": "流式响应暂不支持,请使用 generateContent" - } - })), - ) - .into_response(); - } - - // 非流式响应 - match antigravity - .call_api("generateContent", &antigravity_request) - .await - { - Ok(resp) => { - state.logs.write().await.add( - "info", - &format!( - "[GEMINI] 响应成功: {}", - serde_json::to_string(&resp) - .unwrap_or_default() - .chars() - .take(200) - .collect::() - ), - ); - - // 直接返回 Gemini 格式响应 - Json(resp).into_response() - } - Err(api_err) => { - state.logs.write().await.add( - "error", - &format!( - "[GEMINI] 请求失败 (HTTP {}): {}", - api_err.status_code, api_err.message - ), - ); - - // 直接使用 AntigravityApiError 的状态码构建响应 - build_error_response_with_status(api_err.status_code, &api_err.to_string()) - } - } - } - CredentialData::GeminiOAuth { - creds_file_path, - project_id, - } => { - // 使用 GeminiProvider 处理 Gemini CLI OAuth 凭证 - let mut gemini = GeminiProvider::new(); - if let Err(e) = gemini.load_credentials_from_path(creds_file_path).await { - return ( - StatusCode::INTERNAL_SERVER_ERROR, - Json(serde_json::json!({ - "error": { - "message": format!("加载 Gemini 凭证失败: {}", e) - } - })), - ) - .into_response(); - } - - // 检查并刷新 Token - if !gemini.is_token_valid() { - tracing::info!("[Gemini CLI] Token 需要刷新,开始刷新..."); - match gemini.refresh_token_with_retry(3).await { - Ok(new_token) => { - tracing::info!( - "[Gemini CLI] Token 刷新成功,新 token 长度: {}", - new_token.len() - ); - } - Err(refresh_error) => { - tracing::error!("[Gemini CLI] Token 刷新失败: {:?}", refresh_error); - return ( - StatusCode::UNAUTHORIZED, - Json(serde_json::json!({ - "error": { - "message": format!("Token 刷新失败: {}", refresh_error) - } - })), - ) - .into_response(); - } - } - } - - // 设置项目 ID - if let Some(pid) = project_id { - gemini.project_id = Some(pid.clone()); - } else if gemini.project_id.is_none() { - // 尝试从 API 获取项目 ID - if let Err(e) = gemini.discover_project().await { - tracing::warn!("[Gemini CLI] 获取项目 ID 失败: {},使用随机生成的 ID", e); - let uuid = uuid::Uuid::new_v4(); - let bytes = uuid.as_bytes(); - let adjectives = ["useful", "bright", "swift", "calm", "bold"]; - let nouns = ["fuze", "wave", "spark", "flow", "core"]; - let adj = adjectives[(bytes[0] as usize) % adjectives.len()]; - let noun = nouns[(bytes[1] as usize) % nouns.len()]; - let random_part: String = uuid.to_string()[..5].to_lowercase(); - gemini.project_id = Some(format!("{adj}-{noun}-{random_part}")); - } - } - - let proj_id = gemini.project_id.clone().unwrap_or_else(|| { - let uuid = uuid::Uuid::new_v4(); - format!("proxycast-{}", &uuid.to_string()[..8]) - }); - - state - .logs - .write() - .await - .add("debug", &format!("[GEMINI CLI] 使用 project_id: {proj_id}")); - - // 构建 Gemini CLI 请求体 - // Gemini CLI 使用 Cloud Code Assist 端点,不做模型名称映射 - let gemini_request = build_gemini_cli_request(&request, model, &proj_id); - - state.logs.write().await.add( - "debug", - &format!( - "[GEMINI CLI] 请求体: {}", - serde_json::to_string(&gemini_request).unwrap_or_default() - ), - ); - - if is_stream { - // 流式响应 - 暂不支持 - return ( - StatusCode::NOT_IMPLEMENTED, - Json(serde_json::json!({ - "error": { - "message": "Gemini CLI 流式响应暂不支持,请使用 generateContent" - } - })), - ) - .into_response(); - } - - // 非流式响应 - match gemini.call_api("generateContent", &gemini_request).await { - Ok(resp) => { - state.logs.write().await.add( - "info", - &format!( - "[GEMINI CLI] 响应成功: {}", - serde_json::to_string(&resp) - .unwrap_or_default() - .chars() - .take(200) - .collect::() - ), - ); - - // 直接返回 Gemini 格式响应 - Json(resp).into_response() - } - Err(api_err) => { - state - .logs - .write() - .await - .add("error", &format!("[GEMINI CLI] 请求失败: {api_err}")); - - build_error_response(&api_err.to_string()) - } - } - } - _ => ( - StatusCode::BAD_REQUEST, - Json(serde_json::json!({ - "error": { - "message": "Gemini 原生协议只支持 Antigravity 或 Gemini CLI OAuth 凭证" - } - })), - ) - .into_response(), - } -} - -/// 列出所有可用路由 -async fn list_routes(State(state): State) -> impl IntoResponse { - // 处理 base_url:检查 IP 是否有效(在当前网卡列表中或是特殊地址) - let display_base_url = { - // 从 base_url 中提取 host 部分 - let url_parts: Vec<&str> = state.base_url.split("://").collect(); - let host_port = if url_parts.len() > 1 { - url_parts[1] - } else { - &state.base_url - }; - let host = host_port.split(':').next().unwrap_or("localhost"); - - // 检查是否需要替换 IP - let should_replace = if host == "0.0.0.0" || host == "127.0.0.1" || host == "localhost" { - // 0.0.0.0 需要替换为局域网 IP,127.0.0.1 和 localhost 保持不变 - host == "0.0.0.0" - } else { - // 检查 IP 是否在当前网卡列表中 - if let Ok(network_info) = crate::commands::network_cmd::get_network_info() { - !network_info.all_ips.contains(&host.to_string()) - } else { - false - } - }; - - if should_replace { - // 获取局域网 IP 进行替换 - // 优先选择 192.168.x.x 或 10.x.x.x 开头的 IP(真正的局域网 IP) - if let Ok(network_info) = crate::commands::network_cmd::get_network_info() { - let new_ip = network_info - .all_ips - .iter() - .find(|ip| ip.starts_with("192.168.") || ip.starts_with("10.")) - .or(network_info.lan_ip.as_ref()) - .or_else(|| network_info.all_ips.first()) - .cloned() - .unwrap_or_else(|| "localhost".to_string()); - state.base_url.replace(host, &new_ip) - } else { - state.base_url.replace(host, "localhost") - } - } else { - state.base_url.clone() - } - }; - - let routes = match &state.db { - Some(db) => state - .pool_service - .get_available_routes(db, &display_base_url) - .unwrap_or_default(), - None => Vec::new(), - }; - - // 获取默认 Provider - let default_provider = state.default_provider.read().await.clone(); - - // 添加默认路由 - let mut all_routes = vec![RouteInfo { - selector: "default".to_string(), - provider_type: default_provider.clone(), - credential_count: 1, - endpoints: vec![ - crate::models::route_model::RouteEndpoint { - path: "/v1/messages".to_string(), - protocol: "claude".to_string(), - url: format!("{display_base_url}/v1/messages"), - }, - crate::models::route_model::RouteEndpoint { - path: "/v1/chat/completions".to_string(), - protocol: "openai".to_string(), - url: format!("{display_base_url}/v1/chat/completions"), - }, - ], - tags: vec!["默认".to_string()], - enabled: true, - }]; - all_routes.extend(routes); - - let response = RouteListResponse { - base_url: display_base_url, - default_provider, - routes: all_routes, - }; - - Json(response) -} - -/// 带选择器的 Anthropic messages 处理 -async fn anthropic_messages_with_selector( - State(state): State, - Path(selector): Path, - headers: HeaderMap, - Json(request): Json, -) -> Response { - // 使用 Anthropic 格式的认证验证 - if let Err(e) = handlers::verify_api_key_anthropic(&headers, &state.api_key).await { - state.logs.write().await.add( - "warn", - &format!("Unauthorized request to /{selector}/v1/messages"), - ); - return e.into_response(); - } - - state.logs.write().await.add( - "info", - &format!( - "[REQ] POST /{}/v1/messages model={} stream={}", - selector, request.model, request.stream - ), - ); - - // 尝试解析凭证(不降级,指定什么就用什么) - let credential = match &state.db { - Some(db) => { - // 首先尝试按名称查找 - if let Ok(Some(cred)) = state.pool_service.get_by_name(db, &selector) { - Some(cred) - } - // 然后尝试按 UUID 查找 - else if let Ok(Some(cred)) = state.pool_service.get_by_uuid(db, &selector) { - Some(cred) - } - // 最后尝试按 provider 类型选择(不降级) - else if let Ok(Some(cred)) = - state - .pool_service - .select_credential(db, &selector, Some(&request.model)) - { - Some(cred) - } else { - None - } - } - None => None, - }; - - match credential { - Some(cred) => { - state.logs.write().await.add( - "info", - &format!( - "[ROUTE] Using credential: type={} name={:?} uuid={}", - cred.provider_type, - cred.name, - &cred.uuid[..8] - ), - ); - - // 根据凭证类型调用相应的 Provider - // 注意:这里没有 Flow 捕获,因为是通过 selector 路由的请求 - handlers::call_provider_anthropic(&state, &cred, &request, None).await - } - None => { - // 不再回退到默认 provider,直接返回错误 - state.logs.write().await.add( - "error", - &format!( - "[ROUTE] No available credentials for selector '{selector}', refusing to fallback" - ), - ); - ( - StatusCode::SERVICE_UNAVAILABLE, - Json(serde_json::json!({ - "error": { - "type": "provider_unavailable", - "message": format!("No available credentials for selector '{}'", selector) - } - })), - ) - .into_response() - } - } -} - -/// 带选择器的 OpenAI chat completions 处理 -async fn chat_completions_with_selector( - State(state): State, - Path(selector): Path, - headers: HeaderMap, - Json(request): Json, -) -> Response { - if let Err(e) = handlers::verify_api_key(&headers, &state.api_key).await { - state.logs.write().await.add( - "warn", - &format!("Unauthorized request to /{selector}/v1/chat/completions"), - ); - return e.into_response(); - } - - state.logs.write().await.add( - "info", - &format!( - "[REQ] POST /{}/v1/chat/completions model={} stream={}", - selector, request.model, request.stream - ), - ); - - // 尝试解析凭证(不降级,指定什么就用什么) - let credential = match &state.db { - Some(db) => { - if let Ok(Some(cred)) = state.pool_service.get_by_name(db, &selector) { - Some(cred) - } else if let Ok(Some(cred)) = state.pool_service.get_by_uuid(db, &selector) { - Some(cred) - } else if let Ok(Some(cred)) = - state - .pool_service - .select_credential(db, &selector, Some(&request.model)) - { - Some(cred) - } else { - None - } - } - None => None, - }; - - match credential { - Some(cred) => { - state.logs.write().await.add( - "info", - &format!( - "[ROUTE] Using credential: type={} name={:?} uuid={}", - cred.provider_type, - cred.name, - &cred.uuid[..8] - ), - ); - - // 注意:这里没有 Flow 捕获,因为是通过 selector 路由的请求 - handlers::call_provider_openai(&state, &cred, &request, None).await - } - None => { - // 不再回退到默认 provider,直接返回错误 - state.logs.write().await.add( - "error", - &format!( - "[ROUTE] No available credentials for selector '{selector}', refusing to fallback" - ), - ); - ( - StatusCode::SERVICE_UNAVAILABLE, - Json(serde_json::json!({ - "error": { - "message": format!("No available credentials for selector '{}'", selector), - "type": "provider_unavailable", - "code": "no_credentials" - } - })), - ) - .into_response() - } - } -} - -/// 内部 Anthropic messages 处理 (使用默认 Kiro) -/// 预留:用于内部直接调用 Kiro API -#[allow(dead_code)] -async fn anthropic_messages_internal( - state: &AppState, - request: &AnthropicMessagesRequest, -) -> Response { - // 检查 token - { - let _guard = state.kiro_refresh_lock.lock().await; - let mut kiro = state.kiro.write().await; - let needs_refresh = - kiro.credentials.access_token.is_none() || kiro.is_token_expiring_soon(); - if needs_refresh { - if let Err(e) = kiro.refresh_token().await { - state - .logs - .write() - .await - .add("error", &format!("[AUTH] Token refresh failed: {e}")); - return ( - StatusCode::UNAUTHORIZED, - Json(serde_json::json!({"error": {"message": format!("Token refresh failed: {e}")}})), - ) - .into_response(); - } - } - } - - let openai_request = convert_anthropic_to_openai(request); - let kiro = state.kiro.read().await; - - match kiro.call_api(&openai_request).await { - Ok(resp) => { - let status = resp.status(); - if status.is_success() { - match resp.bytes().await { - Ok(bytes) => { - let body = String::from_utf8_lossy(&bytes).to_string(); - let parsed = parse_cw_response(&body); - if request.stream { - build_anthropic_stream_response(&request.model, &parsed) - } else { - build_anthropic_response(&request.model, &parsed) - } - } - Err(e) => ( - StatusCode::INTERNAL_SERVER_ERROR, - Json(serde_json::json!({"error": {"message": e.to_string()}})), - ) - .into_response(), - } - } else { - let body = resp.text().await.unwrap_or_default(); - ( - StatusCode::from_u16(status.as_u16()).unwrap_or(StatusCode::INTERNAL_SERVER_ERROR), - Json(serde_json::json!({"error": {"message": format!("Upstream error: {}", body)}})), - ) - .into_response() - } - } - Err(e) => ( - StatusCode::INTERNAL_SERVER_ERROR, - Json(serde_json::json!({"error": {"message": e.to_string()}})), - ) - .into_response(), - } -} - -/// 内部 OpenAI chat completions 处理 (使用默认 Kiro) -/// 预留:用于内部直接调用 Kiro API -#[allow(dead_code)] -async fn chat_completions_internal(state: &AppState, request: &ChatCompletionRequest) -> Response { - { - let _guard = state.kiro_refresh_lock.lock().await; - let mut kiro = state.kiro.write().await; - let needs_refresh = - kiro.credentials.access_token.is_none() || kiro.is_token_expiring_soon(); - if needs_refresh { - if let Err(e) = kiro.refresh_token().await { - return ( - StatusCode::UNAUTHORIZED, - Json(serde_json::json!({"error": {"message": format!("Token refresh failed: {e}")}})), - ) - .into_response(); - } - } - } - - let kiro = state.kiro.read().await; - match kiro.call_api(request).await { - Ok(resp) => { - let status = resp.status(); - if status.is_success() { - match resp.text().await { - Ok(body) => { - let parsed = parse_cw_response(&body); - let has_tool_calls = !parsed.tool_calls.is_empty(); - - let message = if has_tool_calls { - serde_json::json!({ - "role": "assistant", - "content": if parsed.content.is_empty() { serde_json::Value::Null } else { serde_json::json!(parsed.content) }, - "tool_calls": parsed.tool_calls.iter().map(|tc| { - serde_json::json!({ - "id": tc.id, - "type": "function", - "function": { - "name": tc.function.name, - "arguments": tc.function.arguments - } - }) - }).collect::>() - }) - } else { - serde_json::json!({ - "role": "assistant", - "content": parsed.content - }) - }; - - let response = serde_json::json!({ - "id": format!("chatcmpl-{}", uuid::Uuid::new_v4()), - "object": "chat.completion", - "created": std::time::SystemTime::now() - .duration_since(std::time::UNIX_EPOCH) - .unwrap_or_default() - .as_secs(), - "model": request.model, - "choices": [{ - "index": 0, - "message": message, - "finish_reason": if has_tool_calls { "tool_calls" } else { "stop" } - }], - "usage": { - "prompt_tokens": 0, - "completion_tokens": 0, - "total_tokens": 0 - } - }); - Json(response).into_response() - } - Err(e) => ( - StatusCode::INTERNAL_SERVER_ERROR, - Json(serde_json::json!({"error": {"message": e.to_string()}})), - ) - .into_response(), - } - } else { - let body = resp.text().await.unwrap_or_default(); - ( - StatusCode::from_u16(status.as_u16()).unwrap_or(StatusCode::INTERNAL_SERVER_ERROR), - Json(serde_json::json!({"error": {"message": format!("Upstream error: {}", body)}})), - ) - .into_response() - } - } - Err(e) => ( - StatusCode::INTERNAL_SERVER_ERROR, - Json(serde_json::json!({"error": {"message": e.to_string()}})), - ) - .into_response(), - } -} +pub use proxycast_server::*;