feat: 添加备份与健康检查能力

This commit is contained in:
jiesen
2025-12-21 08:22:53 +07:00
parent 502096a208
commit baa724823c
15 changed files with 678 additions and 63 deletions
+13
View File
@@ -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",
+4 -1
View File
@@ -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"
+6 -5
View File
@@ -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,
+13
View File
@@ -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
+43 -4
View File
@@ -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<Config, Box<dyn std::error::Error>> {
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<dyn std::error::Error
if let Some(parent) = path.parent() {
std::fs::create_dir_all(parent)?;
}
if path.exists() {
let backup_path = path.with_extension("yaml.backup");
let _ = std::fs::copy(&path, &backup_path);
}
let content = serde_yaml::to_string(config)?;
std::fs::write(&path, content)?;
Ok(())
+1 -7
View File
@@ -19,7 +19,6 @@ pub mod telemetry;
pub mod tray;
pub mod websocket;
use rand::{distributions::Alphanumeric, Rng};
use serde::{Deserialize, Serialize};
use std::sync::Arc;
use tauri::{Manager, Runtime};
@@ -177,12 +176,7 @@ pub type AppState = Arc<RwLock<server::ServerState>>;
pub type LogState = Arc<RwLock<logger::LogStore>>;
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]
+91 -4
View File
@@ -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::<Utc>::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<RwLock<LogStore>>;
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
}
+6 -1
View File
@@ -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<S> ManagementAuthService<S> {
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<S> Service<Request<Body>> for ManagementAuthService<S>
@@ -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);
+4 -3
View File
@@ -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);
}
+12
View File
@@ -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)]
+5
View File
@@ -58,6 +58,8 @@ pub struct RequestProcessor {
pub tokens: Arc<ParkingLotRwLock<TokenTracker>>,
/// 凭证池服务
pub pool_service: Arc<ProviderPoolService>,
/// 热重载协调锁(避免配置更新期间请求读取不一致的配置)
pub reload_lock: Arc<RwLock<()>>,
}
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(())),
}
}
+2 -1
View File
@@ -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())),
}
+360 -37
View File
@@ -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<Arc<crate::telemetry::RequestLogger>>,
/// Amp CLI 路由器
amp_router: Arc<crate::router::AmpRouter>,
/// 备份服务
backup_service: Option<Arc<BackupService>>,
}
/// 启动配置文件监控
@@ -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<String>,
latency_ms: Option<u64>,
}
#[derive(Debug, Serialize)]
struct HealthStatus {
status: String,
timestamp: chrono::DateTime<Utc>,
version: String,
checks: HashMap<String, CheckResult>,
}
async fn health(State(state): State<AppState>) -> impl IntoResponse {
let (db_ready, pool_stats) = match &state.db {
Some(db) => {
let db_ok = db
.lock()
.map(|conn| conn.query_row::<i32, _, _>("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<AppState>) -> 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::<i32, _, _>("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<ChatCompletionRequest>,
) -> 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<AnthropicMessagesRequest>,
) -> 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<serde_json::Value>,
) -> 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<AnthropicMessagesRequest>,
) -> 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<ChatCompletionRequest>,
) -> 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<ChatCompletionRequest>,
) -> 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<AnthropicMessagesRequest>,
) -> 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<String>,
}
#[derive(Debug, Clone, Deserialize)]
pub struct RestoreRequest {
pub backup_path: String,
}
fn snapshot_config(state: &AppState) -> Option<Config> {
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<AppState>) -> impl IntoRespon
Json(response)
}
/// POST /v0/management/backup - 触发数据库备份
pub async fn management_backup(State(state): State<AppState>) -> 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<AppState>,
Json(request): Json<RestoreRequest>,
) -> 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<AppState>) -> impl IntoResponse {
let mut credentials = Vec::new();
+117
View File
@@ -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<Self, String> {
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<Self, String> {
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<PathBuf, String> {
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<PathBuf, String> {
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<fn(rusqlite::backup::Progress)> = 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<fn(rusqlite::backup::Progress)> = None;
conn.restore(DatabaseName::Main, backup_path, progress)
.map_err(|e| format!("恢复失败: {}", e))?;
Ok(())
}
pub fn list_backups(&self) -> Result<Vec<PathBuf>, 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::<Utc>::from(modified);
if modified < cutoff {
let _ = std::fs::remove_file(path);
}
}
Ok(())
}
pub fn backup_dir(&self) -> &PathBuf {
&self.backup_dir
}
}
+1
View File
@@ -1,3 +1,4 @@
pub mod backup_service;
pub mod live_sync;
pub mod mcp_service;
pub mod mcp_sync;