diff --git a/src-tauri/Cargo.lock b/src-tauri/Cargo.lock index ac7bb9e9f..ff94b028e 100644 --- a/src-tauri/Cargo.lock +++ b/src-tauri/Cargo.lock @@ -1289,6 +1289,16 @@ dependencies = [ "tokio", ] +[[package]] +name = "fs2" +version = "0.4.3" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "9564fc758e15025b46aa6643b1b77d047d1a56a1aea6e01002ac0c7026876213" +dependencies = [ + "libc", + "winapi", +] + [[package]] name = "fs_extra" version = "1.3.0" @@ -3378,6 +3388,8 @@ dependencies = [ "chrono", "dashmap", "dirs 5.0.1", + "flate2", + "fs2", "futures", "indexmap 2.12.1", "md5", @@ -3395,6 +3407,7 @@ dependencies = [ "serde_urlencoded", "serde_yaml", "sha2", + "subtle", "tauri", "tauri-build", "tauri-plugin-autostart", diff --git a/src-tauri/Cargo.toml b/src-tauri/Cargo.toml index b0879f7df..3b3f59eb6 100644 --- a/src-tauri/Cargo.toml +++ b/src-tauri/Cargo.toml @@ -38,7 +38,10 @@ async-stream = "0.3" regex = "1" md5 = "0.7" urlencoding = "2" -rusqlite = { version = "0.31", features = ["bundled"] } +subtle = "2.5" +flate2 = "1" +fs2 = "0.4" +rusqlite = { version = "0.31", features = ["bundled", "backup"] } serde_yaml = "0.9" indexmap = { version = "2", features = ["serde"] } zip = "0.6" diff --git a/src-tauri/src/config/mod.rs b/src-tauri/src/config/mod.rs index a8d73a615..fae2fa43d 100644 --- a/src-tauri/src/config/mod.rs +++ b/src-tauri/src/config/mod.rs @@ -21,11 +21,12 @@ pub use hot_reload::{ pub use import::{ImportError, ImportOptions, ImportResult, ImportService, ValidationResult}; pub use path_utils::{collapse_tilde, contains_tilde, expand_tilde}; pub use types::{ - is_default_api_key, AmpConfig, AmpModelMapping, ApiKeyEntry, Config, CredentialEntry, - CredentialPoolConfig, CustomProviderConfig, GeminiApiKeyEntry, IFlowCredentialEntry, - InjectionRuleConfig, InjectionSettings, LoggingConfig, ProviderConfig, ProvidersConfig, - QuotaExceededConfig, RemoteManagementConfig, RetrySettings, RoutingConfig, RoutingRuleConfig, - ServerConfig, TlsConfig, VertexApiKeyEntry, VertexModelAlias, DEFAULT_API_KEY, + generate_secure_api_key, is_default_api_key, AmpConfig, AmpModelMapping, ApiKeyEntry, Config, + CredentialEntry, CredentialPoolConfig, CustomProviderConfig, GeminiApiKeyEntry, + IFlowCredentialEntry, InjectionRuleConfig, InjectionSettings, LoggingConfig, ProviderConfig, + ProvidersConfig, QuotaExceededConfig, RemoteManagementConfig, RetrySettings, RoutingConfig, + RoutingRuleConfig, ServerConfig, TlsConfig, VertexApiKeyEntry, VertexModelAlias, + DEFAULT_API_KEY, }; pub use yaml::{ load_config, save_config, save_config_yaml, ConfigError, ConfigManager, YamlService, diff --git a/src-tauri/src/config/types.rs b/src-tauri/src/config/types.rs index eb4bddc38..5e636855e 100644 --- a/src-tauri/src/config/types.rs +++ b/src-tauri/src/config/types.rs @@ -340,6 +340,19 @@ fn default_api_key() -> String { DEFAULT_API_KEY.to_string() } +/// 生成安全 API Key(32 字节随机) +pub fn generate_secure_api_key() -> String { + use rand::distributions::Alphanumeric; + use rand::Rng; + + let token: String = rand::thread_rng() + .sample_iter(&Alphanumeric) + .take(32) + .map(char::from) + .collect(); + format!("pc_{token}") +} + /// 是否为默认 API Key pub fn is_default_api_key(api_key: &str) -> bool { api_key == DEFAULT_API_KEY diff --git a/src-tauri/src/config/yaml.rs b/src-tauri/src/config/yaml.rs index b2bbda689..71b308ac7 100644 --- a/src-tauri/src/config/yaml.rs +++ b/src-tauri/src/config/yaml.rs @@ -104,6 +104,10 @@ impl ConfigManager { std::fs::create_dir_all(parent).map_err(|e| ConfigError::WriteError(e.to_string()))?; } + if path.exists() { + let backup_path = path.with_extension("yaml.backup"); + let _ = std::fs::copy(path, backup_path); + } let yaml = Self::to_yaml(&self.config)?; std::fs::write(path, yaml).map_err(|e| ConfigError::WriteError(e.to_string())) } @@ -656,26 +660,57 @@ fn json_config_path() -> std::path::PathBuf { /// 加载配置(向后兼容) /// /// 优先加载 YAML 配置,如果不存在则尝试加载 JSON 配置 +/// 首次启动时自动生成强随机 API Key 并保存配置 pub fn load_config() -> Result> { + use super::types::{generate_secure_api_key, is_default_api_key}; + let yaml_path = ConfigManager::default_config_path(); let json_path = json_config_path(); // 优先尝试 YAML 配置 if yaml_path.exists() { let content = std::fs::read_to_string(&yaml_path)?; - let config = serde_yaml::from_str(&content)?; + let mut config: Config = serde_yaml::from_str(&content)?; + // 如果配置中使用默认 API Key,生成强随机 Key 并保存 + if is_default_api_key(&config.server.api_key) { + let new_key = generate_secure_api_key(); + tracing::warn!("[CONFIG] 检测到默认 API Key,已自动生成强随机 Key"); + config.server.api_key = new_key; + // 保存更新后的配置 + if let Err(e) = save_config_yaml(&config) { + tracing::error!("[CONFIG] 保存配置失败: {}", e); + } + } return Ok(config); } // 回退到 JSON 配置 if json_path.exists() { let content = std::fs::read_to_string(&json_path)?; - let config = serde_json::from_str(&content)?; + let mut config: Config = serde_json::from_str(&content)?; + // 如果配置中使用默认 API Key,生成强随机 Key 并保存 + if is_default_api_key(&config.server.api_key) { + let new_key = generate_secure_api_key(); + tracing::warn!("[CONFIG] 检测到默认 API Key,已自动生成强随机 Key"); + config.server.api_key = new_key; + // 保存更新后的配置(迁移到 YAML) + if let Err(e) = save_config_yaml(&config) { + tracing::error!("[CONFIG] 保存配置失败: {}", e); + } + } return Ok(config); } - // 都不存在,返回默认配置 - Ok(Config::default()) + // 都不存在,创建默认配置并生成强随机 API Key + let mut config = Config::default(); + let new_key = generate_secure_api_key(); + tracing::info!("[CONFIG] 首次启动,已生成强随机 API Key"); + config.server.api_key = new_key; + // 保存初始配置 + if let Err(e) = save_config_yaml(&config) { + tracing::error!("[CONFIG] 保存初始配置失败: {}", e); + } + Ok(config) } /// 保存配置(同时写入 YAML 与 JSON,兼容旧版) @@ -699,6 +734,10 @@ pub fn save_config_yaml(config: &Config) -> Result<(), Box>; pub type LogState = Arc>; fn generate_api_key() -> String { - let token: String = rand::thread_rng() - .sample_iter(&Alphanumeric) - .take(32) - .map(char::from) - .collect(); - format!("pc_{token}") + config::generate_secure_api_key() } #[tauri::command] diff --git a/src-tauri/src/logger.rs b/src-tauri/src/logger.rs index 615b89611..88e83afa5 100644 --- a/src-tauri/src/logger.rs +++ b/src-tauri/src/logger.rs @@ -1,9 +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::Write; +use std::io::{Read, Write}; use std::path::PathBuf; use std::sync::Arc; use tokio::sync::RwLock; @@ -79,11 +82,12 @@ impl LogStore { } 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: message.to_string(), + message: sanitized.clone(), }; self.logs.push_back(entry.clone()); @@ -93,7 +97,7 @@ impl LogStore { 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(), message); + 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()); @@ -113,6 +117,7 @@ impl LogStore { 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) @@ -120,7 +125,7 @@ impl LogStore { .write(true) .open(&raw_file) { - let _ = file.write_all(body.as_bytes()); + let _ = file.write_all(sanitized.as_bytes()); } } } @@ -163,6 +168,7 @@ impl LogStore { let Some(dir) = path.parent() else { return; }; + self.archive_old_logs(path); let Ok(entries) = fs::read_dir(dir) else { return; }; @@ -190,7 +196,88 @@ 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); + + 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>; + +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: ***"), + ]; + + 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 +} diff --git a/src-tauri/src/middleware/management_auth.rs b/src-tauri/src/middleware/management_auth.rs index 55c36bc40..5ba4fdd99 100644 --- a/src-tauri/src/middleware/management_auth.rs +++ b/src-tauri/src/middleware/management_auth.rs @@ -24,6 +24,7 @@ use std::{ task::{Context, Poll}, time::{Duration, Instant}, }; +use subtle::ConstantTimeEq; use tower::{Layer, Service}; const MAX_AUTH_FAILURES: u32 = 5; @@ -182,6 +183,10 @@ impl ManagementAuthService { let mut map = failure_map().lock().unwrap(); map.remove(client_id); } + + fn secret_key_matches(provided: &str, expected: &str) -> bool { + provided.as_bytes().ct_eq(expected.as_bytes()).into() + } } impl Service> for ManagementAuthService @@ -240,7 +245,7 @@ where // 3. 验证 secret_key let provided_key = Self::extract_secret_key(&req); match provided_key { - Some(key) if key == secret_key => { + Some(key) if Self::secret_key_matches(&key, &secret_key) => { // 认证成功,继续处理请求 tracing::debug!("[MANAGEMENT_AUTH] Auth successful from {:?}", client_addr); Self::record_success(&client_id); diff --git a/src-tauri/src/middleware/tests.rs b/src-tauri/src/middleware/tests.rs index 5c817138c..4c323cd70 100644 --- a/src-tauri/src/middleware/tests.rs +++ b/src-tauri/src/middleware/tests.rs @@ -134,15 +134,16 @@ fn test_management_auth_rate_limit_after_failures() { let mut service = layer.layer(MockService); let rt = tokio::runtime::Runtime::new().unwrap(); - let client_ip = "203.0.113.10"; + // 使用唯一的 IP 地址避免测试间干扰 + let client_ip = format!("203.0.113.{}", std::process::id() % 256); for _ in 0..5 { let req = - create_request_with_management_key_and_forwarded(Some("invalid"), Some(client_ip)); + create_request_with_management_key_and_forwarded(Some("invalid"), Some(&client_ip)); let response = rt.block_on(async { service.call(req).await.unwrap() }); assert_eq!(response.status(), StatusCode::UNAUTHORIZED); } - let req = create_request_with_management_key_and_forwarded(Some("invalid"), Some(client_ip)); + let req = create_request_with_management_key_and_forwarded(Some("invalid"), Some(&client_ip)); let response = rt.block_on(async { service.call(req).await.unwrap() }); assert_eq!(response.status(), StatusCode::TOO_MANY_REQUESTS); } diff --git a/src-tauri/src/processor/error.rs b/src-tauri/src/processor/error.rs index 7b5cf461b..3bd89675e 100644 --- a/src-tauri/src/processor/error.rs +++ b/src-tauri/src/processor/error.rs @@ -126,6 +126,18 @@ impl ProcessError { ProcessError::Cancelled => "cancelled", } } + + /// 记录带上下文的错误日志 + pub fn log_with_context(&self, request_id: &str, provider: &str, model: &str) { + tracing::error!( + request_id = %request_id, + provider = %provider, + model = %model, + error_type = %self.error_type(), + error_message = %self.to_string(), + "Request processing failed" + ); + } } #[cfg(test)] diff --git a/src-tauri/src/processor/mod.rs b/src-tauri/src/processor/mod.rs index 8ec3dfaa7..2cfb48b5b 100644 --- a/src-tauri/src/processor/mod.rs +++ b/src-tauri/src/processor/mod.rs @@ -58,6 +58,8 @@ pub struct RequestProcessor { pub tokens: Arc>, /// 凭证池服务 pub pool_service: Arc, + /// 热重载协调锁(避免配置更新期间请求读取不一致的配置) + pub reload_lock: Arc>, } impl RequestProcessor { @@ -85,6 +87,7 @@ impl RequestProcessor { stats, tokens, pool_service, + reload_lock: Arc::new(RwLock::new(())), } } @@ -102,6 +105,7 @@ impl RequestProcessor { stats: Arc::new(ParkingLotRwLock::new(StatsAggregator::with_defaults())), tokens: Arc::new(ParkingLotRwLock::new(TokenTracker::with_defaults())), pool_service, + reload_lock: Arc::new(RwLock::new(())), } } @@ -126,6 +130,7 @@ impl RequestProcessor { stats, tokens, pool_service, + reload_lock: Arc::new(RwLock::new(())), } } diff --git a/src-tauri/src/processor/steps/auth.rs b/src-tauri/src/processor/steps/auth.rs index d67559b79..fd302ad8e 100644 --- a/src-tauri/src/processor/steps/auth.rs +++ b/src-tauri/src/processor/steps/auth.rs @@ -5,6 +5,7 @@ use super::traits::{PipelineStep, StepError}; use crate::processor::RequestContext; use async_trait::async_trait; +use subtle::ConstantTimeEq; /// 认证步骤 /// @@ -34,7 +35,7 @@ impl AuthStep { /// 验证 API Key pub fn verify(&self, provided_key: Option<&str>) -> Result<(), StepError> { match provided_key { - Some(key) if key == self.expected_key => Ok(()), + Some(key) if key.as_bytes().ct_eq(self.expected_key.as_bytes()).into() => Ok(()), Some(_) => Err(StepError::Auth("Invalid API key".to_string())), None => Err(StepError::Auth("No API key provided".to_string())), } diff --git a/src-tauri/src/server.rs b/src-tauri/src/server.rs index 1c60e3253..332aff8f2 100644 --- a/src-tauri/src/server.rs +++ b/src-tauri/src/server.rs @@ -23,6 +23,7 @@ use crate::providers::kiro::KiroProvider; use crate::providers::openai_custom::OpenAICustomProvider; use crate::providers::qwen::QwenProvider; use crate::providers::vertex::VertexProvider; +use crate::services::backup_service::BackupService; use crate::services::provider_pool_service::ProviderPoolService; use crate::services::token_cache_service::TokenCacheService; use crate::telemetry::{RequestLog, RequestStatus}; @@ -35,10 +36,14 @@ use axum::{ routing::{get, post}, Json, Router, }; +use chrono::Utc; +use fs2::available_space; use futures::stream; use serde::{Deserialize, Serialize}; +use std::collections::HashMap; use std::path::PathBuf; use std::sync::Arc; +use subtle::ConstantTimeEq; use tokio::sync::{mpsc, oneshot, RwLock}; /// 安全截断字符串到指定字符数,避免 UTF-8 边界问题 @@ -69,6 +74,13 @@ fn message_content_len(content: &crate::models::openai::MessageContent) -> usize } } +fn api_key_matches(provided_key: &str, expected_key: &str) -> bool { + provided_key + .as_bytes() + .ct_eq(expected_key.as_bytes()) + .into() +} + /// 记录请求统计到遥测系统 fn record_request_telemetry( state: &AppState, @@ -290,6 +302,8 @@ impl ServerState { return Err("当前版本未启用 TLS,禁止开启远程管理".into()); } + tracing::warn!("当前未启用 TLS,生产环境请使用反向代理终止 HTTPS"); + // 重新加载凭证 let _ = self.kiro_provider.load_credentials().await; let kiro = self.kiro_provider.clone(); @@ -404,6 +418,8 @@ struct AppState { request_logger: Option>, /// Amp CLI 路由器 amp_router: Arc, + /// 备份服务 + backup_service: Option>, } /// 启动配置文件监控 @@ -553,6 +569,7 @@ async fn start_config_watcher( /// - 更新过程不会阻塞新请求的处理 /// - 现有连接不受影响 async fn update_processor_config(processor: &RequestProcessor, config: &Config) { + let _reload_guard = processor.reload_lock.write().await; // 更新注入器规则 { let mut injector = processor.injector.write().await; @@ -734,6 +751,37 @@ async fn run_server( .unwrap_or_default(), )); + let backup_service = match BackupService::with_defaults() { + Ok(service) => { + tracing::info!( + "[BACKUP] 备份服务初始化成功,备份目录: {:?}", + service.backup_dir() + ); + Some(Arc::new(service)) + } + Err(e) => { + tracing::warn!("[BACKUP] 备份服务初始化失败,自动备份将不可用: {}", e); + None + } + }; + if let Some(service) = backup_service.clone() { + let db_for_backup = db.clone(); + tokio::spawn(async move { + let mut ticker = tokio::time::interval(std::time::Duration::from_secs(24 * 60 * 60)); + loop { + ticker.tick().await; + let result = match &db_for_backup { + Some(db) => service.backup_database_with_connection(db), + None => service.backup_database(), + }; + match result { + Ok(path) => tracing::info!("[BACKUP] 自动备份成功: {:?}", path), + Err(err) => tracing::warn!("[BACKUP] 自动备份失败: {}", err), + } + } + }); + } + let state = AppState { api_key: api_key.to_string(), base_url, @@ -757,6 +805,7 @@ async fn run_server( hot_reload_manager: hot_reload_manager.clone(), request_logger: shared_logger, amp_router, + backup_service, }; // 启动配置文件监控 @@ -785,6 +834,8 @@ async fn run_server( let management_routes = Router::new() .route("/v0/management/status", get(management_status)) + .route("/v0/management/backup", post(management_backup)) + .route("/v0/management/restore", post(management_restore)) .route( "/v0/management/credentials", get(management_list_credentials), @@ -804,6 +855,7 @@ async fn run_server( let app = Router::new() .route("/health", get(health)) + .route("/ready", get(readiness)) .route("/v1/models", get(models)) .route("/v1/routes", get(list_routes)) .route("/v1/chat/completions", post(chat_completions)) @@ -858,43 +910,213 @@ async fn run_server( Ok(()) } +#[derive(Debug, Serialize)] +struct CheckResult { + status: String, + message: Option, + latency_ms: Option, +} + +#[derive(Debug, Serialize)] +struct HealthStatus { + status: String, + timestamp: chrono::DateTime, + version: String, + checks: HashMap, +} + async fn health(State(state): State) -> impl IntoResponse { - let (db_ready, pool_stats) = match &state.db { - Some(db) => { - let db_ok = db - .lock() - .map(|conn| conn.query_row::("SELECT 1", [], |row| row.get(0))) - .is_ok(); - let stats = if db_ok { - state.pool_service.get_overview(db).ok().map(|items| { - let mut total = 0usize; - let mut healthy = 0usize; - let mut disabled = 0usize; - for item in items { - total += item.stats.total_count; - healthy += item.stats.healthy_count; - disabled += item.stats.disabled_count; - } - serde_json::json!({ - "total": total, - "healthy": healthy, - "disabled": disabled - }) - }) - } else { - None - }; - (db_ok, stats) - } - None => (false, None), + let mut checks = HashMap::new(); + + let db_check = check_database(&state).await; + checks.insert("database".to_string(), db_check); + + let pool_check = check_credential_pool(&state).await; + checks.insert("credential_pool".to_string(), pool_check); + + let disk_check = check_disk_space(&state).await; + checks.insert("disk_space".to_string(), disk_check); + + let log_check = check_log_directory(&state).await; + checks.insert("log_directory".to_string(), log_check); + + let overall_status = if checks.values().all(|c| c.status == "healthy") { + "healthy" + } else if checks.values().any(|c| c.status == "unhealthy") { + "unhealthy" + } else { + "degraded" }; - Json(serde_json::json!({ - "status": "healthy", - "version": env!("CARGO_PKG_VERSION"), - "db_ready": db_ready, - "pool": pool_stats - })) + let status_code = match overall_status { + "healthy" | "degraded" => StatusCode::OK, + "unhealthy" => StatusCode::SERVICE_UNAVAILABLE, + _ => StatusCode::INTERNAL_SERVER_ERROR, + }; + + ( + status_code, + Json(HealthStatus { + status: overall_status.to_string(), + timestamp: Utc::now(), + version: env!("CARGO_PKG_VERSION").to_string(), + checks, + }), + ) +} + +async fn readiness(State(state): State) -> impl IntoResponse { + let db_ok = matches!(check_database(&state).await.status.as_str(), "healthy"); + let pool_ok = matches!( + check_credential_pool(&state).await.status.as_str(), + "healthy" + ); + + if !db_ok || !pool_ok { + return ( + StatusCode::SERVICE_UNAVAILABLE, + Json(serde_json::json!({ + "ready": false, + "reason": "Database or credential pool not ready" + })), + ); + } + + ( + StatusCode::OK, + Json(serde_json::json!({ + "ready": true + })), + ) +} + +async fn check_database(state: &AppState) -> CheckResult { + let start = std::time::Instant::now(); + let Some(db) = &state.db else { + return CheckResult { + status: "unhealthy".to_string(), + message: Some("database not initialized".to_string()), + latency_ms: None, + }; + }; + let ok = db + .lock() + .map(|conn| conn.query_row::("SELECT 1", [], |row| row.get(0))) + .is_ok(); + + CheckResult { + status: if ok { "healthy" } else { "unhealthy" }.to_string(), + message: if ok { + None + } else { + Some("database query failed".to_string()) + }, + latency_ms: Some(start.elapsed().as_millis() as u64), + } +} + +async fn check_credential_pool(state: &AppState) -> CheckResult { + let start = std::time::Instant::now(); + let Some(db) = &state.db else { + return CheckResult { + status: "unhealthy".to_string(), + message: Some("database not initialized".to_string()), + latency_ms: None, + }; + }; + + let stats = match state.pool_service.get_overview(db) { + Ok(items) => { + let mut total = 0usize; + let mut healthy = 0usize; + let mut disabled = 0usize; + for item in items { + total += item.stats.total_count; + healthy += item.stats.healthy_count; + disabled += item.stats.disabled_count; + } + (total, healthy, disabled) + } + Err(_) => (0, 0, 0), + }; + + let (total, healthy, disabled) = stats; + let status = if healthy > 0 { + "healthy" + } else if total > 0 { + "degraded" + } else { + "unhealthy" + }; + + CheckResult { + status: status.to_string(), + message: Some(format!( + "total={} healthy={} disabled={}", + total, healthy, disabled + )), + latency_ms: Some(start.elapsed().as_millis() as u64), + } +} + +async fn check_disk_space(state: &AppState) -> CheckResult { + let start = std::time::Instant::now(); + let Some(log_path) = state.logs.read().await.get_log_file_path() else { + return CheckResult { + status: "degraded".to_string(), + message: Some("log path not available".to_string()), + latency_ms: None, + }; + }; + let log_path = PathBuf::from(log_path); + let dir = log_path.parent().unwrap_or(log_path.as_path()); + + let available = available_space(dir).unwrap_or(0); + let available_gb = available / (1024 * 1024 * 1024); + let status = if available_gb >= 10 { + "healthy" + } else if available_gb >= 1 { + "degraded" + } else { + "unhealthy" + }; + + CheckResult { + status: status.to_string(), + message: Some(format!("available_gb={}", available_gb)), + latency_ms: Some(start.elapsed().as_millis() as u64), + } +} + +async fn check_log_directory(state: &AppState) -> CheckResult { + let start = std::time::Instant::now(); + let Some(log_path) = state.logs.read().await.get_log_file_path() else { + return CheckResult { + status: "degraded".to_string(), + message: Some("log path not available".to_string()), + latency_ms: None, + }; + }; + let log_path = PathBuf::from(log_path); + let dir = log_path.parent().unwrap_or(log_path.as_path()); + + let test_file = dir.join(".proxycast_write_check"); + let writable = std::fs::OpenOptions::new() + .create(true) + .write(true) + .open(&test_file) + .and_then(|_| std::fs::remove_file(&test_file)) + .is_ok(); + + CheckResult { + status: if writable { "healthy" } else { "unhealthy" }.to_string(), + message: if writable { + None + } else { + Some("log directory not writable".to_string()) + }, + latency_ms: Some(start.elapsed().as_millis() as u64), + } } async fn models() -> impl IntoResponse { @@ -939,7 +1161,7 @@ async fn verify_api_key( } }; - if key != expected_key { + if !api_key_matches(key, expected_key) { return Err(( StatusCode::UNAUTHORIZED, Json(serde_json::json!({"error": {"message": "Invalid API key"}})), @@ -977,7 +1199,7 @@ async fn verify_api_key_anthropic( } }; - if key != expected_key { + if !api_key_matches(key, expected_key) { return Err(( StatusCode::UNAUTHORIZED, Json(serde_json::json!({ @@ -998,6 +1220,7 @@ async fn chat_completions( headers: HeaderMap, Json(mut request): Json, ) -> Response { + let _reload_guard = state.processor.reload_lock.read().await; if let Err(e) = verify_api_key(&headers, &state.api_key).await { state .logs @@ -1409,6 +1632,7 @@ async fn anthropic_messages( headers: HeaderMap, Json(mut request): Json, ) -> Response { + let _reload_guard = state.processor.reload_lock.read().await; // 使用 Anthropic 格式的认证验证(优先检查 x-api-key) if let Err(e) = verify_api_key_anthropic(&headers, &state.api_key).await { state @@ -2077,6 +2301,7 @@ async fn count_tokens( headers: HeaderMap, Json(_request): Json, ) -> Response { + let _reload_guard = state.processor.reload_lock.read().await; if let Err(e) = verify_api_key(&headers, &state.api_key).await { return e.into_response(); } @@ -2393,6 +2618,7 @@ async fn anthropic_messages_with_selector( headers: HeaderMap, Json(request): Json, ) -> Response { + let _reload_guard = state.processor.reload_lock.read().await; // 使用 Anthropic 格式的认证验证 if let Err(e) = verify_api_key_anthropic(&headers, &state.api_key).await { state.logs.write().await.add( @@ -2472,6 +2698,7 @@ async fn chat_completions_with_selector( headers: HeaderMap, Json(request): Json, ) -> Response { + let _reload_guard = state.processor.reload_lock.read().await; if let Err(e) = verify_api_key(&headers, &state.api_key).await { state.logs.write().await.add( "warn", @@ -2547,6 +2774,7 @@ async fn amp_chat_completions( headers: HeaderMap, Json(mut request): Json, ) -> Response { + let _reload_guard = state.processor.reload_lock.read().await; if let Err(e) = verify_api_key(&headers, &state.api_key).await { state.logs.write().await.add( "warn", @@ -2641,6 +2869,7 @@ async fn amp_messages( headers: HeaderMap, Json(mut request): Json, ) -> Response { + let _reload_guard = state.processor.reload_lock.read().await; // 使用 Anthropic 格式的认证验证 if let Err(e) = verify_api_key_anthropic(&headers, &state.api_key).await { state.logs.write().await.add( @@ -4036,7 +4265,7 @@ async fn ws_upgrade_handler( } }; - if key != state.api_key { + if !api_key_matches(key, &state.api_key) { return axum::http::Response::builder() .status(401) .body(Body::from("Invalid API key")) @@ -4249,6 +4478,7 @@ async fn handle_ws_chat_completions( request_id: &str, mut request: ChatCompletionRequest, ) -> WsProtoMessage { + let _reload_guard = state.processor.reload_lock.read().await; // 创建请求上下文 let mut ctx = RequestContext::new(request.model.clone()).with_stream(request.stream); @@ -4382,6 +4612,7 @@ async fn handle_ws_anthropic_messages( request_id: &str, mut request: AnthropicMessagesRequest, ) -> WsProtoMessage { + let _reload_guard = state.processor.reload_lock.read().await; // 创建请求上下文 let mut ctx = RequestContext::new(request.model.clone()).with_stream(request.stream); @@ -4957,6 +5188,18 @@ pub struct UpdateConfigResponse { pub message: String, } +#[derive(Debug, Clone, Serialize, Deserialize)] +pub struct BackupResponse { + pub success: bool, + pub message: String, + pub backup_path: Option, +} + +#[derive(Debug, Clone, Deserialize)] +pub struct RestoreRequest { + pub backup_path: String, +} + fn snapshot_config(state: &AppState) -> Option { if let Some(manager) = &state.config_manager { if let Ok(guard) = manager.read() { @@ -4989,6 +5232,86 @@ pub async fn management_status(State(state): State) -> impl IntoRespon Json(response) } +/// POST /v0/management/backup - 触发数据库备份 +pub async fn management_backup(State(state): State) -> impl IntoResponse { + let Some(service) = &state.backup_service else { + return ( + StatusCode::SERVICE_UNAVAILABLE, + Json(BackupResponse { + success: false, + message: "Backup service not available".to_string(), + backup_path: None, + }), + ); + }; + + let result = match &state.db { + Some(db) => service.backup_database_with_connection(db), + None => service.backup_database(), + }; + + match result { + Ok(path) => ( + StatusCode::OK, + Json(BackupResponse { + success: true, + message: "Backup created".to_string(), + backup_path: Some(path.to_string_lossy().to_string()), + }), + ), + Err(err) => ( + StatusCode::INTERNAL_SERVER_ERROR, + Json(BackupResponse { + success: false, + message: err, + backup_path: None, + }), + ), + } +} + +/// POST /v0/management/restore - 从备份恢复数据库 +pub async fn management_restore( + State(state): State, + Json(request): Json, +) -> impl IntoResponse { + let Some(service) = &state.backup_service else { + return ( + StatusCode::SERVICE_UNAVAILABLE, + Json(BackupResponse { + success: false, + message: "Backup service not available".to_string(), + backup_path: None, + }), + ); + }; + + let backup_path = PathBuf::from(request.backup_path); + let result = match &state.db { + Some(db) => service.restore_database_with_connection(db, &backup_path), + None => service.restore_database(&backup_path), + }; + + match result { + Ok(()) => ( + StatusCode::OK, + Json(BackupResponse { + success: true, + message: "Restore completed".to_string(), + backup_path: None, + }), + ), + Err(err) => ( + StatusCode::INTERNAL_SERVER_ERROR, + Json(BackupResponse { + success: false, + message: err, + backup_path: None, + }), + ), + } +} + /// GET /v0/management/credentials - 获取凭证列表 pub async fn management_list_credentials(State(state): State) -> impl IntoResponse { let mut credentials = Vec::new(); diff --git a/src-tauri/src/services/backup_service.rs b/src-tauri/src/services/backup_service.rs new file mode 100644 index 000000000..cca7639ae --- /dev/null +++ b/src-tauri/src/services/backup_service.rs @@ -0,0 +1,117 @@ +//! 备份服务 +//! +//! 提供数据库与配置备份的基础能力 + +use crate::database::{get_db_path, DbConnection}; +use chrono::{DateTime, Duration, Utc}; +use rusqlite::DatabaseName; +use std::path::{Path, PathBuf}; + +#[derive(Clone)] +pub struct BackupService { + backup_dir: PathBuf, + retention_days: u32, +} + +impl BackupService { + pub fn new(backup_dir: PathBuf, retention_days: u32) -> Result { + std::fs::create_dir_all(&backup_dir) + .map_err(|e| format!("无法创建备份目录 {:?}: {}", backup_dir, e))?; + Ok(Self { + backup_dir, + retention_days, + }) + } + + pub fn with_defaults() -> Result { + let home = dirs::home_dir().ok_or_else(|| "无法获取主目录".to_string())?; + let backup_dir = home.join(".proxycast").join("backups"); + Self::new(backup_dir, 7) + } + + pub fn backup_database(&self) -> Result { + let db_path = get_db_path()?; + let timestamp = Utc::now().format("%Y%m%d_%H%M%S"); + let backup_path = self.backup_dir.join(format!("proxycast_{}.db", timestamp)); + + std::fs::copy(&db_path, &backup_path).map_err(|e| format!("备份失败: {}", e))?; + + self.cleanup_old_backups()?; + Ok(backup_path) + } + + pub fn backup_database_with_connection(&self, db: &DbConnection) -> Result { + let timestamp = Utc::now().format("%Y%m%d_%H%M%S"); + let backup_path = self.backup_dir.join(format!("proxycast_{}.db", timestamp)); + let conn = db.lock().map_err(|_| "数据库锁已被占用".to_string())?; + let progress: Option = None; + conn.backup(DatabaseName::Main, &backup_path, progress) + .map_err(|e| format!("备份失败: {}", e))?; + + self.cleanup_old_backups()?; + Ok(backup_path) + } + + pub fn restore_database(&self, backup_path: &Path) -> Result<(), String> { + if !backup_path.exists() { + return Err("备份文件不存在".to_string()); + } + let db_path = get_db_path()?; + std::fs::copy(backup_path, db_path).map_err(|e| format!("恢复失败: {}", e))?; + Ok(()) + } + + pub fn restore_database_with_connection( + &self, + db: &DbConnection, + backup_path: &Path, + ) -> Result<(), String> { + if !backup_path.exists() { + return Err("备份文件不存在".to_string()); + } + let mut conn = db.lock().map_err(|_| "数据库锁已被占用".to_string())?; + let progress: Option = None; + conn.restore(DatabaseName::Main, backup_path, progress) + .map_err(|e| format!("恢复失败: {}", e))?; + Ok(()) + } + + pub fn list_backups(&self) -> Result, String> { + let mut backups = Vec::new(); + let entries = + std::fs::read_dir(&self.backup_dir).map_err(|e| format!("无法读取备份目录: {}", e))?; + for entry in entries.flatten() { + let path = entry.path(); + if path.extension().map(|e| e == "db").unwrap_or(false) { + backups.push(path); + } + } + backups.sort(); + Ok(backups) + } + + pub fn cleanup_old_backups(&self) -> Result<(), String> { + let entries = + std::fs::read_dir(&self.backup_dir).map_err(|e| format!("无法读取备份目录: {}", e))?; + let cutoff = Utc::now() - Duration::days(self.retention_days as i64); + + for entry in entries.flatten() { + let path = entry.path(); + let Ok(metadata) = entry.metadata() else { + continue; + }; + let Ok(modified) = metadata.modified() else { + continue; + }; + let modified = DateTime::::from(modified); + if modified < cutoff { + let _ = std::fs::remove_file(path); + } + } + Ok(()) + } + + pub fn backup_dir(&self) -> &PathBuf { + &self.backup_dir + } +} diff --git a/src-tauri/src/services/mod.rs b/src-tauri/src/services/mod.rs index 07768fb05..2cd31f0d4 100644 --- a/src-tauri/src/services/mod.rs +++ b/src-tauri/src/services/mod.rs @@ -1,3 +1,4 @@ +pub mod backup_service; pub mod live_sync; pub mod mcp_service; pub mod mcp_sync;