mirror of
https://github.com/aiclientproxy/proxycast.git
synced 2026-09-24 23:10:56 +08:00
feat: 支持 API Key Provider 在 /v1/chat/completions 和 /v1/messages 端点
- 添加 get_enabled_api_keys_by_type 方法按 Provider 类型获取 API Keys - 添加 get_next_api_key_by_type 方法支持按类型轮询负载均衡 - 在 chat_completions 端点添加 API Key Provider 回退支持 - 在 anthropic_messages 端点添加 API Key Provider 回退支持 - 支持自定义 base_url 的 AnthropicKey 使用 OpenAI 兼容格式调用 - 添加 convert_openai_response_to_anthropic 响应转换函数 - 更新版本号到 v0.28.0 🤖 Generated with [Claude Code](https://claude.com/claude-code) Co-Authored-By: Claude Opus 4.5 <noreply@anthropic.com>
This commit is contained in:
Generated
+2
-1
@@ -3674,7 +3674,7 @@ dependencies = [
|
||||
|
||||
[[package]]
|
||||
name = "proxycast"
|
||||
version = "0.27.0"
|
||||
version = "0.28.0"
|
||||
dependencies = [
|
||||
"anyhow",
|
||||
"arboard",
|
||||
@@ -3691,6 +3691,7 @@ dependencies = [
|
||||
"flate2",
|
||||
"fs2",
|
||||
"futures",
|
||||
"glob",
|
||||
"indexmap 2.12.1",
|
||||
"md5",
|
||||
"notify",
|
||||
|
||||
@@ -1,6 +1,6 @@
|
||||
[package]
|
||||
name = "proxycast"
|
||||
version = "0.27.0"
|
||||
version = "0.28.0"
|
||||
description = "AI API Proxy Desktop App"
|
||||
authors = ["you"]
|
||||
edition = "2021"
|
||||
@@ -64,6 +64,7 @@ url = "2"
|
||||
once_cell = "1"
|
||||
tokio-util = "0.7"
|
||||
arboard = "3"
|
||||
glob = "0.3.3"
|
||||
|
||||
# Platform specific dependencies for browser interceptor
|
||||
|
||||
|
||||
@@ -0,0 +1,439 @@
|
||||
//! 应用启动引导模块
|
||||
//!
|
||||
//! 包含配置验证、状态初始化等启动逻辑。
|
||||
|
||||
use std::sync::Arc;
|
||||
use tokio::sync::RwLock;
|
||||
|
||||
use crate::agent::NativeAgentState;
|
||||
use crate::commands::api_key_provider_cmd::ApiKeyProviderServiceState;
|
||||
use crate::commands::browser_interceptor_cmd::BrowserInterceptorState;
|
||||
use crate::commands::flow_monitor_cmd::{
|
||||
BatchOperationsState, BookmarkManagerState, EnhancedStatsServiceState, FlowInterceptorState,
|
||||
FlowMonitorState, FlowQueryServiceState, FlowReplayerState, QuickFilterManagerState,
|
||||
SessionManagerState,
|
||||
};
|
||||
use crate::commands::machine_id_cmd::MachineIdState;
|
||||
use crate::commands::orchestrator_cmd::OrchestratorState;
|
||||
use crate::commands::plugin_cmd::PluginManagerState;
|
||||
use crate::commands::plugin_install_cmd::PluginInstallerState;
|
||||
use crate::commands::provider_pool_cmd::{CredentialSyncServiceState, ProviderPoolServiceState};
|
||||
use crate::commands::resilience_cmd::ResilienceConfigState;
|
||||
use crate::commands::router_cmd::RouterConfigState;
|
||||
use crate::commands::skill_cmd::SkillServiceState;
|
||||
use crate::config::{self, Config};
|
||||
use crate::database::{self, DbConnection};
|
||||
use crate::flow_monitor::{
|
||||
BatchOperations, BookmarkManager, EnhancedStatsService, FlowFileStore, FlowInterceptor,
|
||||
FlowMonitor, FlowMonitorConfig, FlowQueryService, FlowReplayer, InterceptConfig,
|
||||
QuickFilterManager, RotationConfig, SessionManager,
|
||||
};
|
||||
use crate::logger;
|
||||
use crate::plugin;
|
||||
use crate::server;
|
||||
use crate::services::api_key_provider_service::ApiKeyProviderService;
|
||||
use crate::services::provider_pool_service::ProviderPoolService;
|
||||
use crate::services::skill_service::SkillService;
|
||||
use crate::services::token_cache_service::TokenCacheService;
|
||||
use crate::telemetry;
|
||||
|
||||
use super::types::{AppState, LogState, TokenCacheServiceState};
|
||||
use super::utils::{generate_api_key, is_loopback_host};
|
||||
|
||||
/// 配置验证错误
|
||||
#[derive(Debug)]
|
||||
pub enum ConfigError {
|
||||
LoadFailed(String),
|
||||
SaveFailed(String),
|
||||
InvalidHost,
|
||||
DefaultApiKey,
|
||||
TlsNotSupported,
|
||||
RemoteManagementNotSupported,
|
||||
}
|
||||
|
||||
impl std::fmt::Display for ConfigError {
|
||||
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
|
||||
match self {
|
||||
ConfigError::LoadFailed(e) => write!(f, "配置加载失败: {}", e),
|
||||
ConfigError::SaveFailed(e) => write!(f, "配置保存失败: {}", e),
|
||||
ConfigError::InvalidHost => {
|
||||
write!(f, "当前版本仅支持本地监听,请使用 127.0.0.1/localhost/::1")
|
||||
}
|
||||
ConfigError::DefaultApiKey => write!(f, "检测到使用默认 API key,请配置强密钥"),
|
||||
ConfigError::TlsNotSupported => write!(f, "当前版本尚未支持 TLS"),
|
||||
ConfigError::RemoteManagementNotSupported => {
|
||||
write!(f, "远程管理需要 TLS 支持,当前版本未启用")
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
/// 加载并验证配置
|
||||
pub fn load_and_validate_config() -> Result<Config, ConfigError> {
|
||||
let mut config = config::load_config().map_err(|e| ConfigError::LoadFailed(e.to_string()))?;
|
||||
|
||||
// 自动生成 API key(如果使用默认值)
|
||||
if config.server.api_key == config::DEFAULT_API_KEY {
|
||||
let new_key = generate_api_key();
|
||||
config.server.api_key = new_key;
|
||||
config::save_config(&config).map_err(|e| ConfigError::SaveFailed(e.to_string()))?;
|
||||
tracing::info!("检测到默认 API key,已自动生成并保存新密钥");
|
||||
}
|
||||
|
||||
// 验证主机地址
|
||||
if !is_loopback_host(&config.server.host) {
|
||||
return Err(ConfigError::InvalidHost);
|
||||
}
|
||||
|
||||
// 再次检查 API key(防止保存失败后继续)
|
||||
if config.server.api_key == config::DEFAULT_API_KEY {
|
||||
return Err(ConfigError::DefaultApiKey);
|
||||
}
|
||||
|
||||
// 检查 TLS 配置
|
||||
if config.server.tls.enable {
|
||||
return Err(ConfigError::TlsNotSupported);
|
||||
}
|
||||
|
||||
// 检查远程管理配置
|
||||
if config.remote_management.allow_remote {
|
||||
return Err(ConfigError::RemoteManagementNotSupported);
|
||||
}
|
||||
|
||||
Ok(config)
|
||||
}
|
||||
|
||||
/// 应用状态集合
|
||||
pub struct AppStates {
|
||||
pub state: AppState,
|
||||
pub logs: LogState,
|
||||
pub db: DbConnection,
|
||||
pub skill_service: SkillServiceState,
|
||||
pub provider_pool_service: ProviderPoolServiceState,
|
||||
pub api_key_provider_service: ApiKeyProviderServiceState,
|
||||
pub credential_sync_service: CredentialSyncServiceState,
|
||||
pub token_cache_service: TokenCacheServiceState,
|
||||
pub machine_id_service: MachineIdState,
|
||||
pub router_config: RouterConfigState,
|
||||
pub resilience_config: ResilienceConfigState,
|
||||
pub plugin_manager: PluginManagerState,
|
||||
pub plugin_installer: PluginInstallerState,
|
||||
pub telemetry: crate::commands::telemetry_cmd::TelemetryState,
|
||||
pub flow_monitor: FlowMonitorState,
|
||||
pub flow_query_service: FlowQueryServiceState,
|
||||
pub flow_interceptor: FlowInterceptorState,
|
||||
pub flow_replayer: FlowReplayerState,
|
||||
pub session_manager: SessionManagerState,
|
||||
pub quick_filter_manager: QuickFilterManagerState,
|
||||
pub bookmark_manager: BookmarkManagerState,
|
||||
pub enhanced_stats_service: EnhancedStatsServiceState,
|
||||
pub batch_operations: BatchOperationsState,
|
||||
pub browser_interceptor: BrowserInterceptorState,
|
||||
pub native_agent: NativeAgentState,
|
||||
pub oauth_plugin_manager: crate::commands::oauth_plugin_cmd::OAuthPluginManagerState,
|
||||
pub orchestrator: OrchestratorState,
|
||||
// 用于 setup hook 的共享实例
|
||||
pub shared_stats: Arc<parking_lot::RwLock<telemetry::StatsAggregator>>,
|
||||
pub shared_tokens: Arc<parking_lot::RwLock<telemetry::TokenTracker>>,
|
||||
pub shared_logger: Arc<telemetry::RequestLogger>,
|
||||
pub flow_monitor_arc: Arc<FlowMonitor>,
|
||||
pub flow_interceptor_arc: Arc<FlowInterceptor>,
|
||||
}
|
||||
|
||||
/// 初始化所有应用状态
|
||||
pub fn init_states(config: &Config) -> Result<AppStates, String> {
|
||||
// 核心状态
|
||||
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 db = database::init_database().map_err(|e| format!("数据库初始化失败: {}", e))?;
|
||||
|
||||
// 服务状态
|
||||
let skill_service =
|
||||
SkillService::new().map_err(|e| format!("SkillService 初始化失败: {}", e))?;
|
||||
let skill_service_state = SkillServiceState(Arc::new(skill_service));
|
||||
|
||||
let provider_pool_service = ProviderPoolService::new();
|
||||
let provider_pool_service_state = ProviderPoolServiceState(Arc::new(provider_pool_service));
|
||||
|
||||
let api_key_provider_service = ApiKeyProviderService::new();
|
||||
let api_key_provider_service_state =
|
||||
ApiKeyProviderServiceState(Arc::new(api_key_provider_service));
|
||||
|
||||
let credential_sync_service_state = CredentialSyncServiceState(None);
|
||||
|
||||
let token_cache_service = TokenCacheService::new();
|
||||
let token_cache_service_state = TokenCacheServiceState(Arc::new(token_cache_service));
|
||||
|
||||
let machine_id_service = crate::services::machine_id_service::MachineIdService::new()
|
||||
.map_err(|e| format!("MachineIdService 初始化失败: {}", e))?;
|
||||
let machine_id_service_state: MachineIdState = Arc::new(RwLock::new(machine_id_service));
|
||||
|
||||
let router_config_state = RouterConfigState::default();
|
||||
let resilience_config_state = ResilienceConfigState::default();
|
||||
|
||||
// 插件管理器
|
||||
let plugin_manager = plugin::PluginManager::with_defaults();
|
||||
let plugin_manager_state = PluginManagerState(Arc::new(RwLock::new(plugin_manager)));
|
||||
|
||||
// 插件安装器
|
||||
let plugin_installer_state = init_plugin_installer()?;
|
||||
|
||||
// 遥测系统
|
||||
let (telemetry_state, shared_stats, shared_tokens, shared_logger) = init_telemetry(config)?;
|
||||
|
||||
// Flow Monitor 系统
|
||||
let (
|
||||
flow_monitor_state,
|
||||
flow_query_service_state,
|
||||
flow_interceptor_state,
|
||||
flow_replayer_state,
|
||||
session_manager_state,
|
||||
quick_filter_manager_state,
|
||||
bookmark_manager_state,
|
||||
enhanced_stats_service_state,
|
||||
batch_operations_state,
|
||||
flow_monitor_arc,
|
||||
flow_interceptor_arc,
|
||||
) = init_flow_monitor(&provider_pool_service_state, &db)?;
|
||||
|
||||
// 其他状态
|
||||
let browser_interceptor_state = BrowserInterceptorState::default();
|
||||
let native_agent_state = NativeAgentState::new();
|
||||
let oauth_plugin_manager_state =
|
||||
crate::commands::oauth_plugin_cmd::OAuthPluginManagerState::with_defaults();
|
||||
let orchestrator_state = OrchestratorState::new();
|
||||
|
||||
// 初始化默认技能仓库
|
||||
{
|
||||
let conn = db.lock().expect("Failed to lock database");
|
||||
database::dao::skills::SkillDao::init_default_skill_repos(&conn)
|
||||
.map_err(|e| format!("初始化默认技能仓库失败: {}", e))?;
|
||||
}
|
||||
|
||||
Ok(AppStates {
|
||||
state,
|
||||
logs,
|
||||
db,
|
||||
skill_service: skill_service_state,
|
||||
provider_pool_service: provider_pool_service_state,
|
||||
api_key_provider_service: api_key_provider_service_state,
|
||||
credential_sync_service: credential_sync_service_state,
|
||||
token_cache_service: token_cache_service_state,
|
||||
machine_id_service: machine_id_service_state,
|
||||
router_config: router_config_state,
|
||||
resilience_config: resilience_config_state,
|
||||
plugin_manager: plugin_manager_state,
|
||||
plugin_installer: plugin_installer_state,
|
||||
telemetry: telemetry_state,
|
||||
flow_monitor: flow_monitor_state,
|
||||
flow_query_service: flow_query_service_state,
|
||||
flow_interceptor: flow_interceptor_state,
|
||||
flow_replayer: flow_replayer_state,
|
||||
session_manager: session_manager_state,
|
||||
quick_filter_manager: quick_filter_manager_state,
|
||||
bookmark_manager: bookmark_manager_state,
|
||||
enhanced_stats_service: enhanced_stats_service_state,
|
||||
batch_operations: batch_operations_state,
|
||||
browser_interceptor: browser_interceptor_state,
|
||||
native_agent: native_agent_state,
|
||||
oauth_plugin_manager: oauth_plugin_manager_state,
|
||||
orchestrator: orchestrator_state,
|
||||
shared_stats,
|
||||
shared_tokens,
|
||||
shared_logger,
|
||||
flow_monitor_arc,
|
||||
flow_interceptor_arc,
|
||||
})
|
||||
}
|
||||
|
||||
/// 初始化插件安装器
|
||||
fn init_plugin_installer() -> Result<PluginInstallerState, String> {
|
||||
let db_path = database::get_db_path().map_err(|e| format!("获取数据库路径失败: {}", e))?;
|
||||
let plugins_dir = dirs::data_dir()
|
||||
.unwrap_or_else(|| std::path::PathBuf::from("."))
|
||||
.join("proxycast")
|
||||
.join("plugins");
|
||||
let temp_dir = std::env::temp_dir().join("proxycast_plugin_install");
|
||||
|
||||
let _ = std::fs::create_dir_all(&plugins_dir);
|
||||
let _ = std::fs::create_dir_all(&temp_dir);
|
||||
|
||||
match plugin::installer::PluginInstaller::from_paths(
|
||||
plugins_dir.clone(),
|
||||
temp_dir.clone(),
|
||||
&db_path,
|
||||
) {
|
||||
Ok(installer) => {
|
||||
tracing::info!("[启动] 插件安装器初始化成功");
|
||||
Ok(PluginInstallerState(Arc::new(RwLock::new(installer))))
|
||||
}
|
||||
Err(e) => {
|
||||
tracing::error!("[启动] 插件安装器初始化失败: {}", e);
|
||||
// 使用临时目录作为后备
|
||||
let fallback_plugins_dir = std::env::temp_dir().join("proxycast_plugins_fallback");
|
||||
let fallback_temp_dir = std::env::temp_dir().join("proxycast_plugin_install_fallback");
|
||||
let _ = std::fs::create_dir_all(&fallback_plugins_dir);
|
||||
let _ = std::fs::create_dir_all(&fallback_temp_dir);
|
||||
let installer = plugin::installer::PluginInstaller::from_paths(
|
||||
fallback_plugins_dir,
|
||||
fallback_temp_dir,
|
||||
&db_path,
|
||||
)
|
||||
.map_err(|e| format!("后备插件安装器初始化失败: {}", e))?;
|
||||
Ok(PluginInstallerState(Arc::new(RwLock::new(installer))))
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
/// 初始化遥测系统
|
||||
fn init_telemetry(
|
||||
config: &Config,
|
||||
) -> Result<
|
||||
(
|
||||
crate::commands::telemetry_cmd::TelemetryState,
|
||||
Arc<parking_lot::RwLock<telemetry::StatsAggregator>>,
|
||||
Arc<parking_lot::RwLock<telemetry::TokenTracker>>,
|
||||
Arc<telemetry::RequestLogger>,
|
||||
),
|
||||
String,
|
||||
> {
|
||||
let shared_stats = Arc::new(parking_lot::RwLock::new(
|
||||
telemetry::StatsAggregator::with_defaults(),
|
||||
));
|
||||
let shared_tokens = Arc::new(parking_lot::RwLock::new(
|
||||
telemetry::TokenTracker::with_defaults(),
|
||||
));
|
||||
let log_rotation = telemetry::LogRotationConfig {
|
||||
max_memory_logs: 10000,
|
||||
retention_days: config.logging.retention_days,
|
||||
max_file_size: 10 * 1024 * 1024,
|
||||
enable_file_logging: config.logging.enabled,
|
||||
};
|
||||
let shared_logger = Arc::new(
|
||||
telemetry::RequestLogger::new(log_rotation)
|
||||
.map_err(|e| format!("RequestLogger 初始化失败: {}", e))?,
|
||||
);
|
||||
|
||||
let telemetry_state = crate::commands::telemetry_cmd::TelemetryState::with_shared(
|
||||
shared_stats.clone(),
|
||||
shared_tokens.clone(),
|
||||
Some(shared_logger.clone()),
|
||||
)
|
||||
.map_err(|e| format!("TelemetryState 初始化失败: {}", e))?;
|
||||
|
||||
Ok((telemetry_state, shared_stats, shared_tokens, shared_logger))
|
||||
}
|
||||
|
||||
/// 初始化 Flow Monitor 系统
|
||||
#[allow(clippy::type_complexity)]
|
||||
fn init_flow_monitor(
|
||||
provider_pool_service_state: &ProviderPoolServiceState,
|
||||
db: &DbConnection,
|
||||
) -> Result<
|
||||
(
|
||||
FlowMonitorState,
|
||||
FlowQueryServiceState,
|
||||
FlowInterceptorState,
|
||||
FlowReplayerState,
|
||||
SessionManagerState,
|
||||
QuickFilterManagerState,
|
||||
BookmarkManagerState,
|
||||
EnhancedStatsServiceState,
|
||||
BatchOperationsState,
|
||||
Arc<FlowMonitor>,
|
||||
Arc<FlowInterceptor>,
|
||||
),
|
||||
String,
|
||||
> {
|
||||
let flow_monitor_config = FlowMonitorConfig::default();
|
||||
|
||||
// 初始化文件存储
|
||||
let data_dir = dirs::data_dir()
|
||||
.unwrap_or_else(|| std::path::PathBuf::from("."))
|
||||
.join("proxycast")
|
||||
.join("flows");
|
||||
let _ = std::fs::create_dir_all(&data_dir);
|
||||
|
||||
let rotation_config = RotationConfig::default();
|
||||
let flow_file_store = match FlowFileStore::new(data_dir, rotation_config.clone()) {
|
||||
Ok(store) => Some(Arc::new(store)),
|
||||
Err(e) => {
|
||||
tracing::warn!("无法初始化 Flow 文件存储: {}", e);
|
||||
None
|
||||
}
|
||||
};
|
||||
|
||||
let flow_monitor = Arc::new(FlowMonitor::new(
|
||||
flow_monitor_config,
|
||||
flow_file_store.clone(),
|
||||
));
|
||||
let flow_monitor_state = FlowMonitorState(flow_monitor.clone());
|
||||
|
||||
let flow_interceptor = Arc::new(FlowInterceptor::new(InterceptConfig::default()));
|
||||
let flow_interceptor_state = FlowInterceptorState(flow_interceptor.clone());
|
||||
|
||||
let flow_replayer = Arc::new(FlowReplayer::new(
|
||||
flow_monitor.clone(),
|
||||
provider_pool_service_state.0.clone(),
|
||||
db.clone(),
|
||||
));
|
||||
let flow_replayer_state = FlowReplayerState(flow_replayer);
|
||||
|
||||
let db_path = database::get_db_path().map_err(|e| format!("获取数据库路径失败: {}", e))?;
|
||||
|
||||
let session_manager = Arc::new(
|
||||
SessionManager::new(db_path.clone())
|
||||
.map_err(|e| format!("SessionManager 初始化失败: {}", e))?,
|
||||
);
|
||||
let session_manager_state = SessionManagerState(session_manager.clone());
|
||||
|
||||
let quick_filter_manager = Arc::new(
|
||||
QuickFilterManager::new(db_path.clone())
|
||||
.map_err(|e| format!("QuickFilterManager 初始化失败: {}", e))?,
|
||||
);
|
||||
let quick_filter_manager_state = QuickFilterManagerState(quick_filter_manager);
|
||||
|
||||
let bookmark_manager = Arc::new(
|
||||
BookmarkManager::new(db_path).map_err(|e| format!("BookmarkManager 初始化失败: {}", e))?,
|
||||
);
|
||||
let bookmark_manager_state = BookmarkManagerState(bookmark_manager);
|
||||
|
||||
let enhanced_stats_service = Arc::new(EnhancedStatsService::new(flow_monitor.memory_store()));
|
||||
let enhanced_stats_service_state = EnhancedStatsServiceState(enhanced_stats_service);
|
||||
|
||||
let batch_operations = Arc::new(BatchOperations::new(
|
||||
flow_monitor.clone(),
|
||||
Some(session_manager_state.0.clone()),
|
||||
));
|
||||
let batch_operations_state = BatchOperationsState(batch_operations);
|
||||
|
||||
// FlowQueryService
|
||||
let flow_query_service_state = if let Some(file_store) = flow_file_store {
|
||||
let query_service = FlowQueryService::new(flow_monitor.memory_store(), file_store);
|
||||
FlowQueryServiceState(Arc::new(query_service))
|
||||
} else {
|
||||
let temp_dir = std::env::temp_dir().join("proxycast_flows");
|
||||
let _ = std::fs::create_dir_all(&temp_dir);
|
||||
let temp_store = FlowFileStore::new(temp_dir, rotation_config)
|
||||
.map_err(|e| format!("临时 FlowFileStore 初始化失败: {}", e))?;
|
||||
let query_service =
|
||||
FlowQueryService::new(flow_monitor.memory_store(), Arc::new(temp_store));
|
||||
FlowQueryServiceState(Arc::new(query_service))
|
||||
};
|
||||
|
||||
Ok((
|
||||
flow_monitor_state,
|
||||
flow_query_service_state,
|
||||
flow_interceptor_state,
|
||||
flow_replayer_state,
|
||||
session_manager_state,
|
||||
quick_filter_manager_state,
|
||||
bookmark_manager_state,
|
||||
enhanced_stats_service_state,
|
||||
batch_operations_state,
|
||||
flow_monitor,
|
||||
flow_interceptor,
|
||||
))
|
||||
}
|
||||
@@ -0,0 +1,427 @@
|
||||
//! API 测试和兼容性检查命令
|
||||
//!
|
||||
//! 包含 API 测试、模型列表和兼容性检查命令。
|
||||
|
||||
use crate::app::types::{AppState, LogState, ProviderType};
|
||||
|
||||
/// 测试结果
|
||||
#[derive(serde::Serialize)]
|
||||
pub struct TestResult {
|
||||
pub success: bool,
|
||||
pub status: u16,
|
||||
pub body: String,
|
||||
pub time_ms: u64,
|
||||
}
|
||||
|
||||
/// 模型信息
|
||||
#[derive(serde::Serialize)]
|
||||
pub struct ModelInfo {
|
||||
pub id: String,
|
||||
pub object: String,
|
||||
pub owned_by: String,
|
||||
}
|
||||
|
||||
/// API 检查结果
|
||||
#[derive(serde::Serialize)]
|
||||
pub struct ApiCheckResult {
|
||||
pub model: String,
|
||||
pub available: bool,
|
||||
pub status: u16,
|
||||
pub error_type: Option<String>,
|
||||
pub error_message: Option<String>,
|
||||
pub time_ms: u64,
|
||||
}
|
||||
|
||||
/// API 兼容性结果
|
||||
#[derive(serde::Serialize)]
|
||||
pub struct ApiCompatibilityResult {
|
||||
pub provider: String,
|
||||
pub overall_status: String,
|
||||
pub checked_at: String,
|
||||
pub results: Vec<ApiCheckResult>,
|
||||
pub warnings: Vec<String>,
|
||||
}
|
||||
|
||||
/// 检查 API 兼容性
|
||||
#[tauri::command]
|
||||
pub async fn check_api_compatibility(
|
||||
state: tauri::State<'_, AppState>,
|
||||
logs: tauri::State<'_, LogState>,
|
||||
provider: String,
|
||||
) -> Result<ApiCompatibilityResult, String> {
|
||||
// 使用枚举验证 provider
|
||||
let provider_type: ProviderType = provider.parse().map_err(|e: String| e)?;
|
||||
|
||||
logs.write().await.add(
|
||||
"info",
|
||||
&format!("[API检测] 开始检测 {provider_type} API 兼容性 (Claude Code 功能测试)..."),
|
||||
);
|
||||
|
||||
let s = state.read().await;
|
||||
let mut results: Vec<ApiCheckResult> = Vec::new();
|
||||
let mut warnings: Vec<String> = Vec::new();
|
||||
|
||||
// Claude Code 需要的测试项目
|
||||
let test_cases: Vec<(&str, &str)> = match provider_type {
|
||||
ProviderType::Kiro => vec![
|
||||
("claude-sonnet-4-5", "basic"),
|
||||
("claude-sonnet-4-5", "tool_call"),
|
||||
],
|
||||
ProviderType::Gemini => vec![
|
||||
("gemini-2.5-flash", "basic"),
|
||||
("gemini-2.5-flash", "tool_call"),
|
||||
],
|
||||
ProviderType::Qwen => vec![
|
||||
("qwen3-coder-plus", "basic"),
|
||||
("qwen3-coder-plus", "tool_call"),
|
||||
],
|
||||
ProviderType::Antigravity => vec![
|
||||
("gemini-3-pro-preview", "basic"),
|
||||
("gemini-3-pro-preview", "tool_call"),
|
||||
],
|
||||
ProviderType::Vertex => vec![
|
||||
("gemini-2.0-flash", "basic"),
|
||||
("gemini-2.0-flash", "tool_call"),
|
||||
],
|
||||
ProviderType::GeminiApiKey => vec![
|
||||
("gemini-2.5-flash", "basic"),
|
||||
("gemini-2.5-flash", "tool_call"),
|
||||
],
|
||||
ProviderType::Codex => vec![("gpt-4.1", "basic"), ("gpt-4.1", "tool_call")],
|
||||
ProviderType::ClaudeOAuth => vec![
|
||||
("claude-sonnet-4-5", "basic"),
|
||||
("claude-sonnet-4-5", "tool_call"),
|
||||
],
|
||||
ProviderType::IFlow => vec![("gpt-4o", "basic"), ("gpt-4o", "tool_call")],
|
||||
ProviderType::OpenAI | ProviderType::Claude => vec![],
|
||||
// API Key Provider 类型 - 暂不支持自动测试
|
||||
ProviderType::Anthropic
|
||||
| ProviderType::AzureOpenai
|
||||
| ProviderType::AwsBedrock
|
||||
| ProviderType::Ollama => vec![],
|
||||
};
|
||||
|
||||
for (model, test_type) in test_cases {
|
||||
let start = std::time::Instant::now();
|
||||
let test_name = format!("{model} ({test_type})");
|
||||
|
||||
// 根据测试类型构建不同的请求
|
||||
let test_request = match test_type {
|
||||
"tool_call" => {
|
||||
// 测试 Tool Calls - Claude Code 核心功能
|
||||
crate::models::openai::ChatCompletionRequest {
|
||||
model: model.to_string(),
|
||||
messages: vec![crate::models::openai::ChatMessage {
|
||||
role: "user".to_string(),
|
||||
content: Some(crate::models::openai::MessageContent::Text(
|
||||
"What is 2+2? Use the calculator tool to compute this.".to_string(),
|
||||
)),
|
||||
tool_calls: None,
|
||||
tool_call_id: None,
|
||||
}],
|
||||
temperature: None,
|
||||
max_tokens: Some(100),
|
||||
top_p: None,
|
||||
stream: false,
|
||||
tools: Some(vec![crate::models::openai::Tool::Function {
|
||||
function: crate::models::openai::FunctionDef {
|
||||
name: "calculator".to_string(),
|
||||
description: Some("Perform basic arithmetic calculations".to_string()),
|
||||
parameters: Some(serde_json::json!({
|
||||
"type": "object",
|
||||
"properties": {
|
||||
"expression": {
|
||||
"type": "string",
|
||||
"description": "The math expression to evaluate"
|
||||
}
|
||||
},
|
||||
"required": ["expression"]
|
||||
})),
|
||||
},
|
||||
}]),
|
||||
tool_choice: None,
|
||||
reasoning_effort: None,
|
||||
}
|
||||
}
|
||||
_ => {
|
||||
// 基础对话测试
|
||||
crate::models::openai::ChatCompletionRequest {
|
||||
model: model.to_string(),
|
||||
messages: vec![crate::models::openai::ChatMessage {
|
||||
role: "user".to_string(),
|
||||
content: Some(crate::models::openai::MessageContent::Text(
|
||||
"Say 'OK' only.".to_string(),
|
||||
)),
|
||||
tool_calls: None,
|
||||
tool_call_id: None,
|
||||
}],
|
||||
temperature: None,
|
||||
max_tokens: Some(10),
|
||||
top_p: None,
|
||||
stream: false,
|
||||
tools: None,
|
||||
tool_choice: None,
|
||||
reasoning_effort: None,
|
||||
}
|
||||
}
|
||||
};
|
||||
|
||||
let result = match provider_type {
|
||||
ProviderType::Kiro => s.kiro_provider.call_api(&test_request).await,
|
||||
ProviderType::Gemini => {
|
||||
Err("Gemini API compatibility check not yet implemented".into())
|
||||
}
|
||||
ProviderType::Qwen => Err("Qwen API compatibility check not yet implemented".into()),
|
||||
_ => Err("Provider not supported for direct API check".into()),
|
||||
};
|
||||
|
||||
let time_ms = start.elapsed().as_millis() as u64;
|
||||
|
||||
match result {
|
||||
Ok(resp) => {
|
||||
let status = resp.status().as_u16();
|
||||
let body = resp.text().await.unwrap_or_default();
|
||||
|
||||
let (available, error_type, error_message) = if (200..300).contains(&status) {
|
||||
if test_type == "tool_call" {
|
||||
let has_tool_use =
|
||||
body.contains("\"name\"") && body.contains("\"toolUseId\"");
|
||||
if !has_tool_use {
|
||||
warnings.push(format!(
|
||||
"{test_name}: 响应未包含 tool_use,Claude Code 可能无法正常工作"
|
||||
));
|
||||
}
|
||||
}
|
||||
(true, None, None)
|
||||
} else {
|
||||
let err_type = match status {
|
||||
401 => {
|
||||
warnings.push(format!("{test_name} 返回 401: Token 可能已过期或无效"));
|
||||
Some("AUTH_ERROR".to_string())
|
||||
}
|
||||
403 => {
|
||||
warnings.push(format!(
|
||||
"{test_name} 返回 403: 无权访问,可能需要刷新 Token"
|
||||
));
|
||||
Some("FORBIDDEN".to_string())
|
||||
}
|
||||
400 => {
|
||||
warnings.push(format!("{test_name} 返回 400: 请求格式可能已变更"));
|
||||
Some("BAD_REQUEST".to_string())
|
||||
}
|
||||
404 => {
|
||||
warnings.push(format!("{test_name} 返回 404: 模型或接口可能已下线"));
|
||||
Some("NOT_FOUND".to_string())
|
||||
}
|
||||
429 => {
|
||||
warnings.push(format!("{test_name} 返回 429: 请求过于频繁"));
|
||||
Some("RATE_LIMITED".to_string())
|
||||
}
|
||||
500..=599 => {
|
||||
warnings.push(format!("{test_name} 返回 {status}: 服务端错误"));
|
||||
Some("SERVER_ERROR".to_string())
|
||||
}
|
||||
_ => Some("UNKNOWN_ERROR".to_string()),
|
||||
};
|
||||
(
|
||||
false,
|
||||
err_type,
|
||||
Some(body[..body.len().min(200)].to_string()),
|
||||
)
|
||||
};
|
||||
|
||||
results.push(ApiCheckResult {
|
||||
model: test_name,
|
||||
available,
|
||||
status,
|
||||
error_type,
|
||||
error_message,
|
||||
time_ms,
|
||||
});
|
||||
}
|
||||
Err(e) => {
|
||||
warnings.push(format!("{test_name} 请求失败: {e}"));
|
||||
results.push(ApiCheckResult {
|
||||
model: test_name,
|
||||
available: false,
|
||||
status: 0,
|
||||
error_type: Some("REQUEST_FAILED".to_string()),
|
||||
error_message: Some(e.to_string()),
|
||||
time_ms,
|
||||
});
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
let overall_status = if results.iter().all(|r| r.available) {
|
||||
"healthy".to_string()
|
||||
} else if results.iter().any(|r| r.available) {
|
||||
"partial".to_string()
|
||||
} else {
|
||||
"error".to_string()
|
||||
};
|
||||
|
||||
let checked_at = chrono::Utc::now().to_rfc3339();
|
||||
|
||||
logs.write().await.add(
|
||||
"info",
|
||||
&format!("[API检测] {provider} 检测完成: {overall_status}"),
|
||||
);
|
||||
|
||||
Ok(ApiCompatibilityResult {
|
||||
provider,
|
||||
overall_status,
|
||||
checked_at,
|
||||
results,
|
||||
warnings,
|
||||
})
|
||||
}
|
||||
|
||||
/// 获取可用模型列表
|
||||
#[tauri::command]
|
||||
pub async fn get_available_models() -> Result<Vec<ModelInfo>, String> {
|
||||
Ok(vec![
|
||||
// Kiro/Claude models
|
||||
ModelInfo {
|
||||
id: "claude-sonnet-4-5".to_string(),
|
||||
object: "model".to_string(),
|
||||
owned_by: "anthropic".to_string(),
|
||||
},
|
||||
ModelInfo {
|
||||
id: "claude-sonnet-4-5-20250514".to_string(),
|
||||
object: "model".to_string(),
|
||||
owned_by: "anthropic".to_string(),
|
||||
},
|
||||
ModelInfo {
|
||||
id: "claude-sonnet-4-5-20250929".to_string(),
|
||||
object: "model".to_string(),
|
||||
owned_by: "anthropic".to_string(),
|
||||
},
|
||||
ModelInfo {
|
||||
id: "claude-3-7-sonnet-20250219".to_string(),
|
||||
object: "model".to_string(),
|
||||
owned_by: "anthropic".to_string(),
|
||||
},
|
||||
ModelInfo {
|
||||
id: "claude-3-5-sonnet-latest".to_string(),
|
||||
object: "model".to_string(),
|
||||
owned_by: "anthropic".to_string(),
|
||||
},
|
||||
ModelInfo {
|
||||
id: "claude-opus-4-5-20250514".to_string(),
|
||||
object: "model".to_string(),
|
||||
owned_by: "anthropic".to_string(),
|
||||
},
|
||||
ModelInfo {
|
||||
id: "claude-haiku-4-5-20250514".to_string(),
|
||||
object: "model".to_string(),
|
||||
owned_by: "anthropic".to_string(),
|
||||
},
|
||||
// Gemini models
|
||||
ModelInfo {
|
||||
id: "gemini-2.5-flash".to_string(),
|
||||
object: "model".to_string(),
|
||||
owned_by: "google".to_string(),
|
||||
},
|
||||
ModelInfo {
|
||||
id: "gemini-2.5-flash-lite".to_string(),
|
||||
object: "model".to_string(),
|
||||
owned_by: "google".to_string(),
|
||||
},
|
||||
ModelInfo {
|
||||
id: "gemini-2.5-pro".to_string(),
|
||||
object: "model".to_string(),
|
||||
owned_by: "google".to_string(),
|
||||
},
|
||||
ModelInfo {
|
||||
id: "gemini-2.5-pro-preview-06-05".to_string(),
|
||||
object: "model".to_string(),
|
||||
owned_by: "google".to_string(),
|
||||
},
|
||||
ModelInfo {
|
||||
id: "gemini-3-pro-preview".to_string(),
|
||||
object: "model".to_string(),
|
||||
owned_by: "google".to_string(),
|
||||
},
|
||||
// Qwen models
|
||||
ModelInfo {
|
||||
id: "qwen3-coder-plus".to_string(),
|
||||
object: "model".to_string(),
|
||||
owned_by: "alibaba".to_string(),
|
||||
},
|
||||
ModelInfo {
|
||||
id: "qwen3-coder-flash".to_string(),
|
||||
object: "model".to_string(),
|
||||
owned_by: "alibaba".to_string(),
|
||||
},
|
||||
])
|
||||
}
|
||||
|
||||
/// 测试 API
|
||||
#[tauri::command]
|
||||
pub async fn test_api(
|
||||
state: tauri::State<'_, AppState>,
|
||||
method: String,
|
||||
path: String,
|
||||
body: Option<String>,
|
||||
auth: bool,
|
||||
) -> Result<TestResult, String> {
|
||||
let s = state.read().await;
|
||||
let base_url = format!("http://{}:{}", s.config.server.host, s.config.server.port);
|
||||
let api_key = s
|
||||
.running_api_key
|
||||
.as_ref()
|
||||
.unwrap_or(&s.config.server.api_key);
|
||||
|
||||
let client = reqwest::Client::builder()
|
||||
.no_proxy()
|
||||
.build()
|
||||
.map_err(|e| e.to_string())?;
|
||||
|
||||
let url = format!("{base_url}{path}");
|
||||
|
||||
tracing::info!("Testing API: {} {}", method, url);
|
||||
|
||||
let start = std::time::Instant::now();
|
||||
|
||||
let mut req = match method.as_str() {
|
||||
"GET" => client.get(&url),
|
||||
"POST" => client.post(&url),
|
||||
_ => return Err("Unsupported method".to_string()),
|
||||
};
|
||||
|
||||
req = req.header("Content-Type", "application/json");
|
||||
|
||||
if auth {
|
||||
req = req.header("Authorization", format!("Bearer {api_key}"));
|
||||
}
|
||||
|
||||
if let Some(b) = body {
|
||||
req = req.body(b);
|
||||
}
|
||||
|
||||
match req.send().await {
|
||||
Ok(resp) => {
|
||||
let status = resp.status().as_u16();
|
||||
let body = resp.text().await.unwrap_or_default();
|
||||
let time_ms = start.elapsed().as_millis() as u64;
|
||||
|
||||
tracing::info!(
|
||||
"API test result: status={}, body_len={}",
|
||||
status,
|
||||
body.len()
|
||||
);
|
||||
|
||||
Ok(TestResult {
|
||||
success: (200..300).contains(&status),
|
||||
status,
|
||||
body,
|
||||
time_ms,
|
||||
})
|
||||
}
|
||||
Err(e) => {
|
||||
tracing::error!("API test error: {}", e);
|
||||
Err(e.to_string())
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,128 @@
|
||||
//! 配置管理命令
|
||||
//!
|
||||
//! 包含配置读取、保存、Provider 设置等命令。
|
||||
|
||||
use crate::app::types::{AppState, LogState, ProviderType};
|
||||
use crate::config;
|
||||
|
||||
/// 获取配置
|
||||
#[tauri::command]
|
||||
pub async fn get_config(state: tauri::State<'_, AppState>) -> Result<config::Config, String> {
|
||||
let s = state.read().await;
|
||||
Ok(s.config.clone())
|
||||
}
|
||||
|
||||
/// 保存配置
|
||||
#[tauri::command]
|
||||
pub async fn save_config(
|
||||
state: tauri::State<'_, AppState>,
|
||||
config: config::Config,
|
||||
) -> Result<(), String> {
|
||||
// P0 安全修复:禁止危险的网络配置
|
||||
let host = config.server.host.to_lowercase();
|
||||
if host == "0.0.0.0" || host == "::" {
|
||||
return Err(
|
||||
"安全限制:不允许监听所有网络接口 (0.0.0.0 或 ::)。请使用 127.0.0.1 或 localhost"
|
||||
.to_string(),
|
||||
);
|
||||
}
|
||||
|
||||
// 禁止开启远程管理
|
||||
if config.remote_management.allow_remote {
|
||||
return Err("安全限制:不允许开启远程管理功能".to_string());
|
||||
}
|
||||
|
||||
let mut s = state.write().await;
|
||||
s.config = config.clone();
|
||||
config::save_config(&config).map_err(|e| e.to_string())
|
||||
}
|
||||
|
||||
/// 获取默认 Provider
|
||||
#[tauri::command]
|
||||
pub async fn get_default_provider(state: tauri::State<'_, AppState>) -> Result<String, String> {
|
||||
let s = state.read().await;
|
||||
Ok(s.config.default_provider.clone())
|
||||
}
|
||||
|
||||
/// 设置默认 Provider
|
||||
#[tauri::command]
|
||||
pub async fn set_default_provider(
|
||||
state: tauri::State<'_, AppState>,
|
||||
logs: tauri::State<'_, LogState>,
|
||||
provider: String,
|
||||
) -> Result<String, String> {
|
||||
// 使用枚举验证 provider
|
||||
let provider_type: ProviderType = provider.parse().map_err(|e: String| e)?;
|
||||
|
||||
let mut s = state.write().await;
|
||||
s.config.default_provider = provider.clone();
|
||||
|
||||
// 同时更新运行中服务器的 default_provider_ref
|
||||
{
|
||||
let mut dp = s.default_provider_ref.write().await;
|
||||
*dp = provider.clone();
|
||||
}
|
||||
|
||||
config::save_config(&s.config).map_err(|e| e.to_string())?;
|
||||
logs.write()
|
||||
.await
|
||||
.add("info", &format!("默认 Provider 已切换为: {provider_type}"));
|
||||
Ok(provider)
|
||||
}
|
||||
|
||||
/// 获取端点 Provider 配置
|
||||
#[tauri::command]
|
||||
pub async fn get_endpoint_providers(
|
||||
state: tauri::State<'_, AppState>,
|
||||
) -> Result<serde_json::Value, String> {
|
||||
let s = state.read().await;
|
||||
let ep = &s.config.endpoint_providers;
|
||||
Ok(serde_json::json!({
|
||||
"cursor": ep.cursor.clone(),
|
||||
"claude_code": ep.claude_code.clone(),
|
||||
"codex": ep.codex.clone(),
|
||||
"windsurf": ep.windsurf.clone(),
|
||||
"kiro": ep.kiro.clone(),
|
||||
"other": ep.other.clone()
|
||||
}))
|
||||
}
|
||||
|
||||
/// 设置端点 Provider 配置
|
||||
#[tauri::command]
|
||||
pub async fn set_endpoint_provider(
|
||||
state: tauri::State<'_, AppState>,
|
||||
logs: tauri::State<'_, LogState>,
|
||||
endpoint: String,
|
||||
provider: Option<String>,
|
||||
) -> Result<String, String> {
|
||||
// 验证 provider(如果提供)
|
||||
if let Some(ref p) = provider {
|
||||
if !p.is_empty() {
|
||||
let _: ProviderType = p.parse().map_err(|e: String| e)?;
|
||||
}
|
||||
}
|
||||
|
||||
let mut s = state.write().await;
|
||||
|
||||
// 使用 set_provider 方法设置对应的 provider
|
||||
if !s
|
||||
.config
|
||||
.endpoint_providers
|
||||
.set_provider(&endpoint, provider.clone())
|
||||
{
|
||||
return Err(format!("未知的客户端类型: {}", endpoint));
|
||||
}
|
||||
|
||||
config::save_config(&s.config).map_err(|e| e.to_string())?;
|
||||
|
||||
let provider_display = provider.as_deref().unwrap_or("默认");
|
||||
logs.write().await.add(
|
||||
"info",
|
||||
&format!(
|
||||
"客户端 {} 的 Provider 已设置为: {}",
|
||||
endpoint, provider_display
|
||||
),
|
||||
);
|
||||
|
||||
Ok(provider_display.to_string())
|
||||
}
|
||||
@@ -0,0 +1,89 @@
|
||||
//! 自定义 Provider 命令
|
||||
//!
|
||||
//! 包含 OpenAI Custom 和 Claude Custom Provider 的配置命令。
|
||||
|
||||
use crate::app::types::{AppState, LogState};
|
||||
|
||||
/// OpenAI Custom 状态
|
||||
#[derive(serde::Serialize, serde::Deserialize)]
|
||||
pub struct OpenAICustomStatus {
|
||||
pub enabled: bool,
|
||||
pub has_api_key: bool,
|
||||
pub base_url: String,
|
||||
}
|
||||
|
||||
/// Claude Custom 状态
|
||||
#[derive(serde::Serialize, serde::Deserialize)]
|
||||
pub struct ClaudeCustomStatus {
|
||||
pub enabled: bool,
|
||||
pub has_api_key: bool,
|
||||
pub base_url: String,
|
||||
}
|
||||
|
||||
/// 获取 OpenAI Custom 状态
|
||||
#[tauri::command]
|
||||
pub async fn get_openai_custom_status(
|
||||
state: tauri::State<'_, AppState>,
|
||||
) -> Result<OpenAICustomStatus, String> {
|
||||
let s = state.read().await;
|
||||
let config = &s.openai_custom_provider.config;
|
||||
Ok(OpenAICustomStatus {
|
||||
enabled: config.enabled,
|
||||
has_api_key: config.api_key.is_some(),
|
||||
base_url: s.openai_custom_provider.get_base_url(),
|
||||
})
|
||||
}
|
||||
|
||||
/// 设置 OpenAI Custom 配置
|
||||
#[tauri::command]
|
||||
pub async fn set_openai_custom_config(
|
||||
state: tauri::State<'_, AppState>,
|
||||
logs: tauri::State<'_, LogState>,
|
||||
api_key: Option<String>,
|
||||
base_url: Option<String>,
|
||||
enabled: bool,
|
||||
) -> Result<String, String> {
|
||||
let mut s = state.write().await;
|
||||
s.openai_custom_provider.config.api_key = api_key;
|
||||
s.openai_custom_provider.config.base_url = base_url;
|
||||
s.openai_custom_provider.config.enabled = enabled;
|
||||
logs.write().await.add(
|
||||
"info",
|
||||
&format!("[OpenAI Custom] 配置已更新, enabled={enabled}"),
|
||||
);
|
||||
Ok("OpenAI Custom config updated".to_string())
|
||||
}
|
||||
|
||||
/// 获取 Claude Custom 状态
|
||||
#[tauri::command]
|
||||
pub async fn get_claude_custom_status(
|
||||
state: tauri::State<'_, AppState>,
|
||||
) -> Result<ClaudeCustomStatus, String> {
|
||||
let s = state.read().await;
|
||||
let config = &s.claude_custom_provider.config;
|
||||
Ok(ClaudeCustomStatus {
|
||||
enabled: config.enabled,
|
||||
has_api_key: config.api_key.is_some(),
|
||||
base_url: s.claude_custom_provider.get_base_url(),
|
||||
})
|
||||
}
|
||||
|
||||
/// 设置 Claude Custom 配置
|
||||
#[tauri::command]
|
||||
pub async fn set_claude_custom_config(
|
||||
state: tauri::State<'_, AppState>,
|
||||
logs: tauri::State<'_, LogState>,
|
||||
api_key: Option<String>,
|
||||
base_url: Option<String>,
|
||||
enabled: bool,
|
||||
) -> Result<String, String> {
|
||||
let mut s = state.write().await;
|
||||
s.claude_custom_provider.config.api_key = api_key;
|
||||
s.claude_custom_provider.config.base_url = base_url;
|
||||
s.claude_custom_provider.config.enabled = enabled;
|
||||
logs.write().await.add(
|
||||
"info",
|
||||
&format!("[Claude Custom] 配置已更新, enabled={enabled}"),
|
||||
);
|
||||
Ok("Claude Custom config updated".to_string())
|
||||
}
|
||||
@@ -0,0 +1,184 @@
|
||||
//! Gemini Provider 命令 (Legacy)
|
||||
//!
|
||||
//! 包含 Gemini 凭证管理相关命令。
|
||||
//! 这些命令保留用于向后兼容,新代码应使用统一的 OAuth 命令。
|
||||
|
||||
use crate::app::commands::kiro::{CheckResult, EnvVariable};
|
||||
use crate::app::types::{AppState, LogState};
|
||||
use crate::app::utils::mask_token;
|
||||
use crate::providers;
|
||||
|
||||
/// Gemini 凭证状态
|
||||
#[derive(serde::Serialize)]
|
||||
pub struct GeminiCredentialStatus {
|
||||
pub loaded: bool,
|
||||
pub has_access_token: bool,
|
||||
pub has_refresh_token: bool,
|
||||
pub expiry_date: Option<i64>,
|
||||
pub is_valid: bool,
|
||||
pub creds_path: String,
|
||||
}
|
||||
|
||||
/// 获取 Gemini 凭证状态
|
||||
#[tauri::command]
|
||||
pub async fn get_gemini_credentials(
|
||||
state: tauri::State<'_, AppState>,
|
||||
) -> Result<GeminiCredentialStatus, String> {
|
||||
let s = state.read().await;
|
||||
let creds = &s.gemini_provider.credentials;
|
||||
let path = providers::gemini::GeminiProvider::default_creds_path();
|
||||
|
||||
Ok(GeminiCredentialStatus {
|
||||
loaded: creds.access_token.is_some() || creds.refresh_token.is_some(),
|
||||
has_access_token: creds.access_token.is_some(),
|
||||
has_refresh_token: creds.refresh_token.is_some(),
|
||||
expiry_date: creds.expiry_date,
|
||||
is_valid: s.gemini_provider.is_token_valid(),
|
||||
creds_path: path.to_string_lossy().to_string(),
|
||||
})
|
||||
}
|
||||
|
||||
/// 重新加载 Gemini 凭证
|
||||
#[tauri::command]
|
||||
pub async fn reload_gemini_credentials(
|
||||
state: tauri::State<'_, AppState>,
|
||||
logs: tauri::State<'_, LogState>,
|
||||
) -> Result<String, String> {
|
||||
let mut s = state.write().await;
|
||||
logs.write().await.add("info", "[Gemini] 正在加载凭证...");
|
||||
s.gemini_provider
|
||||
.load_credentials()
|
||||
.await
|
||||
.map_err(|e| e.to_string())?;
|
||||
logs.write().await.add("info", "[Gemini] 凭证加载成功");
|
||||
Ok("Gemini credentials reloaded".to_string())
|
||||
}
|
||||
|
||||
/// 刷新 Gemini Token
|
||||
#[tauri::command]
|
||||
pub async fn refresh_gemini_token(
|
||||
state: tauri::State<'_, AppState>,
|
||||
logs: tauri::State<'_, LogState>,
|
||||
) -> Result<String, String> {
|
||||
let mut s = state.write().await;
|
||||
logs.write().await.add("info", "[Gemini] 正在刷新 Token...");
|
||||
let result = s
|
||||
.gemini_provider
|
||||
.refresh_token()
|
||||
.await
|
||||
.map_err(|e| e.to_string());
|
||||
match &result {
|
||||
Ok(_) => logs.write().await.add("info", "[Gemini] Token 刷新成功"),
|
||||
Err(e) => logs
|
||||
.write()
|
||||
.await
|
||||
.add("error", &format!("[Gemini] Token 刷新失败: {e}")),
|
||||
}
|
||||
result
|
||||
}
|
||||
|
||||
/// 获取 Gemini 环境变量
|
||||
#[tauri::command]
|
||||
pub async fn get_gemini_env_variables(
|
||||
state: tauri::State<'_, AppState>,
|
||||
) -> Result<Vec<EnvVariable>, String> {
|
||||
let s = state.read().await;
|
||||
let creds = &s.gemini_provider.credentials;
|
||||
let mut vars = Vec::new();
|
||||
|
||||
if let Some(token) = &creds.access_token {
|
||||
vars.push(EnvVariable {
|
||||
key: "GEMINI_ACCESS_TOKEN".to_string(),
|
||||
value: token.clone(),
|
||||
masked: mask_token(token),
|
||||
});
|
||||
}
|
||||
if let Some(token) = &creds.refresh_token {
|
||||
vars.push(EnvVariable {
|
||||
key: "GEMINI_REFRESH_TOKEN".to_string(),
|
||||
value: token.clone(),
|
||||
masked: mask_token(token),
|
||||
});
|
||||
}
|
||||
if let Some(expiry) = creds.expiry_date {
|
||||
let expiry_str = expiry.to_string();
|
||||
vars.push(EnvVariable {
|
||||
key: "GEMINI_EXPIRY_DATE".to_string(),
|
||||
value: expiry_str.clone(),
|
||||
masked: expiry_str,
|
||||
});
|
||||
}
|
||||
|
||||
Ok(vars)
|
||||
}
|
||||
|
||||
/// 获取 Gemini Token 文件哈希
|
||||
#[tauri::command]
|
||||
pub async fn get_gemini_token_file_hash() -> Result<String, String> {
|
||||
let path = providers::gemini::GeminiProvider::default_creds_path();
|
||||
if !tokio::fs::try_exists(&path).await.unwrap_or(false) {
|
||||
return Ok("".to_string());
|
||||
}
|
||||
|
||||
let content = tokio::fs::read(&path).await.map_err(|e| e.to_string())?;
|
||||
let hash = format!("{:x}", md5::compute(&content));
|
||||
Ok(hash)
|
||||
}
|
||||
|
||||
/// 检查并重新加载 Gemini 凭证
|
||||
#[tauri::command]
|
||||
pub async fn check_and_reload_gemini_credentials(
|
||||
state: tauri::State<'_, AppState>,
|
||||
logs: tauri::State<'_, LogState>,
|
||||
last_hash: String,
|
||||
) -> Result<CheckResult, String> {
|
||||
let path = providers::gemini::GeminiProvider::default_creds_path();
|
||||
|
||||
if !tokio::fs::try_exists(&path).await.unwrap_or(false) {
|
||||
return Ok(CheckResult {
|
||||
changed: false,
|
||||
new_hash: "".to_string(),
|
||||
reloaded: false,
|
||||
});
|
||||
}
|
||||
|
||||
let content = tokio::fs::read(&path).await.map_err(|e| e.to_string())?;
|
||||
let new_hash = format!("{:x}", md5::compute(&content));
|
||||
|
||||
if !last_hash.is_empty() && new_hash != last_hash {
|
||||
logs.write()
|
||||
.await
|
||||
.add("info", "[Gemini][自动检测] 凭证文件已变化,正在重新加载...");
|
||||
|
||||
let mut s = state.write().await;
|
||||
match s.gemini_provider.load_credentials().await {
|
||||
Ok(_) => {
|
||||
logs.write()
|
||||
.await
|
||||
.add("info", "[Gemini][自动检测] 凭证重新加载成功");
|
||||
Ok(CheckResult {
|
||||
changed: true,
|
||||
new_hash,
|
||||
reloaded: true,
|
||||
})
|
||||
}
|
||||
Err(e) => {
|
||||
logs.write().await.add(
|
||||
"error",
|
||||
&format!("[Gemini][自动检测] 凭证重新加载失败: {e}"),
|
||||
);
|
||||
Ok(CheckResult {
|
||||
changed: true,
|
||||
new_hash,
|
||||
reloaded: false,
|
||||
})
|
||||
}
|
||||
}
|
||||
} else {
|
||||
Ok(CheckResult {
|
||||
changed: false,
|
||||
new_hash,
|
||||
reloaded: false,
|
||||
})
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,231 @@
|
||||
//! Kiro Provider 命令 (Legacy)
|
||||
//!
|
||||
//! 包含 Kiro 凭证管理相关命令。
|
||||
//! 这些命令保留用于向后兼容,新代码应使用统一的 OAuth 命令。
|
||||
|
||||
use crate::app::types::{AppState, LogState};
|
||||
use crate::app::utils::mask_token;
|
||||
use crate::providers;
|
||||
|
||||
/// Kiro 凭证状态
|
||||
#[derive(serde::Serialize)]
|
||||
pub struct KiroCredentialStatus {
|
||||
pub loaded: bool,
|
||||
pub has_access_token: bool,
|
||||
pub has_refresh_token: bool,
|
||||
pub region: Option<String>,
|
||||
pub auth_method: Option<String>,
|
||||
pub expires_at: Option<String>,
|
||||
pub creds_path: String,
|
||||
}
|
||||
|
||||
/// 环境变量
|
||||
#[derive(serde::Serialize)]
|
||||
pub struct EnvVariable {
|
||||
pub key: String,
|
||||
pub value: String,
|
||||
pub masked: String,
|
||||
}
|
||||
|
||||
/// 检查结果
|
||||
#[derive(serde::Serialize)]
|
||||
pub struct CheckResult {
|
||||
pub changed: bool,
|
||||
pub new_hash: String,
|
||||
pub reloaded: bool,
|
||||
}
|
||||
|
||||
/// 刷新 Kiro Token
|
||||
#[tauri::command]
|
||||
pub async fn refresh_kiro_token(
|
||||
state: tauri::State<'_, AppState>,
|
||||
logs: tauri::State<'_, LogState>,
|
||||
) -> Result<String, String> {
|
||||
let mut s = state.write().await;
|
||||
logs.write().await.add("info", "Refreshing Kiro token...");
|
||||
let result = s
|
||||
.kiro_provider
|
||||
.refresh_token()
|
||||
.await
|
||||
.map_err(|e| e.to_string());
|
||||
match &result {
|
||||
Ok(_) => logs
|
||||
.write()
|
||||
.await
|
||||
.add("info", "Token refreshed successfully"),
|
||||
Err(e) => logs
|
||||
.write()
|
||||
.await
|
||||
.add("error", &format!("Token refresh failed: {e}")),
|
||||
}
|
||||
result
|
||||
}
|
||||
|
||||
/// 重新加载凭证
|
||||
#[tauri::command]
|
||||
pub async fn reload_credentials(
|
||||
state: tauri::State<'_, AppState>,
|
||||
logs: tauri::State<'_, LogState>,
|
||||
) -> Result<String, String> {
|
||||
let mut s = state.write().await;
|
||||
logs.write().await.add("info", "Reloading credentials...");
|
||||
s.kiro_provider
|
||||
.load_credentials()
|
||||
.await
|
||||
.map_err(|e| e.to_string())?;
|
||||
logs.write().await.add("info", "Credentials reloaded");
|
||||
Ok("Credentials reloaded".to_string())
|
||||
}
|
||||
|
||||
/// 获取 Kiro 凭证状态
|
||||
#[tauri::command]
|
||||
pub async fn get_kiro_credentials(
|
||||
state: tauri::State<'_, AppState>,
|
||||
) -> Result<KiroCredentialStatus, String> {
|
||||
let s = state.read().await;
|
||||
let creds = &s.kiro_provider.credentials;
|
||||
let path = providers::kiro::KiroProvider::default_creds_path();
|
||||
|
||||
Ok(KiroCredentialStatus {
|
||||
loaded: creds.access_token.is_some() || creds.refresh_token.is_some(),
|
||||
has_access_token: creds.access_token.is_some(),
|
||||
has_refresh_token: creds.refresh_token.is_some(),
|
||||
region: creds.region.clone(),
|
||||
auth_method: creds.auth_method.clone(),
|
||||
expires_at: creds.expires_at.clone(),
|
||||
creds_path: path.to_string_lossy().to_string(),
|
||||
})
|
||||
}
|
||||
|
||||
/// 获取环境变量
|
||||
#[tauri::command]
|
||||
pub async fn get_env_variables(
|
||||
state: tauri::State<'_, AppState>,
|
||||
) -> Result<Vec<EnvVariable>, String> {
|
||||
let s = state.read().await;
|
||||
let creds = &s.kiro_provider.credentials;
|
||||
let mut vars = Vec::new();
|
||||
|
||||
// P0 安全修复:不再返回明文敏感凭证,仅返回 masked 版本
|
||||
if let Some(token) = &creds.access_token {
|
||||
vars.push(EnvVariable {
|
||||
key: "KIRO_ACCESS_TOKEN".to_string(),
|
||||
value: String::new(), // 不返回明文
|
||||
masked: mask_token(token),
|
||||
});
|
||||
}
|
||||
if let Some(token) = &creds.refresh_token {
|
||||
vars.push(EnvVariable {
|
||||
key: "KIRO_REFRESH_TOKEN".to_string(),
|
||||
value: String::new(), // 不返回明文
|
||||
masked: mask_token(token),
|
||||
});
|
||||
}
|
||||
if let Some(id) = &creds.client_id {
|
||||
vars.push(EnvVariable {
|
||||
key: "KIRO_CLIENT_ID".to_string(),
|
||||
value: String::new(), // 不返回明文
|
||||
masked: mask_token(id),
|
||||
});
|
||||
}
|
||||
if let Some(secret) = &creds.client_secret {
|
||||
vars.push(EnvVariable {
|
||||
key: "KIRO_CLIENT_SECRET".to_string(),
|
||||
value: String::new(), // 不返回明文
|
||||
masked: mask_token(secret),
|
||||
});
|
||||
}
|
||||
if let Some(arn) = &creds.profile_arn {
|
||||
vars.push(EnvVariable {
|
||||
key: "KIRO_PROFILE_ARN".to_string(),
|
||||
value: arn.clone(),
|
||||
masked: arn.clone(),
|
||||
});
|
||||
}
|
||||
if let Some(region) = &creds.region {
|
||||
vars.push(EnvVariable {
|
||||
key: "KIRO_REGION".to_string(),
|
||||
value: region.clone(),
|
||||
masked: region.clone(),
|
||||
});
|
||||
}
|
||||
if let Some(method) = &creds.auth_method {
|
||||
vars.push(EnvVariable {
|
||||
key: "KIRO_AUTH_METHOD".to_string(),
|
||||
value: method.clone(),
|
||||
masked: method.clone(),
|
||||
});
|
||||
}
|
||||
|
||||
Ok(vars)
|
||||
}
|
||||
|
||||
/// 获取 Token 文件哈希
|
||||
#[tauri::command]
|
||||
pub async fn get_token_file_hash() -> Result<String, String> {
|
||||
let path = providers::kiro::KiroProvider::default_creds_path();
|
||||
if !tokio::fs::try_exists(&path).await.unwrap_or(false) {
|
||||
return Ok("".to_string());
|
||||
}
|
||||
|
||||
let content = tokio::fs::read(&path).await.map_err(|e| e.to_string())?;
|
||||
let hash = format!("{:x}", md5::compute(&content));
|
||||
Ok(hash)
|
||||
}
|
||||
|
||||
/// 检查凭证文件变化并自动重新加载
|
||||
#[tauri::command]
|
||||
pub async fn check_and_reload_credentials(
|
||||
state: tauri::State<'_, AppState>,
|
||||
logs: tauri::State<'_, LogState>,
|
||||
last_hash: String,
|
||||
) -> Result<CheckResult, String> {
|
||||
let path = providers::kiro::KiroProvider::default_creds_path();
|
||||
|
||||
if !tokio::fs::try_exists(&path).await.unwrap_or(false) {
|
||||
return Ok(CheckResult {
|
||||
changed: false,
|
||||
new_hash: "".to_string(),
|
||||
reloaded: false,
|
||||
});
|
||||
}
|
||||
|
||||
let content = tokio::fs::read(&path).await.map_err(|e| e.to_string())?;
|
||||
let new_hash = format!("{:x}", md5::compute(&content));
|
||||
|
||||
if !last_hash.is_empty() && new_hash != last_hash {
|
||||
logs.write()
|
||||
.await
|
||||
.add("info", "[自动检测] 凭证文件已变化,正在重新加载...");
|
||||
|
||||
let mut s = state.write().await;
|
||||
match s.kiro_provider.load_credentials().await {
|
||||
Ok(_) => {
|
||||
logs.write()
|
||||
.await
|
||||
.add("info", "[自动检测] 凭证重新加载成功");
|
||||
Ok(CheckResult {
|
||||
changed: true,
|
||||
new_hash,
|
||||
reloaded: true,
|
||||
})
|
||||
}
|
||||
Err(e) => {
|
||||
logs.write()
|
||||
.await
|
||||
.add("error", &format!("[自动检测] 凭证重新加载失败: {e}"));
|
||||
Ok(CheckResult {
|
||||
changed: true,
|
||||
new_hash,
|
||||
reloaded: false,
|
||||
})
|
||||
}
|
||||
}
|
||||
} else {
|
||||
Ok(CheckResult {
|
||||
changed: false,
|
||||
new_hash,
|
||||
reloaded: false,
|
||||
})
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,19 @@
|
||||
//! 日志命令
|
||||
//!
|
||||
//! 包含日志查询和清理命令。
|
||||
|
||||
use crate::app::types::LogState;
|
||||
use crate::logger;
|
||||
|
||||
/// 获取日志
|
||||
#[tauri::command]
|
||||
pub async fn get_logs(logs: tauri::State<'_, LogState>) -> Result<Vec<logger::LogEntry>, String> {
|
||||
Ok(logs.read().await.get_logs())
|
||||
}
|
||||
|
||||
/// 清除日志
|
||||
#[tauri::command]
|
||||
pub async fn clear_logs(logs: tauri::State<'_, LogState>) -> Result<(), String> {
|
||||
logs.write().await.clear();
|
||||
Ok(())
|
||||
}
|
||||
@@ -0,0 +1,32 @@
|
||||
//! 内置 Tauri 命令模块
|
||||
//!
|
||||
//! 包含 lib.rs 中定义的所有 Tauri 命令,按功能分类。
|
||||
//!
|
||||
//! ## 模块结构
|
||||
//! - `server` - 服务器控制命令
|
||||
//! - `config` - 配置管理命令
|
||||
//! - `kiro` - Kiro Provider 命令 (legacy)
|
||||
//! - `gemini` - Gemini Provider 命令 (legacy)
|
||||
//! - `qwen` - Qwen Provider 命令 (legacy)
|
||||
//! - `custom_providers` - 自定义 Provider 命令 (OpenAI/Claude Custom)
|
||||
//! - `logs` - 日志命令
|
||||
//! - `api_test` - API 测试和兼容性检查命令
|
||||
|
||||
mod api_test;
|
||||
mod config;
|
||||
mod custom_providers;
|
||||
mod gemini;
|
||||
mod kiro;
|
||||
mod logs;
|
||||
mod qwen;
|
||||
mod server;
|
||||
|
||||
// 重新导出所有命令
|
||||
pub use api_test::*;
|
||||
pub use config::*;
|
||||
pub use custom_providers::*;
|
||||
pub use gemini::*;
|
||||
pub use kiro::*;
|
||||
pub use logs::*;
|
||||
pub use qwen::*;
|
||||
pub use server::*;
|
||||
@@ -0,0 +1,190 @@
|
||||
//! Qwen Provider 命令 (Legacy)
|
||||
//!
|
||||
//! 包含 Qwen 凭证管理相关命令。
|
||||
//! 这些命令保留用于向后兼容,新代码应使用统一的 OAuth 命令。
|
||||
|
||||
use crate::app::commands::kiro::{CheckResult, EnvVariable};
|
||||
use crate::app::types::{AppState, LogState};
|
||||
use crate::app::utils::mask_token;
|
||||
use crate::providers;
|
||||
|
||||
/// Qwen 凭证状态
|
||||
#[derive(serde::Serialize)]
|
||||
pub struct QwenCredentialStatus {
|
||||
pub loaded: bool,
|
||||
pub has_access_token: bool,
|
||||
pub has_refresh_token: bool,
|
||||
pub expiry_date: Option<i64>,
|
||||
pub is_valid: bool,
|
||||
pub creds_path: String,
|
||||
}
|
||||
|
||||
/// 获取 Qwen 凭证状态
|
||||
#[tauri::command]
|
||||
pub async fn get_qwen_credentials(
|
||||
state: tauri::State<'_, AppState>,
|
||||
) -> Result<QwenCredentialStatus, String> {
|
||||
let s = state.read().await;
|
||||
let creds = &s.qwen_provider.credentials;
|
||||
let path = providers::qwen::QwenProvider::default_creds_path();
|
||||
|
||||
Ok(QwenCredentialStatus {
|
||||
loaded: creds.access_token.is_some() || creds.refresh_token.is_some(),
|
||||
has_access_token: creds.access_token.is_some(),
|
||||
has_refresh_token: creds.refresh_token.is_some(),
|
||||
expiry_date: creds.expiry_date,
|
||||
is_valid: s.qwen_provider.is_token_valid(),
|
||||
creds_path: path.to_string_lossy().to_string(),
|
||||
})
|
||||
}
|
||||
|
||||
/// 重新加载 Qwen 凭证
|
||||
#[tauri::command]
|
||||
pub async fn reload_qwen_credentials(
|
||||
state: tauri::State<'_, AppState>,
|
||||
logs: tauri::State<'_, LogState>,
|
||||
) -> Result<String, String> {
|
||||
let mut s = state.write().await;
|
||||
logs.write().await.add("info", "[Qwen] 正在加载凭证...");
|
||||
s.qwen_provider
|
||||
.load_credentials()
|
||||
.await
|
||||
.map_err(|e| e.to_string())?;
|
||||
logs.write().await.add("info", "[Qwen] 凭证加载成功");
|
||||
Ok("Qwen credentials reloaded".to_string())
|
||||
}
|
||||
|
||||
/// 刷新 Qwen Token
|
||||
#[tauri::command]
|
||||
pub async fn refresh_qwen_token(
|
||||
state: tauri::State<'_, AppState>,
|
||||
logs: tauri::State<'_, LogState>,
|
||||
) -> Result<String, String> {
|
||||
let mut s = state.write().await;
|
||||
logs.write().await.add("info", "[Qwen] 正在刷新 Token...");
|
||||
let result = s
|
||||
.qwen_provider
|
||||
.refresh_token()
|
||||
.await
|
||||
.map_err(|e| e.to_string());
|
||||
match &result {
|
||||
Ok(_) => logs.write().await.add("info", "[Qwen] Token 刷新成功"),
|
||||
Err(e) => logs
|
||||
.write()
|
||||
.await
|
||||
.add("error", &format!("[Qwen] Token 刷新失败: {e}")),
|
||||
}
|
||||
result
|
||||
}
|
||||
|
||||
/// 获取 Qwen 环境变量
|
||||
#[tauri::command]
|
||||
pub async fn get_qwen_env_variables(
|
||||
state: tauri::State<'_, AppState>,
|
||||
) -> Result<Vec<EnvVariable>, String> {
|
||||
let s = state.read().await;
|
||||
let creds = &s.qwen_provider.credentials;
|
||||
let mut vars = Vec::new();
|
||||
|
||||
if let Some(token) = &creds.access_token {
|
||||
vars.push(EnvVariable {
|
||||
key: "QWEN_ACCESS_TOKEN".to_string(),
|
||||
value: token.clone(),
|
||||
masked: mask_token(token),
|
||||
});
|
||||
}
|
||||
if let Some(token) = &creds.refresh_token {
|
||||
vars.push(EnvVariable {
|
||||
key: "QWEN_REFRESH_TOKEN".to_string(),
|
||||
value: token.clone(),
|
||||
masked: mask_token(token),
|
||||
});
|
||||
}
|
||||
if let Some(url) = &creds.resource_url {
|
||||
vars.push(EnvVariable {
|
||||
key: "QWEN_RESOURCE_URL".to_string(),
|
||||
value: url.clone(),
|
||||
masked: url.clone(),
|
||||
});
|
||||
}
|
||||
if let Some(expiry) = creds.expiry_date {
|
||||
let expiry_str = expiry.to_string();
|
||||
vars.push(EnvVariable {
|
||||
key: "QWEN_EXPIRY_DATE".to_string(),
|
||||
value: expiry_str.clone(),
|
||||
masked: expiry_str,
|
||||
});
|
||||
}
|
||||
|
||||
Ok(vars)
|
||||
}
|
||||
|
||||
/// 获取 Qwen Token 文件哈希
|
||||
#[tauri::command]
|
||||
pub async fn get_qwen_token_file_hash() -> Result<String, String> {
|
||||
let path = providers::qwen::QwenProvider::default_creds_path();
|
||||
if !tokio::fs::try_exists(&path).await.unwrap_or(false) {
|
||||
return Ok("".to_string());
|
||||
}
|
||||
|
||||
let content = tokio::fs::read(&path).await.map_err(|e| e.to_string())?;
|
||||
let hash = format!("{:x}", md5::compute(&content));
|
||||
Ok(hash)
|
||||
}
|
||||
|
||||
/// 检查并重新加载 Qwen 凭证
|
||||
#[tauri::command]
|
||||
pub async fn check_and_reload_qwen_credentials(
|
||||
state: tauri::State<'_, AppState>,
|
||||
logs: tauri::State<'_, LogState>,
|
||||
last_hash: String,
|
||||
) -> Result<CheckResult, String> {
|
||||
let path = providers::qwen::QwenProvider::default_creds_path();
|
||||
|
||||
if !tokio::fs::try_exists(&path).await.unwrap_or(false) {
|
||||
return Ok(CheckResult {
|
||||
changed: false,
|
||||
new_hash: "".to_string(),
|
||||
reloaded: false,
|
||||
});
|
||||
}
|
||||
|
||||
let content = tokio::fs::read(&path).await.map_err(|e| e.to_string())?;
|
||||
let new_hash = format!("{:x}", md5::compute(&content));
|
||||
|
||||
if !last_hash.is_empty() && new_hash != last_hash {
|
||||
logs.write()
|
||||
.await
|
||||
.add("info", "[Qwen][自动检测] 凭证文件已变化,正在重新加载...");
|
||||
|
||||
let mut s = state.write().await;
|
||||
match s.qwen_provider.load_credentials().await {
|
||||
Ok(_) => {
|
||||
logs.write()
|
||||
.await
|
||||
.add("info", "[Qwen][自动检测] 凭证重新加载成功");
|
||||
Ok(CheckResult {
|
||||
changed: true,
|
||||
new_hash,
|
||||
reloaded: true,
|
||||
})
|
||||
}
|
||||
Err(e) => {
|
||||
logs.write()
|
||||
.await
|
||||
.add("error", &format!("[Qwen][自动检测] 凭证重新加载失败: {e}"));
|
||||
Ok(CheckResult {
|
||||
changed: true,
|
||||
new_hash,
|
||||
reloaded: false,
|
||||
})
|
||||
}
|
||||
}
|
||||
} else {
|
||||
Ok(CheckResult {
|
||||
changed: false,
|
||||
new_hash,
|
||||
reloaded: false,
|
||||
})
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,68 @@
|
||||
//! 服务器控制命令
|
||||
//!
|
||||
//! 包含服务器启动、停止、状态查询等命令。
|
||||
|
||||
use crate::app::types::{AppState, LogState};
|
||||
use crate::app::TokenCacheServiceState;
|
||||
use crate::commands::provider_pool_cmd::ProviderPoolServiceState;
|
||||
use crate::commands::telemetry_cmd::TelemetryState;
|
||||
use crate::database;
|
||||
use crate::server;
|
||||
|
||||
/// 启动服务器
|
||||
#[tauri::command]
|
||||
pub async fn start_server(
|
||||
state: tauri::State<'_, AppState>,
|
||||
logs: tauri::State<'_, LogState>,
|
||||
db: tauri::State<'_, database::DbConnection>,
|
||||
pool_service: tauri::State<'_, ProviderPoolServiceState>,
|
||||
token_cache: tauri::State<'_, TokenCacheServiceState>,
|
||||
) -> Result<String, String> {
|
||||
let mut s = state.write().await;
|
||||
logs.write().await.add("info", "Starting server...");
|
||||
s.start(
|
||||
logs.inner().clone(),
|
||||
pool_service.0.clone(),
|
||||
token_cache.0.clone(),
|
||||
Some(db.inner().clone()),
|
||||
)
|
||||
.await
|
||||
.map_err(|e| e.to_string())?;
|
||||
logs.write().await.add(
|
||||
"info",
|
||||
&format!(
|
||||
"Server started on {}:{}",
|
||||
s.config.server.host, s.config.server.port
|
||||
),
|
||||
);
|
||||
Ok("Server started".to_string())
|
||||
}
|
||||
|
||||
/// 停止服务器
|
||||
#[tauri::command]
|
||||
pub async fn stop_server(
|
||||
state: tauri::State<'_, AppState>,
|
||||
logs: tauri::State<'_, LogState>,
|
||||
) -> Result<String, String> {
|
||||
let mut s = state.write().await;
|
||||
s.stop().await;
|
||||
logs.write().await.add("info", "Server stopped");
|
||||
Ok("Server stopped".to_string())
|
||||
}
|
||||
|
||||
/// 获取服务器状态
|
||||
#[tauri::command]
|
||||
pub async fn get_server_status(
|
||||
state: tauri::State<'_, AppState>,
|
||||
telemetry_state: tauri::State<'_, TelemetryState>,
|
||||
) -> Result<server::ServerStatus, String> {
|
||||
let s = state.read().await;
|
||||
let mut status = s.status();
|
||||
|
||||
// 从遥测系统获取真实的请求计数
|
||||
let stats = telemetry_state.stats.read();
|
||||
let summary = stats.summary(None);
|
||||
status.requests = summary.total_requests;
|
||||
|
||||
Ok(status)
|
||||
}
|
||||
@@ -0,0 +1,26 @@
|
||||
//! 应用核心模块
|
||||
//!
|
||||
//! 包含 Tauri 应用的核心类型、状态管理和启动逻辑。
|
||||
//!
|
||||
//! ## 模块结构
|
||||
//! - `types` - 核心类型定义(ProviderType 等)
|
||||
//! - `state` - 状态类型和初始化
|
||||
//! - `setup` - Tauri setup hook
|
||||
//! - `commands` - 内置 Tauri 命令
|
||||
//! - `utils` - 辅助函数
|
||||
//! - `bootstrap` - 应用启动引导(配置验证、状态初始化)
|
||||
//! - `runner` - 应用运行器(Tauri Builder 配置和命令注册)
|
||||
|
||||
pub mod bootstrap;
|
||||
pub mod commands;
|
||||
pub mod runner;
|
||||
mod setup;
|
||||
mod state;
|
||||
mod types;
|
||||
mod utils;
|
||||
|
||||
pub use runner::run;
|
||||
pub use setup::setup_app;
|
||||
pub use state::*;
|
||||
pub use types::*;
|
||||
pub use utils::*;
|
||||
@@ -0,0 +1,878 @@
|
||||
//! 应用运行器模块
|
||||
//!
|
||||
//! 包含 Tauri 应用的主入口函数和命令注册。
|
||||
|
||||
use std::sync::Arc;
|
||||
use tauri::Manager;
|
||||
|
||||
use crate::commands;
|
||||
use crate::tray::{TrayIconStatus, TrayManager, TrayStateSnapshot};
|
||||
|
||||
use super::bootstrap::{self, AppStates};
|
||||
use super::commands as app_commands;
|
||||
use super::types::{AppState, TrayManagerState};
|
||||
|
||||
/// 运行 Tauri 应用
|
||||
///
|
||||
/// 这是应用的主入口点,负责:
|
||||
/// 1. 加载和验证配置
|
||||
/// 2. 初始化所有应用状态
|
||||
/// 3. 配置 Tauri Builder(插件、状态管理、事件处理)
|
||||
/// 4. 注册所有 Tauri 命令
|
||||
/// 5. 启动应用
|
||||
#[cfg_attr(mobile, tauri::mobile_entry_point)]
|
||||
pub fn run() {
|
||||
// 加载并验证配置
|
||||
let config = match bootstrap::load_and_validate_config() {
|
||||
Ok(cfg) => cfg,
|
||||
Err(err) => {
|
||||
tracing::error!("{}", err);
|
||||
eprintln!("{}", err);
|
||||
return;
|
||||
}
|
||||
};
|
||||
|
||||
// 初始化所有应用状态
|
||||
let states = match bootstrap::init_states(&config) {
|
||||
Ok(s) => s,
|
||||
Err(err) => {
|
||||
tracing::error!("应用状态初始化失败: {}", err);
|
||||
eprintln!("应用状态初始化失败: {}", err);
|
||||
return;
|
||||
}
|
||||
};
|
||||
|
||||
// 解构状态以便使用
|
||||
let AppStates {
|
||||
state,
|
||||
logs,
|
||||
db,
|
||||
skill_service: skill_service_state,
|
||||
provider_pool_service: provider_pool_service_state,
|
||||
api_key_provider_service: api_key_provider_service_state,
|
||||
credential_sync_service: credential_sync_service_state,
|
||||
token_cache_service: token_cache_service_state,
|
||||
machine_id_service: machine_id_service_state,
|
||||
router_config: router_config_state,
|
||||
resilience_config: resilience_config_state,
|
||||
plugin_manager: plugin_manager_state,
|
||||
plugin_installer: plugin_installer_state,
|
||||
telemetry: telemetry_state,
|
||||
flow_monitor: flow_monitor_state,
|
||||
flow_query_service: flow_query_service_state,
|
||||
flow_interceptor: flow_interceptor_state,
|
||||
flow_replayer: flow_replayer_state,
|
||||
session_manager: session_manager_state,
|
||||
quick_filter_manager: quick_filter_manager_state,
|
||||
bookmark_manager: bookmark_manager_state,
|
||||
enhanced_stats_service: enhanced_stats_service_state,
|
||||
batch_operations: batch_operations_state,
|
||||
browser_interceptor: browser_interceptor_state,
|
||||
native_agent: native_agent_state,
|
||||
oauth_plugin_manager: oauth_plugin_manager_state,
|
||||
orchestrator: orchestrator_state,
|
||||
shared_stats,
|
||||
shared_tokens,
|
||||
shared_logger,
|
||||
flow_monitor_arc: flow_monitor,
|
||||
flow_interceptor_arc: flow_interceptor,
|
||||
} = states;
|
||||
|
||||
// Clone for setup hook
|
||||
let state_clone = state.clone();
|
||||
let logs_clone = logs.clone();
|
||||
let db_clone = db.clone();
|
||||
let pool_service_clone = provider_pool_service_state.0.clone();
|
||||
let token_cache_clone = token_cache_service_state.0.clone();
|
||||
let shared_stats_clone = shared_stats.clone();
|
||||
let shared_tokens_clone = shared_tokens.clone();
|
||||
let shared_logger_clone = shared_logger.clone();
|
||||
let flow_monitor_clone = flow_monitor.clone();
|
||||
let flow_interceptor_clone = flow_interceptor.clone();
|
||||
|
||||
let mut builder = tauri::Builder::default()
|
||||
.plugin(tauri_plugin_shell::init())
|
||||
.plugin(tauri_plugin_dialog::init())
|
||||
.plugin(tauri_plugin_autostart::init(
|
||||
tauri_plugin_autostart::MacosLauncher::LaunchAgent,
|
||||
Some(vec!["--minimized"]),
|
||||
))
|
||||
// 单实例插件:当第二个实例启动时,将 URL 传递给第一个实例
|
||||
.plugin(tauri_plugin_single_instance::init(|app, args, _cwd| {
|
||||
tracing::info!("[单实例] 收到来自新实例的参数: {:?}", args);
|
||||
|
||||
// 处理传入的 URL 参数
|
||||
for arg in args.iter().skip(1) {
|
||||
// 跳过第一个参数(程序路径)
|
||||
if arg.starts_with("http://") || arg.starts_with("https://") {
|
||||
tracing::info!("[单实例] 收到 URL: {}", arg);
|
||||
|
||||
#[cfg(target_os = "macos")]
|
||||
{
|
||||
crate::browser_interceptor::platform::macos::handle_deep_link_url(
|
||||
arg.clone(),
|
||||
);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// 将窗口带到前台
|
||||
if let Some(window) = app.get_webview_window("main") {
|
||||
let _ = window.show();
|
||||
let _ = window.set_focus();
|
||||
}
|
||||
}));
|
||||
|
||||
// 添加 Deep Link 插件(用于浏览器拦截)
|
||||
#[cfg(target_os = "macos")]
|
||||
{
|
||||
builder = builder.plugin(tauri_plugin_deep_link::init());
|
||||
}
|
||||
|
||||
builder
|
||||
.manage(state)
|
||||
.manage(logs)
|
||||
.manage(db)
|
||||
.manage(skill_service_state)
|
||||
.manage(provider_pool_service_state)
|
||||
.manage(api_key_provider_service_state)
|
||||
.manage(credential_sync_service_state)
|
||||
.manage(token_cache_service_state)
|
||||
.manage(machine_id_service_state)
|
||||
.manage(router_config_state)
|
||||
.manage(resilience_config_state)
|
||||
.manage(telemetry_state)
|
||||
.manage(plugin_manager_state)
|
||||
.manage(plugin_installer_state)
|
||||
.manage(flow_monitor_state)
|
||||
.manage(flow_query_service_state)
|
||||
.manage(flow_interceptor_state)
|
||||
.manage(flow_replayer_state)
|
||||
.manage(session_manager_state)
|
||||
.manage(quick_filter_manager_state)
|
||||
.manage(bookmark_manager_state)
|
||||
.manage(enhanced_stats_service_state)
|
||||
.manage(batch_operations_state)
|
||||
.manage(browser_interceptor_state)
|
||||
.manage(native_agent_state)
|
||||
.manage(oauth_plugin_manager_state)
|
||||
.manage(orchestrator_state)
|
||||
.on_window_event(move |window, event| {
|
||||
// 处理窗口关闭事件
|
||||
if let tauri::WindowEvent::CloseRequested { api, .. } = event {
|
||||
// 获取配置,检查是否启用最小化到托盘
|
||||
let app_handle = window.app_handle();
|
||||
if let Some(app_state) = app_handle.try_state::<AppState>() {
|
||||
// 使用 block_on 同步获取配置
|
||||
let minimize_to_tray = tauri::async_runtime::block_on(async {
|
||||
let state = app_state.read().await;
|
||||
state.config.minimize_to_tray
|
||||
});
|
||||
|
||||
if minimize_to_tray {
|
||||
// 阻止默认关闭行为
|
||||
api.prevent_close();
|
||||
// 隐藏窗口而不是关闭
|
||||
if let Err(e) = window.hide() {
|
||||
tracing::error!("[窗口] 隐藏窗口失败: {}", e);
|
||||
} else {
|
||||
tracing::info!("[窗口] 窗口已最小化到托盘");
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
})
|
||||
.setup(move |app| {
|
||||
// 设置 deep-link 事件监听(用于浏览器拦截)
|
||||
#[cfg(target_os = "macos")]
|
||||
{
|
||||
use tauri_plugin_deep_link::DeepLinkExt;
|
||||
let _listener_id = app.deep_link().on_open_url(|event| {
|
||||
for url in event.urls() {
|
||||
tracing::info!("[Deep Link] 收到 URL: {}", url);
|
||||
crate::browser_interceptor::platform::macos::handle_deep_link_url(
|
||||
url.to_string(),
|
||||
);
|
||||
}
|
||||
});
|
||||
tracing::info!("[启动] Deep Link 事件监听已设置");
|
||||
}
|
||||
|
||||
// 初始化托盘管理器
|
||||
// Requirements 1.4: 应用启动时显示停止状态图标
|
||||
match TrayManager::new(app.handle()) {
|
||||
Ok(tray_manager) => {
|
||||
tracing::info!("[启动] 托盘管理器初始化成功");
|
||||
// 将托盘管理器存储到应用状态中
|
||||
let tray_state: TrayManagerState<tauri::Wry> =
|
||||
TrayManagerState(Arc::new(tokio::sync::RwLock::new(Some(tray_manager))));
|
||||
app.manage(tray_state);
|
||||
}
|
||||
Err(e) => {
|
||||
tracing::error!("[启动] 托盘管理器初始化失败: {}", e);
|
||||
// 即使托盘初始化失败,应用仍然可以运行
|
||||
let tray_state: TrayManagerState<tauri::Wry> =
|
||||
TrayManagerState(Arc::new(tokio::sync::RwLock::new(None)));
|
||||
app.manage(tray_state);
|
||||
}
|
||||
}
|
||||
// 自动启动服务器
|
||||
let state = state_clone.clone();
|
||||
let logs = logs_clone.clone();
|
||||
let db = db_clone.clone();
|
||||
let pool_service = pool_service_clone.clone();
|
||||
let token_cache = token_cache_clone.clone();
|
||||
let shared_stats = shared_stats_clone.clone();
|
||||
let shared_tokens = shared_tokens_clone.clone();
|
||||
let shared_logger = shared_logger_clone.clone();
|
||||
let shared_flow_monitor = flow_monitor_clone.clone();
|
||||
let app_handle = app.handle().clone();
|
||||
tauri::async_runtime::spawn(async move {
|
||||
// 先加载凭证池中的凭证
|
||||
{
|
||||
logs.write().await.add("info", "[启动] 正在加载凭证池...");
|
||||
|
||||
// 获取凭证池概览信息
|
||||
match pool_service.get_overview(&db) {
|
||||
Ok(overview) => {
|
||||
let mut loaded_types = Vec::new();
|
||||
let mut total_credentials = 0;
|
||||
|
||||
for provider_overview in overview {
|
||||
let count = provider_overview.stats.total_count;
|
||||
if count > 0 {
|
||||
total_credentials += count;
|
||||
let provider_name =
|
||||
match provider_overview.provider_type.as_str() {
|
||||
"kiro" => "Kiro",
|
||||
"gemini" => "Gemini",
|
||||
"qwen" => "通义千问",
|
||||
"antigravity" => "Antigravity",
|
||||
"openai" => "OpenAI",
|
||||
"claude" => "Claude",
|
||||
"codex" => "Codex",
|
||||
"claude_oauth" => "Claude OAuth",
|
||||
"iflow" => "iFlow",
|
||||
_ => &provider_overview.provider_type,
|
||||
};
|
||||
loaded_types.push(format!("{} ({} 个)", provider_name, count));
|
||||
}
|
||||
}
|
||||
|
||||
if loaded_types.is_empty() {
|
||||
logs.write().await.add("warn", "[启动] 未找到任何可用凭证");
|
||||
} else {
|
||||
let message = format!(
|
||||
"[启动] 凭证已加载: {} (共 {} 个)",
|
||||
loaded_types.join(", "),
|
||||
total_credentials
|
||||
);
|
||||
logs.write().await.add("info", &message);
|
||||
}
|
||||
}
|
||||
Err(e) => {
|
||||
logs.write()
|
||||
.await
|
||||
.add("warn", &format!("[启动] 获取凭证池信息失败: {}", e));
|
||||
}
|
||||
}
|
||||
|
||||
// 兼容性:仍然尝试加载旧的 Kiro 凭证(如果存在)
|
||||
let mut s = state.write().await;
|
||||
if let Err(e) = s.kiro_provider.load_credentials().await {
|
||||
logs.write()
|
||||
.await
|
||||
.add("debug", &format!("[启动] 旧版 Kiro 凭证加载失败: {e}"));
|
||||
}
|
||||
}
|
||||
// 启动服务器(使用共享的遥测实例和 Flow Monitor)
|
||||
let server_started;
|
||||
let server_address;
|
||||
{
|
||||
let mut s = state.write().await;
|
||||
logs.write()
|
||||
.await
|
||||
.add("info", "[启动] 正在自动启动服务器...");
|
||||
match s
|
||||
.start_with_telemetry_and_flow_monitor(
|
||||
logs.clone(),
|
||||
pool_service,
|
||||
token_cache,
|
||||
Some(db),
|
||||
Some(shared_stats),
|
||||
Some(shared_tokens),
|
||||
Some(shared_logger),
|
||||
Some(shared_flow_monitor),
|
||||
Some(flow_interceptor_clone),
|
||||
)
|
||||
.await
|
||||
{
|
||||
Ok(_) => {
|
||||
let host = s.config.server.host.clone();
|
||||
let port = s.config.server.port;
|
||||
logs.write()
|
||||
.await
|
||||
.add("info", &format!("[启动] 服务器已启动: {host}:{port}"));
|
||||
server_started = true;
|
||||
server_address = format!("{}:{}", host, port);
|
||||
}
|
||||
Err(e) => {
|
||||
logs.write()
|
||||
.await
|
||||
.add("error", &format!("[启动] 服务器启动失败: {e}"));
|
||||
server_started = false;
|
||||
server_address = String::new();
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// 更新托盘状态
|
||||
// Requirements 7.1: API 服务器状态变化时更新托盘图标
|
||||
if let Some(tray_state) = app_handle.try_state::<TrayManagerState<tauri::Wry>>() {
|
||||
let tray_guard = tray_state.0.read().await;
|
||||
if let Some(tray_manager) = tray_guard.as_ref() {
|
||||
// 计算初始图标状态
|
||||
// 服务器刚启动时,假设凭证健康(后续会通过状态同步更新)
|
||||
let icon_status = if server_started {
|
||||
TrayIconStatus::Running
|
||||
} else {
|
||||
TrayIconStatus::Stopped
|
||||
};
|
||||
|
||||
let snapshot = TrayStateSnapshot {
|
||||
icon_status,
|
||||
server_running: server_started,
|
||||
server_address,
|
||||
available_credentials: 0, // 初始值,后续通过状态同步更新
|
||||
total_credentials: 0,
|
||||
today_requests: 0,
|
||||
auto_start_enabled: false, // 后续通过状态同步更新
|
||||
};
|
||||
|
||||
if let Err(e) = tray_manager.update_state(snapshot).await {
|
||||
tracing::error!("[启动] 更新托盘状态失败: {}", e);
|
||||
} else {
|
||||
tracing::info!("[启动] 托盘状态已更新");
|
||||
}
|
||||
}
|
||||
}
|
||||
});
|
||||
Ok(())
|
||||
})
|
||||
.invoke_handler(tauri::generate_handler![
|
||||
// Server commands (from app::commands)
|
||||
app_commands::start_server,
|
||||
app_commands::stop_server,
|
||||
app_commands::get_server_status,
|
||||
// Config commands (from app::commands)
|
||||
app_commands::get_config,
|
||||
app_commands::save_config,
|
||||
app_commands::get_default_provider,
|
||||
app_commands::set_default_provider,
|
||||
app_commands::get_endpoint_providers,
|
||||
app_commands::set_endpoint_provider,
|
||||
// Unified OAuth commands (new)
|
||||
commands::oauth_cmd::get_oauth_credentials,
|
||||
commands::oauth_cmd::reload_oauth_credentials,
|
||||
commands::oauth_cmd::refresh_oauth_token,
|
||||
commands::oauth_cmd::get_oauth_env_variables,
|
||||
commands::oauth_cmd::get_oauth_token_file_hash,
|
||||
commands::oauth_cmd::check_and_reload_oauth_credentials,
|
||||
commands::oauth_cmd::get_all_oauth_credentials,
|
||||
// Legacy Kiro commands (from app::commands, deprecated)
|
||||
app_commands::refresh_kiro_token,
|
||||
app_commands::reload_credentials,
|
||||
app_commands::get_kiro_credentials,
|
||||
app_commands::get_env_variables,
|
||||
app_commands::get_token_file_hash,
|
||||
app_commands::check_and_reload_credentials,
|
||||
// Legacy Gemini commands (from app::commands, deprecated)
|
||||
app_commands::get_gemini_credentials,
|
||||
app_commands::reload_gemini_credentials,
|
||||
app_commands::refresh_gemini_token,
|
||||
app_commands::get_gemini_env_variables,
|
||||
app_commands::get_gemini_token_file_hash,
|
||||
app_commands::check_and_reload_gemini_credentials,
|
||||
// Legacy Qwen commands (from app::commands, deprecated)
|
||||
app_commands::get_qwen_credentials,
|
||||
app_commands::reload_qwen_credentials,
|
||||
app_commands::refresh_qwen_token,
|
||||
app_commands::get_qwen_env_variables,
|
||||
app_commands::get_qwen_token_file_hash,
|
||||
app_commands::check_and_reload_qwen_credentials,
|
||||
// OpenAI Custom commands (from app::commands)
|
||||
app_commands::get_openai_custom_status,
|
||||
app_commands::set_openai_custom_config,
|
||||
// Claude Custom commands (from app::commands)
|
||||
app_commands::get_claude_custom_status,
|
||||
app_commands::set_claude_custom_config,
|
||||
// Log commands (from app::commands)
|
||||
app_commands::get_logs,
|
||||
app_commands::clear_logs,
|
||||
// API test commands (from app::commands)
|
||||
app_commands::test_api,
|
||||
app_commands::get_available_models,
|
||||
app_commands::check_api_compatibility,
|
||||
// Switch commands
|
||||
commands::switch_cmd::get_switch_providers,
|
||||
commands::switch_cmd::get_current_switch_provider,
|
||||
commands::switch_cmd::add_switch_provider,
|
||||
commands::switch_cmd::update_switch_provider,
|
||||
commands::switch_cmd::delete_switch_provider,
|
||||
commands::switch_cmd::switch_provider,
|
||||
commands::switch_cmd::import_default_config,
|
||||
commands::switch_cmd::read_live_provider_settings,
|
||||
commands::switch_cmd::check_config_sync_status,
|
||||
commands::switch_cmd::sync_from_external_config,
|
||||
// Config commands
|
||||
commands::config_cmd::get_config_status,
|
||||
commands::config_cmd::get_config_dir_path,
|
||||
commands::config_cmd::open_config_folder,
|
||||
commands::config_cmd::get_tool_versions,
|
||||
commands::config_cmd::get_auto_launch_status,
|
||||
commands::config_cmd::set_auto_launch,
|
||||
// Config import/export commands
|
||||
commands::config_cmd::export_config,
|
||||
commands::config_cmd::validate_config_yaml,
|
||||
commands::config_cmd::import_config,
|
||||
commands::config_cmd::get_config_paths,
|
||||
// Enhanced export/import commands (using ExportService/ImportService)
|
||||
commands::config_cmd::export_bundle,
|
||||
commands::config_cmd::export_config_yaml,
|
||||
commands::config_cmd::validate_import,
|
||||
commands::config_cmd::import_bundle,
|
||||
// Path utility commands
|
||||
commands::config_cmd::expand_path,
|
||||
commands::config_cmd::open_auth_dir,
|
||||
commands::config_cmd::check_for_updates,
|
||||
commands::config_cmd::download_update,
|
||||
// MCP commands
|
||||
commands::mcp_cmd::get_mcp_servers,
|
||||
commands::mcp_cmd::add_mcp_server,
|
||||
commands::mcp_cmd::update_mcp_server,
|
||||
commands::mcp_cmd::delete_mcp_server,
|
||||
commands::mcp_cmd::toggle_mcp_server,
|
||||
commands::mcp_cmd::import_mcp_from_app,
|
||||
commands::mcp_cmd::sync_all_mcp_to_live,
|
||||
// Prompt commands
|
||||
commands::prompt_cmd::get_prompts,
|
||||
commands::prompt_cmd::upsert_prompt,
|
||||
commands::prompt_cmd::add_prompt,
|
||||
commands::prompt_cmd::update_prompt,
|
||||
commands::prompt_cmd::delete_prompt,
|
||||
commands::prompt_cmd::enable_prompt,
|
||||
commands::prompt_cmd::import_prompt_from_file,
|
||||
commands::prompt_cmd::get_current_prompt_file_content,
|
||||
commands::prompt_cmd::auto_import_prompt,
|
||||
commands::prompt_cmd::switch_prompt,
|
||||
// Skill commands
|
||||
commands::skill_cmd::get_skills,
|
||||
commands::skill_cmd::get_skills_for_app,
|
||||
commands::skill_cmd::install_skill,
|
||||
commands::skill_cmd::install_skill_for_app,
|
||||
commands::skill_cmd::uninstall_skill,
|
||||
commands::skill_cmd::uninstall_skill_for_app,
|
||||
commands::skill_cmd::get_skill_repos,
|
||||
commands::skill_cmd::add_skill_repo,
|
||||
commands::skill_cmd::remove_skill_repo,
|
||||
commands::skill_cmd::get_installed_proxycast_skills,
|
||||
// Provider Pool commands
|
||||
commands::provider_pool_cmd::get_provider_pool_overview,
|
||||
commands::provider_pool_cmd::get_provider_pool_credentials,
|
||||
commands::provider_pool_cmd::add_provider_pool_credential,
|
||||
commands::provider_pool_cmd::update_provider_pool_credential,
|
||||
commands::provider_pool_cmd::delete_provider_pool_credential,
|
||||
commands::provider_pool_cmd::toggle_provider_pool_credential,
|
||||
commands::provider_pool_cmd::reset_provider_pool_credential,
|
||||
commands::provider_pool_cmd::reset_provider_pool_health,
|
||||
commands::provider_pool_cmd::check_provider_pool_credential_health,
|
||||
commands::provider_pool_cmd::check_provider_pool_type_health,
|
||||
commands::provider_pool_cmd::add_kiro_oauth_credential,
|
||||
commands::provider_pool_cmd::add_kiro_from_json,
|
||||
commands::provider_pool_cmd::add_gemini_oauth_credential,
|
||||
commands::provider_pool_cmd::add_qwen_oauth_credential,
|
||||
commands::provider_pool_cmd::add_antigravity_oauth_credential,
|
||||
commands::provider_pool_cmd::add_openai_key_credential,
|
||||
commands::provider_pool_cmd::add_claude_key_credential,
|
||||
commands::provider_pool_cmd::add_gemini_api_key_credential,
|
||||
commands::provider_pool_cmd::add_codex_oauth_credential,
|
||||
commands::provider_pool_cmd::add_claude_oauth_credential,
|
||||
commands::provider_pool_cmd::add_iflow_oauth_credential,
|
||||
commands::provider_pool_cmd::add_iflow_cookie_credential,
|
||||
commands::provider_pool_cmd::refresh_pool_credential_token,
|
||||
commands::provider_pool_cmd::get_pool_credential_oauth_status,
|
||||
commands::provider_pool_cmd::debug_kiro_credentials,
|
||||
commands::provider_pool_cmd::test_user_credentials,
|
||||
commands::provider_pool_cmd::migrate_private_config_to_pool,
|
||||
commands::provider_pool_cmd::start_antigravity_oauth_login,
|
||||
commands::provider_pool_cmd::get_antigravity_auth_url_and_wait,
|
||||
commands::provider_pool_cmd::get_codex_auth_url_and_wait,
|
||||
commands::provider_pool_cmd::start_codex_oauth_login,
|
||||
commands::provider_pool_cmd::get_claude_oauth_auth_url_and_wait,
|
||||
commands::provider_pool_cmd::start_claude_oauth_login,
|
||||
commands::provider_pool_cmd::exchange_claude_oauth_code,
|
||||
commands::provider_pool_cmd::claude_oauth_with_cookie,
|
||||
commands::provider_pool_cmd::get_qwen_device_code_and_wait,
|
||||
commands::provider_pool_cmd::start_qwen_device_code_login,
|
||||
commands::provider_pool_cmd::get_iflow_auth_url_and_wait,
|
||||
commands::provider_pool_cmd::start_iflow_oauth_login,
|
||||
commands::provider_pool_cmd::get_gemini_auth_url_and_wait,
|
||||
commands::provider_pool_cmd::start_gemini_oauth_login,
|
||||
commands::provider_pool_cmd::exchange_gemini_code,
|
||||
commands::provider_pool_cmd::get_kiro_credential_fingerprint,
|
||||
commands::provider_pool_cmd::get_credential_health,
|
||||
commands::provider_pool_cmd::get_all_credential_health,
|
||||
// Kiro Builder ID 登录命令
|
||||
commands::provider_pool_cmd::start_kiro_builder_id_login,
|
||||
commands::provider_pool_cmd::poll_kiro_builder_id_auth,
|
||||
commands::provider_pool_cmd::cancel_kiro_builder_id_login,
|
||||
commands::provider_pool_cmd::add_kiro_from_builder_id_auth,
|
||||
// Kiro Social Auth 登录命令 (Google/GitHub)
|
||||
commands::provider_pool_cmd::start_kiro_social_auth_login,
|
||||
commands::provider_pool_cmd::exchange_kiro_social_auth_token,
|
||||
commands::provider_pool_cmd::cancel_kiro_social_auth_login,
|
||||
commands::provider_pool_cmd::start_kiro_social_auth_callback_server,
|
||||
// Playwright 指纹浏览器登录命令
|
||||
commands::provider_pool_cmd::check_playwright_available,
|
||||
commands::provider_pool_cmd::install_playwright,
|
||||
commands::provider_pool_cmd::start_kiro_playwright_login,
|
||||
commands::provider_pool_cmd::cancel_kiro_playwright_login,
|
||||
// API Key Provider commands
|
||||
commands::api_key_provider_cmd::get_api_key_providers,
|
||||
commands::api_key_provider_cmd::get_api_key_provider,
|
||||
commands::api_key_provider_cmd::add_custom_api_key_provider,
|
||||
commands::api_key_provider_cmd::update_api_key_provider,
|
||||
commands::api_key_provider_cmd::delete_custom_api_key_provider,
|
||||
commands::api_key_provider_cmd::add_api_key,
|
||||
commands::api_key_provider_cmd::delete_api_key,
|
||||
commands::api_key_provider_cmd::toggle_api_key,
|
||||
commands::api_key_provider_cmd::update_api_key_alias,
|
||||
commands::api_key_provider_cmd::get_next_api_key,
|
||||
commands::api_key_provider_cmd::record_api_key_usage,
|
||||
commands::api_key_provider_cmd::record_api_key_error,
|
||||
commands::api_key_provider_cmd::get_provider_ui_state,
|
||||
commands::api_key_provider_cmd::set_provider_ui_state,
|
||||
commands::api_key_provider_cmd::update_provider_sort_orders,
|
||||
commands::api_key_provider_cmd::export_api_key_providers,
|
||||
commands::api_key_provider_cmd::import_api_key_providers,
|
||||
// Legacy API Key migration commands
|
||||
commands::api_key_provider_cmd::get_legacy_api_key_credentials,
|
||||
commands::api_key_provider_cmd::migrate_legacy_api_key_credentials,
|
||||
commands::api_key_provider_cmd::delete_legacy_api_key_credential,
|
||||
// Route commands
|
||||
commands::route_cmd::get_available_routes,
|
||||
commands::route_cmd::get_route_curl_examples,
|
||||
// Router config commands
|
||||
commands::router_cmd::get_model_aliases,
|
||||
commands::router_cmd::add_model_alias,
|
||||
commands::router_cmd::remove_model_alias,
|
||||
commands::router_cmd::get_routing_rules,
|
||||
commands::router_cmd::add_routing_rule,
|
||||
commands::router_cmd::remove_routing_rule,
|
||||
commands::router_cmd::update_routing_rule,
|
||||
commands::router_cmd::get_exclusions,
|
||||
commands::router_cmd::add_exclusion,
|
||||
commands::router_cmd::remove_exclusion,
|
||||
commands::router_cmd::set_router_default_provider,
|
||||
commands::router_cmd::get_recommended_presets,
|
||||
commands::router_cmd::apply_recommended_preset,
|
||||
commands::router_cmd::clear_all_routing_config,
|
||||
// Resilience config commands
|
||||
commands::resilience_cmd::get_retry_config,
|
||||
commands::resilience_cmd::update_retry_config,
|
||||
commands::resilience_cmd::get_failover_config,
|
||||
commands::resilience_cmd::update_failover_config,
|
||||
commands::resilience_cmd::get_switch_log,
|
||||
commands::resilience_cmd::clear_switch_log,
|
||||
// Telemetry commands
|
||||
commands::telemetry_cmd::get_request_logs,
|
||||
commands::telemetry_cmd::get_request_log_detail,
|
||||
commands::telemetry_cmd::clear_request_logs,
|
||||
commands::telemetry_cmd::get_stats_summary,
|
||||
commands::telemetry_cmd::get_stats_by_provider,
|
||||
commands::telemetry_cmd::get_stats_by_model,
|
||||
commands::telemetry_cmd::get_token_summary,
|
||||
commands::telemetry_cmd::get_token_stats_by_provider,
|
||||
commands::telemetry_cmd::get_token_stats_by_model,
|
||||
commands::telemetry_cmd::get_token_stats_by_day,
|
||||
// Injection commands
|
||||
commands::injection_cmd::get_injection_config,
|
||||
commands::injection_cmd::set_injection_enabled,
|
||||
commands::injection_cmd::get_injection_rules,
|
||||
commands::injection_cmd::add_injection_rule,
|
||||
commands::injection_cmd::remove_injection_rule,
|
||||
commands::injection_cmd::update_injection_rule,
|
||||
// Usage commands
|
||||
commands::usage_cmd::get_kiro_usage,
|
||||
// Tray commands
|
||||
commands::tray_cmd::sync_tray_state,
|
||||
commands::tray_cmd::update_tray_server_status,
|
||||
commands::tray_cmd::update_tray_credential_status,
|
||||
commands::tray_cmd::get_tray_state,
|
||||
commands::tray_cmd::refresh_tray_menu,
|
||||
commands::tray_cmd::refresh_tray_with_stats,
|
||||
// Plugin commands
|
||||
commands::plugin_cmd::get_plugin_status,
|
||||
commands::plugin_cmd::get_plugins,
|
||||
commands::plugin_cmd::get_plugin_info,
|
||||
commands::plugin_cmd::enable_plugin,
|
||||
commands::plugin_cmd::disable_plugin,
|
||||
commands::plugin_cmd::update_plugin_config,
|
||||
commands::plugin_cmd::get_plugin_config,
|
||||
commands::plugin_cmd::reload_plugins,
|
||||
commands::plugin_cmd::unload_plugin,
|
||||
commands::plugin_cmd::get_plugins_dir,
|
||||
// Plugin Install commands
|
||||
commands::plugin_install_cmd::install_plugin_from_file,
|
||||
commands::plugin_install_cmd::install_plugin_from_url,
|
||||
commands::plugin_install_cmd::uninstall_plugin,
|
||||
commands::plugin_install_cmd::list_installed_plugins,
|
||||
commands::plugin_install_cmd::get_installed_plugin,
|
||||
commands::plugin_install_cmd::is_plugin_installed,
|
||||
// Plugin UI commands
|
||||
commands::plugin_cmd::get_plugins_with_ui,
|
||||
// Flow Monitor commands
|
||||
commands::flow_monitor_cmd::query_flows,
|
||||
commands::flow_monitor_cmd::get_flow_detail,
|
||||
commands::flow_monitor_cmd::search_flows,
|
||||
commands::flow_monitor_cmd::get_flow_stats,
|
||||
commands::flow_monitor_cmd::export_flows,
|
||||
commands::flow_monitor_cmd::update_flow_annotations,
|
||||
commands::flow_monitor_cmd::toggle_flow_starred,
|
||||
commands::flow_monitor_cmd::add_flow_comment,
|
||||
commands::flow_monitor_cmd::add_flow_tag,
|
||||
commands::flow_monitor_cmd::remove_flow_tag,
|
||||
commands::flow_monitor_cmd::set_flow_marker,
|
||||
commands::flow_monitor_cmd::cleanup_flows,
|
||||
commands::flow_monitor_cmd::get_recent_flows,
|
||||
commands::flow_monitor_cmd::get_flow_monitor_status,
|
||||
commands::flow_monitor_cmd::get_flow_monitor_debug_info,
|
||||
commands::flow_monitor_cmd::create_test_flows,
|
||||
commands::flow_monitor_cmd::enable_flow_monitor,
|
||||
commands::flow_monitor_cmd::disable_flow_monitor,
|
||||
commands::flow_monitor_cmd::subscribe_flow_events,
|
||||
commands::flow_monitor_cmd::get_all_flow_tags,
|
||||
// Flow Monitor filter expression commands
|
||||
commands::flow_monitor_cmd::parse_filter,
|
||||
commands::flow_monitor_cmd::validate_filter,
|
||||
commands::flow_monitor_cmd::get_filter_help_items,
|
||||
commands::flow_monitor_cmd::get_filter_help_text,
|
||||
commands::flow_monitor_cmd::query_flows_with_expression,
|
||||
// Flow Interceptor commands
|
||||
commands::flow_monitor_cmd::intercept_config_get,
|
||||
commands::flow_monitor_cmd::intercept_config_set,
|
||||
commands::flow_monitor_cmd::intercept_continue,
|
||||
commands::flow_monitor_cmd::intercept_cancel,
|
||||
commands::flow_monitor_cmd::intercept_get_flow,
|
||||
commands::flow_monitor_cmd::intercept_list_flows,
|
||||
commands::flow_monitor_cmd::intercept_count,
|
||||
commands::flow_monitor_cmd::intercept_is_enabled,
|
||||
commands::flow_monitor_cmd::intercept_enable,
|
||||
commands::flow_monitor_cmd::intercept_disable,
|
||||
commands::flow_monitor_cmd::intercept_set_editing,
|
||||
commands::flow_monitor_cmd::subscribe_intercept_events,
|
||||
// Flow Monitor realtime enhancement commands
|
||||
commands::flow_monitor_cmd::get_threshold_config,
|
||||
commands::flow_monitor_cmd::update_threshold_config,
|
||||
commands::flow_monitor_cmd::get_request_rate,
|
||||
commands::flow_monitor_cmd::set_rate_window,
|
||||
// Flow Replayer commands
|
||||
commands::flow_monitor_cmd::replay_flow,
|
||||
commands::flow_monitor_cmd::replay_flows_batch,
|
||||
// Flow Diff commands
|
||||
commands::flow_monitor_cmd::diff_flows,
|
||||
// Session Management commands
|
||||
commands::flow_monitor_cmd::create_session,
|
||||
commands::flow_monitor_cmd::get_session,
|
||||
commands::flow_monitor_cmd::list_sessions,
|
||||
commands::flow_monitor_cmd::add_flow_to_session,
|
||||
commands::flow_monitor_cmd::remove_flow_from_session,
|
||||
commands::flow_monitor_cmd::update_session,
|
||||
commands::flow_monitor_cmd::archive_session,
|
||||
commands::flow_monitor_cmd::unarchive_session,
|
||||
commands::flow_monitor_cmd::delete_session,
|
||||
commands::flow_monitor_cmd::export_session,
|
||||
commands::flow_monitor_cmd::get_session_flow_count,
|
||||
commands::flow_monitor_cmd::is_flow_in_session,
|
||||
commands::flow_monitor_cmd::get_sessions_for_flow,
|
||||
commands::flow_monitor_cmd::get_auto_session_config,
|
||||
commands::flow_monitor_cmd::set_auto_session_config,
|
||||
commands::flow_monitor_cmd::register_active_session,
|
||||
// Quick Filter commands
|
||||
commands::flow_monitor_cmd::save_quick_filter,
|
||||
commands::flow_monitor_cmd::get_quick_filter,
|
||||
commands::flow_monitor_cmd::update_quick_filter,
|
||||
commands::flow_monitor_cmd::delete_quick_filter,
|
||||
commands::flow_monitor_cmd::list_quick_filters,
|
||||
commands::flow_monitor_cmd::list_quick_filters_by_group,
|
||||
commands::flow_monitor_cmd::list_quick_filter_groups,
|
||||
commands::flow_monitor_cmd::export_quick_filters,
|
||||
commands::flow_monitor_cmd::import_quick_filters,
|
||||
commands::flow_monitor_cmd::find_quick_filter_by_name,
|
||||
// Code Export commands
|
||||
commands::flow_monitor_cmd::export_flow_as_code,
|
||||
commands::flow_monitor_cmd::export_flows_as_code,
|
||||
commands::flow_monitor_cmd::get_code_export_formats,
|
||||
// Bookmark Management commands
|
||||
commands::flow_monitor_cmd::add_bookmark,
|
||||
commands::flow_monitor_cmd::get_bookmark,
|
||||
commands::flow_monitor_cmd::get_bookmark_by_flow_id,
|
||||
commands::flow_monitor_cmd::remove_bookmark,
|
||||
commands::flow_monitor_cmd::remove_bookmark_by_flow_id,
|
||||
commands::flow_monitor_cmd::update_bookmark,
|
||||
commands::flow_monitor_cmd::list_bookmarks,
|
||||
commands::flow_monitor_cmd::list_bookmark_groups,
|
||||
commands::flow_monitor_cmd::is_flow_bookmarked,
|
||||
commands::flow_monitor_cmd::get_bookmark_count,
|
||||
commands::flow_monitor_cmd::export_bookmarks,
|
||||
commands::flow_monitor_cmd::import_bookmarks,
|
||||
commands::flow_monitor_cmd::toggle_bookmark,
|
||||
// Enhanced Stats commands
|
||||
commands::flow_monitor_cmd::get_enhanced_stats,
|
||||
commands::flow_monitor_cmd::get_request_trend,
|
||||
commands::flow_monitor_cmd::get_token_distribution,
|
||||
commands::flow_monitor_cmd::get_latency_histogram,
|
||||
commands::flow_monitor_cmd::export_stats_report,
|
||||
// Batch Operations commands
|
||||
commands::flow_monitor_cmd::batch_star_flows,
|
||||
commands::flow_monitor_cmd::batch_unstar_flows,
|
||||
commands::flow_monitor_cmd::batch_add_tags,
|
||||
commands::flow_monitor_cmd::batch_remove_tags,
|
||||
commands::flow_monitor_cmd::batch_export_flows,
|
||||
commands::flow_monitor_cmd::batch_delete_flows,
|
||||
commands::flow_monitor_cmd::batch_add_to_session,
|
||||
// Window control commands
|
||||
commands::window_cmd::get_window_size,
|
||||
commands::window_cmd::set_window_size,
|
||||
commands::window_cmd::resize_for_flow_monitor,
|
||||
commands::window_cmd::restore_window_size,
|
||||
commands::window_cmd::toggle_window_size,
|
||||
commands::window_cmd::center_window,
|
||||
commands::window_cmd::get_window_size_options,
|
||||
commands::window_cmd::set_window_size_by_option,
|
||||
commands::window_cmd::toggle_fullscreen,
|
||||
commands::window_cmd::is_fullscreen,
|
||||
// Browser Interceptor commands
|
||||
commands::browser_interceptor_cmd::get_browser_interceptor_state,
|
||||
commands::browser_interceptor_cmd::start_browser_interceptor,
|
||||
commands::browser_interceptor_cmd::stop_browser_interceptor,
|
||||
commands::browser_interceptor_cmd::restore_normal_browser_behavior,
|
||||
commands::browser_interceptor_cmd::temporary_disable_interceptor,
|
||||
commands::browser_interceptor_cmd::get_intercepted_urls,
|
||||
commands::browser_interceptor_cmd::get_interceptor_history,
|
||||
commands::browser_interceptor_cmd::copy_intercepted_url_to_clipboard,
|
||||
commands::browser_interceptor_cmd::open_url_in_fingerprint_browser,
|
||||
commands::browser_interceptor_cmd::dismiss_intercepted_url,
|
||||
commands::browser_interceptor_cmd::update_browser_interceptor_config,
|
||||
commands::browser_interceptor_cmd::get_default_browser_interceptor_config,
|
||||
commands::browser_interceptor_cmd::validate_browser_interceptor_config,
|
||||
commands::browser_interceptor_cmd::is_browser_interceptor_running,
|
||||
commands::browser_interceptor_cmd::get_browser_interceptor_statistics,
|
||||
// Browser Interceptor notification commands
|
||||
commands::browser_interceptor_cmd::show_notification,
|
||||
commands::browser_interceptor_cmd::show_url_intercept_notification,
|
||||
commands::browser_interceptor_cmd::show_status_notification,
|
||||
// Auto fix commands
|
||||
commands::auto_fix_cmd::auto_fix_configuration,
|
||||
// Machine ID commands
|
||||
commands::machine_id_cmd::get_current_machine_id,
|
||||
commands::machine_id_cmd::set_machine_id,
|
||||
commands::machine_id_cmd::generate_random_machine_id,
|
||||
commands::machine_id_cmd::validate_machine_id,
|
||||
commands::machine_id_cmd::check_admin_privileges,
|
||||
commands::machine_id_cmd::get_os_type,
|
||||
commands::machine_id_cmd::backup_machine_id_to_file,
|
||||
commands::machine_id_cmd::restore_machine_id_from_file,
|
||||
commands::machine_id_cmd::format_machine_id,
|
||||
commands::machine_id_cmd::detect_machine_id_format,
|
||||
commands::machine_id_cmd::convert_machine_id_format,
|
||||
commands::machine_id_cmd::get_machine_id_history,
|
||||
commands::machine_id_cmd::clear_machine_id_override,
|
||||
commands::machine_id_cmd::copy_machine_id_to_clipboard,
|
||||
commands::machine_id_cmd::paste_machine_id_from_clipboard,
|
||||
commands::machine_id_cmd::get_system_info,
|
||||
// Kiro Local commands
|
||||
commands::kiro_local::switch_kiro_to_local,
|
||||
commands::kiro_local::get_kiro_fingerprint_info,
|
||||
commands::kiro_local::get_local_kiro_credential_uuid,
|
||||
// Agent commands
|
||||
commands::agent_cmd::agent_start_process,
|
||||
commands::agent_cmd::agent_stop_process,
|
||||
commands::agent_cmd::agent_get_process_status,
|
||||
commands::agent_cmd::agent_create_session,
|
||||
commands::agent_cmd::agent_send_message,
|
||||
commands::agent_cmd::agent_list_sessions,
|
||||
commands::agent_cmd::agent_get_session,
|
||||
commands::agent_cmd::agent_delete_session,
|
||||
// Native Agent commands
|
||||
commands::native_agent_cmd::native_agent_init,
|
||||
commands::native_agent_cmd::native_agent_status,
|
||||
commands::native_agent_cmd::native_agent_reset,
|
||||
commands::native_agent_cmd::native_agent_chat,
|
||||
commands::native_agent_cmd::native_agent_chat_stream,
|
||||
commands::native_agent_cmd::native_agent_create_session,
|
||||
commands::native_agent_cmd::native_agent_get_session,
|
||||
commands::native_agent_cmd::native_agent_delete_session,
|
||||
commands::native_agent_cmd::native_agent_list_sessions,
|
||||
// Network commands
|
||||
commands::network_cmd::get_network_info,
|
||||
// OAuth Plugin commands
|
||||
commands::oauth_plugin_cmd::init_oauth_plugin_system,
|
||||
commands::oauth_plugin_cmd::list_oauth_plugins,
|
||||
commands::oauth_plugin_cmd::get_oauth_plugin,
|
||||
commands::oauth_plugin_cmd::enable_oauth_plugin,
|
||||
commands::oauth_plugin_cmd::disable_oauth_plugin,
|
||||
commands::oauth_plugin_cmd::install_oauth_plugin,
|
||||
commands::oauth_plugin_cmd::uninstall_oauth_plugin,
|
||||
commands::oauth_plugin_cmd::check_oauth_plugin_updates,
|
||||
commands::oauth_plugin_cmd::update_oauth_plugin,
|
||||
commands::oauth_plugin_cmd::reload_oauth_plugins,
|
||||
commands::oauth_plugin_cmd::get_oauth_plugin_config,
|
||||
commands::oauth_plugin_cmd::update_oauth_plugin_config,
|
||||
commands::oauth_plugin_cmd::scan_oauth_plugin_directory,
|
||||
// OAuth Plugin credential commands
|
||||
commands::oauth_plugin_cmd::plugin_credential_list,
|
||||
commands::oauth_plugin_cmd::plugin_credential_get,
|
||||
commands::oauth_plugin_cmd::plugin_credential_create,
|
||||
commands::oauth_plugin_cmd::plugin_credential_update,
|
||||
commands::oauth_plugin_cmd::plugin_credential_delete,
|
||||
commands::oauth_plugin_cmd::plugin_credential_validate,
|
||||
commands::oauth_plugin_cmd::plugin_credential_refresh,
|
||||
// OAuth Plugin SDK commands
|
||||
commands::oauth_plugin_cmd::plugin_database_query,
|
||||
commands::oauth_plugin_cmd::plugin_database_execute,
|
||||
commands::oauth_plugin_cmd::plugin_http_request,
|
||||
commands::oauth_plugin_cmd::plugin_crypto_encrypt,
|
||||
commands::oauth_plugin_cmd::plugin_crypto_decrypt,
|
||||
commands::oauth_plugin_cmd::plugin_notification,
|
||||
commands::oauth_plugin_cmd::plugin_storage_get,
|
||||
commands::oauth_plugin_cmd::plugin_storage_set,
|
||||
commands::oauth_plugin_cmd::plugin_storage_delete,
|
||||
commands::oauth_plugin_cmd::plugin_storage_keys,
|
||||
commands::oauth_plugin_cmd::plugin_config_get,
|
||||
commands::oauth_plugin_cmd::plugin_config_set,
|
||||
// OAuth Plugin UI commands
|
||||
commands::oauth_plugin_cmd::read_plugin_ui_file,
|
||||
// Orchestrator commands
|
||||
commands::orchestrator_cmd::init_orchestrator,
|
||||
commands::orchestrator_cmd::get_orchestrator_config,
|
||||
commands::orchestrator_cmd::update_orchestrator_config,
|
||||
commands::orchestrator_cmd::get_pool_stats,
|
||||
commands::orchestrator_cmd::get_tier_models,
|
||||
commands::orchestrator_cmd::get_all_models,
|
||||
commands::orchestrator_cmd::update_orchestrator_credentials,
|
||||
commands::orchestrator_cmd::add_orchestrator_credential,
|
||||
commands::orchestrator_cmd::remove_orchestrator_credential,
|
||||
commands::orchestrator_cmd::mark_credential_unhealthy,
|
||||
commands::orchestrator_cmd::mark_credential_healthy,
|
||||
commands::orchestrator_cmd::update_credential_load,
|
||||
commands::orchestrator_cmd::select_model,
|
||||
commands::orchestrator_cmd::quick_select_model,
|
||||
commands::orchestrator_cmd::select_model_for_task,
|
||||
commands::orchestrator_cmd::list_strategies,
|
||||
commands::orchestrator_cmd::list_service_tiers,
|
||||
commands::orchestrator_cmd::list_task_hints,
|
||||
])
|
||||
.run(tauri::generate_context!())
|
||||
.expect("error while running tauri application");
|
||||
}
|
||||
@@ -0,0 +1,244 @@
|
||||
//! Tauri Setup Hook
|
||||
//!
|
||||
//! 包含应用启动时的初始化逻辑。
|
||||
|
||||
use std::sync::Arc;
|
||||
use tauri::{App, Manager};
|
||||
|
||||
use crate::agent::NativeAgentState;
|
||||
use crate::commands::browser_interceptor_cmd::BrowserInterceptorState;
|
||||
use crate::commands::oauth_plugin_cmd::OAuthPluginManagerState;
|
||||
use crate::database;
|
||||
use crate::flow_monitor::FlowInterceptor;
|
||||
use crate::services::provider_pool_service::ProviderPoolService;
|
||||
use crate::services::token_cache_service::TokenCacheService;
|
||||
use crate::telemetry;
|
||||
use crate::tray::{TrayIconStatus, TrayManager, TrayStateSnapshot};
|
||||
|
||||
use super::types::{AppState, LogState, TrayManagerState};
|
||||
|
||||
/// Tauri setup hook
|
||||
///
|
||||
/// 在应用启动时执行初始化逻辑
|
||||
pub fn setup_app(
|
||||
app: &mut App,
|
||||
state: AppState,
|
||||
logs: LogState,
|
||||
db: database::DbConnection,
|
||||
pool_service: Arc<ProviderPoolService>,
|
||||
token_cache: Arc<TokenCacheService>,
|
||||
shared_stats: Arc<parking_lot::RwLock<telemetry::StatsAggregator>>,
|
||||
shared_tokens: Arc<parking_lot::RwLock<telemetry::TokenTracker>>,
|
||||
shared_logger: Arc<telemetry::RequestLogger>,
|
||||
flow_monitor: Arc<crate::flow_monitor::FlowMonitor>,
|
||||
flow_interceptor: Arc<FlowInterceptor>,
|
||||
) -> Result<(), Box<dyn std::error::Error>> {
|
||||
// 设置 deep-link 事件监听(用于浏览器拦截)
|
||||
#[cfg(target_os = "macos")]
|
||||
{
|
||||
use tauri_plugin_deep_link::DeepLinkExt;
|
||||
let _listener_id = app.deep_link().on_open_url(|event| {
|
||||
for url in event.urls() {
|
||||
tracing::info!("[Deep Link] 收到 URL: {}", url);
|
||||
crate::browser_interceptor::platform::macos::handle_deep_link_url(url.to_string());
|
||||
}
|
||||
});
|
||||
tracing::info!("[启动] Deep Link 事件监听已设置");
|
||||
}
|
||||
|
||||
// 初始化托盘管理器
|
||||
match TrayManager::new(app.handle()) {
|
||||
Ok(tray_manager) => {
|
||||
tracing::info!("[启动] 托盘管理器初始化成功");
|
||||
let tray_state: TrayManagerState<tauri::Wry> =
|
||||
TrayManagerState(Arc::new(tokio::sync::RwLock::new(Some(tray_manager))));
|
||||
app.manage(tray_state);
|
||||
}
|
||||
Err(e) => {
|
||||
tracing::error!("[启动] 托盘管理器初始化失败: {}", e);
|
||||
let tray_state: TrayManagerState<tauri::Wry> =
|
||||
TrayManagerState(Arc::new(tokio::sync::RwLock::new(None)));
|
||||
app.manage(tray_state);
|
||||
}
|
||||
}
|
||||
|
||||
// 初始化 BrowserInterceptorState
|
||||
let browser_interceptor_state = BrowserInterceptorState::default();
|
||||
app.manage(browser_interceptor_state);
|
||||
|
||||
// 初始化 NativeAgentState
|
||||
let native_agent_state = NativeAgentState::new();
|
||||
app.manage(native_agent_state);
|
||||
|
||||
// 初始化 OAuth Plugin Manager State
|
||||
let oauth_plugin_manager_state = OAuthPluginManagerState::with_defaults();
|
||||
app.manage(oauth_plugin_manager_state);
|
||||
|
||||
// 初始化默认 skill repos
|
||||
{
|
||||
let conn = db.lock().expect("Failed to lock database");
|
||||
database::dao::skills::SkillDao::init_default_skill_repos(&conn)
|
||||
.expect("Failed to initialize default skill repos");
|
||||
}
|
||||
|
||||
// 自动启动服务器
|
||||
let app_handle = app.handle().clone();
|
||||
tauri::async_runtime::spawn(async move {
|
||||
start_server_async(
|
||||
state,
|
||||
logs,
|
||||
db,
|
||||
pool_service,
|
||||
token_cache,
|
||||
shared_stats,
|
||||
shared_tokens,
|
||||
shared_logger,
|
||||
flow_monitor,
|
||||
flow_interceptor,
|
||||
app_handle,
|
||||
)
|
||||
.await;
|
||||
});
|
||||
|
||||
Ok(())
|
||||
}
|
||||
|
||||
/// 异步启动服务器
|
||||
async fn start_server_async(
|
||||
state: AppState,
|
||||
logs: LogState,
|
||||
db: database::DbConnection,
|
||||
pool_service: Arc<ProviderPoolService>,
|
||||
token_cache: Arc<TokenCacheService>,
|
||||
shared_stats: Arc<parking_lot::RwLock<telemetry::StatsAggregator>>,
|
||||
shared_tokens: Arc<parking_lot::RwLock<telemetry::TokenTracker>>,
|
||||
shared_logger: Arc<telemetry::RequestLogger>,
|
||||
shared_flow_monitor: Arc<crate::flow_monitor::FlowMonitor>,
|
||||
flow_interceptor: Arc<FlowInterceptor>,
|
||||
app_handle: tauri::AppHandle,
|
||||
) {
|
||||
// 先加载凭证池中的凭证
|
||||
{
|
||||
logs.write().await.add("info", "[启动] 正在加载凭证池...");
|
||||
|
||||
match pool_service.get_overview(&db) {
|
||||
Ok(overview) => {
|
||||
let mut loaded_types = Vec::new();
|
||||
let mut total_credentials = 0;
|
||||
|
||||
for provider_overview in overview {
|
||||
let count = provider_overview.stats.total_count;
|
||||
if count > 0 {
|
||||
total_credentials += count;
|
||||
let provider_name = match provider_overview.provider_type.as_str() {
|
||||
"kiro" => "Kiro",
|
||||
"gemini" => "Gemini",
|
||||
"qwen" => "通义千问",
|
||||
"antigravity" => "Antigravity",
|
||||
"openai" => "OpenAI",
|
||||
"claude" => "Claude",
|
||||
"codex" => "Codex",
|
||||
"claude_oauth" => "Claude OAuth",
|
||||
"iflow" => "iFlow",
|
||||
_ => &provider_overview.provider_type,
|
||||
};
|
||||
loaded_types.push(format!("{} ({} 个)", provider_name, count));
|
||||
}
|
||||
}
|
||||
|
||||
if loaded_types.is_empty() {
|
||||
logs.write().await.add("warn", "[启动] 未找到任何可用凭证");
|
||||
} else {
|
||||
let message = format!(
|
||||
"[启动] 凭证已加载: {} (共 {} 个)",
|
||||
loaded_types.join(", "),
|
||||
total_credentials
|
||||
);
|
||||
logs.write().await.add("info", &message);
|
||||
}
|
||||
}
|
||||
Err(e) => {
|
||||
logs.write()
|
||||
.await
|
||||
.add("warn", &format!("[启动] 获取凭证池信息失败: {}", e));
|
||||
}
|
||||
}
|
||||
|
||||
// 兼容性:仍然尝试加载旧的 Kiro 凭证(如果存在)
|
||||
let mut s = state.write().await;
|
||||
if let Err(e) = s.kiro_provider.load_credentials().await {
|
||||
logs.write()
|
||||
.await
|
||||
.add("debug", &format!("[启动] 旧版 Kiro 凭证加载失败: {e}"));
|
||||
}
|
||||
}
|
||||
|
||||
// 启动服务器
|
||||
let server_started;
|
||||
let server_address;
|
||||
{
|
||||
let mut s = state.write().await;
|
||||
logs.write()
|
||||
.await
|
||||
.add("info", "[启动] 正在自动启动服务器...");
|
||||
match s
|
||||
.start_with_telemetry_and_flow_monitor(
|
||||
logs.clone(),
|
||||
pool_service,
|
||||
token_cache,
|
||||
Some(db),
|
||||
Some(shared_stats),
|
||||
Some(shared_tokens),
|
||||
Some(shared_logger),
|
||||
Some(shared_flow_monitor),
|
||||
Some(flow_interceptor),
|
||||
)
|
||||
.await
|
||||
{
|
||||
Ok(_) => {
|
||||
let host = s.config.server.host.clone();
|
||||
let port = s.config.server.port;
|
||||
logs.write()
|
||||
.await
|
||||
.add("info", &format!("[启动] 服务器已启动: {host}:{port}"));
|
||||
server_started = true;
|
||||
server_address = format!("{}:{}", host, port);
|
||||
}
|
||||
Err(e) => {
|
||||
logs.write()
|
||||
.await
|
||||
.add("error", &format!("[启动] 服务器启动失败: {e}"));
|
||||
server_started = false;
|
||||
server_address = String::new();
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// 更新托盘状态
|
||||
if let Some(tray_state) = app_handle.try_state::<TrayManagerState<tauri::Wry>>() {
|
||||
let tray_guard = tray_state.0.read().await;
|
||||
if let Some(tray_manager) = tray_guard.as_ref() {
|
||||
let icon_status = if server_started {
|
||||
TrayIconStatus::Running
|
||||
} else {
|
||||
TrayIconStatus::Stopped
|
||||
};
|
||||
|
||||
let snapshot = TrayStateSnapshot {
|
||||
icon_status,
|
||||
server_running: server_started,
|
||||
server_address,
|
||||
available_credentials: 0,
|
||||
total_credentials: 0,
|
||||
today_requests: 0,
|
||||
auto_start_enabled: false,
|
||||
};
|
||||
|
||||
if let Err(e) = tray_manager.update_state(snapshot).await {
|
||||
tracing::error!("[启动] 更新托盘状态失败: {}", e);
|
||||
} else {
|
||||
tracing::info!("[启动] 托盘状态已更新");
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,319 @@
|
||||
//! 状态初始化模块
|
||||
//!
|
||||
//! 包含应用状态的初始化逻辑。
|
||||
|
||||
use std::sync::Arc;
|
||||
use tokio::sync::RwLock;
|
||||
|
||||
use crate::commands::api_key_provider_cmd::ApiKeyProviderServiceState;
|
||||
use crate::commands::flow_monitor_cmd::{
|
||||
BatchOperationsState, BookmarkManagerState, EnhancedStatsServiceState, FlowInterceptorState,
|
||||
FlowMonitorState, FlowQueryServiceState, FlowReplayerState, QuickFilterManagerState,
|
||||
SessionManagerState,
|
||||
};
|
||||
use crate::commands::machine_id_cmd::MachineIdState;
|
||||
use crate::commands::orchestrator_cmd::OrchestratorState;
|
||||
use crate::commands::plugin_cmd::PluginManagerState;
|
||||
use crate::commands::plugin_install_cmd::PluginInstallerState;
|
||||
use crate::commands::provider_pool_cmd::{CredentialSyncServiceState, ProviderPoolServiceState};
|
||||
use crate::commands::resilience_cmd::ResilienceConfigState;
|
||||
use crate::commands::router_cmd::RouterConfigState;
|
||||
use crate::commands::skill_cmd::SkillServiceState;
|
||||
use crate::config::Config;
|
||||
use crate::database;
|
||||
use crate::flow_monitor::{
|
||||
BatchOperations, BookmarkManager, EnhancedStatsService, FlowFileStore, FlowInterceptor,
|
||||
FlowMonitor, FlowMonitorConfig, FlowQueryService, FlowReplayer, InterceptConfig,
|
||||
QuickFilterManager, RotationConfig, SessionManager,
|
||||
};
|
||||
use crate::plugin;
|
||||
use crate::services::api_key_provider_service::ApiKeyProviderService;
|
||||
use crate::services::provider_pool_service::ProviderPoolService;
|
||||
use crate::services::skill_service::SkillService;
|
||||
use crate::services::token_cache_service::TokenCacheService;
|
||||
use crate::telemetry;
|
||||
|
||||
use super::types::{AppState, LogState, TokenCacheServiceState};
|
||||
use crate::logger;
|
||||
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)));
|
||||
(state, logs)
|
||||
}
|
||||
|
||||
/// 初始化服务状态
|
||||
pub struct ServiceStates {
|
||||
pub skill_service: SkillServiceState,
|
||||
pub provider_pool_service: ProviderPoolServiceState,
|
||||
pub api_key_provider_service: ApiKeyProviderServiceState,
|
||||
pub credential_sync_service: CredentialSyncServiceState,
|
||||
pub token_cache_service: TokenCacheServiceState,
|
||||
pub machine_id_service: MachineIdState,
|
||||
pub router_config: RouterConfigState,
|
||||
pub resilience_config: ResilienceConfigState,
|
||||
pub plugin_manager: PluginManagerState,
|
||||
pub plugin_installer: PluginInstallerState,
|
||||
pub orchestrator: OrchestratorState,
|
||||
}
|
||||
|
||||
/// 初始化所有服务状态
|
||||
pub fn init_service_states() -> ServiceStates {
|
||||
// Initialize SkillService
|
||||
let skill_service = SkillService::new().expect("Failed to initialize SkillService");
|
||||
let skill_service_state = SkillServiceState(Arc::new(skill_service));
|
||||
|
||||
// Initialize ProviderPoolService
|
||||
let provider_pool_service = ProviderPoolService::new();
|
||||
let provider_pool_service_state = ProviderPoolServiceState(Arc::new(provider_pool_service));
|
||||
|
||||
// Initialize ApiKeyProviderService
|
||||
let api_key_provider_service = ApiKeyProviderService::new();
|
||||
let api_key_provider_service_state =
|
||||
ApiKeyProviderServiceState(Arc::new(api_key_provider_service));
|
||||
|
||||
// Initialize CredentialSyncService (optional)
|
||||
let credential_sync_service_state = CredentialSyncServiceState(None);
|
||||
|
||||
// Initialize TokenCacheService
|
||||
let token_cache_service = TokenCacheService::new();
|
||||
let token_cache_service_state = TokenCacheServiceState(Arc::new(token_cache_service));
|
||||
|
||||
// Initialize MachineIdService
|
||||
let machine_id_service = crate::services::machine_id_service::MachineIdService::new()
|
||||
.expect("Failed to initialize MachineIdService");
|
||||
let machine_id_service_state: MachineIdState = Arc::new(RwLock::new(machine_id_service));
|
||||
|
||||
// Initialize RouterConfigState
|
||||
let router_config_state = RouterConfigState::default();
|
||||
|
||||
// Initialize ResilienceConfigState
|
||||
let resilience_config_state = ResilienceConfigState::default();
|
||||
|
||||
// Initialize PluginManager
|
||||
let plugin_manager = plugin::PluginManager::with_defaults();
|
||||
let plugin_manager_state = PluginManagerState(Arc::new(RwLock::new(plugin_manager)));
|
||||
|
||||
// Initialize PluginInstaller
|
||||
let plugin_installer_state = init_plugin_installer();
|
||||
|
||||
// Initialize Orchestrator State
|
||||
let orchestrator_state = OrchestratorState::new();
|
||||
|
||||
ServiceStates {
|
||||
skill_service: skill_service_state,
|
||||
provider_pool_service: provider_pool_service_state,
|
||||
api_key_provider_service: api_key_provider_service_state,
|
||||
credential_sync_service: credential_sync_service_state,
|
||||
token_cache_service: token_cache_service_state,
|
||||
machine_id_service: machine_id_service_state,
|
||||
router_config: router_config_state,
|
||||
resilience_config: resilience_config_state,
|
||||
plugin_manager: plugin_manager_state,
|
||||
plugin_installer: plugin_installer_state,
|
||||
orchestrator: orchestrator_state,
|
||||
}
|
||||
}
|
||||
|
||||
/// 初始化插件安装器
|
||||
fn init_plugin_installer() -> PluginInstallerState {
|
||||
let db_path = database::get_db_path().expect("Failed to get database path for PluginInstaller");
|
||||
let plugins_dir = dirs::data_dir()
|
||||
.unwrap_or_else(|| std::path::PathBuf::from("."))
|
||||
.join("proxycast")
|
||||
.join("plugins");
|
||||
let temp_dir = std::env::temp_dir().join("proxycast_plugin_install");
|
||||
|
||||
// 创建目录(如果不存在)
|
||||
if let Err(e) = std::fs::create_dir_all(&plugins_dir) {
|
||||
tracing::warn!("无法创建插件目录: {}", e);
|
||||
}
|
||||
if let Err(e) = std::fs::create_dir_all(&temp_dir) {
|
||||
tracing::warn!("无法创建插件临时目录: {}", e);
|
||||
}
|
||||
|
||||
match plugin::installer::PluginInstaller::from_paths(plugins_dir, temp_dir, &db_path) {
|
||||
Ok(installer) => {
|
||||
tracing::info!("[启动] 插件安装器初始化成功");
|
||||
PluginInstallerState(Arc::new(RwLock::new(installer)))
|
||||
}
|
||||
Err(e) => {
|
||||
tracing::error!("[启动] 插件安装器初始化失败: {}", e);
|
||||
// 创建一个默认的安装器(使用临时目录)
|
||||
let fallback_plugins_dir = std::env::temp_dir().join("proxycast_plugins_fallback");
|
||||
let fallback_temp_dir = std::env::temp_dir().join("proxycast_plugin_install_fallback");
|
||||
let _ = std::fs::create_dir_all(&fallback_plugins_dir);
|
||||
let _ = std::fs::create_dir_all(&fallback_temp_dir);
|
||||
let installer = plugin::installer::PluginInstaller::from_paths(
|
||||
fallback_plugins_dir,
|
||||
fallback_temp_dir,
|
||||
&db_path,
|
||||
)
|
||||
.expect("Failed to create fallback PluginInstaller");
|
||||
PluginInstallerState(Arc::new(RwLock::new(installer)))
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
/// 遥测状态
|
||||
pub struct TelemetryStates {
|
||||
pub stats: Arc<parking_lot::RwLock<telemetry::StatsAggregator>>,
|
||||
pub tokens: Arc<parking_lot::RwLock<telemetry::TokenTracker>>,
|
||||
pub logger: Arc<telemetry::RequestLogger>,
|
||||
pub telemetry_state: crate::commands::telemetry_cmd::TelemetryState,
|
||||
}
|
||||
|
||||
/// 初始化遥测状态
|
||||
pub fn init_telemetry_states(config: &Config) -> TelemetryStates {
|
||||
let shared_stats = Arc::new(parking_lot::RwLock::new(
|
||||
telemetry::StatsAggregator::with_defaults(),
|
||||
));
|
||||
let shared_tokens = Arc::new(parking_lot::RwLock::new(
|
||||
telemetry::TokenTracker::with_defaults(),
|
||||
));
|
||||
let log_rotation = telemetry::LogRotationConfig {
|
||||
max_memory_logs: 10000,
|
||||
retention_days: config.logging.retention_days,
|
||||
max_file_size: 10 * 1024 * 1024,
|
||||
enable_file_logging: config.logging.enabled,
|
||||
};
|
||||
let shared_logger = Arc::new(
|
||||
telemetry::RequestLogger::new(log_rotation).expect("Failed to create RequestLogger"),
|
||||
);
|
||||
|
||||
let telemetry_state = crate::commands::telemetry_cmd::TelemetryState::with_shared(
|
||||
shared_stats.clone(),
|
||||
shared_tokens.clone(),
|
||||
Some(shared_logger.clone()),
|
||||
)
|
||||
.expect("Failed to create TelemetryState");
|
||||
|
||||
TelemetryStates {
|
||||
stats: shared_stats,
|
||||
tokens: shared_tokens,
|
||||
logger: shared_logger,
|
||||
telemetry_state,
|
||||
}
|
||||
}
|
||||
|
||||
/// Flow Monitor 状态
|
||||
pub struct FlowMonitorStates {
|
||||
pub flow_monitor: Arc<FlowMonitor>,
|
||||
pub flow_monitor_state: FlowMonitorState,
|
||||
pub flow_interceptor: Arc<FlowInterceptor>,
|
||||
pub flow_interceptor_state: FlowInterceptorState,
|
||||
pub flow_replayer_state: FlowReplayerState,
|
||||
pub flow_query_service_state: FlowQueryServiceState,
|
||||
pub session_manager_state: SessionManagerState,
|
||||
pub quick_filter_manager_state: QuickFilterManagerState,
|
||||
pub bookmark_manager_state: BookmarkManagerState,
|
||||
pub enhanced_stats_service_state: EnhancedStatsServiceState,
|
||||
pub batch_operations_state: BatchOperationsState,
|
||||
}
|
||||
|
||||
/// 初始化 Flow Monitor 状态
|
||||
pub fn init_flow_monitor_states(
|
||||
provider_pool_service: Arc<ProviderPoolService>,
|
||||
db: database::DbConnection,
|
||||
) -> FlowMonitorStates {
|
||||
let flow_monitor_config = FlowMonitorConfig::default();
|
||||
let flow_file_store = init_flow_file_store();
|
||||
|
||||
let flow_monitor = Arc::new(FlowMonitor::new(
|
||||
flow_monitor_config,
|
||||
flow_file_store.clone(),
|
||||
));
|
||||
let flow_monitor_state = FlowMonitorState(flow_monitor.clone());
|
||||
|
||||
// 初始化 Flow 拦截器
|
||||
let flow_interceptor = Arc::new(FlowInterceptor::new(InterceptConfig::default()));
|
||||
let flow_interceptor_state = FlowInterceptorState(flow_interceptor.clone());
|
||||
|
||||
// 初始化 Flow 重放器
|
||||
let flow_replayer = Arc::new(FlowReplayer::new(
|
||||
flow_monitor.clone(),
|
||||
provider_pool_service,
|
||||
db,
|
||||
));
|
||||
let flow_replayer_state = FlowReplayerState(flow_replayer);
|
||||
|
||||
// 初始化会话管理器
|
||||
let db_path = database::get_db_path().expect("Failed to get database path");
|
||||
let session_manager =
|
||||
Arc::new(SessionManager::new(db_path.clone()).expect("Failed to create SessionManager"));
|
||||
let session_manager_state = SessionManagerState(session_manager.clone());
|
||||
|
||||
// 初始化快速过滤器管理器
|
||||
let quick_filter_manager = Arc::new(
|
||||
QuickFilterManager::new(db_path.clone()).expect("Failed to create QuickFilterManager"),
|
||||
);
|
||||
let quick_filter_manager_state = QuickFilterManagerState(quick_filter_manager);
|
||||
|
||||
// 初始化书签管理器
|
||||
let bookmark_manager =
|
||||
Arc::new(BookmarkManager::new(db_path).expect("Failed to create BookmarkManager"));
|
||||
let bookmark_manager_state = BookmarkManagerState(bookmark_manager);
|
||||
|
||||
// 初始化增强统计服务
|
||||
let enhanced_stats_service = Arc::new(EnhancedStatsService::new(flow_monitor.memory_store()));
|
||||
let enhanced_stats_service_state = EnhancedStatsServiceState(enhanced_stats_service);
|
||||
|
||||
// 初始化批量操作服务
|
||||
let batch_operations = Arc::new(BatchOperations::new(
|
||||
flow_monitor.clone(),
|
||||
Some(session_manager_state.0.clone()),
|
||||
));
|
||||
let batch_operations_state = BatchOperationsState(batch_operations);
|
||||
|
||||
// FlowQueryService
|
||||
let flow_query_service_state = if let Some(file_store) = flow_file_store {
|
||||
let query_service = FlowQueryService::new(flow_monitor.memory_store(), file_store);
|
||||
FlowQueryServiceState(Arc::new(query_service))
|
||||
} else {
|
||||
let temp_dir = std::env::temp_dir().join("proxycast_flows");
|
||||
let _ = std::fs::create_dir_all(&temp_dir);
|
||||
let rotation_config = RotationConfig::default();
|
||||
let temp_store = FlowFileStore::new(temp_dir, rotation_config)
|
||||
.expect("Failed to create temp FlowFileStore");
|
||||
let query_service =
|
||||
FlowQueryService::new(flow_monitor.memory_store(), Arc::new(temp_store));
|
||||
FlowQueryServiceState(Arc::new(query_service))
|
||||
};
|
||||
|
||||
FlowMonitorStates {
|
||||
flow_monitor,
|
||||
flow_monitor_state,
|
||||
flow_interceptor,
|
||||
flow_interceptor_state,
|
||||
flow_replayer_state,
|
||||
flow_query_service_state,
|
||||
session_manager_state,
|
||||
quick_filter_manager_state,
|
||||
bookmark_manager_state,
|
||||
enhanced_stats_service_state,
|
||||
batch_operations_state,
|
||||
}
|
||||
}
|
||||
|
||||
/// 初始化 Flow 文件存储
|
||||
fn init_flow_file_store() -> Option<Arc<FlowFileStore>> {
|
||||
let data_dir = dirs::data_dir()
|
||||
.unwrap_or_else(|| std::path::PathBuf::from("."))
|
||||
.join("proxycast")
|
||||
.join("flows");
|
||||
|
||||
if let Err(e) = std::fs::create_dir_all(&data_dir) {
|
||||
tracing::warn!("无法创建 Flow 存储目录: {}", e);
|
||||
}
|
||||
|
||||
let rotation_config = RotationConfig::default();
|
||||
match FlowFileStore::new(data_dir, rotation_config) {
|
||||
Ok(store) => Some(Arc::new(store)),
|
||||
Err(e) => {
|
||||
tracing::warn!("无法初始化 Flow 文件存储: {}", e);
|
||||
None
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,172 @@
|
||||
//! 核心类型定义
|
||||
//!
|
||||
//! 包含 Provider 类型枚举和相关实现。
|
||||
|
||||
use serde::{Deserialize, Serialize};
|
||||
use std::sync::Arc;
|
||||
use tauri::Runtime;
|
||||
use tokio::sync::RwLock;
|
||||
|
||||
use crate::logger;
|
||||
use crate::server;
|
||||
use crate::services::token_cache_service::TokenCacheService;
|
||||
use crate::tray::TrayManager;
|
||||
|
||||
/// Provider 类型枚举
|
||||
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash, Serialize, Deserialize)]
|
||||
#[serde(rename_all = "lowercase")]
|
||||
pub enum ProviderType {
|
||||
Kiro,
|
||||
Gemini,
|
||||
Qwen,
|
||||
#[serde(rename = "openai")]
|
||||
OpenAI,
|
||||
Claude,
|
||||
Antigravity,
|
||||
Vertex,
|
||||
#[serde(rename = "gemini_api_key")]
|
||||
GeminiApiKey,
|
||||
Codex,
|
||||
#[serde(rename = "claude_oauth")]
|
||||
ClaudeOAuth,
|
||||
#[serde(rename = "iflow")]
|
||||
IFlow,
|
||||
// API Key Provider 类型
|
||||
Anthropic,
|
||||
#[serde(rename = "azure_openai")]
|
||||
AzureOpenai,
|
||||
#[serde(rename = "aws_bedrock")]
|
||||
AwsBedrock,
|
||||
Ollama,
|
||||
}
|
||||
|
||||
impl std::fmt::Display for ProviderType {
|
||||
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
|
||||
match self {
|
||||
ProviderType::Kiro => write!(f, "kiro"),
|
||||
ProviderType::Gemini => write!(f, "gemini"),
|
||||
ProviderType::Qwen => write!(f, "qwen"),
|
||||
ProviderType::OpenAI => write!(f, "openai"),
|
||||
ProviderType::Claude => write!(f, "claude"),
|
||||
ProviderType::Antigravity => write!(f, "antigravity"),
|
||||
ProviderType::Vertex => write!(f, "vertex"),
|
||||
ProviderType::GeminiApiKey => write!(f, "gemini_api_key"),
|
||||
ProviderType::Codex => write!(f, "codex"),
|
||||
ProviderType::ClaudeOAuth => write!(f, "claude_oauth"),
|
||||
ProviderType::IFlow => write!(f, "iflow"),
|
||||
ProviderType::Anthropic => write!(f, "anthropic"),
|
||||
ProviderType::AzureOpenai => write!(f, "azure_openai"),
|
||||
ProviderType::AwsBedrock => write!(f, "aws_bedrock"),
|
||||
ProviderType::Ollama => write!(f, "ollama"),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
impl std::str::FromStr for ProviderType {
|
||||
type Err = String;
|
||||
|
||||
fn from_str(s: &str) -> Result<Self, Self::Err> {
|
||||
match s.to_lowercase().as_str() {
|
||||
"kiro" => Ok(ProviderType::Kiro),
|
||||
"gemini" => Ok(ProviderType::Gemini),
|
||||
"qwen" => Ok(ProviderType::Qwen),
|
||||
"openai" => Ok(ProviderType::OpenAI),
|
||||
"claude" => Ok(ProviderType::Claude),
|
||||
"antigravity" => Ok(ProviderType::Antigravity),
|
||||
"vertex" => Ok(ProviderType::Vertex),
|
||||
"gemini_api_key" => Ok(ProviderType::GeminiApiKey),
|
||||
"codex" => Ok(ProviderType::Codex),
|
||||
"claude_oauth" => Ok(ProviderType::ClaudeOAuth),
|
||||
"iflow" => Ok(ProviderType::IFlow),
|
||||
"anthropic" => Ok(ProviderType::Anthropic),
|
||||
"azure_openai" | "azure-openai" => Ok(ProviderType::AzureOpenai),
|
||||
"aws_bedrock" | "aws-bedrock" => Ok(ProviderType::AwsBedrock),
|
||||
"ollama" => Ok(ProviderType::Ollama),
|
||||
_ => Err(format!("Invalid provider: {s}")),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
/// 应用状态类型别名
|
||||
pub type AppState = Arc<RwLock<server::ServerState>>;
|
||||
|
||||
/// 日志状态类型别名
|
||||
pub type LogState = Arc<RwLock<logger::LogStore>>;
|
||||
|
||||
/// TokenCacheService 状态封装
|
||||
pub struct TokenCacheServiceState(pub Arc<TokenCacheService>);
|
||||
|
||||
/// TrayManager 状态封装
|
||||
pub struct TrayManagerState<R: Runtime>(pub Arc<tokio::sync::RwLock<Option<TrayManager<R>>>>);
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
|
||||
#[test]
|
||||
fn test_provider_type_from_str() {
|
||||
assert_eq!("kiro".parse::<ProviderType>().unwrap(), ProviderType::Kiro);
|
||||
assert_eq!(
|
||||
"gemini".parse::<ProviderType>().unwrap(),
|
||||
ProviderType::Gemini
|
||||
);
|
||||
assert_eq!("qwen".parse::<ProviderType>().unwrap(), ProviderType::Qwen);
|
||||
assert_eq!(
|
||||
"openai".parse::<ProviderType>().unwrap(),
|
||||
ProviderType::OpenAI
|
||||
);
|
||||
assert_eq!(
|
||||
"claude".parse::<ProviderType>().unwrap(),
|
||||
ProviderType::Claude
|
||||
);
|
||||
assert_eq!(
|
||||
"vertex".parse::<ProviderType>().unwrap(),
|
||||
ProviderType::Vertex
|
||||
);
|
||||
assert_eq!(
|
||||
"gemini_api_key".parse::<ProviderType>().unwrap(),
|
||||
ProviderType::GeminiApiKey
|
||||
);
|
||||
assert_eq!("KIRO".parse::<ProviderType>().unwrap(), ProviderType::Kiro);
|
||||
assert_eq!(
|
||||
"Gemini".parse::<ProviderType>().unwrap(),
|
||||
ProviderType::Gemini
|
||||
);
|
||||
assert_eq!(
|
||||
"VERTEX".parse::<ProviderType>().unwrap(),
|
||||
ProviderType::Vertex
|
||||
);
|
||||
assert!("invalid".parse::<ProviderType>().is_err());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_provider_type_display() {
|
||||
assert_eq!(ProviderType::Kiro.to_string(), "kiro");
|
||||
assert_eq!(ProviderType::Gemini.to_string(), "gemini");
|
||||
assert_eq!(ProviderType::Qwen.to_string(), "qwen");
|
||||
assert_eq!(ProviderType::OpenAI.to_string(), "openai");
|
||||
assert_eq!(ProviderType::Claude.to_string(), "claude");
|
||||
assert_eq!(ProviderType::Vertex.to_string(), "vertex");
|
||||
assert_eq!(ProviderType::GeminiApiKey.to_string(), "gemini_api_key");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_provider_type_serde() {
|
||||
assert_eq!(
|
||||
serde_json::to_string(&ProviderType::Kiro).unwrap(),
|
||||
"\"kiro\""
|
||||
);
|
||||
assert_eq!(
|
||||
serde_json::to_string(&ProviderType::OpenAI).unwrap(),
|
||||
"\"openai\""
|
||||
);
|
||||
assert_eq!(
|
||||
serde_json::from_str::<ProviderType>("\"kiro\"").unwrap(),
|
||||
ProviderType::Kiro
|
||||
);
|
||||
assert_eq!(
|
||||
serde_json::from_str::<ProviderType>("\"openai\"").unwrap(),
|
||||
ProviderType::OpenAI
|
||||
);
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,53 @@
|
||||
//! 辅助函数
|
||||
//!
|
||||
//! 包含通用工具函数。
|
||||
|
||||
use crate::config;
|
||||
|
||||
/// 生成安全的 API Key
|
||||
pub fn generate_api_key() -> String {
|
||||
config::generate_secure_api_key()
|
||||
}
|
||||
|
||||
/// 检查是否为回环地址
|
||||
pub fn is_loopback_host(host: &str) -> bool {
|
||||
if host == "localhost" {
|
||||
return true;
|
||||
}
|
||||
match host.parse::<std::net::IpAddr>() {
|
||||
Ok(addr) => addr.is_loopback(),
|
||||
Err(_) => false,
|
||||
}
|
||||
}
|
||||
|
||||
/// 掩码敏感 Token
|
||||
pub fn mask_token(token: &str) -> String {
|
||||
let chars: Vec<char> = token.chars().collect();
|
||||
if chars.len() <= 12 {
|
||||
"****".to_string()
|
||||
} else {
|
||||
let prefix: String = chars[..6].iter().collect();
|
||||
let suffix: String = chars[chars.len() - 4..].iter().collect();
|
||||
format!("{prefix}****{suffix}")
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
|
||||
#[test]
|
||||
fn test_is_loopback_host() {
|
||||
assert!(is_loopback_host("localhost"));
|
||||
assert!(is_loopback_host("127.0.0.1"));
|
||||
assert!(is_loopback_host("::1"));
|
||||
assert!(!is_loopback_host("0.0.0.0"));
|
||||
assert!(!is_loopback_host("192.168.1.1"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_mask_token() {
|
||||
assert_eq!(mask_token("short"), "****");
|
||||
assert_eq!(mask_token("abcdefghijklmnop"), "abcdef****mnop");
|
||||
}
|
||||
}
|
||||
@@ -396,3 +396,174 @@ pub fn import_api_key_providers(
|
||||
) -> Result<ImportResult, String> {
|
||||
service.0.import_config(&db, &config_json)
|
||||
}
|
||||
|
||||
// ============================================================================
|
||||
// 迁移命令 - 将旧的凭证池 API Key 迁移到新的 API Key Provider 系统
|
||||
// ============================================================================
|
||||
|
||||
use crate::commands::provider_pool_cmd::ProviderPoolServiceState;
|
||||
use crate::database::dao::provider_pool::ProviderPoolDao;
|
||||
use crate::models::provider_pool_model::CredentialData;
|
||||
|
||||
/// 迁移结果
|
||||
#[derive(Debug, Clone, Serialize, Deserialize)]
|
||||
pub struct MigrationResult {
|
||||
/// 迁移成功的凭证数量
|
||||
pub migrated_count: usize,
|
||||
/// 跳过的凭证数量(已存在或不支持)
|
||||
pub skipped_count: usize,
|
||||
/// 删除的旧凭证数量
|
||||
pub deleted_count: usize,
|
||||
/// 错误信息
|
||||
pub errors: Vec<String>,
|
||||
}
|
||||
|
||||
/// 获取需要迁移的旧 API Key 凭证列表
|
||||
#[tauri::command]
|
||||
pub fn get_legacy_api_key_credentials(
|
||||
db: State<'_, DbConnection>,
|
||||
) -> Result<Vec<LegacyApiKeyCredential>, String> {
|
||||
let conn = db.lock().map_err(|e| e.to_string())?;
|
||||
let all_credentials = ProviderPoolDao::get_all(&conn).map_err(|e| e.to_string())?;
|
||||
|
||||
let legacy_credentials: Vec<LegacyApiKeyCredential> = all_credentials
|
||||
.into_iter()
|
||||
.filter_map(|cred| match &cred.credential {
|
||||
CredentialData::OpenAIKey { api_key, base_url } => Some(LegacyApiKeyCredential {
|
||||
uuid: cred.uuid.to_string(),
|
||||
provider_type: "openai".to_string(),
|
||||
name: cred.name.clone(),
|
||||
api_key_masked: mask_api_key(api_key),
|
||||
base_url: base_url.clone(),
|
||||
usage_count: cred.usage_count as i64,
|
||||
error_count: cred.error_count as i64,
|
||||
created_at: cred.created_at.to_rfc3339(),
|
||||
}),
|
||||
CredentialData::ClaudeKey { api_key, base_url } => Some(LegacyApiKeyCredential {
|
||||
uuid: cred.uuid.to_string(),
|
||||
provider_type: "anthropic".to_string(),
|
||||
name: cred.name.clone(),
|
||||
api_key_masked: mask_api_key(api_key),
|
||||
base_url: base_url.clone(),
|
||||
usage_count: cred.usage_count as i64,
|
||||
error_count: cred.error_count as i64,
|
||||
created_at: cred.created_at.to_rfc3339(),
|
||||
}),
|
||||
_ => None,
|
||||
})
|
||||
.collect();
|
||||
|
||||
Ok(legacy_credentials)
|
||||
}
|
||||
|
||||
/// 旧的 API Key 凭证信息(用于前端显示)
|
||||
#[derive(Debug, Clone, Serialize, Deserialize)]
|
||||
pub struct LegacyApiKeyCredential {
|
||||
pub uuid: String,
|
||||
pub provider_type: String,
|
||||
pub name: Option<String>,
|
||||
pub api_key_masked: String,
|
||||
pub base_url: Option<String>,
|
||||
pub usage_count: i64,
|
||||
pub error_count: i64,
|
||||
pub created_at: String,
|
||||
}
|
||||
|
||||
/// 迁移旧的 API Key 凭证到新的 API Key Provider 系统
|
||||
#[tauri::command]
|
||||
pub fn migrate_legacy_api_key_credentials(
|
||||
db: State<'_, DbConnection>,
|
||||
api_key_service: State<'_, ApiKeyProviderServiceState>,
|
||||
pool_service: State<'_, ProviderPoolServiceState>,
|
||||
delete_after_migration: bool,
|
||||
) -> Result<MigrationResult, String> {
|
||||
let conn = db.lock().map_err(|e| e.to_string())?;
|
||||
let all_credentials = ProviderPoolDao::get_all(&conn).map_err(|e| e.to_string())?;
|
||||
drop(conn);
|
||||
|
||||
let mut migrated_count = 0;
|
||||
let mut skipped_count = 0;
|
||||
let mut deleted_count = 0;
|
||||
let mut errors = Vec::new();
|
||||
|
||||
for cred in all_credentials {
|
||||
let (provider_id, api_key, base_url) = match &cred.credential {
|
||||
CredentialData::OpenAIKey { api_key, base_url } => {
|
||||
("openai".to_string(), api_key.clone(), base_url.clone())
|
||||
}
|
||||
CredentialData::ClaudeKey { api_key, base_url } => {
|
||||
("anthropic".to_string(), api_key.clone(), base_url.clone())
|
||||
}
|
||||
_ => {
|
||||
// 不是 API Key 类型,跳过
|
||||
continue;
|
||||
}
|
||||
};
|
||||
|
||||
// 尝试添加到新的 API Key Provider 系统
|
||||
let alias = cred.name.clone();
|
||||
match api_key_service
|
||||
.0
|
||||
.add_api_key(&db, &provider_id, &api_key, alias)
|
||||
{
|
||||
Ok(_) => {
|
||||
migrated_count += 1;
|
||||
tracing::info!(
|
||||
"迁移成功: {} -> {} ({})",
|
||||
cred.uuid,
|
||||
provider_id,
|
||||
cred.name.as_deref().unwrap_or("未命名")
|
||||
);
|
||||
|
||||
// 如果需要删除旧凭证
|
||||
if delete_after_migration {
|
||||
match pool_service
|
||||
.0
|
||||
.delete_credential(&db, &cred.uuid.to_string())
|
||||
{
|
||||
Ok(_) => {
|
||||
deleted_count += 1;
|
||||
tracing::info!("删除旧凭证: {}", cred.uuid);
|
||||
}
|
||||
Err(e) => {
|
||||
errors.push(format!("删除旧凭证 {} 失败: {}", cred.uuid, e));
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
Err(e) => {
|
||||
// 可能是重复的 API Key,跳过
|
||||
skipped_count += 1;
|
||||
tracing::warn!("迁移跳过: {} - {}", cred.uuid, e);
|
||||
}
|
||||
}
|
||||
|
||||
// 如果有自定义 base_url,记录警告(新系统可能需要手动配置)
|
||||
if let Some(url) = base_url {
|
||||
if !url.is_empty() {
|
||||
errors.push(format!(
|
||||
"凭证 {} 有自定义 base_url ({}),请在新系统中手动配置",
|
||||
cred.name.as_deref().unwrap_or(&cred.uuid.to_string()),
|
||||
url
|
||||
));
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
Ok(MigrationResult {
|
||||
migrated_count,
|
||||
skipped_count,
|
||||
deleted_count,
|
||||
errors,
|
||||
})
|
||||
}
|
||||
|
||||
/// 删除单个旧的 API Key 凭证
|
||||
#[tauri::command]
|
||||
pub fn delete_legacy_api_key_credential(
|
||||
db: State<'_, DbConnection>,
|
||||
pool_service: State<'_, ProviderPoolServiceState>,
|
||||
uuid: String,
|
||||
) -> Result<bool, String> {
|
||||
pool_service.0.delete_credential(&db, &uuid)
|
||||
}
|
||||
|
||||
@@ -11,6 +11,8 @@ pub mod mcp_cmd;
|
||||
pub mod native_agent_cmd;
|
||||
pub mod network_cmd;
|
||||
pub mod oauth_cmd;
|
||||
pub mod oauth_plugin_cmd;
|
||||
pub mod orchestrator_cmd;
|
||||
pub mod plugin_cmd;
|
||||
pub mod plugin_install_cmd;
|
||||
pub mod prompt_cmd;
|
||||
|
||||
@@ -0,0 +1,933 @@
|
||||
//! OAuth Provider 插件命令
|
||||
//!
|
||||
//! 提供 OAuth Provider 插件管理的 Tauri 命令:
|
||||
//! - list_oauth_plugins: 获取所有已安装的 OAuth Provider 插件
|
||||
//! - get_oauth_plugin: 获取单个插件信息
|
||||
//! - enable_oauth_plugin: 启用插件
|
||||
//! - disable_oauth_plugin: 禁用插件
|
||||
//! - install_oauth_plugin: 安装插件
|
||||
//! - uninstall_oauth_plugin: 卸载插件
|
||||
//! - 插件 SDK 命令
|
||||
|
||||
use crate::credential::{
|
||||
get_global_registry, init_global_registry, OAuthPluginLoader, PluginPermission,
|
||||
PluginSdkContext, PluginSource,
|
||||
};
|
||||
use crate::database::DbConnection;
|
||||
use rusqlite::params;
|
||||
use serde::{Deserialize, Serialize};
|
||||
use std::path::PathBuf;
|
||||
use std::sync::Arc;
|
||||
use tokio::sync::RwLock;
|
||||
use tracing::{error, info};
|
||||
|
||||
// ============================================================================
|
||||
// 状态管理
|
||||
// ============================================================================
|
||||
|
||||
/// OAuth 插件管理器状态
|
||||
pub struct OAuthPluginManagerState {
|
||||
/// 插件加载器
|
||||
pub loader: Arc<RwLock<OAuthPluginLoader>>,
|
||||
/// 是否已初始化
|
||||
pub initialized: Arc<RwLock<bool>>,
|
||||
}
|
||||
|
||||
impl OAuthPluginManagerState {
|
||||
/// 创建新状态
|
||||
pub fn new(plugins_dir: PathBuf) -> Self {
|
||||
Self {
|
||||
loader: Arc::new(RwLock::new(OAuthPluginLoader::new(plugins_dir))),
|
||||
initialized: Arc::new(RwLock::new(false)),
|
||||
}
|
||||
}
|
||||
|
||||
/// 使用默认配置创建
|
||||
pub fn with_defaults() -> Self {
|
||||
Self::new(OAuthPluginLoader::default_plugins_dir())
|
||||
}
|
||||
}
|
||||
|
||||
// ============================================================================
|
||||
// 响应类型
|
||||
// ============================================================================
|
||||
|
||||
/// 认证类型信息
|
||||
#[derive(Debug, Clone, Serialize, Deserialize)]
|
||||
#[serde(rename_all = "camelCase")]
|
||||
pub struct AuthTypeInfoResponse {
|
||||
pub id: String,
|
||||
pub display_name: String,
|
||||
pub description: String,
|
||||
pub category: String,
|
||||
pub icon: Option<String>,
|
||||
}
|
||||
|
||||
/// 模型家族信息
|
||||
#[derive(Debug, Clone, Serialize, Deserialize)]
|
||||
#[serde(rename_all = "camelCase")]
|
||||
pub struct ModelFamilyResponse {
|
||||
pub name: String,
|
||||
pub pattern: String,
|
||||
pub tier: Option<String>,
|
||||
pub description: Option<String>,
|
||||
}
|
||||
|
||||
/// OAuth 插件信息响应
|
||||
#[derive(Debug, Clone, Serialize, Deserialize)]
|
||||
#[serde(rename_all = "camelCase")]
|
||||
pub struct OAuthPluginInfoResponse {
|
||||
pub id: String,
|
||||
pub display_name: String,
|
||||
pub version: String,
|
||||
pub description: String,
|
||||
pub target_protocol: String,
|
||||
pub category: String,
|
||||
pub enabled: bool,
|
||||
pub install_path: String,
|
||||
pub installed_at: String,
|
||||
pub last_used_at: Option<String>,
|
||||
pub credential_count: u32,
|
||||
pub healthy_credential_count: u32,
|
||||
pub auth_types: Vec<AuthTypeInfoResponse>,
|
||||
}
|
||||
|
||||
/// 插件安装来源
|
||||
#[derive(Debug, Clone, Serialize, Deserialize)]
|
||||
#[serde(tag = "type", rename_all = "snake_case")]
|
||||
pub enum PluginSourceRequest {
|
||||
GitHub {
|
||||
owner: String,
|
||||
repo: String,
|
||||
version: Option<String>,
|
||||
},
|
||||
LocalFile {
|
||||
path: String,
|
||||
},
|
||||
Builtin {
|
||||
id: String,
|
||||
},
|
||||
}
|
||||
|
||||
impl From<PluginSourceRequest> for PluginSource {
|
||||
fn from(req: PluginSourceRequest) -> Self {
|
||||
match req {
|
||||
PluginSourceRequest::GitHub {
|
||||
owner,
|
||||
repo,
|
||||
version,
|
||||
} => PluginSource::GitHub {
|
||||
owner,
|
||||
repo,
|
||||
version,
|
||||
},
|
||||
PluginSourceRequest::LocalFile { path } => PluginSource::LocalFile {
|
||||
path: PathBuf::from(path),
|
||||
},
|
||||
PluginSourceRequest::Builtin { id } => PluginSource::Builtin { id },
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
/// 安装结果
|
||||
#[derive(Debug, Clone, Serialize, Deserialize)]
|
||||
#[serde(rename_all = "camelCase")]
|
||||
pub struct InstallResultResponse {
|
||||
pub success: bool,
|
||||
pub plugin_id: Option<String>,
|
||||
pub error: Option<String>,
|
||||
}
|
||||
|
||||
/// 插件更新信息
|
||||
#[derive(Debug, Clone, Serialize, Deserialize)]
|
||||
#[serde(rename_all = "camelCase")]
|
||||
pub struct PluginUpdateResponse {
|
||||
pub plugin_id: String,
|
||||
pub current_version: String,
|
||||
pub latest_version: String,
|
||||
pub changelog: Option<String>,
|
||||
}
|
||||
|
||||
// ============================================================================
|
||||
// 插件管理命令
|
||||
// ============================================================================
|
||||
|
||||
/// 初始化 OAuth 插件系统
|
||||
#[tauri::command]
|
||||
pub async fn init_oauth_plugin_system(
|
||||
state: tauri::State<'_, OAuthPluginManagerState>,
|
||||
) -> Result<(), String> {
|
||||
let mut initialized = state.initialized.write().await;
|
||||
if *initialized {
|
||||
return Ok(());
|
||||
}
|
||||
|
||||
let loader = state.loader.read().await;
|
||||
|
||||
// 初始化全局注册表
|
||||
let registry = init_global_registry(loader.plugins_dir().to_path_buf());
|
||||
|
||||
// 加载所有插件
|
||||
match loader.load_all(®istry).await {
|
||||
Ok(loaded) => {
|
||||
info!("已加载 {} 个 OAuth Provider 插件", loaded.len());
|
||||
*initialized = true;
|
||||
Ok(())
|
||||
}
|
||||
Err(e) => {
|
||||
error!("加载 OAuth Provider 插件失败: {}", e);
|
||||
Err(e.to_string())
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
/// 获取所有已安装的 OAuth Provider 插件
|
||||
#[tauri::command]
|
||||
pub async fn list_oauth_plugins(
|
||||
db: tauri::State<'_, DbConnection>,
|
||||
) -> Result<Vec<OAuthPluginInfoResponse>, String> {
|
||||
let registry = get_global_registry().ok_or("OAuth 插件系统未初始化")?;
|
||||
|
||||
let infos = registry.get_plugin_infos();
|
||||
|
||||
// 查询每个插件的凭证数量
|
||||
let conn = db.lock().map_err(|e| e.to_string())?;
|
||||
|
||||
let plugins: Vec<OAuthPluginInfoResponse> = infos
|
||||
.into_iter()
|
||||
.map(|info| {
|
||||
// 查询凭证数量
|
||||
let credential_count: u32 = conn
|
||||
.query_row(
|
||||
"SELECT COUNT(*) FROM plugin_credentials WHERE plugin_id = ?",
|
||||
params![info.id],
|
||||
|row| row.get(0),
|
||||
)
|
||||
.unwrap_or(0);
|
||||
|
||||
OAuthPluginInfoResponse {
|
||||
id: info.id.clone(),
|
||||
display_name: info.display_name,
|
||||
version: info.version,
|
||||
description: info.description,
|
||||
target_protocol: info.target_protocol,
|
||||
category: format!("{:?}", info.category),
|
||||
enabled: info.enabled,
|
||||
install_path: registry
|
||||
.plugins_dir()
|
||||
.join(&info.id)
|
||||
.to_string_lossy()
|
||||
.to_string(),
|
||||
installed_at: chrono::Utc::now().to_rfc3339(), // TODO: 从数据库获取
|
||||
last_used_at: None,
|
||||
credential_count,
|
||||
healthy_credential_count: info.healthy_credential_count,
|
||||
auth_types: info
|
||||
.auth_types
|
||||
.into_iter()
|
||||
.map(|a| AuthTypeInfoResponse {
|
||||
id: a.id,
|
||||
display_name: a.display_name,
|
||||
description: a.description,
|
||||
category: format!("{:?}", a.category),
|
||||
icon: a.icon,
|
||||
})
|
||||
.collect(),
|
||||
}
|
||||
})
|
||||
.collect();
|
||||
|
||||
Ok(plugins)
|
||||
}
|
||||
|
||||
/// 获取单个插件信息
|
||||
#[tauri::command]
|
||||
pub async fn get_oauth_plugin(
|
||||
db: tauri::State<'_, DbConnection>,
|
||||
plugin_id: String,
|
||||
) -> Result<Option<OAuthPluginInfoResponse>, String> {
|
||||
let plugins = list_oauth_plugins(db).await?;
|
||||
Ok(plugins.into_iter().find(|p| p.id == plugin_id))
|
||||
}
|
||||
|
||||
/// 启用插件
|
||||
#[tauri::command]
|
||||
pub async fn enable_oauth_plugin(
|
||||
db: tauri::State<'_, DbConnection>,
|
||||
plugin_id: String,
|
||||
) -> Result<(), String> {
|
||||
let registry = get_global_registry().ok_or("OAuth 插件系统未初始化")?;
|
||||
|
||||
if registry.enable_plugin(&plugin_id) {
|
||||
info!("已启用 OAuth 插件: {}", plugin_id);
|
||||
|
||||
// 更新数据库
|
||||
let conn = db.lock().map_err(|e| e.to_string())?;
|
||||
conn.execute(
|
||||
"UPDATE credential_provider_plugins SET enabled = 1, updated_at = ? WHERE id = ?",
|
||||
params![chrono::Utc::now().to_rfc3339(), plugin_id],
|
||||
)
|
||||
.map_err(|e| e.to_string())?;
|
||||
|
||||
Ok(())
|
||||
} else {
|
||||
Err(format!("插件不存在: {}", plugin_id))
|
||||
}
|
||||
}
|
||||
|
||||
/// 禁用插件
|
||||
#[tauri::command]
|
||||
pub async fn disable_oauth_plugin(
|
||||
db: tauri::State<'_, DbConnection>,
|
||||
plugin_id: String,
|
||||
) -> Result<(), String> {
|
||||
let registry = get_global_registry().ok_or("OAuth 插件系统未初始化")?;
|
||||
|
||||
if registry.disable_plugin(&plugin_id) {
|
||||
info!("已禁用 OAuth 插件: {}", plugin_id);
|
||||
|
||||
// 更新数据库
|
||||
let conn = db.lock().map_err(|e| e.to_string())?;
|
||||
conn.execute(
|
||||
"UPDATE credential_provider_plugins SET enabled = 0, updated_at = ? WHERE id = ?",
|
||||
params![chrono::Utc::now().to_rfc3339(), plugin_id],
|
||||
)
|
||||
.map_err(|e| e.to_string())?;
|
||||
|
||||
Ok(())
|
||||
} else {
|
||||
Err(format!("插件不存在: {}", plugin_id))
|
||||
}
|
||||
}
|
||||
|
||||
/// 安装插件
|
||||
#[tauri::command]
|
||||
pub async fn install_oauth_plugin(
|
||||
_state: tauri::State<'_, OAuthPluginManagerState>,
|
||||
source: PluginSourceRequest,
|
||||
) -> Result<InstallResultResponse, String> {
|
||||
let registry = get_global_registry().ok_or("OAuth 插件系统未初始化")?;
|
||||
|
||||
let plugin_source: PluginSource = source.into();
|
||||
|
||||
match registry.install_plugin(plugin_source).await {
|
||||
Ok(plugin_id) => {
|
||||
info!("已安装 OAuth 插件: {}", plugin_id);
|
||||
Ok(InstallResultResponse {
|
||||
success: true,
|
||||
plugin_id: Some(plugin_id),
|
||||
error: None,
|
||||
})
|
||||
}
|
||||
Err(e) => {
|
||||
error!("安装 OAuth 插件失败: {}", e);
|
||||
Ok(InstallResultResponse {
|
||||
success: false,
|
||||
plugin_id: None,
|
||||
error: Some(e.to_string()),
|
||||
})
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
/// 卸载插件
|
||||
#[tauri::command]
|
||||
pub async fn uninstall_oauth_plugin(
|
||||
db: tauri::State<'_, DbConnection>,
|
||||
plugin_id: String,
|
||||
) -> Result<(), String> {
|
||||
let registry = get_global_registry().ok_or("OAuth 插件系统未初始化")?;
|
||||
|
||||
// 使用块作用域确保 MutexGuard 在 await 之前释放
|
||||
{
|
||||
let conn = db.lock().map_err(|e| e.to_string())?;
|
||||
|
||||
// 删除插件凭证
|
||||
conn.execute(
|
||||
"DELETE FROM plugin_credentials WHERE plugin_id = ?",
|
||||
params![plugin_id],
|
||||
)
|
||||
.map_err(|e| e.to_string())?;
|
||||
|
||||
// 删除插件存储
|
||||
conn.execute(
|
||||
"DELETE FROM plugin_storage WHERE plugin_id = ?",
|
||||
params![plugin_id],
|
||||
)
|
||||
.map_err(|e| e.to_string())?;
|
||||
|
||||
// 删除插件记录
|
||||
conn.execute(
|
||||
"DELETE FROM credential_provider_plugins WHERE id = ?",
|
||||
params![plugin_id],
|
||||
)
|
||||
.map_err(|e| e.to_string())?;
|
||||
} // conn 在此处释放
|
||||
|
||||
// 从注册表卸载
|
||||
registry
|
||||
.uninstall_plugin(&plugin_id)
|
||||
.await
|
||||
.map_err(|e| e.to_string())?;
|
||||
|
||||
info!("已卸载 OAuth 插件: {}", plugin_id);
|
||||
Ok(())
|
||||
}
|
||||
|
||||
/// 检查插件更新
|
||||
#[tauri::command]
|
||||
pub async fn check_oauth_plugin_updates() -> Result<Vec<PluginUpdateResponse>, String> {
|
||||
let registry = get_global_registry().ok_or("OAuth 插件系统未初始化")?;
|
||||
|
||||
let updates = registry.check_updates().await.map_err(|e| e.to_string())?;
|
||||
|
||||
Ok(updates
|
||||
.into_iter()
|
||||
.map(|u| PluginUpdateResponse {
|
||||
plugin_id: u.plugin_id,
|
||||
current_version: u.current_version,
|
||||
latest_version: u.latest_version,
|
||||
changelog: u.changelog,
|
||||
})
|
||||
.collect())
|
||||
}
|
||||
|
||||
/// 更新插件
|
||||
#[tauri::command]
|
||||
pub async fn update_oauth_plugin(_plugin_id: String) -> Result<(), String> {
|
||||
// TODO: 实现插件更新逻辑
|
||||
Err("插件更新功能尚未实现".to_string())
|
||||
}
|
||||
|
||||
/// 重新加载所有插件
|
||||
#[tauri::command]
|
||||
pub fn reload_oauth_plugins(
|
||||
_state: tauri::State<'_, OAuthPluginManagerState>,
|
||||
) -> Result<(), String> {
|
||||
// TODO: 由于 DashMap 生命周期限制,暂时不支持热重载
|
||||
// 需要重启应用来重新加载插件
|
||||
info!("请重启应用以重新加载 OAuth 插件");
|
||||
Err("请重启应用以重新加载 OAuth 插件".to_string())
|
||||
}
|
||||
|
||||
/// 获取插件配置
|
||||
#[tauri::command]
|
||||
pub async fn get_oauth_plugin_config(plugin_id: String) -> Result<serde_json::Value, String> {
|
||||
let registry = get_global_registry().ok_or("OAuth 插件系统未初始化")?;
|
||||
|
||||
let state = registry
|
||||
.get_plugin_state(&plugin_id)
|
||||
.ok_or(format!("插件不存在: {}", plugin_id))?;
|
||||
|
||||
Ok(state.config)
|
||||
}
|
||||
|
||||
/// 更新插件配置
|
||||
#[tauri::command]
|
||||
pub async fn update_oauth_plugin_config(
|
||||
db: tauri::State<'_, DbConnection>,
|
||||
plugin_id: String,
|
||||
config: serde_json::Value,
|
||||
) -> Result<(), String> {
|
||||
let registry = get_global_registry().ok_or("OAuth 插件系统未初始化")?;
|
||||
|
||||
registry
|
||||
.update_plugin_config(&plugin_id, config.clone())
|
||||
.await
|
||||
.map_err(|e| e.to_string())?;
|
||||
|
||||
// 更新数据库
|
||||
let conn = db.lock().map_err(|e| e.to_string())?;
|
||||
conn.execute(
|
||||
"UPDATE credential_provider_plugins SET config = ?, updated_at = ? WHERE id = ?",
|
||||
params![
|
||||
config.to_string(),
|
||||
chrono::Utc::now().to_rfc3339(),
|
||||
plugin_id
|
||||
],
|
||||
)
|
||||
.map_err(|e| e.to_string())?;
|
||||
|
||||
info!("已更新 OAuth 插件配置: {}", plugin_id);
|
||||
Ok(())
|
||||
}
|
||||
|
||||
/// 扫描插件目录
|
||||
#[tauri::command]
|
||||
pub async fn scan_oauth_plugin_directory(
|
||||
state: tauri::State<'_, OAuthPluginManagerState>,
|
||||
) -> Result<Vec<String>, String> {
|
||||
let loader = state.loader.read().await;
|
||||
|
||||
let paths = loader.scan().await.map_err(|e| e.to_string())?;
|
||||
|
||||
Ok(paths
|
||||
.into_iter()
|
||||
.map(|p| p.to_string_lossy().to_string())
|
||||
.collect())
|
||||
}
|
||||
|
||||
// ============================================================================
|
||||
// 插件凭证命令
|
||||
// ============================================================================
|
||||
|
||||
/// 凭证信息响应
|
||||
#[derive(Debug, Clone, Serialize, Deserialize)]
|
||||
#[serde(rename_all = "camelCase")]
|
||||
pub struct CredentialInfoResponse {
|
||||
pub id: String,
|
||||
pub plugin_id: String,
|
||||
pub auth_type: String,
|
||||
pub display_name: Option<String>,
|
||||
pub status: String,
|
||||
pub created_at: String,
|
||||
pub updated_at: String,
|
||||
pub last_used_at: Option<String>,
|
||||
pub config: serde_json::Value,
|
||||
}
|
||||
|
||||
/// 获取插件凭证列表
|
||||
#[tauri::command]
|
||||
pub async fn plugin_credential_list(
|
||||
db: tauri::State<'_, DbConnection>,
|
||||
plugin_id: String,
|
||||
) -> Result<PluginCredentialListResponse, String> {
|
||||
let conn = db.lock().map_err(|e| e.to_string())?;
|
||||
|
||||
let mut stmt = conn
|
||||
.prepare(
|
||||
"SELECT id, plugin_id, auth_type, display_name, status,
|
||||
config_encrypted, created_at, updated_at, last_used_at
|
||||
FROM plugin_credentials WHERE plugin_id = ?",
|
||||
)
|
||||
.map_err(|e| e.to_string())?;
|
||||
|
||||
let rows = stmt
|
||||
.query_map(params![plugin_id], |row| {
|
||||
Ok(CredentialInfoResponse {
|
||||
id: row.get(0)?,
|
||||
plugin_id: row.get(1)?,
|
||||
auth_type: row.get(2)?,
|
||||
display_name: row.get(3)?,
|
||||
status: row.get(4)?,
|
||||
config: serde_json::json!({}), // 不返回加密配置
|
||||
created_at: row.get(6)?,
|
||||
updated_at: row.get(7)?,
|
||||
last_used_at: row.get(8)?,
|
||||
})
|
||||
})
|
||||
.map_err(|e| e.to_string())?;
|
||||
|
||||
let credentials: Vec<CredentialInfoResponse> = rows.filter_map(|r| r.ok()).collect();
|
||||
|
||||
Ok(PluginCredentialListResponse { credentials })
|
||||
}
|
||||
|
||||
/// 凭证列表响应
|
||||
#[derive(Debug, Clone, Serialize, Deserialize)]
|
||||
#[serde(rename_all = "camelCase")]
|
||||
pub struct PluginCredentialListResponse {
|
||||
pub credentials: Vec<CredentialInfoResponse>,
|
||||
}
|
||||
|
||||
/// 获取单个凭证
|
||||
#[tauri::command]
|
||||
pub async fn plugin_credential_get(
|
||||
db: tauri::State<'_, DbConnection>,
|
||||
plugin_id: String,
|
||||
credential_id: String,
|
||||
) -> Result<PluginCredentialGetResponse, String> {
|
||||
let result = plugin_credential_list(db, plugin_id).await?;
|
||||
let credential = result
|
||||
.credentials
|
||||
.into_iter()
|
||||
.find(|c| c.id == credential_id);
|
||||
Ok(PluginCredentialGetResponse { credential })
|
||||
}
|
||||
|
||||
/// 单个凭证响应
|
||||
#[derive(Debug, Clone, Serialize, Deserialize)]
|
||||
#[serde(rename_all = "camelCase")]
|
||||
pub struct PluginCredentialGetResponse {
|
||||
pub credential: Option<CredentialInfoResponse>,
|
||||
}
|
||||
|
||||
/// 创建凭证
|
||||
#[tauri::command]
|
||||
pub async fn plugin_credential_create(
|
||||
db: tauri::State<'_, DbConnection>,
|
||||
plugin_id: String,
|
||||
auth_type: String,
|
||||
config: serde_json::Value,
|
||||
) -> Result<String, String> {
|
||||
let conn = db.lock().map_err(|e| e.to_string())?;
|
||||
|
||||
let credential_id = uuid::Uuid::new_v4().to_string();
|
||||
let now = chrono::Utc::now().to_rfc3339();
|
||||
|
||||
// TODO: 加密配置
|
||||
let config_encrypted = config.to_string();
|
||||
|
||||
conn.execute(
|
||||
"INSERT INTO plugin_credentials
|
||||
(id, plugin_id, auth_type, status, config_encrypted, created_at, updated_at)
|
||||
VALUES (?, ?, ?, 'active', ?, ?, ?)",
|
||||
params![
|
||||
credential_id,
|
||||
plugin_id,
|
||||
auth_type,
|
||||
config_encrypted,
|
||||
now,
|
||||
now
|
||||
],
|
||||
)
|
||||
.map_err(|e| e.to_string())?;
|
||||
|
||||
info!("已创建插件凭证: {} (插件: {})", credential_id, plugin_id);
|
||||
Ok(credential_id)
|
||||
}
|
||||
|
||||
/// 更新凭证
|
||||
#[tauri::command]
|
||||
pub async fn plugin_credential_update(
|
||||
db: tauri::State<'_, DbConnection>,
|
||||
plugin_id: String,
|
||||
credential_id: String,
|
||||
config: serde_json::Value,
|
||||
) -> Result<(), String> {
|
||||
let conn = db.lock().map_err(|e| e.to_string())?;
|
||||
|
||||
let now = chrono::Utc::now().to_rfc3339();
|
||||
let config_encrypted = config.to_string();
|
||||
|
||||
let affected = conn
|
||||
.execute(
|
||||
"UPDATE plugin_credentials SET config_encrypted = ?, updated_at = ?
|
||||
WHERE id = ? AND plugin_id = ?",
|
||||
params![config_encrypted, now, credential_id, plugin_id],
|
||||
)
|
||||
.map_err(|e| e.to_string())?;
|
||||
|
||||
if affected == 0 {
|
||||
return Err(format!("凭证不存在: {}", credential_id));
|
||||
}
|
||||
|
||||
info!("已更新插件凭证: {}", credential_id);
|
||||
Ok(())
|
||||
}
|
||||
|
||||
/// 删除凭证
|
||||
#[tauri::command]
|
||||
pub async fn plugin_credential_delete(
|
||||
db: tauri::State<'_, DbConnection>,
|
||||
plugin_id: String,
|
||||
credential_id: String,
|
||||
) -> Result<(), String> {
|
||||
let conn = db.lock().map_err(|e| e.to_string())?;
|
||||
|
||||
let affected = conn
|
||||
.execute(
|
||||
"DELETE FROM plugin_credentials WHERE id = ? AND plugin_id = ?",
|
||||
params![credential_id, plugin_id],
|
||||
)
|
||||
.map_err(|e| e.to_string())?;
|
||||
|
||||
if affected == 0 {
|
||||
return Err(format!("凭证不存在: {}", credential_id));
|
||||
}
|
||||
|
||||
info!("已删除插件凭证: {}", credential_id);
|
||||
Ok(())
|
||||
}
|
||||
|
||||
/// 验证凭证
|
||||
#[tauri::command]
|
||||
pub async fn plugin_credential_validate(
|
||||
plugin_id: String,
|
||||
credential_id: String,
|
||||
) -> Result<serde_json::Value, String> {
|
||||
let registry = get_global_registry().ok_or("OAuth 插件系统未初始化")?;
|
||||
|
||||
let plugin = registry
|
||||
.get(&plugin_id)
|
||||
.ok_or(format!("插件不存在: {}", plugin_id))?;
|
||||
|
||||
let result = plugin
|
||||
.validate_credential(&credential_id)
|
||||
.await
|
||||
.map_err(|e| e.to_string())?;
|
||||
|
||||
Ok(serde_json::json!({
|
||||
"valid": result.valid,
|
||||
"message": result.message
|
||||
}))
|
||||
}
|
||||
|
||||
/// 刷新凭证
|
||||
#[tauri::command]
|
||||
pub async fn plugin_credential_refresh(
|
||||
db: tauri::State<'_, DbConnection>,
|
||||
plugin_id: String,
|
||||
credential_id: String,
|
||||
) -> Result<(), String> {
|
||||
let registry = get_global_registry().ok_or("OAuth 插件系统未初始化")?;
|
||||
|
||||
let plugin = registry
|
||||
.get(&plugin_id)
|
||||
.ok_or(format!("插件不存在: {}", plugin_id))?;
|
||||
|
||||
plugin
|
||||
.refresh_token(&credential_id)
|
||||
.await
|
||||
.map_err(|e| e.to_string())?;
|
||||
|
||||
// 更新最后使用时间
|
||||
let conn = db.lock().map_err(|e| e.to_string())?;
|
||||
conn.execute(
|
||||
"UPDATE plugin_credentials SET last_used_at = ? WHERE id = ?",
|
||||
params![chrono::Utc::now().to_rfc3339(), credential_id],
|
||||
)
|
||||
.map_err(|e| e.to_string())?;
|
||||
|
||||
info!("已刷新插件凭证: {}", credential_id);
|
||||
Ok(())
|
||||
}
|
||||
|
||||
// ============================================================================
|
||||
// 插件 SDK 命令
|
||||
// ============================================================================
|
||||
|
||||
/// 插件数据库查询
|
||||
#[tauri::command]
|
||||
pub async fn plugin_database_query(
|
||||
plugin_id: String,
|
||||
sql: String,
|
||||
params: Vec<serde_json::Value>,
|
||||
) -> Result<serde_json::Value, String> {
|
||||
// 创建 SDK 上下文
|
||||
let context = PluginSdkContext::new(plugin_id.clone(), vec![PluginPermission::DatabaseRead]);
|
||||
|
||||
// 执行查询
|
||||
let result = context
|
||||
.database_query(&sql, params)
|
||||
.await
|
||||
.map_err(|e| e.to_string())?;
|
||||
|
||||
Ok(serde_json::to_value(result).unwrap())
|
||||
}
|
||||
|
||||
/// 插件数据库执行
|
||||
#[tauri::command]
|
||||
pub async fn plugin_database_execute(
|
||||
plugin_id: String,
|
||||
sql: String,
|
||||
params: Vec<serde_json::Value>,
|
||||
) -> Result<serde_json::Value, String> {
|
||||
let context = PluginSdkContext::new(plugin_id.clone(), vec![PluginPermission::DatabaseWrite]);
|
||||
|
||||
let affected = context
|
||||
.database_execute(&sql, params)
|
||||
.await
|
||||
.map_err(|e| e.to_string())?;
|
||||
|
||||
Ok(serde_json::json!({ "affected": affected }))
|
||||
}
|
||||
|
||||
/// 插件 HTTP 请求
|
||||
#[tauri::command]
|
||||
pub async fn plugin_http_request(
|
||||
plugin_id: String,
|
||||
url: String,
|
||||
method: String,
|
||||
headers: std::collections::HashMap<String, String>,
|
||||
body: Option<String>,
|
||||
timeout_ms: u64,
|
||||
) -> Result<serde_json::Value, String> {
|
||||
use crate::credential::HttpRequestOptions;
|
||||
|
||||
let context = PluginSdkContext::new(plugin_id.clone(), vec![PluginPermission::HttpRequest]);
|
||||
|
||||
let options = HttpRequestOptions {
|
||||
method,
|
||||
headers,
|
||||
body,
|
||||
timeout_ms,
|
||||
};
|
||||
|
||||
let response = context
|
||||
.http_request(&url, options)
|
||||
.await
|
||||
.map_err(|e| e.to_string())?;
|
||||
|
||||
Ok(serde_json::to_value(response).unwrap())
|
||||
}
|
||||
|
||||
/// 插件加密
|
||||
#[tauri::command]
|
||||
pub async fn plugin_crypto_encrypt(
|
||||
plugin_id: String,
|
||||
data: String,
|
||||
) -> Result<serde_json::Value, String> {
|
||||
let context = PluginSdkContext::new(plugin_id.clone(), vec![PluginPermission::CryptoEncrypt]);
|
||||
|
||||
let encrypted = context
|
||||
.crypto_encrypt(&data)
|
||||
.await
|
||||
.map_err(|e| e.to_string())?;
|
||||
|
||||
Ok(serde_json::json!({ "encrypted": encrypted }))
|
||||
}
|
||||
|
||||
/// 插件解密
|
||||
#[tauri::command]
|
||||
pub async fn plugin_crypto_decrypt(
|
||||
plugin_id: String,
|
||||
data: String,
|
||||
) -> Result<serde_json::Value, String> {
|
||||
let context = PluginSdkContext::new(plugin_id.clone(), vec![PluginPermission::CryptoDecrypt]);
|
||||
|
||||
let decrypted = context
|
||||
.crypto_decrypt(&data)
|
||||
.await
|
||||
.map_err(|e| e.to_string())?;
|
||||
|
||||
Ok(serde_json::json!({ "decrypted": decrypted }))
|
||||
}
|
||||
|
||||
/// 插件通知
|
||||
#[tauri::command]
|
||||
pub async fn plugin_notification(
|
||||
plugin_id: String,
|
||||
level: String,
|
||||
message: String,
|
||||
) -> Result<(), String> {
|
||||
let context = PluginSdkContext::new(plugin_id.clone(), vec![PluginPermission::Notification]);
|
||||
|
||||
match level.as_str() {
|
||||
"success" => context
|
||||
.notification_success(&message)
|
||||
.map_err(|e| e.to_string())?,
|
||||
"error" => context
|
||||
.notification_error(&message)
|
||||
.map_err(|e| e.to_string())?,
|
||||
"info" => context
|
||||
.notification_info(&message)
|
||||
.map_err(|e| e.to_string())?,
|
||||
_ => context
|
||||
.notification_info(&message)
|
||||
.map_err(|e| e.to_string())?,
|
||||
}
|
||||
|
||||
Ok(())
|
||||
}
|
||||
|
||||
/// 插件存储获取
|
||||
#[tauri::command]
|
||||
pub async fn plugin_storage_get(
|
||||
db: tauri::State<'_, DbConnection>,
|
||||
plugin_id: String,
|
||||
key: String,
|
||||
) -> Result<serde_json::Value, String> {
|
||||
let conn = db.lock().map_err(|e| e.to_string())?;
|
||||
|
||||
let value: Option<String> = conn
|
||||
.query_row(
|
||||
"SELECT value FROM plugin_storage WHERE plugin_id = ? AND key = ?",
|
||||
params![plugin_id, key],
|
||||
|row| row.get(0),
|
||||
)
|
||||
.ok();
|
||||
|
||||
Ok(serde_json::json!({ "value": value }))
|
||||
}
|
||||
|
||||
/// 插件存储设置
|
||||
#[tauri::command]
|
||||
pub async fn plugin_storage_set(
|
||||
db: tauri::State<'_, DbConnection>,
|
||||
plugin_id: String,
|
||||
key: String,
|
||||
value: String,
|
||||
) -> Result<(), String> {
|
||||
let conn = db.lock().map_err(|e| e.to_string())?;
|
||||
let now = chrono::Utc::now().to_rfc3339();
|
||||
|
||||
conn.execute(
|
||||
"INSERT OR REPLACE INTO plugin_storage (plugin_id, key, value, created_at, updated_at)
|
||||
VALUES (?, ?, ?, ?, ?)",
|
||||
params![plugin_id, key, value, now, now],
|
||||
)
|
||||
.map_err(|e| e.to_string())?;
|
||||
|
||||
Ok(())
|
||||
}
|
||||
|
||||
/// 插件存储删除
|
||||
#[tauri::command]
|
||||
pub async fn plugin_storage_delete(
|
||||
db: tauri::State<'_, DbConnection>,
|
||||
plugin_id: String,
|
||||
key: String,
|
||||
) -> Result<(), String> {
|
||||
let conn = db.lock().map_err(|e| e.to_string())?;
|
||||
|
||||
conn.execute(
|
||||
"DELETE FROM plugin_storage WHERE plugin_id = ? AND key = ?",
|
||||
params![plugin_id, key],
|
||||
)
|
||||
.map_err(|e| e.to_string())?;
|
||||
|
||||
Ok(())
|
||||
}
|
||||
|
||||
/// 插件存储键列表
|
||||
#[tauri::command]
|
||||
pub async fn plugin_storage_keys(
|
||||
db: tauri::State<'_, DbConnection>,
|
||||
plugin_id: String,
|
||||
) -> Result<serde_json::Value, String> {
|
||||
let conn = db.lock().map_err(|e| e.to_string())?;
|
||||
|
||||
let mut stmt = conn
|
||||
.prepare("SELECT key FROM plugin_storage WHERE plugin_id = ?")
|
||||
.map_err(|e| e.to_string())?;
|
||||
|
||||
let keys: Vec<String> = stmt
|
||||
.query_map(params![plugin_id], |row| row.get(0))
|
||||
.map_err(|e| e.to_string())?
|
||||
.filter_map(|r| r.ok())
|
||||
.collect();
|
||||
|
||||
Ok(serde_json::json!({ "keys": keys }))
|
||||
}
|
||||
|
||||
/// 插件配置获取
|
||||
#[tauri::command]
|
||||
pub async fn plugin_config_get(plugin_id: String) -> Result<serde_json::Value, String> {
|
||||
get_oauth_plugin_config(plugin_id)
|
||||
.await
|
||||
.map(|c| serde_json::json!({ "config": c }))
|
||||
}
|
||||
|
||||
/// 插件配置设置
|
||||
#[tauri::command]
|
||||
pub async fn plugin_config_set(
|
||||
db: tauri::State<'_, DbConnection>,
|
||||
plugin_id: String,
|
||||
config: serde_json::Value,
|
||||
) -> Result<(), String> {
|
||||
update_oauth_plugin_config(db, plugin_id, config).await
|
||||
}
|
||||
|
||||
/// 读取插件 UI 文件
|
||||
/// 用于前端动态加载插件的 React 组件
|
||||
#[tauri::command]
|
||||
pub async fn read_plugin_ui_file(path: String) -> Result<String, String> {
|
||||
use std::fs;
|
||||
|
||||
// 安全检查:确保路径在插件目录内
|
||||
let path = std::path::PathBuf::from(&path);
|
||||
|
||||
// 读取文件内容
|
||||
fs::read_to_string(&path).map_err(|e| format!("读取插件 UI 文件失败: {}", e))
|
||||
}
|
||||
@@ -0,0 +1,493 @@
|
||||
//! 模型编排器 Tauri 命令
|
||||
//!
|
||||
//! 提供前端访问模型编排器的接口。
|
||||
|
||||
use crate::database::dao::provider_pool::ProviderPoolDao;
|
||||
use crate::database::DbConnection;
|
||||
use crate::orchestrator::{
|
||||
get_global_orchestrator, init_global_orchestrator, AvailableModel, CredentialInfo,
|
||||
OrchestratorConfig, PoolStats, ProviderType, SelectionContext, SelectionResult, ServiceTier,
|
||||
StrategyInfo, TaskHint,
|
||||
};
|
||||
use serde::{Deserialize, Serialize};
|
||||
use tauri::State;
|
||||
use tokio::sync::RwLock;
|
||||
|
||||
/// 编排器状态
|
||||
pub struct OrchestratorState {
|
||||
initialized: RwLock<bool>,
|
||||
}
|
||||
|
||||
impl OrchestratorState {
|
||||
pub fn new() -> Self {
|
||||
Self {
|
||||
initialized: RwLock::new(false),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
impl Default for OrchestratorState {
|
||||
fn default() -> Self {
|
||||
Self::new()
|
||||
}
|
||||
}
|
||||
|
||||
// ============================================================================
|
||||
// 初始化命令
|
||||
// ============================================================================
|
||||
|
||||
/// 初始化编排器
|
||||
#[tauri::command]
|
||||
pub async fn init_orchestrator(
|
||||
state: State<'_, OrchestratorState>,
|
||||
db: State<'_, DbConnection>,
|
||||
) -> Result<(), String> {
|
||||
let mut initialized = state.initialized.write().await;
|
||||
if *initialized {
|
||||
return Ok(());
|
||||
}
|
||||
|
||||
let orchestrator = init_global_orchestrator();
|
||||
*initialized = true;
|
||||
|
||||
// 从数据库加载凭证并同步到 orchestrator
|
||||
let credentials = {
|
||||
let conn = db
|
||||
.lock()
|
||||
.map_err(|e| format!("获取数据库连接失败: {}", e))?;
|
||||
ProviderPoolDao::get_all(&conn).map_err(|e| format!("获取凭证列表失败: {}", e))?
|
||||
};
|
||||
|
||||
// 转换凭证格式
|
||||
let cred_infos: Vec<CredentialInfo> = credentials
|
||||
.iter()
|
||||
.filter(|c| !c.is_disabled && c.is_healthy)
|
||||
.map(|c| {
|
||||
// 从 credential 中提取支持的模型列表
|
||||
let supported_models = extract_supported_models(&c.credential);
|
||||
|
||||
CredentialInfo {
|
||||
id: c.uuid.clone(),
|
||||
provider_type: map_pool_provider_type(&c.provider_type.to_string()),
|
||||
supported_models,
|
||||
is_healthy: c.is_healthy,
|
||||
current_load: None,
|
||||
}
|
||||
})
|
||||
.collect();
|
||||
|
||||
if !cred_infos.is_empty() {
|
||||
orchestrator.update_credentials(cred_infos).await;
|
||||
tracing::info!("已从凭证池同步 {} 个凭证到编排器", credentials.len());
|
||||
}
|
||||
|
||||
tracing::info!("模型编排器已初始化");
|
||||
Ok(())
|
||||
}
|
||||
|
||||
/// 从 credential 提取支持的模型列表
|
||||
fn extract_supported_models(
|
||||
credential: &crate::models::provider_pool_model::CredentialData,
|
||||
) -> Vec<String> {
|
||||
use crate::models::provider_pool_model::CredentialData;
|
||||
|
||||
match credential {
|
||||
CredentialData::ClaudeKey { .. } | CredentialData::ClaudeOAuth { .. } => {
|
||||
vec![
|
||||
"claude-opus-4-5-20251101".to_string(),
|
||||
"claude-sonnet-4-5-20250929".to_string(),
|
||||
"claude-sonnet-4-20250514".to_string(),
|
||||
"claude-3-5-haiku-20241022".to_string(),
|
||||
]
|
||||
}
|
||||
CredentialData::OpenAIKey { .. } => {
|
||||
vec![
|
||||
"gpt-4o".to_string(),
|
||||
"gpt-4o-mini".to_string(),
|
||||
"gpt-4-turbo".to_string(),
|
||||
"o1".to_string(),
|
||||
"o1-mini".to_string(),
|
||||
]
|
||||
}
|
||||
CredentialData::GeminiOAuth { .. } => {
|
||||
vec![
|
||||
"gemini-2.0-flash-exp".to_string(),
|
||||
"gemini-1.5-pro".to_string(),
|
||||
"gemini-1.5-flash".to_string(),
|
||||
]
|
||||
}
|
||||
CredentialData::GeminiApiKey {
|
||||
excluded_models, ..
|
||||
} => {
|
||||
let all_models = vec![
|
||||
"gemini-2.0-flash-exp".to_string(),
|
||||
"gemini-1.5-pro".to_string(),
|
||||
"gemini-1.5-flash".to_string(),
|
||||
];
|
||||
all_models
|
||||
.into_iter()
|
||||
.filter(|m| !excluded_models.contains(m))
|
||||
.collect()
|
||||
}
|
||||
CredentialData::KiroOAuth { .. } => {
|
||||
vec![
|
||||
"claude-sonnet-4-5-20250929".to_string(),
|
||||
"claude-sonnet-4-20250514".to_string(),
|
||||
]
|
||||
}
|
||||
CredentialData::CodexOAuth { .. } => {
|
||||
vec!["codex-mini-latest".to_string()]
|
||||
}
|
||||
CredentialData::QwenOAuth { .. } => {
|
||||
vec![
|
||||
"qwen-max".to_string(),
|
||||
"qwen-plus".to_string(),
|
||||
"qwen-turbo".to_string(),
|
||||
]
|
||||
}
|
||||
CredentialData::AntigravityOAuth { .. } => {
|
||||
vec![
|
||||
"gemini-claude-sonnet-4-5".to_string(),
|
||||
"gemini-claude-sonnet-4-5-thinking".to_string(),
|
||||
"gemini-claude-opus-4-5-thinking".to_string(),
|
||||
]
|
||||
}
|
||||
_ => vec![],
|
||||
}
|
||||
}
|
||||
|
||||
/// 映射 PoolProviderType 到 orchestrator 的 ProviderType
|
||||
fn map_pool_provider_type(pool_type: &str) -> ProviderType {
|
||||
match pool_type.to_lowercase().as_str() {
|
||||
"claude" | "claude_oauth" => ProviderType::Anthropic,
|
||||
"openai" => ProviderType::OpenAI,
|
||||
"gemini" | "gemini_api_key" | "gemini_oauth" => ProviderType::Google,
|
||||
"kiro" => ProviderType::Kiro,
|
||||
"codex" => ProviderType::OpenAI,
|
||||
"qwen" => ProviderType::Custom,
|
||||
_ => ProviderType::Custom,
|
||||
}
|
||||
}
|
||||
|
||||
/// 获取编排器配置
|
||||
#[tauri::command]
|
||||
pub async fn get_orchestrator_config() -> Result<OrchestratorConfig, String> {
|
||||
let orchestrator = get_global_orchestrator().ok_or("编排器未初始化")?;
|
||||
|
||||
Ok(orchestrator.get_config().await)
|
||||
}
|
||||
|
||||
/// 更新编排器配置
|
||||
#[tauri::command]
|
||||
pub async fn update_orchestrator_config(config: OrchestratorConfig) -> Result<(), String> {
|
||||
let orchestrator = get_global_orchestrator().ok_or("编排器未初始化")?;
|
||||
|
||||
orchestrator.update_config(config).await;
|
||||
Ok(())
|
||||
}
|
||||
|
||||
// ============================================================================
|
||||
// 模型池命令
|
||||
// ============================================================================
|
||||
|
||||
/// 获取模型池统计
|
||||
#[tauri::command]
|
||||
pub async fn get_pool_stats() -> Result<PoolStats, String> {
|
||||
let orchestrator = get_global_orchestrator().ok_or("编排器未初始化")?;
|
||||
|
||||
Ok(orchestrator.get_pool_stats().await)
|
||||
}
|
||||
|
||||
/// 获取指定等级的模型列表
|
||||
#[tauri::command]
|
||||
pub async fn get_tier_models(tier: String) -> Result<Vec<AvailableModel>, String> {
|
||||
let orchestrator = get_global_orchestrator().ok_or("编排器未初始化")?;
|
||||
|
||||
let service_tier =
|
||||
ServiceTier::from_str(&tier).ok_or_else(|| format!("无效的服务等级: {}", tier))?;
|
||||
|
||||
Ok(orchestrator.get_models(service_tier).await)
|
||||
}
|
||||
|
||||
/// 获取所有可用模型
|
||||
#[tauri::command]
|
||||
pub async fn get_all_models() -> Result<Vec<AvailableModel>, String> {
|
||||
let orchestrator = get_global_orchestrator().ok_or("编排器未初始化")?;
|
||||
|
||||
Ok(orchestrator.get_all_models().await)
|
||||
}
|
||||
|
||||
// ============================================================================
|
||||
// 凭证管理命令
|
||||
// ============================================================================
|
||||
|
||||
/// 凭证信息请求
|
||||
#[derive(Debug, Clone, Serialize, Deserialize)]
|
||||
pub struct CredentialInfoRequest {
|
||||
pub id: String,
|
||||
pub provider_type: String,
|
||||
pub supported_models: Vec<String>,
|
||||
pub is_healthy: bool,
|
||||
pub current_load: Option<u8>,
|
||||
}
|
||||
|
||||
impl From<CredentialInfoRequest> for CredentialInfo {
|
||||
fn from(req: CredentialInfoRequest) -> Self {
|
||||
CredentialInfo {
|
||||
id: req.id,
|
||||
provider_type: ProviderType::from_str(&req.provider_type)
|
||||
.unwrap_or(ProviderType::Custom),
|
||||
supported_models: req.supported_models,
|
||||
is_healthy: req.is_healthy,
|
||||
current_load: req.current_load,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
/// 更新凭证列表
|
||||
#[tauri::command]
|
||||
pub async fn update_orchestrator_credentials(
|
||||
credentials: Vec<CredentialInfoRequest>,
|
||||
) -> Result<(), String> {
|
||||
let orchestrator = get_global_orchestrator().ok_or("编排器未初始化")?;
|
||||
|
||||
let creds: Vec<CredentialInfo> = credentials.into_iter().map(Into::into).collect();
|
||||
orchestrator.update_credentials(creds).await;
|
||||
|
||||
Ok(())
|
||||
}
|
||||
|
||||
/// 添加凭证
|
||||
#[tauri::command]
|
||||
pub async fn add_orchestrator_credential(credential: CredentialInfoRequest) -> Result<(), String> {
|
||||
let orchestrator = get_global_orchestrator().ok_or("编排器未初始化")?;
|
||||
|
||||
orchestrator.add_credential(credential.into()).await;
|
||||
Ok(())
|
||||
}
|
||||
|
||||
/// 移除凭证
|
||||
#[tauri::command]
|
||||
pub async fn remove_orchestrator_credential(credential_id: String) -> Result<(), String> {
|
||||
let orchestrator = get_global_orchestrator().ok_or("编排器未初始化")?;
|
||||
|
||||
orchestrator.remove_credential(&credential_id).await;
|
||||
Ok(())
|
||||
}
|
||||
|
||||
/// 标记凭证为不健康
|
||||
#[tauri::command]
|
||||
pub async fn mark_credential_unhealthy(
|
||||
model_id: String,
|
||||
credential_id: String,
|
||||
) -> Result<(), String> {
|
||||
let orchestrator = get_global_orchestrator().ok_or("编排器未初始化")?;
|
||||
|
||||
orchestrator.mark_unhealthy(&model_id, &credential_id).await;
|
||||
Ok(())
|
||||
}
|
||||
|
||||
/// 标记凭证为健康
|
||||
#[tauri::command]
|
||||
pub async fn mark_credential_healthy(credential_id: String) -> Result<(), String> {
|
||||
let orchestrator = get_global_orchestrator().ok_or("编排器未初始化")?;
|
||||
|
||||
orchestrator.mark_healthy(&credential_id).await;
|
||||
Ok(())
|
||||
}
|
||||
|
||||
/// 更新凭证负载
|
||||
#[tauri::command]
|
||||
pub async fn update_credential_load(credential_id: String, load: u8) -> Result<(), String> {
|
||||
let orchestrator = get_global_orchestrator().ok_or("编排器未初始化")?;
|
||||
|
||||
orchestrator.update_load(&credential_id, load).await;
|
||||
Ok(())
|
||||
}
|
||||
|
||||
// ============================================================================
|
||||
// 模型选择命令
|
||||
// ============================================================================
|
||||
|
||||
/// 选择请求
|
||||
#[derive(Debug, Clone, Serialize, Deserialize)]
|
||||
pub struct SelectionRequest {
|
||||
pub tier: String,
|
||||
pub task_hint: Option<String>,
|
||||
pub requires_vision: Option<bool>,
|
||||
pub requires_tools: Option<bool>,
|
||||
pub preferred_provider: Option<String>,
|
||||
pub excluded_models: Option<Vec<String>>,
|
||||
pub strategy_id: Option<String>,
|
||||
}
|
||||
|
||||
/// 选择模型
|
||||
#[tauri::command]
|
||||
pub async fn select_model(request: SelectionRequest) -> Result<SelectionResult, String> {
|
||||
let orchestrator = get_global_orchestrator().ok_or("编排器未初始化")?;
|
||||
|
||||
let tier = ServiceTier::from_str(&request.tier)
|
||||
.ok_or_else(|| format!("无效的服务等级: {}", request.tier))?;
|
||||
|
||||
let mut ctx = SelectionContext::new(tier);
|
||||
|
||||
if let Some(hint) = &request.task_hint {
|
||||
ctx.task_hint = match hint.to_lowercase().as_str() {
|
||||
"coding" => Some(TaskHint::Coding),
|
||||
"writing" => Some(TaskHint::Writing),
|
||||
"analysis" => Some(TaskHint::Analysis),
|
||||
"chat" => Some(TaskHint::Chat),
|
||||
"translation" => Some(TaskHint::Translation),
|
||||
"summarization" => Some(TaskHint::Summarization),
|
||||
"math" => Some(TaskHint::Math),
|
||||
_ => Some(TaskHint::Other),
|
||||
};
|
||||
}
|
||||
|
||||
if let Some(vision) = request.requires_vision {
|
||||
ctx.requires_vision = vision;
|
||||
}
|
||||
|
||||
if let Some(tools) = request.requires_tools {
|
||||
ctx.requires_tools = tools;
|
||||
}
|
||||
|
||||
if let Some(provider) = request.preferred_provider {
|
||||
ctx.preferred_provider = Some(provider);
|
||||
}
|
||||
|
||||
if let Some(excluded) = request.excluded_models {
|
||||
ctx.excluded_models = excluded;
|
||||
}
|
||||
|
||||
let result = if let Some(strategy_id) = &request.strategy_id {
|
||||
orchestrator.select_with_strategy(strategy_id, &ctx).await
|
||||
} else {
|
||||
orchestrator.select(&ctx).await
|
||||
};
|
||||
|
||||
result.map_err(|e| e.to_string())
|
||||
}
|
||||
|
||||
/// 快速选择模型
|
||||
#[tauri::command]
|
||||
pub async fn quick_select_model() -> Result<SelectionResult, String> {
|
||||
let orchestrator = get_global_orchestrator().ok_or("编排器未初始化")?;
|
||||
|
||||
orchestrator.quick_select().await.map_err(|e| e.to_string())
|
||||
}
|
||||
|
||||
/// 为特定任务选择模型
|
||||
#[tauri::command]
|
||||
pub async fn select_model_for_task(tier: String, task: String) -> Result<SelectionResult, String> {
|
||||
let orchestrator = get_global_orchestrator().ok_or("编排器未初始化")?;
|
||||
|
||||
let service_tier =
|
||||
ServiceTier::from_str(&tier).ok_or_else(|| format!("无效的服务等级: {}", tier))?;
|
||||
|
||||
let task_hint = match task.to_lowercase().as_str() {
|
||||
"coding" => TaskHint::Coding,
|
||||
"writing" => TaskHint::Writing,
|
||||
"analysis" => TaskHint::Analysis,
|
||||
"chat" => TaskHint::Chat,
|
||||
"translation" => TaskHint::Translation,
|
||||
"summarization" => TaskHint::Summarization,
|
||||
"math" => TaskHint::Math,
|
||||
_ => TaskHint::Other,
|
||||
};
|
||||
|
||||
orchestrator
|
||||
.select_for_task(service_tier, task_hint)
|
||||
.await
|
||||
.map_err(|e| e.to_string())
|
||||
}
|
||||
|
||||
// ============================================================================
|
||||
// 策略命令
|
||||
// ============================================================================
|
||||
|
||||
/// 列出所有可用策略
|
||||
#[tauri::command]
|
||||
pub async fn list_strategies() -> Result<Vec<StrategyInfo>, String> {
|
||||
let orchestrator = get_global_orchestrator().ok_or("编排器未初始化")?;
|
||||
|
||||
Ok(orchestrator.list_strategies().await)
|
||||
}
|
||||
|
||||
/// 获取服务等级列表
|
||||
#[tauri::command]
|
||||
pub fn list_service_tiers() -> Vec<ServiceTierInfo> {
|
||||
ServiceTier::all()
|
||||
.iter()
|
||||
.map(|t| ServiceTierInfo {
|
||||
id: format!("{:?}", t).to_lowercase(),
|
||||
display_name: t.display_name().to_string(),
|
||||
description: t.description().to_string(),
|
||||
level: t.level(),
|
||||
})
|
||||
.collect()
|
||||
}
|
||||
|
||||
/// 服务等级信息
|
||||
#[derive(Debug, Clone, Serialize, Deserialize)]
|
||||
pub struct ServiceTierInfo {
|
||||
pub id: String,
|
||||
pub display_name: String,
|
||||
pub description: String,
|
||||
pub level: u8,
|
||||
}
|
||||
|
||||
/// 获取任务类型列表
|
||||
#[tauri::command]
|
||||
pub fn list_task_hints() -> Vec<TaskHintInfo> {
|
||||
vec![
|
||||
TaskHintInfo {
|
||||
id: "coding".to_string(),
|
||||
display_name: "代码".to_string(),
|
||||
description: "代码生成、编辑、调试".to_string(),
|
||||
},
|
||||
TaskHintInfo {
|
||||
id: "writing".to_string(),
|
||||
display_name: "写作".to_string(),
|
||||
description: "文章、报告、创意写作".to_string(),
|
||||
},
|
||||
TaskHintInfo {
|
||||
id: "analysis".to_string(),
|
||||
display_name: "分析".to_string(),
|
||||
description: "数据分析、推理、研究".to_string(),
|
||||
},
|
||||
TaskHintInfo {
|
||||
id: "chat".to_string(),
|
||||
display_name: "对话".to_string(),
|
||||
description: "日常对话、问答".to_string(),
|
||||
},
|
||||
TaskHintInfo {
|
||||
id: "translation".to_string(),
|
||||
display_name: "翻译".to_string(),
|
||||
description: "语言翻译".to_string(),
|
||||
},
|
||||
TaskHintInfo {
|
||||
id: "summarization".to_string(),
|
||||
display_name: "摘要".to_string(),
|
||||
description: "文本摘要、总结".to_string(),
|
||||
},
|
||||
TaskHintInfo {
|
||||
id: "math".to_string(),
|
||||
display_name: "数学".to_string(),
|
||||
description: "数学计算、推理".to_string(),
|
||||
},
|
||||
TaskHintInfo {
|
||||
id: "other".to_string(),
|
||||
display_name: "其他".to_string(),
|
||||
description: "其他任务".to_string(),
|
||||
},
|
||||
]
|
||||
}
|
||||
|
||||
/// 任务类型信息
|
||||
#[derive(Debug, Clone, Serialize, Deserialize)]
|
||||
pub struct TaskHintInfo {
|
||||
pub id: String,
|
||||
pub display_name: String,
|
||||
pub description: String,
|
||||
}
|
||||
@@ -64,6 +64,11 @@ impl ProtocolSelector {
|
||||
PoolProviderType::Codex => Protocol::OpenAI, // Codex uses OpenAI protocol
|
||||
PoolProviderType::ClaudeOAuth => Protocol::Anthropic, // Claude OAuth uses Anthropic protocol
|
||||
PoolProviderType::IFlow => Protocol::OpenAI, // iFlow uses OpenAI protocol
|
||||
// API Key Provider 类型
|
||||
PoolProviderType::Anthropic => Protocol::Anthropic,
|
||||
PoolProviderType::AzureOpenai => Protocol::OpenAI,
|
||||
PoolProviderType::AwsBedrock => Protocol::Anthropic,
|
||||
PoolProviderType::Ollama => Protocol::OpenAI,
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
@@ -1,23 +1,67 @@
|
||||
//! 凭证池管理模块
|
||||
//!
|
||||
//! 提供多凭证管理、负载均衡和健康检查功能
|
||||
//!
|
||||
//! ## 模块结构
|
||||
//!
|
||||
//! - `types` - 凭证相关类型定义
|
||||
//! - `pool` - 凭证池管理
|
||||
//! - `balancer` - 负载均衡策略
|
||||
//! - `health` - 健康检查
|
||||
//! - `quota` - 配额管理
|
||||
//! - `sync` - 数据库同步
|
||||
//! - `plugin` - OAuth Provider 插件 Trait
|
||||
//! - `registry` - 插件注册表
|
||||
//! - `oauth_plugin_loader` - OAuth Provider 插件加载器
|
||||
//! - `sdk` - ProxyCast Plugin SDK
|
||||
//! - `risk` - 风控模块(限流检测、冷却期管理)
|
||||
//! - `unified` - 统一凭证管理器
|
||||
|
||||
mod balancer;
|
||||
mod health;
|
||||
pub mod oauth_plugin_loader;
|
||||
pub mod plugin;
|
||||
mod pool;
|
||||
mod quota;
|
||||
pub mod registry;
|
||||
pub mod risk;
|
||||
pub mod sdk;
|
||||
mod sync;
|
||||
mod types;
|
||||
mod unified;
|
||||
|
||||
pub use balancer::{BalanceStrategy, CooldownInfo, CredentialSelection, LoadBalancer};
|
||||
pub use health::{HealthCheckConfig, HealthCheckResult, HealthChecker, HealthStatus};
|
||||
pub use oauth_plugin_loader::{
|
||||
BinaryManifest, ExternalOAuthPlugin, OAuthPluginLoader, OAuthPluginManifest, ProviderManifest,
|
||||
UiManifest,
|
||||
};
|
||||
pub use plugin::{
|
||||
AcquiredCredential, AuthTypeInfo, CredentialCategory, CredentialConfig,
|
||||
CredentialProviderPlugin, ModelFamily, ModelInfo, OAuthPluginError, OAuthPluginInfo,
|
||||
OAuthPluginResult, PluginInstance, ProviderError, ProviderErrorType, StandardProtocol,
|
||||
TokenRefreshResult, UsageResult, ValidationResult,
|
||||
};
|
||||
pub use pool::{CredentialPool, PoolError, PoolStatus};
|
||||
pub use quota::{
|
||||
create_shared_quota_manager, start_quota_cleanup_task, AllCredentialsExhaustedError,
|
||||
QuotaAutoSwitchResult, QuotaExceededRecord, QuotaManager,
|
||||
};
|
||||
pub use registry::{
|
||||
get_global_registry, init_global_registry, CredentialProviderRegistry, PluginSource,
|
||||
PluginState, PluginUpdate,
|
||||
};
|
||||
pub use risk::{CooldownConfig, RateLimitEvent, RateLimitStats, RiskController, RiskLevel};
|
||||
pub use sdk::{
|
||||
DatabaseCallback, HttpRequestOptions, HttpResponse, JsonRpcError, JsonRpcRequest,
|
||||
JsonRpcResponse, PluginPermission, PluginSdkContext, QueryResult, SdkError, SdkMethodHandler,
|
||||
SdkResult,
|
||||
};
|
||||
pub use sync::{CredentialSyncService, SyncError};
|
||||
pub use types::{Credential, CredentialData, CredentialStats, CredentialStatus};
|
||||
pub use unified::{
|
||||
get_global_unified_manager, init_global_unified_manager, UnifiedCredentialManager,
|
||||
};
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests;
|
||||
|
||||
@@ -0,0 +1,723 @@
|
||||
//! OAuth Provider 插件加载器
|
||||
//!
|
||||
//! 负责从外部目录加载 OAuth Provider 插件。
|
||||
//! 与通用插件加载器不同,此加载器专门处理 oauth_provider 类型的插件。
|
||||
|
||||
use super::plugin::{
|
||||
AcquiredCredential, AuthTypeInfo, CredentialCategory, CredentialConfig,
|
||||
CredentialProviderPlugin, ModelFamily, ModelInfo, OAuthPluginError, OAuthPluginResult,
|
||||
ProviderError, StandardProtocol, TokenRefreshResult, UsageResult, ValidationResult,
|
||||
};
|
||||
use super::registry::CredentialProviderRegistry;
|
||||
use async_trait::async_trait;
|
||||
use serde::{Deserialize, Serialize};
|
||||
use std::collections::HashMap;
|
||||
use std::path::{Path, PathBuf};
|
||||
use std::process::Stdio;
|
||||
use std::sync::Arc;
|
||||
use tokio::fs;
|
||||
use tokio::process::{Child, Command};
|
||||
use tokio::sync::Mutex;
|
||||
use tracing::{debug, info, warn};
|
||||
|
||||
/// OAuth Provider 插件的 plugin.json 结构
|
||||
#[derive(Debug, Clone, Serialize, Deserialize)]
|
||||
pub struct OAuthPluginManifest {
|
||||
/// 插件名称
|
||||
pub name: String,
|
||||
/// 版本
|
||||
pub version: String,
|
||||
/// 描述
|
||||
#[serde(default)]
|
||||
pub description: String,
|
||||
/// 作者
|
||||
#[serde(default)]
|
||||
pub author: Option<String>,
|
||||
/// 主页
|
||||
#[serde(default)]
|
||||
pub homepage: Option<String>,
|
||||
/// 许可证
|
||||
#[serde(default)]
|
||||
pub license: Option<String>,
|
||||
/// 插件类型(必须是 "oauth_provider")
|
||||
pub plugin_type: String,
|
||||
/// 入口(二进制名称)
|
||||
pub entry: String,
|
||||
/// 最低 ProxyCast 版本
|
||||
#[serde(default)]
|
||||
pub min_proxycast_version: Option<String>,
|
||||
/// Provider 配置
|
||||
pub provider: ProviderManifest,
|
||||
/// 二进制配置
|
||||
#[serde(default)]
|
||||
pub binary: Option<BinaryManifest>,
|
||||
/// UI 配置
|
||||
#[serde(default)]
|
||||
pub ui: Option<UiManifest>,
|
||||
}
|
||||
|
||||
/// Provider 配置
|
||||
#[derive(Debug, Clone, Serialize, Deserialize)]
|
||||
pub struct ProviderManifest {
|
||||
/// Provider ID
|
||||
pub id: String,
|
||||
/// 显示名称
|
||||
pub display_name: String,
|
||||
/// 目标协议
|
||||
pub target_protocol: String,
|
||||
/// 支持的模型模式
|
||||
#[serde(default)]
|
||||
pub supported_models: Vec<String>,
|
||||
/// 认证类型
|
||||
#[serde(default)]
|
||||
pub auth_types: Vec<String>,
|
||||
/// 凭证 Schema
|
||||
#[serde(default)]
|
||||
pub credential_schemas: HashMap<String, serde_json::Value>,
|
||||
}
|
||||
|
||||
/// 二进制配置
|
||||
#[derive(Debug, Clone, Serialize, Deserialize)]
|
||||
pub struct BinaryManifest {
|
||||
/// 二进制名称
|
||||
pub binary_name: String,
|
||||
/// GitHub owner
|
||||
pub github_owner: String,
|
||||
/// GitHub repo
|
||||
pub github_repo: String,
|
||||
/// 平台二进制映射
|
||||
pub platform_binaries: HashMap<String, String>,
|
||||
/// 校验文件
|
||||
#[serde(default)]
|
||||
pub checksum_file: Option<String>,
|
||||
}
|
||||
|
||||
/// UI 配置
|
||||
#[derive(Debug, Clone, Serialize, Deserialize)]
|
||||
pub struct UiManifest {
|
||||
/// 显示位置
|
||||
#[serde(default)]
|
||||
pub surfaces: Vec<String>,
|
||||
/// 图标
|
||||
#[serde(default)]
|
||||
pub icon: Option<String>,
|
||||
/// 标题
|
||||
#[serde(default)]
|
||||
pub title: Option<String>,
|
||||
/// UI 入口文件
|
||||
#[serde(default)]
|
||||
pub entry: Option<String>,
|
||||
/// 样式文件
|
||||
#[serde(default)]
|
||||
pub styles: Option<String>,
|
||||
/// 默认宽度
|
||||
#[serde(default)]
|
||||
pub default_width: Option<u32>,
|
||||
/// 默认高度
|
||||
#[serde(default)]
|
||||
pub default_height: Option<u32>,
|
||||
/// 权限列表
|
||||
#[serde(default)]
|
||||
pub permissions: Vec<String>,
|
||||
}
|
||||
|
||||
/// OAuth Provider 插件加载器
|
||||
pub struct OAuthPluginLoader {
|
||||
/// 插件目录
|
||||
plugins_dir: PathBuf,
|
||||
}
|
||||
|
||||
impl OAuthPluginLoader {
|
||||
/// 创建新的加载器
|
||||
pub fn new(plugins_dir: PathBuf) -> Self {
|
||||
Self { plugins_dir }
|
||||
}
|
||||
|
||||
/// 默认插件目录
|
||||
pub fn default_plugins_dir() -> PathBuf {
|
||||
dirs::config_dir()
|
||||
.unwrap_or_else(|| PathBuf::from("."))
|
||||
.join("proxycast")
|
||||
.join("plugins")
|
||||
}
|
||||
|
||||
/// 使用默认配置创建
|
||||
pub fn with_defaults() -> Self {
|
||||
Self::new(Self::default_plugins_dir())
|
||||
}
|
||||
|
||||
/// 确保插件目录存在
|
||||
pub async fn ensure_plugins_dir(&self) -> OAuthPluginResult<()> {
|
||||
if !self.plugins_dir.exists() {
|
||||
fs::create_dir_all(&self.plugins_dir)
|
||||
.await
|
||||
.map_err(|e| OAuthPluginError::IoError(e))?;
|
||||
}
|
||||
Ok(())
|
||||
}
|
||||
|
||||
/// 扫描所有 OAuth Provider 插件
|
||||
pub async fn scan(&self) -> OAuthPluginResult<Vec<PathBuf>> {
|
||||
self.ensure_plugins_dir().await?;
|
||||
|
||||
let mut plugins = Vec::new();
|
||||
let mut entries = fs::read_dir(&self.plugins_dir)
|
||||
.await
|
||||
.map_err(|e| OAuthPluginError::IoError(e))?;
|
||||
|
||||
while let Some(entry) = entries
|
||||
.next_entry()
|
||||
.await
|
||||
.map_err(|e| OAuthPluginError::IoError(e))?
|
||||
{
|
||||
let path = entry.path();
|
||||
|
||||
// 检查是否是目录且包含 plugin.json
|
||||
if path.is_dir() && path.join("plugin.json").exists() {
|
||||
// 读取 plugin.json 检查类型
|
||||
let manifest_path = path.join("plugin.json");
|
||||
if let Ok(content) = fs::read_to_string(&manifest_path).await {
|
||||
if let Ok(manifest) = serde_json::from_str::<OAuthPluginManifest>(&content) {
|
||||
if manifest.plugin_type == "oauth_provider" {
|
||||
plugins.push(path);
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
Ok(plugins)
|
||||
}
|
||||
|
||||
/// 加载插件清单
|
||||
pub async fn load_manifest(&self, plugin_dir: &Path) -> OAuthPluginResult<OAuthPluginManifest> {
|
||||
let manifest_path = plugin_dir.join("plugin.json");
|
||||
|
||||
let content = fs::read_to_string(&manifest_path)
|
||||
.await
|
||||
.map_err(|e| OAuthPluginError::InitError(format!("无法读取 plugin.json: {}", e)))?;
|
||||
|
||||
let manifest: OAuthPluginManifest = serde_json::from_str(&content)?;
|
||||
|
||||
// 验证插件类型
|
||||
if manifest.plugin_type != "oauth_provider" {
|
||||
return Err(OAuthPluginError::InitError(format!(
|
||||
"无效的插件类型: {} (期望 oauth_provider)",
|
||||
manifest.plugin_type
|
||||
)));
|
||||
}
|
||||
|
||||
Ok(manifest)
|
||||
}
|
||||
|
||||
/// 加载单个插件
|
||||
pub async fn load(
|
||||
&self,
|
||||
plugin_dir: &Path,
|
||||
) -> OAuthPluginResult<Arc<dyn CredentialProviderPlugin>> {
|
||||
let manifest = self.load_manifest(plugin_dir).await?;
|
||||
|
||||
info!(
|
||||
"Loading OAuth provider plugin: {} v{}",
|
||||
manifest.provider.id, manifest.version
|
||||
);
|
||||
|
||||
// 查找二进制文件
|
||||
let binary_path = self.find_binary(plugin_dir, &manifest)?;
|
||||
|
||||
// 加载配置
|
||||
let config_path = plugin_dir.join("config.json");
|
||||
let config = if config_path.exists() {
|
||||
let content = fs::read_to_string(&config_path)
|
||||
.await
|
||||
.map_err(|e| OAuthPluginError::IoError(e))?;
|
||||
serde_json::from_str(&content).unwrap_or_default()
|
||||
} else {
|
||||
serde_json::json!({})
|
||||
};
|
||||
|
||||
// 创建外部插件实例
|
||||
let plugin = ExternalOAuthPlugin::new(manifest, binary_path, config);
|
||||
|
||||
Ok(Arc::new(plugin))
|
||||
}
|
||||
|
||||
/// 查找二进制文件
|
||||
fn find_binary(
|
||||
&self,
|
||||
plugin_dir: &Path,
|
||||
manifest: &OAuthPluginManifest,
|
||||
) -> OAuthPluginResult<PathBuf> {
|
||||
let bin_dir = plugin_dir.join("bin");
|
||||
|
||||
// 获取当前平台的二进制名称
|
||||
let platform_key = get_platform_key();
|
||||
|
||||
let binary_name = if let Some(binary) = &manifest.binary {
|
||||
binary
|
||||
.platform_binaries
|
||||
.get(&platform_key)
|
||||
.cloned()
|
||||
.unwrap_or_else(|| manifest.entry.clone())
|
||||
} else {
|
||||
manifest.entry.clone()
|
||||
};
|
||||
|
||||
// 尝试几个可能的位置
|
||||
let candidates = vec![
|
||||
bin_dir.join(&binary_name),
|
||||
plugin_dir.join(&binary_name),
|
||||
plugin_dir.join("bin").join(&manifest.entry),
|
||||
];
|
||||
|
||||
for path in candidates {
|
||||
if path.exists() {
|
||||
return Ok(path);
|
||||
}
|
||||
}
|
||||
|
||||
Err(OAuthPluginError::InitError(format!(
|
||||
"找不到二进制文件: {} (平台: {})",
|
||||
binary_name, platform_key
|
||||
)))
|
||||
}
|
||||
|
||||
/// 加载所有插件到注册表
|
||||
pub async fn load_all(
|
||||
&self,
|
||||
registry: &CredentialProviderRegistry,
|
||||
) -> OAuthPluginResult<Vec<String>> {
|
||||
let plugin_dirs = self.scan().await?;
|
||||
let mut loaded = Vec::new();
|
||||
|
||||
for dir in plugin_dirs {
|
||||
match self.load(&dir).await {
|
||||
Ok(plugin) => {
|
||||
let id = plugin.id().to_string();
|
||||
if let Err(e) = registry.register(plugin).await {
|
||||
warn!("注册插件失败 {}: {}", id, e);
|
||||
} else {
|
||||
loaded.push(id);
|
||||
}
|
||||
}
|
||||
Err(e) => {
|
||||
warn!("加载插件失败 {:?}: {}", dir, e);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
Ok(loaded)
|
||||
}
|
||||
|
||||
/// 获取插件目录
|
||||
pub fn plugins_dir(&self) -> &Path {
|
||||
&self.plugins_dir
|
||||
}
|
||||
}
|
||||
|
||||
/// 获取当前平台的 key
|
||||
fn get_platform_key() -> String {
|
||||
match (std::env::consts::ARCH, std::env::consts::OS) {
|
||||
("aarch64", "macos") => "macos-arm64".to_string(),
|
||||
("x86_64", "macos") => "macos-x64".to_string(),
|
||||
("x86_64", "linux") => "linux-x64".to_string(),
|
||||
("aarch64", "linux") => "linux-arm64".to_string(),
|
||||
("x86_64", "windows") => "windows-x64".to_string(),
|
||||
(arch, os) => format!("{}-{}", os, arch),
|
||||
}
|
||||
}
|
||||
|
||||
// ============================================================================
|
||||
// 外部 OAuth 插件(通过二进制调用)
|
||||
// ============================================================================
|
||||
|
||||
/// 外部 OAuth 插件
|
||||
///
|
||||
/// 通过调用外部二进制实现 CredentialProviderPlugin trait。
|
||||
/// 使用 JSON-RPC 或 stdin/stdout 通信。
|
||||
pub struct ExternalOAuthPlugin {
|
||||
/// 插件清单
|
||||
manifest: OAuthPluginManifest,
|
||||
/// 二进制路径
|
||||
binary_path: PathBuf,
|
||||
/// 插件配置
|
||||
config: serde_json::Value,
|
||||
/// 进程句柄
|
||||
process: Mutex<Option<Child>>,
|
||||
}
|
||||
|
||||
impl ExternalOAuthPlugin {
|
||||
/// 创建新的外部插件
|
||||
pub fn new(
|
||||
manifest: OAuthPluginManifest,
|
||||
binary_path: PathBuf,
|
||||
config: serde_json::Value,
|
||||
) -> Self {
|
||||
Self {
|
||||
manifest,
|
||||
binary_path,
|
||||
config,
|
||||
process: Mutex::new(None),
|
||||
}
|
||||
}
|
||||
|
||||
/// 调用插件命令
|
||||
async fn call_command(
|
||||
&self,
|
||||
method: &str,
|
||||
params: serde_json::Value,
|
||||
) -> OAuthPluginResult<serde_json::Value> {
|
||||
let _request = serde_json::json!({
|
||||
"jsonrpc": "2.0",
|
||||
"method": method,
|
||||
"params": params,
|
||||
"id": 1
|
||||
});
|
||||
|
||||
let _output = Command::new(&self.binary_path)
|
||||
.arg("--json-rpc")
|
||||
.stdin(Stdio::piped())
|
||||
.stdout(Stdio::piped())
|
||||
.stderr(Stdio::piped())
|
||||
.spawn()
|
||||
.map_err(|e| OAuthPluginError::InitError(format!("启动插件进程失败: {}", e)))?
|
||||
.wait_with_output()
|
||||
.await
|
||||
.map_err(|e| OAuthPluginError::InitError(format!("等待插件进程失败: {}", e)))?;
|
||||
|
||||
// TODO: 实现完整的 JSON-RPC 通信
|
||||
// 目前返回模拟数据
|
||||
|
||||
debug!(
|
||||
"Plugin {} called method {} (simulated)",
|
||||
self.manifest.provider.id, method
|
||||
);
|
||||
|
||||
Ok(serde_json::json!({}))
|
||||
}
|
||||
}
|
||||
|
||||
#[async_trait]
|
||||
impl CredentialProviderPlugin for ExternalOAuthPlugin {
|
||||
fn id(&self) -> &str {
|
||||
&self.manifest.provider.id
|
||||
}
|
||||
|
||||
fn display_name(&self) -> &str {
|
||||
&self.manifest.provider.display_name
|
||||
}
|
||||
|
||||
fn version(&self) -> &str {
|
||||
&self.manifest.version
|
||||
}
|
||||
|
||||
fn description(&self) -> &str {
|
||||
&self.manifest.description
|
||||
}
|
||||
|
||||
fn target_protocol(&self) -> StandardProtocol {
|
||||
StandardProtocol::from_str(&self.manifest.provider.target_protocol)
|
||||
.unwrap_or(StandardProtocol::Anthropic)
|
||||
}
|
||||
|
||||
fn ui_category(&self) -> CredentialCategory {
|
||||
CredentialCategory::OAuth
|
||||
}
|
||||
|
||||
fn supported_auth_types(&self) -> Vec<AuthTypeInfo> {
|
||||
self.manifest
|
||||
.provider
|
||||
.auth_types
|
||||
.iter()
|
||||
.map(|id| {
|
||||
let schema = self
|
||||
.manifest
|
||||
.provider
|
||||
.credential_schemas
|
||||
.get(id)
|
||||
.cloned()
|
||||
.unwrap_or_default();
|
||||
|
||||
AuthTypeInfo {
|
||||
id: id.clone(),
|
||||
display_name: id.clone(),
|
||||
description: schema
|
||||
.get("description")
|
||||
.and_then(|v| v.as_str())
|
||||
.unwrap_or("")
|
||||
.to_string(),
|
||||
category: CredentialCategory::OAuth,
|
||||
icon: None,
|
||||
}
|
||||
})
|
||||
.collect()
|
||||
}
|
||||
|
||||
fn credential_schema_for_auth(&self, auth_type: &str) -> serde_json::Value {
|
||||
self.manifest
|
||||
.provider
|
||||
.credential_schemas
|
||||
.get(auth_type)
|
||||
.cloned()
|
||||
.unwrap_or_default()
|
||||
}
|
||||
|
||||
fn parse_credential_config(
|
||||
&self,
|
||||
_auth_type: &str,
|
||||
_config: serde_json::Value,
|
||||
) -> OAuthPluginResult<Box<dyn CredentialConfig>> {
|
||||
// TODO: 调用外部二进制解析配置
|
||||
Err(OAuthPluginError::ConfigParseError(
|
||||
"外部插件配置解析未实现".to_string(),
|
||||
))
|
||||
}
|
||||
|
||||
async fn create_credential(
|
||||
&self,
|
||||
auth_type: &str,
|
||||
config: serde_json::Value,
|
||||
) -> OAuthPluginResult<String> {
|
||||
let result = self
|
||||
.call_command(
|
||||
"create_credential",
|
||||
serde_json::json!({
|
||||
"auth_type": auth_type,
|
||||
"config": config
|
||||
}),
|
||||
)
|
||||
.await?;
|
||||
|
||||
result
|
||||
.get("credential_id")
|
||||
.and_then(|v| v.as_str())
|
||||
.map(|s| s.to_string())
|
||||
.ok_or_else(|| OAuthPluginError::ConfigParseError("无效的凭证 ID".to_string()))
|
||||
}
|
||||
|
||||
fn model_families(&self) -> Vec<ModelFamily> {
|
||||
self.manifest
|
||||
.provider
|
||||
.supported_models
|
||||
.iter()
|
||||
.map(|pattern| ModelFamily {
|
||||
name: pattern.clone(),
|
||||
pattern: pattern.clone(),
|
||||
tier: None,
|
||||
description: None,
|
||||
})
|
||||
.collect()
|
||||
}
|
||||
|
||||
async fn list_models(&self) -> OAuthPluginResult<Vec<ModelInfo>> {
|
||||
// TODO: 调用外部二进制获取模型列表
|
||||
Ok(vec![])
|
||||
}
|
||||
|
||||
fn supports_model(&self, model: &str) -> bool {
|
||||
for pattern in &self.manifest.provider.supported_models {
|
||||
if let Ok(glob) = glob::Pattern::new(pattern) {
|
||||
if glob.matches(model) {
|
||||
return true;
|
||||
}
|
||||
}
|
||||
}
|
||||
false
|
||||
}
|
||||
|
||||
async fn acquire_credential(&self, model: &str) -> OAuthPluginResult<AcquiredCredential> {
|
||||
let result = self
|
||||
.call_command(
|
||||
"acquire_credential",
|
||||
serde_json::json!({
|
||||
"model": model
|
||||
}),
|
||||
)
|
||||
.await?;
|
||||
|
||||
serde_json::from_value(result)
|
||||
.map_err(|e| OAuthPluginError::ConfigParseError(format!("解析凭证失败: {}", e)))
|
||||
}
|
||||
|
||||
async fn release_credential(
|
||||
&self,
|
||||
credential_id: &str,
|
||||
result: UsageResult,
|
||||
) -> OAuthPluginResult<()> {
|
||||
self.call_command(
|
||||
"release_credential",
|
||||
serde_json::json!({
|
||||
"credential_id": credential_id,
|
||||
"result": result
|
||||
}),
|
||||
)
|
||||
.await?;
|
||||
Ok(())
|
||||
}
|
||||
|
||||
async fn validate_credential(
|
||||
&self,
|
||||
credential_id: &str,
|
||||
) -> OAuthPluginResult<ValidationResult> {
|
||||
let result = self
|
||||
.call_command(
|
||||
"validate_credential",
|
||||
serde_json::json!({
|
||||
"credential_id": credential_id
|
||||
}),
|
||||
)
|
||||
.await?;
|
||||
|
||||
serde_json::from_value(result)
|
||||
.map_err(|e| OAuthPluginError::ValidationError(format!("解析验证结果失败: {}", e)))
|
||||
}
|
||||
|
||||
async fn refresh_token(&self, credential_id: &str) -> OAuthPluginResult<TokenRefreshResult> {
|
||||
let result = self
|
||||
.call_command(
|
||||
"refresh_token",
|
||||
serde_json::json!({
|
||||
"credential_id": credential_id
|
||||
}),
|
||||
)
|
||||
.await?;
|
||||
|
||||
serde_json::from_value(result)
|
||||
.map_err(|e| OAuthPluginError::TokenRefreshError(format!("解析刷新结果失败: {}", e)))
|
||||
}
|
||||
|
||||
async fn transform_request(&self, request: &mut serde_json::Value) -> OAuthPluginResult<()> {
|
||||
let result = self
|
||||
.call_command(
|
||||
"transform_request",
|
||||
serde_json::json!({
|
||||
"request": request.clone()
|
||||
}),
|
||||
)
|
||||
.await?;
|
||||
|
||||
if let Some(transformed) = result.get("request") {
|
||||
*request = transformed.clone();
|
||||
}
|
||||
|
||||
Ok(())
|
||||
}
|
||||
|
||||
async fn transform_response(&self, response: &mut serde_json::Value) -> OAuthPluginResult<()> {
|
||||
let result = self
|
||||
.call_command(
|
||||
"transform_response",
|
||||
serde_json::json!({
|
||||
"response": response.clone()
|
||||
}),
|
||||
)
|
||||
.await?;
|
||||
|
||||
if let Some(transformed) = result.get("response") {
|
||||
*response = transformed.clone();
|
||||
}
|
||||
|
||||
Ok(())
|
||||
}
|
||||
|
||||
async fn apply_risk_control(
|
||||
&self,
|
||||
request: &mut serde_json::Value,
|
||||
credential_id: &str,
|
||||
) -> OAuthPluginResult<()> {
|
||||
let result = self
|
||||
.call_command(
|
||||
"apply_risk_control",
|
||||
serde_json::json!({
|
||||
"request": request.clone(),
|
||||
"credential_id": credential_id
|
||||
}),
|
||||
)
|
||||
.await?;
|
||||
|
||||
if let Some(modified) = result.get("request") {
|
||||
*request = modified.clone();
|
||||
}
|
||||
|
||||
Ok(())
|
||||
}
|
||||
|
||||
fn parse_error(&self, _status: u16, _body: &str) -> Option<ProviderError> {
|
||||
// TODO: 调用外部二进制解析错误
|
||||
None
|
||||
}
|
||||
|
||||
fn get_plugin_config(&self) -> serde_json::Value {
|
||||
self.config.clone()
|
||||
}
|
||||
|
||||
async fn update_plugin_config(&self, _config: serde_json::Value) -> OAuthPluginResult<()> {
|
||||
// TODO: 持久化配置更新
|
||||
Ok(())
|
||||
}
|
||||
|
||||
async fn init(&self) -> OAuthPluginResult<()> {
|
||||
info!(
|
||||
"Initializing external OAuth plugin: {} ({})",
|
||||
self.manifest.provider.id,
|
||||
self.binary_path.display()
|
||||
);
|
||||
Ok(())
|
||||
}
|
||||
|
||||
async fn shutdown(&self) -> OAuthPluginResult<()> {
|
||||
info!(
|
||||
"Shutting down external OAuth plugin: {}",
|
||||
self.manifest.provider.id
|
||||
);
|
||||
|
||||
// 终止进程(如果有)
|
||||
let mut process = self.process.lock().await;
|
||||
if let Some(mut child) = process.take() {
|
||||
let _ = child.kill().await;
|
||||
}
|
||||
|
||||
Ok(())
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
use std::env::temp_dir;
|
||||
|
||||
#[test]
|
||||
fn test_platform_key() {
|
||||
let key = get_platform_key();
|
||||
assert!(!key.is_empty());
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_loader_creation() {
|
||||
let loader = OAuthPluginLoader::new(temp_dir().join("test_oauth_plugins"));
|
||||
assert!(loader.plugins_dir().exists() || true); // 目录可能不存在
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_manifest_parsing() {
|
||||
let json = r#"{
|
||||
"name": "test-provider",
|
||||
"version": "1.0.0",
|
||||
"description": "Test OAuth Provider",
|
||||
"plugin_type": "oauth_provider",
|
||||
"entry": "test-provider-cli",
|
||||
"provider": {
|
||||
"id": "test",
|
||||
"display_name": "Test Provider",
|
||||
"target_protocol": "anthropic",
|
||||
"supported_models": ["test-*"],
|
||||
"auth_types": ["oauth"]
|
||||
}
|
||||
}"#;
|
||||
|
||||
let manifest: OAuthPluginManifest = serde_json::from_str(json).unwrap();
|
||||
assert_eq!(manifest.name, "test-provider");
|
||||
assert_eq!(manifest.provider.id, "test");
|
||||
assert_eq!(manifest.plugin_type, "oauth_provider");
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,568 @@
|
||||
//! OAuth 凭证提供商插件 Trait
|
||||
//!
|
||||
//! 定义 OAuth Provider 插件必须实现的接口,支持动态注册和独立更新。
|
||||
//! 设计原则:
|
||||
//! - 不依赖任何硬编码枚举
|
||||
//! - 新增 Provider 只需实现此 trait 并注册
|
||||
//! - 凭证配置由插件自己定义 Schema
|
||||
//! - 一个插件可支持多种认证方式(OAuth、API Key、第三方中转)
|
||||
|
||||
use async_trait::async_trait;
|
||||
use chrono::{DateTime, Utc};
|
||||
use serde::{Deserialize, Serialize};
|
||||
use std::any::Any;
|
||||
use std::collections::HashMap;
|
||||
use std::sync::Arc;
|
||||
use thiserror::Error;
|
||||
|
||||
/// OAuth Provider 插件错误类型
|
||||
#[derive(Error, Debug)]
|
||||
pub enum OAuthPluginError {
|
||||
#[error("凭证获取失败: {0}")]
|
||||
AcquireError(String),
|
||||
|
||||
#[error("凭证释放失败: {0}")]
|
||||
ReleaseError(String),
|
||||
|
||||
#[error("Token 刷新失败: {0}")]
|
||||
TokenRefreshError(String),
|
||||
|
||||
#[error("凭证验证失败: {0}")]
|
||||
ValidationError(String),
|
||||
|
||||
#[error("配置解析失败: {0}")]
|
||||
ConfigParseError(String),
|
||||
|
||||
#[error("协议转换失败: {0}")]
|
||||
TransformError(String),
|
||||
|
||||
#[error("风控检查失败: {0}")]
|
||||
RiskControlError(String),
|
||||
|
||||
#[error("模型不支持: {0}")]
|
||||
UnsupportedModel(String),
|
||||
|
||||
#[error("插件初始化失败: {0}")]
|
||||
InitError(String),
|
||||
|
||||
#[error("IO 错误: {0}")]
|
||||
IoError(#[from] std::io::Error),
|
||||
|
||||
#[error("JSON 解析错误: {0}")]
|
||||
JsonError(#[from] serde_json::Error),
|
||||
}
|
||||
|
||||
pub type OAuthPluginResult<T> = Result<T, OAuthPluginError>;
|
||||
|
||||
// ============================================================================
|
||||
// 认证类型信息
|
||||
// ============================================================================
|
||||
|
||||
/// 凭证分组(用于 UI Tab 展示)
|
||||
#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize, Default)]
|
||||
#[serde(rename_all = "lowercase")]
|
||||
pub enum CredentialCategory {
|
||||
/// OAuth 凭证 Tab
|
||||
#[default]
|
||||
OAuth,
|
||||
/// API Key Tab
|
||||
ApiKey,
|
||||
/// 其他配置 Tab(第三方中转、Cookie 等)
|
||||
Other,
|
||||
}
|
||||
|
||||
/// 认证方式信息
|
||||
#[derive(Debug, Clone, Serialize, Deserialize)]
|
||||
pub struct AuthTypeInfo {
|
||||
/// 认证方式 ID(如 "oauth", "api_key", "third_party")
|
||||
pub id: String,
|
||||
/// 显示名称(如 "OAuth 登录", "官方 API Key", "第三方中转")
|
||||
pub display_name: String,
|
||||
/// 描述(如 "使用官方 OAuth 授权")
|
||||
pub description: String,
|
||||
/// UI 分组(显示在哪个 Tab)
|
||||
pub category: CredentialCategory,
|
||||
/// 图标名称 (Lucide icon)
|
||||
#[serde(default)]
|
||||
pub icon: Option<String>,
|
||||
}
|
||||
|
||||
// ============================================================================
|
||||
// 模型家族定义
|
||||
// ============================================================================
|
||||
|
||||
/// 模型家族(用于 Mini/Pro/Max 分层)
|
||||
#[derive(Debug, Clone, Serialize, Deserialize)]
|
||||
pub struct ModelFamily {
|
||||
/// 家族名称(如 "opus", "sonnet", "haiku")
|
||||
pub name: String,
|
||||
/// 匹配模式(如 "claude-opus-*", "claude-*-sonnet")
|
||||
pub pattern: String,
|
||||
/// 服务等级(1=Mini, 2=Pro, 3=Max)
|
||||
#[serde(default)]
|
||||
pub tier: Option<u8>,
|
||||
/// 描述
|
||||
#[serde(default)]
|
||||
pub description: Option<String>,
|
||||
}
|
||||
|
||||
/// 模型信息
|
||||
#[derive(Debug, Clone, Serialize, Deserialize)]
|
||||
pub struct ModelInfo {
|
||||
/// 模型 ID
|
||||
pub id: String,
|
||||
/// 显示名称
|
||||
pub display_name: String,
|
||||
/// 模型家族
|
||||
#[serde(default)]
|
||||
pub family: Option<String>,
|
||||
/// 上下文长度
|
||||
#[serde(default)]
|
||||
pub context_length: Option<u32>,
|
||||
/// 是否支持视觉
|
||||
#[serde(default)]
|
||||
pub supports_vision: bool,
|
||||
/// 是否支持工具调用
|
||||
#[serde(default)]
|
||||
pub supports_tools: bool,
|
||||
/// 输入价格(每 1M tokens)
|
||||
#[serde(default)]
|
||||
pub input_cost_per_million: Option<f64>,
|
||||
/// 输出价格(每 1M tokens)
|
||||
#[serde(default)]
|
||||
pub output_cost_per_million: Option<f64>,
|
||||
}
|
||||
|
||||
// ============================================================================
|
||||
// 凭证配置 Trait
|
||||
// ============================================================================
|
||||
|
||||
/// 凭证配置 trait(代替 CredentialData 枚举)
|
||||
///
|
||||
/// 每个插件自己定义凭证配置结构
|
||||
pub trait CredentialConfig: Send + Sync + Any {
|
||||
/// 转换为 Any,用于向下转型
|
||||
fn as_any(&self) -> &dyn Any;
|
||||
|
||||
/// 凭证类型(如 "oauth", "api_key", "third_party")
|
||||
fn credential_type(&self) -> &str;
|
||||
|
||||
/// 序列化为 JSON
|
||||
fn to_json(&self) -> serde_json::Value;
|
||||
|
||||
/// 克隆
|
||||
fn clone_box(&self) -> Box<dyn CredentialConfig>;
|
||||
}
|
||||
|
||||
// ============================================================================
|
||||
// 凭证获取结果
|
||||
// ============================================================================
|
||||
|
||||
/// 获取的凭证
|
||||
#[derive(Debug, Clone, Serialize, Deserialize)]
|
||||
pub struct AcquiredCredential {
|
||||
/// 凭证 ID
|
||||
pub id: String,
|
||||
/// 凭证名称
|
||||
#[serde(default)]
|
||||
pub name: Option<String>,
|
||||
/// 认证方式
|
||||
pub auth_type: String,
|
||||
/// Base URL(如果有)
|
||||
#[serde(default)]
|
||||
pub base_url: Option<String>,
|
||||
/// 请求头(Key-Value 对)
|
||||
#[serde(default)]
|
||||
pub headers: HashMap<String, String>,
|
||||
/// 额外元数据
|
||||
#[serde(default)]
|
||||
pub metadata: HashMap<String, serde_json::Value>,
|
||||
}
|
||||
|
||||
/// 凭证使用结果
|
||||
#[derive(Debug, Clone, Serialize, Deserialize)]
|
||||
pub enum UsageResult {
|
||||
/// 成功使用
|
||||
Success {
|
||||
/// 延迟(毫秒)
|
||||
latency_ms: u64,
|
||||
/// 输入 tokens
|
||||
input_tokens: Option<u32>,
|
||||
/// 输出 tokens
|
||||
output_tokens: Option<u32>,
|
||||
},
|
||||
/// 使用失败
|
||||
Error {
|
||||
/// 错误类型
|
||||
error_type: String,
|
||||
/// 错误消息
|
||||
message: String,
|
||||
/// 是否应标记为不健康
|
||||
mark_unhealthy: bool,
|
||||
/// 冷却时间(秒)
|
||||
cooldown_seconds: Option<u64>,
|
||||
},
|
||||
}
|
||||
|
||||
/// Token 刷新结果
|
||||
#[derive(Debug, Clone, Serialize, Deserialize)]
|
||||
pub struct TokenRefreshResult {
|
||||
/// 新的 access_token
|
||||
pub access_token: String,
|
||||
/// 新的 refresh_token(如果更新了)
|
||||
#[serde(default)]
|
||||
pub refresh_token: Option<String>,
|
||||
/// 过期时间
|
||||
#[serde(default)]
|
||||
pub expires_at: Option<DateTime<Utc>>,
|
||||
}
|
||||
|
||||
/// 凭证验证结果
|
||||
#[derive(Debug, Clone, Serialize, Deserialize)]
|
||||
pub struct ValidationResult {
|
||||
/// 是否有效
|
||||
pub valid: bool,
|
||||
/// 消息
|
||||
#[serde(default)]
|
||||
pub message: Option<String>,
|
||||
/// 额外信息
|
||||
#[serde(default)]
|
||||
pub details: HashMap<String, serde_json::Value>,
|
||||
}
|
||||
|
||||
// ============================================================================
|
||||
// Provider 错误解析
|
||||
// ============================================================================
|
||||
|
||||
/// Provider 错误
|
||||
#[derive(Debug, Clone, Serialize, Deserialize)]
|
||||
pub struct ProviderError {
|
||||
/// 错误类型
|
||||
pub error_type: ProviderErrorType,
|
||||
/// 错误消息
|
||||
pub message: String,
|
||||
/// HTTP 状态码
|
||||
#[serde(default)]
|
||||
pub status_code: Option<u16>,
|
||||
/// 是否可重试
|
||||
#[serde(default)]
|
||||
pub retryable: bool,
|
||||
/// 建议的冷却时间(秒)
|
||||
#[serde(default)]
|
||||
pub cooldown_seconds: Option<u64>,
|
||||
}
|
||||
|
||||
/// Provider 错误类型
|
||||
#[derive(Debug, Clone, Serialize, Deserialize, PartialEq, Eq)]
|
||||
#[serde(rename_all = "snake_case")]
|
||||
pub enum ProviderErrorType {
|
||||
/// 认证错误(Token 无效、过期等)
|
||||
Authentication,
|
||||
/// 授权错误(无权限)
|
||||
Authorization,
|
||||
/// 限流
|
||||
RateLimit,
|
||||
/// 配额超限
|
||||
QuotaExceeded,
|
||||
/// 模型不可用
|
||||
ModelUnavailable,
|
||||
/// 内容安全过滤
|
||||
ContentFiltered,
|
||||
/// 服务器错误
|
||||
ServerError,
|
||||
/// 网络错误
|
||||
NetworkError,
|
||||
/// 未知错误
|
||||
Unknown,
|
||||
}
|
||||
|
||||
// ============================================================================
|
||||
// 输出协议
|
||||
// ============================================================================
|
||||
|
||||
/// 目标标准协议
|
||||
#[derive(Debug, Clone, Serialize, Deserialize, PartialEq, Eq)]
|
||||
#[serde(rename_all = "lowercase")]
|
||||
pub enum StandardProtocol {
|
||||
/// Anthropic Claude API
|
||||
Anthropic,
|
||||
/// OpenAI Chat Completions API
|
||||
OpenAI,
|
||||
/// Google Gemini API
|
||||
Gemini,
|
||||
/// 通义千问
|
||||
Qwen,
|
||||
/// 其他 OpenAI 兼容
|
||||
OpenAICompat,
|
||||
}
|
||||
|
||||
impl StandardProtocol {
|
||||
pub fn as_str(&self) -> &'static str {
|
||||
match self {
|
||||
StandardProtocol::Anthropic => "anthropic",
|
||||
StandardProtocol::OpenAI => "openai",
|
||||
StandardProtocol::Gemini => "gemini",
|
||||
StandardProtocol::Qwen => "qwen",
|
||||
StandardProtocol::OpenAICompat => "openai_compat",
|
||||
}
|
||||
}
|
||||
|
||||
pub fn from_str(s: &str) -> Option<Self> {
|
||||
match s.to_lowercase().as_str() {
|
||||
"anthropic" => Some(StandardProtocol::Anthropic),
|
||||
"openai" => Some(StandardProtocol::OpenAI),
|
||||
"gemini" => Some(StandardProtocol::Gemini),
|
||||
"qwen" => Some(StandardProtocol::Qwen),
|
||||
"openai_compat" => Some(StandardProtocol::OpenAICompat),
|
||||
_ => None,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// ============================================================================
|
||||
// 主 Trait:CredentialProviderPlugin
|
||||
// ============================================================================
|
||||
|
||||
/// 凭证提供商插件 - 核心 Trait
|
||||
///
|
||||
/// 设计原则:
|
||||
/// - 不依赖任何硬编码枚举
|
||||
/// - 新增 Provider 只需实现此 trait 并注册
|
||||
/// - 凭证配置由插件自己定义 Schema
|
||||
/// - 一个插件可支持多种认证方式(OAuth、API Key、第三方中转)
|
||||
#[async_trait]
|
||||
pub trait CredentialProviderPlugin: Send + Sync {
|
||||
// ========== 基础信息 ==========
|
||||
|
||||
/// 插件唯一标识(代替 ProviderType 枚举)
|
||||
fn id(&self) -> &str;
|
||||
|
||||
/// 显示名称
|
||||
fn display_name(&self) -> &str;
|
||||
|
||||
/// 插件版本
|
||||
fn version(&self) -> &str;
|
||||
|
||||
/// 插件描述
|
||||
fn description(&self) -> &str {
|
||||
""
|
||||
}
|
||||
|
||||
/// 默认目标标准协议
|
||||
fn target_protocol(&self) -> StandardProtocol;
|
||||
|
||||
/// 根据模型动态返回目标协议(用于 Antigravity 等多协议 Provider)
|
||||
fn target_protocol_for_model(&self, _model: &str) -> StandardProtocol {
|
||||
self.target_protocol() // 默认返回固定协议
|
||||
}
|
||||
|
||||
/// UI 分组
|
||||
fn ui_category(&self) -> CredentialCategory {
|
||||
CredentialCategory::OAuth
|
||||
}
|
||||
|
||||
// ========== 多认证方式支持 ==========
|
||||
|
||||
/// 支持的认证方式(一个插件可支持多种)
|
||||
/// 例如 Anthropic 同时支持 OAuth、API Key、第三方中转
|
||||
fn supported_auth_types(&self) -> Vec<AuthTypeInfo>;
|
||||
|
||||
/// 根据认证方式返回对应的凭证配置 Schema
|
||||
fn credential_schema_for_auth(&self, auth_type: &str) -> serde_json::Value;
|
||||
|
||||
/// 解析凭证配置(从 JSON 解析成插件内部结构)
|
||||
fn parse_credential_config(
|
||||
&self,
|
||||
auth_type: &str,
|
||||
config: serde_json::Value,
|
||||
) -> OAuthPluginResult<Box<dyn CredentialConfig>>;
|
||||
|
||||
/// 创建凭证(从用户输入创建)
|
||||
async fn create_credential(
|
||||
&self,
|
||||
auth_type: &str,
|
||||
config: serde_json::Value,
|
||||
) -> OAuthPluginResult<String>;
|
||||
|
||||
// ========== 模型能力 ==========
|
||||
|
||||
/// 模型家族定义(用于 Mini/Pro/Max 分层)
|
||||
fn model_families(&self) -> Vec<ModelFamily>;
|
||||
|
||||
/// 获取支持的模型列表
|
||||
async fn list_models(&self) -> OAuthPluginResult<Vec<ModelInfo>>;
|
||||
|
||||
/// 检查是否支持某个模型
|
||||
fn supports_model(&self, model: &str) -> bool;
|
||||
|
||||
// ========== 凭证管理 ==========
|
||||
|
||||
/// 获取可用凭证
|
||||
async fn acquire_credential(&self, model: &str) -> OAuthPluginResult<AcquiredCredential>;
|
||||
|
||||
/// 释放凭证(报告使用结果)
|
||||
async fn release_credential(
|
||||
&self,
|
||||
credential_id: &str,
|
||||
result: UsageResult,
|
||||
) -> OAuthPluginResult<()>;
|
||||
|
||||
/// 验证凭证有效性
|
||||
async fn validate_credential(&self, credential_id: &str)
|
||||
-> OAuthPluginResult<ValidationResult>;
|
||||
|
||||
/// 刷新 Token(OAuth 类型)
|
||||
async fn refresh_token(&self, credential_id: &str) -> OAuthPluginResult<TokenRefreshResult>;
|
||||
|
||||
// ========== 协议转换 ==========
|
||||
|
||||
/// 将输入请求转换成标准协议
|
||||
async fn transform_request(&self, request: &mut serde_json::Value) -> OAuthPluginResult<()>;
|
||||
|
||||
/// 将响应转换回来(如果需要)
|
||||
async fn transform_response(&self, response: &mut serde_json::Value) -> OAuthPluginResult<()>;
|
||||
|
||||
// ========== 风控适配 ==========
|
||||
|
||||
/// 应用特有的风控逻辑
|
||||
async fn apply_risk_control(
|
||||
&self,
|
||||
request: &mut serde_json::Value,
|
||||
credential_id: &str,
|
||||
) -> OAuthPluginResult<()>;
|
||||
|
||||
/// 解析特有的错误码
|
||||
fn parse_error(&self, status: u16, body: &str) -> Option<ProviderError>;
|
||||
|
||||
// ========== 插件配置(非凭证配置)==========
|
||||
|
||||
/// 插件配置 Schema(用于 UI 动态生成表单)
|
||||
fn plugin_config_schema(&self) -> serde_json::Value {
|
||||
serde_json::json!({})
|
||||
}
|
||||
|
||||
/// 获取插件配置
|
||||
fn get_plugin_config(&self) -> serde_json::Value {
|
||||
serde_json::json!({})
|
||||
}
|
||||
|
||||
/// 更新插件配置
|
||||
async fn update_plugin_config(&self, _config: serde_json::Value) -> OAuthPluginResult<()> {
|
||||
Ok(())
|
||||
}
|
||||
|
||||
// ========== 生命周期 ==========
|
||||
|
||||
/// 初始化插件
|
||||
async fn init(&self) -> OAuthPluginResult<()>;
|
||||
|
||||
/// 关闭插件
|
||||
async fn shutdown(&self) -> OAuthPluginResult<()>;
|
||||
}
|
||||
|
||||
/// 插件实例类型别名
|
||||
pub type PluginInstance = Arc<dyn CredentialProviderPlugin>;
|
||||
|
||||
// ============================================================================
|
||||
// 插件信息(用于 UI 显示)
|
||||
// ============================================================================
|
||||
|
||||
/// 插件信息(用于 UI 显示)
|
||||
#[derive(Debug, Clone, Serialize, Deserialize)]
|
||||
pub struct OAuthPluginInfo {
|
||||
/// 插件 ID
|
||||
pub id: String,
|
||||
/// 显示名称
|
||||
pub display_name: String,
|
||||
/// 版本
|
||||
pub version: String,
|
||||
/// 描述
|
||||
pub description: String,
|
||||
/// 目标协议
|
||||
pub target_protocol: String,
|
||||
/// UI 分组
|
||||
pub category: CredentialCategory,
|
||||
/// 支持的认证方式
|
||||
pub auth_types: Vec<AuthTypeInfo>,
|
||||
/// 是否启用
|
||||
pub enabled: bool,
|
||||
/// 凭证数量
|
||||
pub credential_count: u32,
|
||||
/// 健康凭证数量
|
||||
pub healthy_credential_count: u32,
|
||||
}
|
||||
|
||||
impl OAuthPluginInfo {
|
||||
/// 从插件实例创建信息
|
||||
pub fn from_plugin(plugin: &dyn CredentialProviderPlugin) -> Self {
|
||||
Self {
|
||||
id: plugin.id().to_string(),
|
||||
display_name: plugin.display_name().to_string(),
|
||||
version: plugin.version().to_string(),
|
||||
description: plugin.description().to_string(),
|
||||
target_protocol: plugin.target_protocol().as_str().to_string(),
|
||||
category: plugin.ui_category(),
|
||||
auth_types: plugin.supported_auth_types(),
|
||||
enabled: true,
|
||||
credential_count: 0,
|
||||
healthy_credential_count: 0,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
|
||||
#[test]
|
||||
fn test_standard_protocol_conversion() {
|
||||
assert_eq!(
|
||||
StandardProtocol::from_str("anthropic"),
|
||||
Some(StandardProtocol::Anthropic)
|
||||
);
|
||||
assert_eq!(
|
||||
StandardProtocol::from_str("OPENAI"),
|
||||
Some(StandardProtocol::OpenAI)
|
||||
);
|
||||
assert_eq!(StandardProtocol::from_str("unknown"), None);
|
||||
|
||||
assert_eq!(StandardProtocol::Anthropic.as_str(), "anthropic");
|
||||
assert_eq!(StandardProtocol::OpenAI.as_str(), "openai");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_auth_type_info_serialization() {
|
||||
let info = AuthTypeInfo {
|
||||
id: "oauth".to_string(),
|
||||
display_name: "OAuth 登录".to_string(),
|
||||
description: "使用官方 OAuth 授权".to_string(),
|
||||
category: CredentialCategory::OAuth,
|
||||
icon: Some("Key".to_string()),
|
||||
};
|
||||
|
||||
let json = serde_json::to_string(&info).unwrap();
|
||||
assert!(json.contains("oauth"));
|
||||
assert!(json.contains("OAuth 登录"));
|
||||
|
||||
let parsed: AuthTypeInfo = serde_json::from_str(&json).unwrap();
|
||||
assert_eq!(parsed.id, "oauth");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_provider_error_serialization() {
|
||||
let error = ProviderError {
|
||||
error_type: ProviderErrorType::RateLimit,
|
||||
message: "Too many requests".to_string(),
|
||||
status_code: Some(429),
|
||||
retryable: true,
|
||||
cooldown_seconds: Some(60),
|
||||
};
|
||||
|
||||
let json = serde_json::to_string(&error).unwrap();
|
||||
let parsed: ProviderError = serde_json::from_str(&json).unwrap();
|
||||
|
||||
assert_eq!(parsed.error_type, ProviderErrorType::RateLimit);
|
||||
assert_eq!(parsed.status_code, Some(429));
|
||||
assert!(parsed.retryable);
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,802 @@
|
||||
//! OAuth Provider 插件注册表
|
||||
//!
|
||||
//! 管理所有 OAuth Provider 插件的注册、发现和生命周期。
|
||||
//! 支持从外部目录动态加载插件。
|
||||
|
||||
use super::plugin::{OAuthPluginError, OAuthPluginInfo, OAuthPluginResult, PluginInstance};
|
||||
use dashmap::DashMap;
|
||||
use glob::Pattern;
|
||||
use serde::{Deserialize, Serialize};
|
||||
use std::path::{Path, PathBuf};
|
||||
use std::sync::Arc;
|
||||
use tokio::sync::RwLock;
|
||||
use tracing::{debug, error, info, warn};
|
||||
|
||||
/// 插件来源
|
||||
#[derive(Debug, Clone, Serialize, Deserialize)]
|
||||
#[serde(tag = "type", rename_all = "snake_case")]
|
||||
pub enum PluginSource {
|
||||
/// 从 GitHub Release 安装
|
||||
GitHub {
|
||||
owner: String,
|
||||
repo: String,
|
||||
version: Option<String>,
|
||||
},
|
||||
/// 从本地文件安装
|
||||
LocalFile { path: PathBuf },
|
||||
/// 内置插件(编译时包含)
|
||||
Builtin { id: String },
|
||||
}
|
||||
|
||||
/// 插件更新信息
|
||||
#[derive(Debug, Clone, Serialize, Deserialize)]
|
||||
pub struct PluginUpdate {
|
||||
/// 插件 ID
|
||||
pub plugin_id: String,
|
||||
/// 当前版本
|
||||
pub current_version: String,
|
||||
/// 最新版本
|
||||
pub latest_version: String,
|
||||
/// 更新说明
|
||||
pub changelog: Option<String>,
|
||||
}
|
||||
|
||||
/// 插件状态
|
||||
#[derive(Debug, Clone, Serialize, Deserialize, Default)]
|
||||
pub struct PluginState {
|
||||
/// 是否启用
|
||||
pub enabled: bool,
|
||||
/// 插件配置
|
||||
pub config: serde_json::Value,
|
||||
/// 安装时间
|
||||
pub installed_at: Option<String>,
|
||||
/// 最后使用时间
|
||||
pub last_used_at: Option<String>,
|
||||
}
|
||||
|
||||
/// OAuth Provider 插件注册表
|
||||
///
|
||||
/// 负责管理所有 OAuth Provider 插件的注册、发现和路由。
|
||||
/// 支持运行时动态加载和卸载插件。
|
||||
pub struct CredentialProviderRegistry {
|
||||
/// 已注册的插件(id -> 插件实例)
|
||||
providers: DashMap<String, PluginInstance>,
|
||||
|
||||
/// 模型到插件的映射(用于快速查找)
|
||||
/// 键是模型模式(如 "claude-*"),值是插件 ID
|
||||
model_patterns: RwLock<Vec<(Pattern, String)>>,
|
||||
|
||||
/// 插件状态
|
||||
plugin_states: DashMap<String, PluginState>,
|
||||
|
||||
/// 插件目录
|
||||
plugins_dir: PathBuf,
|
||||
}
|
||||
|
||||
impl std::fmt::Debug for CredentialProviderRegistry {
|
||||
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
|
||||
f.debug_struct("CredentialProviderRegistry")
|
||||
.field("plugins_dir", &self.plugins_dir)
|
||||
.field("provider_count", &self.providers.len())
|
||||
.finish()
|
||||
}
|
||||
}
|
||||
|
||||
impl CredentialProviderRegistry {
|
||||
/// 创建新的注册表
|
||||
pub fn new(plugins_dir: PathBuf) -> Self {
|
||||
Self {
|
||||
providers: DashMap::new(),
|
||||
model_patterns: RwLock::new(Vec::new()),
|
||||
plugin_states: DashMap::new(),
|
||||
plugins_dir,
|
||||
}
|
||||
}
|
||||
|
||||
/// 获取插件目录
|
||||
pub fn plugins_dir(&self) -> &Path {
|
||||
&self.plugins_dir
|
||||
}
|
||||
|
||||
// ========================================================================
|
||||
// 插件注册
|
||||
// ========================================================================
|
||||
|
||||
/// 注册插件
|
||||
pub async fn register(&self, plugin: PluginInstance) -> OAuthPluginResult<()> {
|
||||
let id = plugin.id().to_string();
|
||||
let display_name = plugin.display_name().to_string();
|
||||
|
||||
info!(
|
||||
"Registering OAuth provider plugin: {} ({})",
|
||||
id, display_name
|
||||
);
|
||||
|
||||
// 初始化插件
|
||||
plugin.init().await?;
|
||||
|
||||
// 注册模型模式
|
||||
let families = plugin.model_families();
|
||||
let mut patterns = self.model_patterns.write().await;
|
||||
|
||||
for family in families {
|
||||
match Pattern::new(&family.pattern) {
|
||||
Ok(pattern) => {
|
||||
patterns.push((pattern, id.clone()));
|
||||
debug!(
|
||||
"Registered model pattern '{}' for plugin '{}'",
|
||||
family.pattern, id
|
||||
);
|
||||
}
|
||||
Err(e) => {
|
||||
warn!(
|
||||
"Invalid model pattern '{}' for plugin '{}': {}",
|
||||
family.pattern, id, e
|
||||
);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// 创建默认状态
|
||||
if !self.plugin_states.contains_key(&id) {
|
||||
self.plugin_states.insert(
|
||||
id.clone(),
|
||||
PluginState {
|
||||
enabled: true,
|
||||
config: serde_json::json!({}),
|
||||
installed_at: Some(chrono::Utc::now().to_rfc3339()),
|
||||
last_used_at: None,
|
||||
},
|
||||
);
|
||||
}
|
||||
|
||||
// 注册插件
|
||||
self.providers.insert(id.clone(), plugin);
|
||||
|
||||
info!("Successfully registered OAuth provider plugin: {}", id);
|
||||
Ok(())
|
||||
}
|
||||
|
||||
/// 注销插件
|
||||
pub async fn unregister(&self, plugin_id: &str) -> OAuthPluginResult<()> {
|
||||
info!("Unregistering OAuth provider plugin: {}", plugin_id);
|
||||
|
||||
// 移除插件
|
||||
if let Some((_, plugin)) = self.providers.remove(plugin_id) {
|
||||
// 关闭插件
|
||||
if let Err(e) = plugin.shutdown().await {
|
||||
warn!("Error shutting down plugin {}: {}", plugin_id, e);
|
||||
}
|
||||
}
|
||||
|
||||
// 移除模型模式
|
||||
let mut patterns = self.model_patterns.write().await;
|
||||
patterns.retain(|(_, id)| id != plugin_id);
|
||||
|
||||
// 移除状态
|
||||
self.plugin_states.remove(plugin_id);
|
||||
|
||||
info!(
|
||||
"Successfully unregistered OAuth provider plugin: {}",
|
||||
plugin_id
|
||||
);
|
||||
Ok(())
|
||||
}
|
||||
|
||||
// ========================================================================
|
||||
// 插件查找
|
||||
// ========================================================================
|
||||
|
||||
/// 根据 ID 获取插件
|
||||
pub fn get(&self, plugin_id: &str) -> Option<PluginInstance> {
|
||||
self.providers.get(plugin_id).map(|r| r.value().clone())
|
||||
}
|
||||
|
||||
/// 根据模型名称查找插件
|
||||
pub async fn find_by_model(&self, model: &str) -> Option<PluginInstance> {
|
||||
let patterns = self.model_patterns.read().await;
|
||||
|
||||
// 按注册顺序查找匹配的模式
|
||||
for (pattern, plugin_id) in patterns.iter() {
|
||||
if pattern.matches(model) {
|
||||
// 检查插件是否启用
|
||||
if let Some(state) = self.plugin_states.get(plugin_id) {
|
||||
if !state.enabled {
|
||||
continue;
|
||||
}
|
||||
}
|
||||
|
||||
if let Some(plugin) = self.providers.get(plugin_id) {
|
||||
// 更新最后使用时间
|
||||
if let Some(mut state) = self.plugin_states.get_mut(plugin_id) {
|
||||
state.last_used_at = Some(chrono::Utc::now().to_rfc3339());
|
||||
}
|
||||
return Some(plugin.value().clone());
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
None
|
||||
}
|
||||
|
||||
/// 获取所有已注册的插件
|
||||
pub fn get_all(&self) -> Vec<PluginInstance> {
|
||||
self.providers.iter().map(|r| r.value().clone()).collect()
|
||||
}
|
||||
|
||||
/// 获取所有已启用的插件
|
||||
pub fn get_enabled(&self) -> Vec<PluginInstance> {
|
||||
self.providers
|
||||
.iter()
|
||||
.filter(|r| {
|
||||
self.plugin_states
|
||||
.get(r.key())
|
||||
.map(|s| s.enabled)
|
||||
.unwrap_or(true)
|
||||
})
|
||||
.map(|r| r.value().clone())
|
||||
.collect()
|
||||
}
|
||||
|
||||
/// 获取所有插件信息(用于 UI 显示)
|
||||
pub fn get_plugin_infos(&self) -> Vec<OAuthPluginInfo> {
|
||||
use super::plugin::CredentialCategory;
|
||||
|
||||
let mut infos: Vec<OAuthPluginInfo> = self
|
||||
.providers
|
||||
.iter()
|
||||
.map(|r| {
|
||||
let plugin = r.value();
|
||||
let state = self
|
||||
.plugin_states
|
||||
.get(r.key())
|
||||
.map(|s| s.value().clone())
|
||||
.unwrap_or_default();
|
||||
|
||||
let mut info = OAuthPluginInfo::from_plugin(plugin.as_ref());
|
||||
info.enabled = state.enabled;
|
||||
info
|
||||
})
|
||||
.collect();
|
||||
|
||||
// 同时扫描插件目录中已安装但未加载的插件
|
||||
if let Ok(entries) = std::fs::read_dir(&self.plugins_dir) {
|
||||
for entry in entries.flatten() {
|
||||
let path = entry.path();
|
||||
if path.is_dir() {
|
||||
let plugin_json = path.join("plugin.json");
|
||||
if plugin_json.exists() {
|
||||
if let Ok(content) = std::fs::read_to_string(&plugin_json) {
|
||||
if let Ok(manifest) =
|
||||
serde_json::from_str::<serde_json::Value>(&content)
|
||||
{
|
||||
// 只处理 oauth_provider 类型的插件
|
||||
let plugin_type = manifest["plugin_type"].as_str().unwrap_or("");
|
||||
if plugin_type != "oauth_provider" {
|
||||
continue;
|
||||
}
|
||||
|
||||
let plugin_id = manifest["name"].as_str().unwrap_or_default();
|
||||
|
||||
// 跳过已经在 providers 中的插件
|
||||
if infos.iter().any(|i| i.id == plugin_id) {
|
||||
continue;
|
||||
}
|
||||
|
||||
let state = self
|
||||
.plugin_states
|
||||
.get(plugin_id)
|
||||
.map(|s| s.value().clone())
|
||||
.unwrap_or_else(|| PluginState {
|
||||
enabled: true, // 默认启用
|
||||
config: serde_json::json!({}),
|
||||
installed_at: None,
|
||||
last_used_at: None,
|
||||
});
|
||||
|
||||
let info = OAuthPluginInfo {
|
||||
id: plugin_id.to_string(),
|
||||
display_name: manifest["provider"]["display_name"]
|
||||
.as_str()
|
||||
.or_else(|| manifest["name"].as_str())
|
||||
.unwrap_or(plugin_id)
|
||||
.to_string(),
|
||||
version: manifest["version"]
|
||||
.as_str()
|
||||
.unwrap_or("0.0.0")
|
||||
.to_string(),
|
||||
description: manifest["description"]
|
||||
.as_str()
|
||||
.unwrap_or("")
|
||||
.to_string(),
|
||||
target_protocol: manifest["provider"]["target_protocol"]
|
||||
.as_str()
|
||||
.unwrap_or("unknown")
|
||||
.to_string(),
|
||||
category: CredentialCategory::OAuth,
|
||||
auth_types: vec![],
|
||||
enabled: state.enabled,
|
||||
credential_count: 0,
|
||||
healthy_credential_count: 0,
|
||||
};
|
||||
infos.push(info);
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
infos
|
||||
}
|
||||
|
||||
// ========================================================================
|
||||
// 插件状态管理
|
||||
// ========================================================================
|
||||
|
||||
/// 启用插件
|
||||
pub fn enable_plugin(&self, plugin_id: &str) -> bool {
|
||||
// 如果状态不存在,先检查插件目录是否存在
|
||||
if !self.plugin_states.contains_key(plugin_id) {
|
||||
let plugin_dir = self.plugins_dir.join(plugin_id);
|
||||
if plugin_dir.join("plugin.json").exists() {
|
||||
// 创建默认状态
|
||||
let state = PluginState {
|
||||
enabled: true,
|
||||
config: serde_json::json!({}),
|
||||
installed_at: Some(chrono::Utc::now().to_rfc3339()),
|
||||
last_used_at: None,
|
||||
};
|
||||
self.plugin_states.insert(plugin_id.to_string(), state);
|
||||
info!("Enabled OAuth provider plugin: {}", plugin_id);
|
||||
return true;
|
||||
}
|
||||
return false;
|
||||
}
|
||||
|
||||
if let Some(mut state) = self.plugin_states.get_mut(plugin_id) {
|
||||
state.enabled = true;
|
||||
info!("Enabled OAuth provider plugin: {}", plugin_id);
|
||||
true
|
||||
} else {
|
||||
false
|
||||
}
|
||||
}
|
||||
|
||||
/// 禁用插件
|
||||
pub fn disable_plugin(&self, plugin_id: &str) -> bool {
|
||||
// 如果状态不存在,先检查插件目录是否存在
|
||||
if !self.plugin_states.contains_key(plugin_id) {
|
||||
let plugin_dir = self.plugins_dir.join(plugin_id);
|
||||
if plugin_dir.join("plugin.json").exists() {
|
||||
// 创建默认状态(禁用)
|
||||
let state = PluginState {
|
||||
enabled: false,
|
||||
config: serde_json::json!({}),
|
||||
installed_at: Some(chrono::Utc::now().to_rfc3339()),
|
||||
last_used_at: None,
|
||||
};
|
||||
self.plugin_states.insert(plugin_id.to_string(), state);
|
||||
info!("Disabled OAuth provider plugin: {}", plugin_id);
|
||||
return true;
|
||||
}
|
||||
return false;
|
||||
}
|
||||
|
||||
if let Some(mut state) = self.plugin_states.get_mut(plugin_id) {
|
||||
state.enabled = false;
|
||||
info!("Disabled OAuth provider plugin: {}", plugin_id);
|
||||
true
|
||||
} else {
|
||||
false
|
||||
}
|
||||
}
|
||||
|
||||
/// 获取插件状态
|
||||
pub fn get_plugin_state(&self, plugin_id: &str) -> Option<PluginState> {
|
||||
self.plugin_states.get(plugin_id).map(|r| r.value().clone())
|
||||
}
|
||||
|
||||
/// 更新插件配置
|
||||
pub async fn update_plugin_config(
|
||||
&self,
|
||||
plugin_id: &str,
|
||||
config: serde_json::Value,
|
||||
) -> OAuthPluginResult<()> {
|
||||
// 更新状态中的配置
|
||||
if let Some(mut state) = self.plugin_states.get_mut(plugin_id) {
|
||||
state.config = config.clone();
|
||||
}
|
||||
|
||||
// 通知插件配置更新
|
||||
if let Some(plugin) = self.providers.get(plugin_id) {
|
||||
plugin.update_plugin_config(config).await?;
|
||||
}
|
||||
|
||||
Ok(())
|
||||
}
|
||||
|
||||
// ========================================================================
|
||||
// 插件安装管理
|
||||
// ========================================================================
|
||||
|
||||
/// 安装插件(从外部来源)
|
||||
pub async fn install_plugin(&self, source: PluginSource) -> OAuthPluginResult<String> {
|
||||
match source {
|
||||
PluginSource::GitHub {
|
||||
owner,
|
||||
repo,
|
||||
version,
|
||||
} => {
|
||||
self.install_from_github(&owner, &repo, version.as_deref())
|
||||
.await
|
||||
}
|
||||
PluginSource::LocalFile { path } => self.install_from_local(&path).await,
|
||||
PluginSource::Builtin { id } => Err(OAuthPluginError::InitError(format!(
|
||||
"Builtin plugin '{}' cannot be installed manually",
|
||||
id
|
||||
))),
|
||||
}
|
||||
}
|
||||
|
||||
/// 从 GitHub Release 安装插件
|
||||
async fn install_from_github(
|
||||
&self,
|
||||
owner: &str,
|
||||
repo: &str,
|
||||
version: Option<&str>,
|
||||
) -> OAuthPluginResult<String> {
|
||||
let version_tag = version.unwrap_or("latest");
|
||||
|
||||
// 构建下载 URL
|
||||
let download_url = if version_tag == "latest" {
|
||||
format!(
|
||||
"https://github.com/{}/{}/releases/latest/download/{}-plugin.zip",
|
||||
owner, repo, repo
|
||||
)
|
||||
} else {
|
||||
format!(
|
||||
"https://github.com/{}/{}/releases/download/{}/{}-plugin.zip",
|
||||
owner, repo, version_tag, repo
|
||||
)
|
||||
};
|
||||
|
||||
info!("Downloading plugin from: {}", download_url);
|
||||
|
||||
// 下载插件包
|
||||
let client = reqwest::Client::new();
|
||||
let response = client.get(&download_url).send().await.map_err(|e| {
|
||||
OAuthPluginError::InitError(format!("Failed to download plugin: {}", e))
|
||||
})?;
|
||||
|
||||
if !response.status().is_success() {
|
||||
return Err(OAuthPluginError::InitError(format!(
|
||||
"Failed to download plugin: HTTP {}",
|
||||
response.status()
|
||||
)));
|
||||
}
|
||||
|
||||
let bytes = response
|
||||
.bytes()
|
||||
.await
|
||||
.map_err(|e| OAuthPluginError::InitError(format!("Failed to read response: {}", e)))?;
|
||||
|
||||
// 创建临时目录解压
|
||||
let temp_dir = std::env::temp_dir().join(format!("oauth_plugin_{}", uuid::Uuid::new_v4()));
|
||||
std::fs::create_dir_all(&temp_dir)?;
|
||||
|
||||
// 解压 ZIP 文件
|
||||
let cursor = std::io::Cursor::new(bytes);
|
||||
let mut archive = zip::ZipArchive::new(cursor)
|
||||
.map_err(|e| OAuthPluginError::InitError(format!("Failed to open zip: {}", e)))?;
|
||||
|
||||
for i in 0..archive.len() {
|
||||
let mut file = archive.by_index(i).map_err(|e| {
|
||||
OAuthPluginError::InitError(format!("Failed to read zip entry: {}", e))
|
||||
})?;
|
||||
|
||||
let outpath = temp_dir.join(file.name());
|
||||
|
||||
if file.name().ends_with('/') {
|
||||
std::fs::create_dir_all(&outpath)?;
|
||||
} else {
|
||||
if let Some(p) = outpath.parent() {
|
||||
if !p.exists() {
|
||||
std::fs::create_dir_all(p)?;
|
||||
}
|
||||
}
|
||||
let mut outfile = std::fs::File::create(&outpath)?;
|
||||
std::io::copy(&mut file, &mut outfile)?;
|
||||
}
|
||||
}
|
||||
|
||||
// 读取 plugin.json 获取插件 ID
|
||||
let plugin_json_path = temp_dir.join("plugin.json");
|
||||
if !plugin_json_path.exists() {
|
||||
std::fs::remove_dir_all(&temp_dir)?;
|
||||
return Err(OAuthPluginError::InitError(
|
||||
"Invalid plugin package: missing plugin.json".to_string(),
|
||||
));
|
||||
}
|
||||
|
||||
let plugin_json = std::fs::read_to_string(&plugin_json_path)?;
|
||||
let plugin_info: serde_json::Value = serde_json::from_str(&plugin_json)
|
||||
.map_err(|e| OAuthPluginError::InitError(format!("Invalid plugin.json: {}", e)))?;
|
||||
|
||||
let plugin_id = plugin_info["name"]
|
||||
.as_str()
|
||||
.ok_or_else(|| {
|
||||
OAuthPluginError::InitError("Missing 'name' in plugin.json".to_string())
|
||||
})?
|
||||
.to_string();
|
||||
|
||||
// 移动到插件目录
|
||||
let target_dir = self.plugins_dir.join(&plugin_id);
|
||||
if target_dir.exists() {
|
||||
std::fs::remove_dir_all(&target_dir)?;
|
||||
}
|
||||
std::fs::rename(&temp_dir, &target_dir)?;
|
||||
|
||||
info!("Plugin installed to: {:?}", target_dir);
|
||||
|
||||
// 注册插件(创建 PluginInstance)
|
||||
self.register_from_dir(&target_dir, &plugin_id).await?;
|
||||
|
||||
Ok(plugin_id)
|
||||
}
|
||||
|
||||
/// 从本地文件安装插件
|
||||
async fn install_from_local(&self, path: &Path) -> OAuthPluginResult<String> {
|
||||
// 检查路径是否存在
|
||||
if !path.exists() {
|
||||
return Err(OAuthPluginError::InitError(format!(
|
||||
"Path does not exist: {:?}",
|
||||
path
|
||||
)));
|
||||
}
|
||||
|
||||
// 如果是目录,直接复制
|
||||
if path.is_dir() {
|
||||
let plugin_json_path = path.join("plugin.json");
|
||||
if !plugin_json_path.exists() {
|
||||
return Err(OAuthPluginError::InitError(
|
||||
"Invalid plugin directory: missing plugin.json".to_string(),
|
||||
));
|
||||
}
|
||||
|
||||
let plugin_json = std::fs::read_to_string(&plugin_json_path)?;
|
||||
let plugin_info: serde_json::Value = serde_json::from_str(&plugin_json)
|
||||
.map_err(|e| OAuthPluginError::InitError(format!("Invalid plugin.json: {}", e)))?;
|
||||
|
||||
let plugin_id = plugin_info["name"]
|
||||
.as_str()
|
||||
.ok_or_else(|| {
|
||||
OAuthPluginError::InitError("Missing 'name' in plugin.json".to_string())
|
||||
})?
|
||||
.to_string();
|
||||
|
||||
let target_dir = self.plugins_dir.join(&plugin_id);
|
||||
if target_dir.exists() {
|
||||
std::fs::remove_dir_all(&target_dir)?;
|
||||
}
|
||||
|
||||
// 复制目录
|
||||
copy_dir_all(path, &target_dir)?;
|
||||
|
||||
info!("Plugin installed from local directory to: {:?}", target_dir);
|
||||
|
||||
// 注册插件
|
||||
self.register_from_dir(&target_dir, &plugin_id).await?;
|
||||
|
||||
return Ok(plugin_id);
|
||||
}
|
||||
|
||||
// 如果是 ZIP 文件
|
||||
if path.extension().map_or(false, |ext| ext == "zip") {
|
||||
let file = std::fs::File::open(path)?;
|
||||
let mut archive = zip::ZipArchive::new(file)
|
||||
.map_err(|e| OAuthPluginError::InitError(format!("Failed to open zip: {}", e)))?;
|
||||
|
||||
let temp_dir =
|
||||
std::env::temp_dir().join(format!("oauth_plugin_{}", uuid::Uuid::new_v4()));
|
||||
std::fs::create_dir_all(&temp_dir)?;
|
||||
|
||||
for i in 0..archive.len() {
|
||||
let mut file = archive.by_index(i).map_err(|e| {
|
||||
OAuthPluginError::InitError(format!("Failed to read zip entry: {}", e))
|
||||
})?;
|
||||
|
||||
let outpath = temp_dir.join(file.name());
|
||||
|
||||
if file.name().ends_with('/') {
|
||||
std::fs::create_dir_all(&outpath)?;
|
||||
} else {
|
||||
if let Some(p) = outpath.parent() {
|
||||
if !p.exists() {
|
||||
std::fs::create_dir_all(p)?;
|
||||
}
|
||||
}
|
||||
let mut outfile = std::fs::File::create(&outpath)?;
|
||||
std::io::copy(&mut file, &mut outfile)?;
|
||||
}
|
||||
}
|
||||
|
||||
let plugin_json_path = temp_dir.join("plugin.json");
|
||||
if !plugin_json_path.exists() {
|
||||
std::fs::remove_dir_all(&temp_dir)?;
|
||||
return Err(OAuthPluginError::InitError(
|
||||
"Invalid plugin package: missing plugin.json".to_string(),
|
||||
));
|
||||
}
|
||||
|
||||
let plugin_json = std::fs::read_to_string(&plugin_json_path)?;
|
||||
let plugin_info: serde_json::Value = serde_json::from_str(&plugin_json)
|
||||
.map_err(|e| OAuthPluginError::InitError(format!("Invalid plugin.json: {}", e)))?;
|
||||
|
||||
let plugin_id = plugin_info["name"]
|
||||
.as_str()
|
||||
.ok_or_else(|| {
|
||||
OAuthPluginError::InitError("Missing 'name' in plugin.json".to_string())
|
||||
})?
|
||||
.to_string();
|
||||
|
||||
let target_dir = self.plugins_dir.join(&plugin_id);
|
||||
if target_dir.exists() {
|
||||
std::fs::remove_dir_all(&target_dir)?;
|
||||
}
|
||||
std::fs::rename(&temp_dir, &target_dir)?;
|
||||
|
||||
info!("Plugin installed from zip to: {:?}", target_dir);
|
||||
|
||||
self.register_from_dir(&target_dir, &plugin_id).await?;
|
||||
|
||||
return Ok(plugin_id);
|
||||
}
|
||||
|
||||
Err(OAuthPluginError::InitError(format!(
|
||||
"Unsupported file type: {:?}",
|
||||
path
|
||||
)))
|
||||
}
|
||||
|
||||
/// 从目录注册插件
|
||||
async fn register_from_dir(&self, plugin_dir: &Path, plugin_id: &str) -> OAuthPluginResult<()> {
|
||||
// 读取 plugin.json
|
||||
let plugin_json_path = plugin_dir.join("plugin.json");
|
||||
let plugin_json = std::fs::read_to_string(&plugin_json_path)?;
|
||||
let manifest: serde_json::Value = serde_json::from_str(&plugin_json)
|
||||
.map_err(|e| OAuthPluginError::InitError(format!("Invalid plugin.json: {}", e)))?;
|
||||
|
||||
// 设置初始状态
|
||||
let state = PluginState {
|
||||
enabled: true,
|
||||
config: serde_json::json!({}),
|
||||
installed_at: Some(chrono::Utc::now().to_rfc3339()),
|
||||
last_used_at: None,
|
||||
};
|
||||
self.plugin_states.insert(plugin_id.to_string(), state);
|
||||
|
||||
info!(
|
||||
"Registered plugin: {} ({})",
|
||||
plugin_id,
|
||||
manifest["version"].as_str().unwrap_or("unknown")
|
||||
);
|
||||
|
||||
Ok(())
|
||||
}
|
||||
|
||||
/// 卸载插件
|
||||
pub async fn uninstall_plugin(&self, plugin_id: &str) -> OAuthPluginResult<()> {
|
||||
// 1. 注销插件
|
||||
self.unregister(plugin_id).await?;
|
||||
|
||||
// 2. 删除插件目录
|
||||
let plugin_dir = self.plugins_dir.join(plugin_id);
|
||||
if plugin_dir.exists() {
|
||||
std::fs::remove_dir_all(&plugin_dir)?;
|
||||
info!("Removed plugin directory: {:?}", plugin_dir);
|
||||
}
|
||||
|
||||
Ok(())
|
||||
}
|
||||
|
||||
/// 检查插件更新
|
||||
pub async fn check_updates(&self) -> OAuthPluginResult<Vec<PluginUpdate>> {
|
||||
// TODO: 实现更新检查逻辑
|
||||
// 1. 遍历所有插件
|
||||
// 2. 检查 GitHub Release 或其他来源
|
||||
// 3. 比较版本号
|
||||
// 4. 返回有更新的插件列表
|
||||
|
||||
Ok(vec![])
|
||||
}
|
||||
|
||||
// ========================================================================
|
||||
// 生命周期管理
|
||||
// ========================================================================
|
||||
|
||||
/// 关闭所有插件
|
||||
pub async fn shutdown_all(&self) -> OAuthPluginResult<()> {
|
||||
info!("Shutting down all OAuth provider plugins...");
|
||||
|
||||
for entry in self.providers.iter() {
|
||||
let plugin_id = entry.key();
|
||||
let plugin = entry.value();
|
||||
|
||||
if let Err(e) = plugin.shutdown().await {
|
||||
error!("Error shutting down plugin {}: {}", plugin_id, e);
|
||||
} else {
|
||||
debug!("Successfully shut down plugin: {}", plugin_id);
|
||||
}
|
||||
}
|
||||
|
||||
info!("All OAuth provider plugins shut down");
|
||||
Ok(())
|
||||
}
|
||||
}
|
||||
|
||||
// ============================================================================
|
||||
// 全局注册表
|
||||
// ============================================================================
|
||||
|
||||
use once_cell::sync::OnceCell;
|
||||
|
||||
static GLOBAL_REGISTRY: OnceCell<Arc<CredentialProviderRegistry>> = OnceCell::new();
|
||||
|
||||
/// 初始化全局注册表
|
||||
pub fn init_global_registry(plugins_dir: PathBuf) -> Arc<CredentialProviderRegistry> {
|
||||
let registry = Arc::new(CredentialProviderRegistry::new(plugins_dir));
|
||||
GLOBAL_REGISTRY
|
||||
.set(registry.clone())
|
||||
.expect("Global registry already initialized");
|
||||
registry
|
||||
}
|
||||
|
||||
/// 获取全局注册表
|
||||
pub fn get_global_registry() -> Option<Arc<CredentialProviderRegistry>> {
|
||||
GLOBAL_REGISTRY.get().cloned()
|
||||
}
|
||||
|
||||
/// 递归复制目录
|
||||
fn copy_dir_all(src: &Path, dst: &Path) -> std::io::Result<()> {
|
||||
std::fs::create_dir_all(dst)?;
|
||||
for entry in std::fs::read_dir(src)? {
|
||||
let entry = entry?;
|
||||
let ty = entry.file_type()?;
|
||||
let src_path = entry.path();
|
||||
let dst_path = dst.join(entry.file_name());
|
||||
|
||||
if ty.is_dir() {
|
||||
copy_dir_all(&src_path, &dst_path)?;
|
||||
} else {
|
||||
std::fs::copy(&src_path, &dst_path)?;
|
||||
}
|
||||
}
|
||||
Ok(())
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
use std::env::temp_dir;
|
||||
|
||||
#[test]
|
||||
fn test_registry_creation() {
|
||||
let registry = CredentialProviderRegistry::new(temp_dir().join("test_plugins"));
|
||||
assert_eq!(registry.get_all().len(), 0);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_find_by_model_empty() {
|
||||
let registry = CredentialProviderRegistry::new(temp_dir().join("test_plugins"));
|
||||
let result = registry.find_by_model("claude-opus-4").await;
|
||||
assert!(result.is_none());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_plugin_state_default() {
|
||||
let state = PluginState::default();
|
||||
assert!(!state.enabled);
|
||||
assert!(state.installed_at.is_none());
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,586 @@
|
||||
//! 风控模块
|
||||
//!
|
||||
//! 提供限流检测、冷却期管理和风险评估功能。
|
||||
//!
|
||||
//! ## 功能
|
||||
//!
|
||||
//! - **限流检测**: 检测 API 返回的限流错误(429、rate limit)
|
||||
//! - **冷却期管理**: 自动计算和管理凭证冷却时间
|
||||
//! - **风险评估**: 根据历史数据评估凭证风险等级
|
||||
|
||||
use chrono::{DateTime, Duration, Utc};
|
||||
use dashmap::DashMap;
|
||||
use serde::{Deserialize, Serialize};
|
||||
use std::collections::VecDeque;
|
||||
use std::sync::atomic::{AtomicU64, Ordering};
|
||||
|
||||
/// 风险等级
|
||||
#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
|
||||
#[serde(rename_all = "snake_case")]
|
||||
pub enum RiskLevel {
|
||||
/// 低风险 - 正常使用
|
||||
Low,
|
||||
/// 中风险 - 接近限流阈值
|
||||
Medium,
|
||||
/// 高风险 - 频繁触发限流
|
||||
High,
|
||||
/// 危险 - 需要立即冷却
|
||||
Critical,
|
||||
}
|
||||
|
||||
impl RiskLevel {
|
||||
/// 获取风险等级对应的冷却时间倍数
|
||||
pub fn cooldown_multiplier(&self) -> f64 {
|
||||
match self {
|
||||
RiskLevel::Low => 1.0,
|
||||
RiskLevel::Medium => 1.5,
|
||||
RiskLevel::High => 2.0,
|
||||
RiskLevel::Critical => 3.0,
|
||||
}
|
||||
}
|
||||
|
||||
/// 获取风险等级描述
|
||||
pub fn description(&self) -> &'static str {
|
||||
match self {
|
||||
RiskLevel::Low => "正常",
|
||||
RiskLevel::Medium => "接近限流",
|
||||
RiskLevel::High => "频繁限流",
|
||||
RiskLevel::Critical => "需要冷却",
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
/// 限流事件
|
||||
#[derive(Debug, Clone, Serialize, Deserialize)]
|
||||
pub struct RateLimitEvent {
|
||||
/// 凭证 ID
|
||||
pub credential_id: String,
|
||||
/// 事件时间
|
||||
pub timestamp: DateTime<Utc>,
|
||||
/// HTTP 状态码
|
||||
pub status_code: Option<u16>,
|
||||
/// 错误消息
|
||||
pub error_message: Option<String>,
|
||||
/// 建议的重试时间(秒)
|
||||
pub retry_after_secs: Option<u64>,
|
||||
}
|
||||
|
||||
impl RateLimitEvent {
|
||||
/// 创建新的限流事件
|
||||
pub fn new(credential_id: String) -> Self {
|
||||
Self {
|
||||
credential_id,
|
||||
timestamp: Utc::now(),
|
||||
status_code: None,
|
||||
error_message: None,
|
||||
retry_after_secs: None,
|
||||
}
|
||||
}
|
||||
|
||||
/// 设置状态码
|
||||
pub fn with_status_code(mut self, code: u16) -> Self {
|
||||
self.status_code = Some(code);
|
||||
self
|
||||
}
|
||||
|
||||
/// 设置错误消息
|
||||
pub fn with_error_message(mut self, message: String) -> Self {
|
||||
self.error_message = Some(message);
|
||||
self
|
||||
}
|
||||
|
||||
/// 设置重试时间
|
||||
pub fn with_retry_after(mut self, secs: u64) -> Self {
|
||||
self.retry_after_secs = Some(secs);
|
||||
self
|
||||
}
|
||||
}
|
||||
|
||||
/// 冷却配置
|
||||
#[derive(Debug, Clone, Serialize, Deserialize)]
|
||||
pub struct CooldownConfig {
|
||||
/// 基础冷却时间(秒)
|
||||
pub base_cooldown_secs: u64,
|
||||
/// 最大冷却时间(秒)
|
||||
pub max_cooldown_secs: u64,
|
||||
/// 冷却时间增长因子(指数退避)
|
||||
pub backoff_factor: f64,
|
||||
/// 限流事件窗口大小(保留最近 N 个事件)
|
||||
pub event_window_size: usize,
|
||||
/// 限流事件时间窗口(秒)- 只统计此时间内的事件
|
||||
pub event_time_window_secs: u64,
|
||||
/// 触发中风险的限流次数阈值
|
||||
pub medium_risk_threshold: u32,
|
||||
/// 触发高风险的限流次数阈值
|
||||
pub high_risk_threshold: u32,
|
||||
/// 触发危险的限流次数阈值
|
||||
pub critical_risk_threshold: u32,
|
||||
}
|
||||
|
||||
impl Default for CooldownConfig {
|
||||
fn default() -> Self {
|
||||
Self {
|
||||
base_cooldown_secs: 60, // 1 分钟
|
||||
max_cooldown_secs: 3600, // 1 小时
|
||||
backoff_factor: 2.0, // 指数退避因子
|
||||
event_window_size: 100, // 保留最近 100 个事件
|
||||
event_time_window_secs: 3600, // 1 小时内的事件
|
||||
medium_risk_threshold: 3, // 3 次限流 -> 中风险
|
||||
high_risk_threshold: 5, // 5 次限流 -> 高风险
|
||||
critical_risk_threshold: 10, // 10 次限流 -> 危险
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
/// 凭证风控状态
|
||||
#[derive(Debug)]
|
||||
struct CredentialRiskState {
|
||||
/// 限流事件历史
|
||||
events: VecDeque<RateLimitEvent>,
|
||||
/// 连续限流次数
|
||||
consecutive_rate_limits: AtomicU64,
|
||||
/// 当前冷却结束时间
|
||||
cooldown_until: Option<DateTime<Utc>>,
|
||||
/// 上次限流时间
|
||||
last_rate_limit: Option<DateTime<Utc>>,
|
||||
}
|
||||
|
||||
impl CredentialRiskState {
|
||||
fn new() -> Self {
|
||||
Self {
|
||||
events: VecDeque::new(),
|
||||
consecutive_rate_limits: AtomicU64::new(0),
|
||||
cooldown_until: None,
|
||||
last_rate_limit: None,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
/// 风控控制器
|
||||
///
|
||||
/// 管理凭证的限流检测和冷却期
|
||||
pub struct RiskController {
|
||||
/// 配置
|
||||
config: CooldownConfig,
|
||||
/// 各凭证的风控状态
|
||||
states: DashMap<String, CredentialRiskState>,
|
||||
}
|
||||
|
||||
impl RiskController {
|
||||
/// 创建新的风控控制器
|
||||
pub fn new(config: CooldownConfig) -> Self {
|
||||
Self {
|
||||
config,
|
||||
states: DashMap::new(),
|
||||
}
|
||||
}
|
||||
|
||||
/// 使用默认配置创建
|
||||
pub fn with_defaults() -> Self {
|
||||
Self::new(CooldownConfig::default())
|
||||
}
|
||||
|
||||
/// 获取配置
|
||||
pub fn config(&self) -> &CooldownConfig {
|
||||
&self.config
|
||||
}
|
||||
|
||||
/// 记录限流事件
|
||||
///
|
||||
/// # 返回
|
||||
/// 建议的冷却时间(秒)
|
||||
pub fn record_rate_limit(&self, event: RateLimitEvent) -> u64 {
|
||||
let credential_id = event.credential_id.clone();
|
||||
let retry_after = event.retry_after_secs;
|
||||
|
||||
let mut state = self
|
||||
.states
|
||||
.entry(credential_id.clone())
|
||||
.or_insert_with(CredentialRiskState::new);
|
||||
|
||||
// 更新连续限流次数
|
||||
state.consecutive_rate_limits.fetch_add(1, Ordering::SeqCst);
|
||||
state.last_rate_limit = Some(Utc::now());
|
||||
|
||||
// 添加事件到历史
|
||||
state.events.push_back(event);
|
||||
|
||||
// 清理过期事件
|
||||
self.cleanup_old_events(&mut state);
|
||||
|
||||
// 计算冷却时间
|
||||
let cooldown_secs = self.calculate_cooldown(&state, retry_after);
|
||||
|
||||
// 设置冷却结束时间
|
||||
state.cooldown_until = Some(Utc::now() + Duration::seconds(cooldown_secs as i64));
|
||||
|
||||
cooldown_secs
|
||||
}
|
||||
|
||||
/// 记录成功请求(重置连续限流计数)
|
||||
pub fn record_success(&self, credential_id: &str) {
|
||||
if let Some(state) = self.states.get_mut(credential_id) {
|
||||
state.consecutive_rate_limits.store(0, Ordering::SeqCst);
|
||||
}
|
||||
}
|
||||
|
||||
/// 获取凭证的风险等级
|
||||
pub fn get_risk_level(&self, credential_id: &str) -> RiskLevel {
|
||||
let state = match self.states.get(credential_id) {
|
||||
Some(s) => s,
|
||||
None => return RiskLevel::Low,
|
||||
};
|
||||
|
||||
let recent_count = self.count_recent_events(&state);
|
||||
|
||||
if recent_count >= self.config.critical_risk_threshold {
|
||||
RiskLevel::Critical
|
||||
} else if recent_count >= self.config.high_risk_threshold {
|
||||
RiskLevel::High
|
||||
} else if recent_count >= self.config.medium_risk_threshold {
|
||||
RiskLevel::Medium
|
||||
} else {
|
||||
RiskLevel::Low
|
||||
}
|
||||
}
|
||||
|
||||
/// 检查凭证是否在冷却中
|
||||
pub fn is_in_cooldown(&self, credential_id: &str) -> bool {
|
||||
self.states
|
||||
.get(credential_id)
|
||||
.and_then(|state| state.cooldown_until)
|
||||
.map(|until| Utc::now() < until)
|
||||
.unwrap_or(false)
|
||||
}
|
||||
|
||||
/// 获取凭证的冷却结束时间
|
||||
pub fn get_cooldown_until(&self, credential_id: &str) -> Option<DateTime<Utc>> {
|
||||
self.states
|
||||
.get(credential_id)
|
||||
.and_then(|state| state.cooldown_until)
|
||||
.filter(|until| Utc::now() < *until)
|
||||
}
|
||||
|
||||
/// 获取凭证的剩余冷却时间(秒)
|
||||
pub fn get_remaining_cooldown_secs(&self, credential_id: &str) -> Option<u64> {
|
||||
self.get_cooldown_until(credential_id).map(|until| {
|
||||
let remaining = until - Utc::now();
|
||||
remaining.num_seconds().max(0) as u64
|
||||
})
|
||||
}
|
||||
|
||||
/// 手动清除凭证的冷却状态
|
||||
pub fn clear_cooldown(&self, credential_id: &str) {
|
||||
if let Some(mut state) = self.states.get_mut(credential_id) {
|
||||
state.cooldown_until = None;
|
||||
state.consecutive_rate_limits.store(0, Ordering::SeqCst);
|
||||
}
|
||||
}
|
||||
|
||||
/// 获取所有处于冷却中的凭证 ID
|
||||
pub fn get_cooling_credentials(&self) -> Vec<String> {
|
||||
let now = Utc::now();
|
||||
self.states
|
||||
.iter()
|
||||
.filter(|entry| {
|
||||
entry
|
||||
.value()
|
||||
.cooldown_until
|
||||
.map(|until| now < until)
|
||||
.unwrap_or(false)
|
||||
})
|
||||
.map(|entry| entry.key().clone())
|
||||
.collect()
|
||||
}
|
||||
|
||||
/// 获取凭证的限流事件统计
|
||||
pub fn get_event_stats(&self, credential_id: &str) -> Option<RateLimitStats> {
|
||||
self.states.get(credential_id).map(|state| {
|
||||
let recent_count = self.count_recent_events(&state);
|
||||
let consecutive = state.consecutive_rate_limits.load(Ordering::SeqCst);
|
||||
|
||||
RateLimitStats {
|
||||
total_events: state.events.len(),
|
||||
recent_events: recent_count as usize,
|
||||
consecutive_rate_limits: consecutive,
|
||||
last_rate_limit: state.last_rate_limit,
|
||||
cooldown_until: state.cooldown_until,
|
||||
risk_level: self.get_risk_level(credential_id),
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
/// 检测响应是否为限流错误
|
||||
pub fn is_rate_limit_error(status_code: u16, body: Option<&str>) -> bool {
|
||||
// HTTP 429 Too Many Requests
|
||||
if status_code == 429 {
|
||||
return true;
|
||||
}
|
||||
|
||||
// 检查响应体中的限流关键词
|
||||
if let Some(body) = body {
|
||||
let body_lower = body.to_lowercase();
|
||||
if body_lower.contains("rate limit")
|
||||
|| body_lower.contains("rate_limit")
|
||||
|| body_lower.contains("ratelimit")
|
||||
|| body_lower.contains("too many requests")
|
||||
|| body_lower.contains("quota exceeded")
|
||||
|| body_lower.contains("resource_exhausted")
|
||||
{
|
||||
return true;
|
||||
}
|
||||
}
|
||||
|
||||
false
|
||||
}
|
||||
|
||||
/// 从响应头解析 Retry-After
|
||||
pub fn parse_retry_after(header_value: &str) -> Option<u64> {
|
||||
// 尝试解析为秒数
|
||||
if let Ok(secs) = header_value.parse::<u64>() {
|
||||
return Some(secs);
|
||||
}
|
||||
|
||||
// 尝试解析为 HTTP 日期格式
|
||||
if let Ok(date) = DateTime::parse_from_rfc2822(header_value) {
|
||||
let until = date.with_timezone(&Utc);
|
||||
let now = Utc::now();
|
||||
if until > now {
|
||||
return Some((until - now).num_seconds() as u64);
|
||||
}
|
||||
}
|
||||
|
||||
None
|
||||
}
|
||||
|
||||
/// 清理过期事件
|
||||
fn cleanup_old_events(&self, state: &mut CredentialRiskState) {
|
||||
let cutoff = Utc::now() - Duration::seconds(self.config.event_time_window_secs as i64);
|
||||
|
||||
// 移除过期事件
|
||||
while let Some(front) = state.events.front() {
|
||||
if front.timestamp < cutoff {
|
||||
state.events.pop_front();
|
||||
} else {
|
||||
break;
|
||||
}
|
||||
}
|
||||
|
||||
// 限制事件数量
|
||||
while state.events.len() > self.config.event_window_size {
|
||||
state.events.pop_front();
|
||||
}
|
||||
}
|
||||
|
||||
/// 统计最近的限流事件数
|
||||
fn count_recent_events(&self, state: &CredentialRiskState) -> u32 {
|
||||
let cutoff = Utc::now() - Duration::seconds(self.config.event_time_window_secs as i64);
|
||||
state
|
||||
.events
|
||||
.iter()
|
||||
.filter(|e| e.timestamp >= cutoff)
|
||||
.count() as u32
|
||||
}
|
||||
|
||||
/// 计算冷却时间
|
||||
fn calculate_cooldown(&self, state: &CredentialRiskState, retry_after: Option<u64>) -> u64 {
|
||||
// 如果有 Retry-After,优先使用
|
||||
if let Some(retry) = retry_after {
|
||||
return retry.min(self.config.max_cooldown_secs);
|
||||
}
|
||||
|
||||
// 使用指数退避计算冷却时间
|
||||
let consecutive = state.consecutive_rate_limits.load(Ordering::SeqCst);
|
||||
let base = self.config.base_cooldown_secs as f64;
|
||||
let factor = self.config.backoff_factor;
|
||||
|
||||
// cooldown = base * factor^(consecutive - 1)
|
||||
let cooldown = if consecutive > 0 {
|
||||
base * factor.powi((consecutive - 1) as i32)
|
||||
} else {
|
||||
base
|
||||
};
|
||||
|
||||
// 根据风险等级调整
|
||||
let risk_level = self.get_risk_level_from_state(state);
|
||||
let adjusted = cooldown * risk_level.cooldown_multiplier();
|
||||
|
||||
// 限制在最大值内
|
||||
(adjusted as u64).min(self.config.max_cooldown_secs)
|
||||
}
|
||||
|
||||
/// 从状态计算风险等级
|
||||
fn get_risk_level_from_state(&self, state: &CredentialRiskState) -> RiskLevel {
|
||||
let recent_count = self.count_recent_events(state);
|
||||
|
||||
if recent_count >= self.config.critical_risk_threshold {
|
||||
RiskLevel::Critical
|
||||
} else if recent_count >= self.config.high_risk_threshold {
|
||||
RiskLevel::High
|
||||
} else if recent_count >= self.config.medium_risk_threshold {
|
||||
RiskLevel::Medium
|
||||
} else {
|
||||
RiskLevel::Low
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
impl Default for RiskController {
|
||||
fn default() -> Self {
|
||||
Self::with_defaults()
|
||||
}
|
||||
}
|
||||
|
||||
/// 限流事件统计
|
||||
#[derive(Debug, Clone, Serialize, Deserialize)]
|
||||
pub struct RateLimitStats {
|
||||
/// 总事件数
|
||||
pub total_events: usize,
|
||||
/// 最近事件数(时间窗口内)
|
||||
pub recent_events: usize,
|
||||
/// 连续限流次数
|
||||
pub consecutive_rate_limits: u64,
|
||||
/// 上次限流时间
|
||||
pub last_rate_limit: Option<DateTime<Utc>>,
|
||||
/// 冷却结束时间
|
||||
pub cooldown_until: Option<DateTime<Utc>>,
|
||||
/// 风险等级
|
||||
pub risk_level: RiskLevel,
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
|
||||
#[test]
|
||||
fn test_risk_controller_new() {
|
||||
let controller = RiskController::with_defaults();
|
||||
assert_eq!(controller.config().base_cooldown_secs, 60);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_record_rate_limit() {
|
||||
let controller = RiskController::with_defaults();
|
||||
let event = RateLimitEvent::new("cred-1".to_string()).with_status_code(429);
|
||||
|
||||
let cooldown = controller.record_rate_limit(event);
|
||||
assert!(cooldown >= 60); // 至少是基础冷却时间
|
||||
|
||||
assert!(controller.is_in_cooldown("cred-1"));
|
||||
assert_eq!(controller.get_risk_level("cred-1"), RiskLevel::Low);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_risk_level_escalation() {
|
||||
let controller = RiskController::with_defaults();
|
||||
|
||||
// 记录多次限流事件
|
||||
for i in 0..5 {
|
||||
let event = RateLimitEvent::new("cred-1".to_string())
|
||||
.with_status_code(429)
|
||||
.with_error_message(format!("Rate limit {}", i));
|
||||
controller.record_rate_limit(event);
|
||||
}
|
||||
|
||||
// 应该达到高风险
|
||||
assert_eq!(controller.get_risk_level("cred-1"), RiskLevel::High);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_record_success_resets_consecutive() {
|
||||
let controller = RiskController::with_defaults();
|
||||
|
||||
// 记录限流
|
||||
let event = RateLimitEvent::new("cred-1".to_string());
|
||||
controller.record_rate_limit(event);
|
||||
|
||||
// 记录成功
|
||||
controller.record_success("cred-1");
|
||||
|
||||
// 连续计数应该重置
|
||||
let stats = controller.get_event_stats("cred-1").unwrap();
|
||||
assert_eq!(stats.consecutive_rate_limits, 0);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_clear_cooldown() {
|
||||
let controller = RiskController::with_defaults();
|
||||
|
||||
let event = RateLimitEvent::new("cred-1".to_string());
|
||||
controller.record_rate_limit(event);
|
||||
|
||||
assert!(controller.is_in_cooldown("cred-1"));
|
||||
|
||||
controller.clear_cooldown("cred-1");
|
||||
|
||||
assert!(!controller.is_in_cooldown("cred-1"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_is_rate_limit_error() {
|
||||
assert!(RiskController::is_rate_limit_error(429, None));
|
||||
assert!(RiskController::is_rate_limit_error(
|
||||
200,
|
||||
Some("rate limit exceeded")
|
||||
));
|
||||
assert!(RiskController::is_rate_limit_error(
|
||||
500,
|
||||
Some("RESOURCE_EXHAUSTED")
|
||||
));
|
||||
assert!(!RiskController::is_rate_limit_error(200, Some("success")));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_parse_retry_after() {
|
||||
assert_eq!(RiskController::parse_retry_after("60"), Some(60));
|
||||
assert_eq!(RiskController::parse_retry_after("3600"), Some(3600));
|
||||
assert!(RiskController::parse_retry_after("invalid").is_none());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_retry_after_priority() {
|
||||
let controller = RiskController::with_defaults();
|
||||
|
||||
// 使用 retry_after 的事件
|
||||
let event = RateLimitEvent::new("cred-1".to_string()).with_retry_after(120);
|
||||
|
||||
let cooldown = controller.record_rate_limit(event);
|
||||
assert_eq!(cooldown, 120); // 应该使用 retry_after 的值
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_exponential_backoff() {
|
||||
let controller = RiskController::with_defaults();
|
||||
|
||||
// 第一次限流
|
||||
let event1 = RateLimitEvent::new("cred-1".to_string());
|
||||
let cooldown1 = controller.record_rate_limit(event1);
|
||||
|
||||
// 第二次限流(应该更长)
|
||||
let event2 = RateLimitEvent::new("cred-1".to_string());
|
||||
let cooldown2 = controller.record_rate_limit(event2);
|
||||
|
||||
assert!(cooldown2 > cooldown1);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_get_cooling_credentials() {
|
||||
let controller = RiskController::with_defaults();
|
||||
|
||||
controller.record_rate_limit(RateLimitEvent::new("cred-1".to_string()));
|
||||
controller.record_rate_limit(RateLimitEvent::new("cred-2".to_string()));
|
||||
|
||||
let cooling = controller.get_cooling_credentials();
|
||||
assert_eq!(cooling.len(), 2);
|
||||
assert!(cooling.contains(&"cred-1".to_string()));
|
||||
assert!(cooling.contains(&"cred-2".to_string()));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_risk_level_cooldown_multiplier() {
|
||||
assert_eq!(RiskLevel::Low.cooldown_multiplier(), 1.0);
|
||||
assert_eq!(RiskLevel::Medium.cooldown_multiplier(), 1.5);
|
||||
assert_eq!(RiskLevel::High.cooldown_multiplier(), 2.0);
|
||||
assert_eq!(RiskLevel::Critical.cooldown_multiplier(), 3.0);
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,778 @@
|
||||
//! ProxyCast Plugin SDK
|
||||
//!
|
||||
//! 提供给 OAuth Provider 插件使用的 SDK 接口。
|
||||
//! 插件可以通过这些接口访问 ProxyCast 的核心功能。
|
||||
|
||||
use serde::{Deserialize, Serialize};
|
||||
use std::collections::HashMap;
|
||||
use std::sync::Arc;
|
||||
|
||||
/// SDK 错误类型
|
||||
#[derive(Debug, Clone, Serialize, Deserialize)]
|
||||
pub enum SdkError {
|
||||
/// 数据库错误
|
||||
DatabaseError(String),
|
||||
/// HTTP 错误
|
||||
HttpError(String),
|
||||
/// 加密错误
|
||||
CryptoError(String),
|
||||
/// 权限错误
|
||||
PermissionDenied(String),
|
||||
/// 未找到
|
||||
NotFound(String),
|
||||
/// 参数错误
|
||||
InvalidArgument(String),
|
||||
/// 内部错误
|
||||
InternalError(String),
|
||||
}
|
||||
|
||||
impl std::fmt::Display for SdkError {
|
||||
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
|
||||
match self {
|
||||
SdkError::DatabaseError(msg) => write!(f, "Database error: {}", msg),
|
||||
SdkError::HttpError(msg) => write!(f, "HTTP error: {}", msg),
|
||||
SdkError::CryptoError(msg) => write!(f, "Crypto error: {}", msg),
|
||||
SdkError::PermissionDenied(msg) => write!(f, "Permission denied: {}", msg),
|
||||
SdkError::NotFound(msg) => write!(f, "Not found: {}", msg),
|
||||
SdkError::InvalidArgument(msg) => write!(f, "Invalid argument: {}", msg),
|
||||
SdkError::InternalError(msg) => write!(f, "Internal error: {}", msg),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
impl std::error::Error for SdkError {}
|
||||
|
||||
/// SDK 结果类型
|
||||
pub type SdkResult<T> = Result<T, SdkError>;
|
||||
|
||||
/// HTTP 请求选项
|
||||
#[derive(Debug, Clone, Serialize, Deserialize, Default)]
|
||||
pub struct HttpRequestOptions {
|
||||
/// HTTP 方法
|
||||
#[serde(default = "default_method")]
|
||||
pub method: String,
|
||||
/// 请求头
|
||||
#[serde(default)]
|
||||
pub headers: HashMap<String, String>,
|
||||
/// 请求体
|
||||
#[serde(default)]
|
||||
pub body: Option<String>,
|
||||
/// 超时(毫秒)
|
||||
#[serde(default = "default_timeout")]
|
||||
pub timeout_ms: u64,
|
||||
}
|
||||
|
||||
fn default_method() -> String {
|
||||
"GET".to_string()
|
||||
}
|
||||
|
||||
fn default_timeout() -> u64 {
|
||||
30000
|
||||
}
|
||||
|
||||
/// HTTP 响应
|
||||
#[derive(Debug, Clone, Serialize, Deserialize)]
|
||||
pub struct HttpResponse {
|
||||
/// 状态码
|
||||
pub status: u16,
|
||||
/// 响应头
|
||||
pub headers: HashMap<String, String>,
|
||||
/// 响应体
|
||||
pub body: String,
|
||||
}
|
||||
|
||||
/// 数据库查询结果
|
||||
#[derive(Debug, Clone, Serialize, Deserialize)]
|
||||
pub struct QueryResult {
|
||||
/// 列名
|
||||
pub columns: Vec<String>,
|
||||
/// 行数据
|
||||
pub rows: Vec<Vec<serde_json::Value>>,
|
||||
}
|
||||
|
||||
/// 插件权限
|
||||
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash, Serialize, Deserialize)]
|
||||
pub enum PluginPermission {
|
||||
/// 读取数据库
|
||||
DatabaseRead,
|
||||
/// 写入数据库
|
||||
DatabaseWrite,
|
||||
/// 发送 HTTP 请求
|
||||
HttpRequest,
|
||||
/// 加密数据
|
||||
CryptoEncrypt,
|
||||
/// 解密数据
|
||||
CryptoDecrypt,
|
||||
/// 发送通知
|
||||
Notification,
|
||||
/// 发布事件
|
||||
EventEmit,
|
||||
/// 订阅事件
|
||||
EventSubscribe,
|
||||
/// 访问文件系统
|
||||
FileSystemRead,
|
||||
/// 写入文件系统
|
||||
FileSystemWrite,
|
||||
}
|
||||
|
||||
/// 数据库连接包装
|
||||
///
|
||||
/// 由于 rusqlite::Connection 不是 Send + Sync,我们使用回调模式
|
||||
pub type DatabaseCallback =
|
||||
Box<dyn Fn(&str, Vec<serde_json::Value>) -> Result<QueryResult, String> + Send + Sync>;
|
||||
|
||||
/// 插件 SDK 上下文
|
||||
///
|
||||
/// 提供给插件的 SDK 接口,包含所有可用的功能。
|
||||
pub struct PluginSdkContext {
|
||||
/// 插件 ID
|
||||
pub plugin_id: String,
|
||||
/// 授予的权限
|
||||
pub permissions: Vec<PluginPermission>,
|
||||
/// 数据库查询回调
|
||||
db_query_callback: Option<Arc<DatabaseCallback>>,
|
||||
/// HTTP 客户端
|
||||
http_client: reqwest::Client,
|
||||
}
|
||||
|
||||
impl PluginSdkContext {
|
||||
/// 创建新的 SDK 上下文
|
||||
pub fn new(plugin_id: String, permissions: Vec<PluginPermission>) -> Self {
|
||||
Self {
|
||||
plugin_id,
|
||||
permissions,
|
||||
db_query_callback: None,
|
||||
http_client: reqwest::Client::new(),
|
||||
}
|
||||
}
|
||||
|
||||
/// 设置数据库查询回调
|
||||
pub fn with_database_callback(mut self, callback: DatabaseCallback) -> Self {
|
||||
self.db_query_callback = Some(Arc::new(callback));
|
||||
self
|
||||
}
|
||||
|
||||
/// 检查权限
|
||||
fn check_permission(&self, required: PluginPermission) -> SdkResult<()> {
|
||||
if self.permissions.contains(&required) {
|
||||
Ok(())
|
||||
} else {
|
||||
Err(SdkError::PermissionDenied(format!(
|
||||
"Plugin '{}' does not have {:?} permission",
|
||||
self.plugin_id, required
|
||||
)))
|
||||
}
|
||||
}
|
||||
|
||||
// ========================================================================
|
||||
// 数据库操作
|
||||
// ========================================================================
|
||||
|
||||
/// 执行数据库查询
|
||||
pub async fn database_query(
|
||||
&self,
|
||||
sql: &str,
|
||||
params: Vec<serde_json::Value>,
|
||||
) -> SdkResult<QueryResult> {
|
||||
self.check_permission(PluginPermission::DatabaseRead)?;
|
||||
|
||||
let callback = self
|
||||
.db_query_callback
|
||||
.as_ref()
|
||||
.ok_or_else(|| SdkError::DatabaseError("Database not initialized".to_string()))?;
|
||||
|
||||
// 安全检查:只允许 SELECT 语句
|
||||
let sql_upper = sql.trim().to_uppercase();
|
||||
if !sql_upper.starts_with("SELECT") {
|
||||
return Err(SdkError::PermissionDenied(
|
||||
"Only SELECT queries are allowed for database_query".to_string(),
|
||||
));
|
||||
}
|
||||
|
||||
// 限制只能查询插件自己的表或公共表
|
||||
if !self.is_allowed_table(sql) {
|
||||
return Err(SdkError::PermissionDenied(
|
||||
"Access to this table is not allowed".to_string(),
|
||||
));
|
||||
}
|
||||
|
||||
// 执行数据库查询
|
||||
callback(sql, params).map_err(|e| SdkError::DatabaseError(e))
|
||||
}
|
||||
|
||||
/// 执行数据库写入
|
||||
pub async fn database_execute(
|
||||
&self,
|
||||
sql: &str,
|
||||
_params: Vec<serde_json::Value>,
|
||||
) -> SdkResult<u64> {
|
||||
self.check_permission(PluginPermission::DatabaseWrite)?;
|
||||
|
||||
let _callback = self
|
||||
.db_query_callback
|
||||
.as_ref()
|
||||
.ok_or_else(|| SdkError::DatabaseError("Database not initialized".to_string()))?;
|
||||
|
||||
// 限制只能操作插件自己的表
|
||||
if !self.is_plugin_table(sql) {
|
||||
return Err(SdkError::PermissionDenied(
|
||||
"Can only modify plugin-owned tables".to_string(),
|
||||
));
|
||||
}
|
||||
|
||||
// TODO: 执行实际的数据库写入
|
||||
Ok(0)
|
||||
}
|
||||
|
||||
/// 检查是否是允许访问的表
|
||||
fn is_allowed_table(&self, sql: &str) -> bool {
|
||||
let sql_lower = sql.to_lowercase();
|
||||
|
||||
// 允许访问的公共表
|
||||
let public_tables = ["credential_provider_plugins", "plugin_credentials"];
|
||||
|
||||
// 检查是否访问公共表
|
||||
for table in public_tables {
|
||||
if sql_lower.contains(table) {
|
||||
return true;
|
||||
}
|
||||
}
|
||||
|
||||
// 检查是否访问插件自己的表(以 plugin_{plugin_id}_ 为前缀)
|
||||
let plugin_prefix = format!("plugin_{}.", self.plugin_id.replace('-', "_"));
|
||||
sql_lower.contains(&plugin_prefix)
|
||||
}
|
||||
|
||||
/// 检查是否是插件自己的表
|
||||
fn is_plugin_table(&self, sql: &str) -> bool {
|
||||
let sql_lower = sql.to_lowercase();
|
||||
let plugin_prefix = format!("plugin_{}.", self.plugin_id.replace('-', "_"));
|
||||
sql_lower.contains(&plugin_prefix)
|
||||
}
|
||||
|
||||
// ========================================================================
|
||||
// HTTP 操作
|
||||
// ========================================================================
|
||||
|
||||
/// 发送 HTTP 请求
|
||||
pub async fn http_request(
|
||||
&self,
|
||||
url: &str,
|
||||
options: HttpRequestOptions,
|
||||
) -> SdkResult<HttpResponse> {
|
||||
self.check_permission(PluginPermission::HttpRequest)?;
|
||||
|
||||
let method = options.method.to_uppercase();
|
||||
let mut request = match method.as_str() {
|
||||
"GET" => self.http_client.get(url),
|
||||
"POST" => self.http_client.post(url),
|
||||
"PUT" => self.http_client.put(url),
|
||||
"DELETE" => self.http_client.delete(url),
|
||||
"PATCH" => self.http_client.patch(url),
|
||||
"HEAD" => self.http_client.head(url),
|
||||
_ => {
|
||||
return Err(SdkError::InvalidArgument(format!(
|
||||
"Unsupported HTTP method: {}",
|
||||
method
|
||||
)))
|
||||
}
|
||||
};
|
||||
|
||||
// 添加请求头
|
||||
for (key, value) in options.headers {
|
||||
request = request.header(&key, &value);
|
||||
}
|
||||
|
||||
// 添加请求体
|
||||
if let Some(body) = options.body {
|
||||
request = request.body(body);
|
||||
}
|
||||
|
||||
// 设置超时
|
||||
request = request.timeout(std::time::Duration::from_millis(options.timeout_ms));
|
||||
|
||||
// 发送请求
|
||||
let response = request
|
||||
.send()
|
||||
.await
|
||||
.map_err(|e| SdkError::HttpError(e.to_string()))?;
|
||||
|
||||
let status = response.status().as_u16();
|
||||
let headers: HashMap<String, String> = response
|
||||
.headers()
|
||||
.iter()
|
||||
.filter_map(|(k, v)| {
|
||||
v.to_str()
|
||||
.ok()
|
||||
.map(|v| (k.as_str().to_string(), v.to_string()))
|
||||
})
|
||||
.collect();
|
||||
|
||||
let body = response
|
||||
.text()
|
||||
.await
|
||||
.map_err(|e| SdkError::HttpError(e.to_string()))?;
|
||||
|
||||
Ok(HttpResponse {
|
||||
status,
|
||||
headers,
|
||||
body,
|
||||
})
|
||||
}
|
||||
|
||||
// ========================================================================
|
||||
// 加密操作
|
||||
// ========================================================================
|
||||
|
||||
/// 加密数据
|
||||
pub async fn crypto_encrypt(&self, data: &str) -> SdkResult<String> {
|
||||
self.check_permission(PluginPermission::CryptoEncrypt)?;
|
||||
|
||||
// TODO: 使用 ProxyCast 的加密服务
|
||||
// 暂时使用 base64 编码作为占位符
|
||||
use base64::Engine;
|
||||
Ok(base64::engine::general_purpose::STANDARD.encode(data.as_bytes()))
|
||||
}
|
||||
|
||||
/// 解密数据
|
||||
pub async fn crypto_decrypt(&self, data: &str) -> SdkResult<String> {
|
||||
self.check_permission(PluginPermission::CryptoDecrypt)?;
|
||||
|
||||
// TODO: 使用 ProxyCast 的解密服务
|
||||
// 暂时使用 base64 解码作为占位符
|
||||
use base64::Engine;
|
||||
let bytes = base64::engine::general_purpose::STANDARD
|
||||
.decode(data)
|
||||
.map_err(|e| SdkError::CryptoError(e.to_string()))?;
|
||||
|
||||
String::from_utf8(bytes).map_err(|e| SdkError::CryptoError(e.to_string()))
|
||||
}
|
||||
|
||||
// ========================================================================
|
||||
// 通知操作
|
||||
// ========================================================================
|
||||
|
||||
/// 发送成功通知
|
||||
pub fn notification_success(&self, message: &str) -> SdkResult<()> {
|
||||
self.check_permission(PluginPermission::Notification)?;
|
||||
tracing::info!("[Plugin {}] Success: {}", self.plugin_id, message);
|
||||
// TODO: 发送到前端通知系统
|
||||
Ok(())
|
||||
}
|
||||
|
||||
/// 发送错误通知
|
||||
pub fn notification_error(&self, message: &str) -> SdkResult<()> {
|
||||
self.check_permission(PluginPermission::Notification)?;
|
||||
tracing::error!("[Plugin {}] Error: {}", self.plugin_id, message);
|
||||
// TODO: 发送到前端通知系统
|
||||
Ok(())
|
||||
}
|
||||
|
||||
/// 发送信息通知
|
||||
pub fn notification_info(&self, message: &str) -> SdkResult<()> {
|
||||
self.check_permission(PluginPermission::Notification)?;
|
||||
tracing::info!("[Plugin {}] Info: {}", self.plugin_id, message);
|
||||
// TODO: 发送到前端通知系统
|
||||
Ok(())
|
||||
}
|
||||
|
||||
// ========================================================================
|
||||
// 事件操作
|
||||
// ========================================================================
|
||||
|
||||
/// 发布事件
|
||||
pub fn event_emit(&self, event: &str, data: serde_json::Value) -> SdkResult<()> {
|
||||
self.check_permission(PluginPermission::EventEmit)?;
|
||||
tracing::debug!(
|
||||
"[Plugin {}] Emitting event '{}': {:?}",
|
||||
self.plugin_id,
|
||||
event,
|
||||
data
|
||||
);
|
||||
// TODO: 通过事件总线发布事件
|
||||
Ok(())
|
||||
}
|
||||
|
||||
// ========================================================================
|
||||
// 插件存储
|
||||
// ========================================================================
|
||||
|
||||
/// 获取插件存储的值
|
||||
pub async fn storage_get(&self, _key: &str) -> SdkResult<Option<String>> {
|
||||
self.check_permission(PluginPermission::DatabaseRead)?;
|
||||
|
||||
// TODO: 从插件存储表读取
|
||||
Ok(None)
|
||||
}
|
||||
|
||||
/// 设置插件存储的值
|
||||
pub async fn storage_set(&self, _key: &str, _value: &str) -> SdkResult<()> {
|
||||
self.check_permission(PluginPermission::DatabaseWrite)?;
|
||||
|
||||
// TODO: 写入插件存储表
|
||||
Ok(())
|
||||
}
|
||||
|
||||
/// 删除插件存储的值
|
||||
pub async fn storage_delete(&self, _key: &str) -> SdkResult<()> {
|
||||
self.check_permission(PluginPermission::DatabaseWrite)?;
|
||||
|
||||
// TODO: 从插件存储表删除
|
||||
Ok(())
|
||||
}
|
||||
}
|
||||
|
||||
/// JSON-RPC 请求格式
|
||||
#[derive(Debug, Clone, Serialize, Deserialize)]
|
||||
pub struct JsonRpcRequest {
|
||||
/// JSON-RPC 版本
|
||||
pub jsonrpc: String,
|
||||
/// 方法名
|
||||
pub method: String,
|
||||
/// 参数
|
||||
pub params: serde_json::Value,
|
||||
/// 请求 ID
|
||||
pub id: serde_json::Value,
|
||||
}
|
||||
|
||||
/// JSON-RPC 响应格式
|
||||
#[derive(Debug, Clone, Serialize, Deserialize)]
|
||||
pub struct JsonRpcResponse {
|
||||
/// JSON-RPC 版本
|
||||
pub jsonrpc: String,
|
||||
/// 结果(成功时)
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub result: Option<serde_json::Value>,
|
||||
/// 错误(失败时)
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub error: Option<JsonRpcError>,
|
||||
/// 请求 ID
|
||||
pub id: serde_json::Value,
|
||||
}
|
||||
|
||||
/// JSON-RPC 错误
|
||||
#[derive(Debug, Clone, Serialize, Deserialize)]
|
||||
pub struct JsonRpcError {
|
||||
/// 错误码
|
||||
pub code: i32,
|
||||
/// 错误消息
|
||||
pub message: String,
|
||||
/// 附加数据
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub data: Option<serde_json::Value>,
|
||||
}
|
||||
|
||||
impl JsonRpcResponse {
|
||||
/// 创建成功响应
|
||||
pub fn success(id: serde_json::Value, result: serde_json::Value) -> Self {
|
||||
Self {
|
||||
jsonrpc: "2.0".to_string(),
|
||||
result: Some(result),
|
||||
error: None,
|
||||
id,
|
||||
}
|
||||
}
|
||||
|
||||
/// 创建错误响应
|
||||
pub fn error(id: serde_json::Value, code: i32, message: String) -> Self {
|
||||
Self {
|
||||
jsonrpc: "2.0".to_string(),
|
||||
result: None,
|
||||
error: Some(JsonRpcError {
|
||||
code,
|
||||
message,
|
||||
data: None,
|
||||
}),
|
||||
id,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
/// SDK 方法处理器
|
||||
///
|
||||
/// 处理来自外部插件的 SDK 调用请求
|
||||
pub struct SdkMethodHandler {
|
||||
context: PluginSdkContext,
|
||||
}
|
||||
|
||||
impl SdkMethodHandler {
|
||||
/// 创建新的处理器
|
||||
pub fn new(context: PluginSdkContext) -> Self {
|
||||
Self { context }
|
||||
}
|
||||
|
||||
/// 处理 JSON-RPC 请求
|
||||
pub async fn handle(&self, request: JsonRpcRequest) -> JsonRpcResponse {
|
||||
match request.method.as_str() {
|
||||
// 数据库方法
|
||||
"database.query" => self.handle_database_query(request).await,
|
||||
"database.execute" => self.handle_database_execute(request).await,
|
||||
|
||||
// HTTP 方法
|
||||
"http.request" => self.handle_http_request(request).await,
|
||||
|
||||
// 加密方法
|
||||
"crypto.encrypt" => self.handle_crypto_encrypt(request).await,
|
||||
"crypto.decrypt" => self.handle_crypto_decrypt(request).await,
|
||||
|
||||
// 通知方法
|
||||
"notification.success" => self.handle_notification(request, "success"),
|
||||
"notification.error" => self.handle_notification(request, "error"),
|
||||
"notification.info" => self.handle_notification(request, "info"),
|
||||
|
||||
// 事件方法
|
||||
"event.emit" => self.handle_event_emit(request),
|
||||
|
||||
// 存储方法
|
||||
"storage.get" => self.handle_storage_get(request).await,
|
||||
"storage.set" => self.handle_storage_set(request).await,
|
||||
"storage.delete" => self.handle_storage_delete(request).await,
|
||||
|
||||
// 未知方法
|
||||
_ => JsonRpcResponse::error(
|
||||
request.id,
|
||||
-32601,
|
||||
format!("Method not found: {}", request.method),
|
||||
),
|
||||
}
|
||||
}
|
||||
|
||||
async fn handle_database_query(&self, request: JsonRpcRequest) -> JsonRpcResponse {
|
||||
#[derive(Deserialize)]
|
||||
struct Params {
|
||||
sql: String,
|
||||
#[serde(default)]
|
||||
params: Vec<serde_json::Value>,
|
||||
}
|
||||
|
||||
match serde_json::from_value::<Params>(request.params.clone()) {
|
||||
Ok(params) => match self
|
||||
.context
|
||||
.database_query(¶ms.sql, params.params)
|
||||
.await
|
||||
{
|
||||
Ok(result) => {
|
||||
JsonRpcResponse::success(request.id, serde_json::to_value(result).unwrap())
|
||||
}
|
||||
Err(e) => JsonRpcResponse::error(request.id, -32000, e.to_string()),
|
||||
},
|
||||
Err(e) => JsonRpcResponse::error(request.id, -32602, format!("Invalid params: {}", e)),
|
||||
}
|
||||
}
|
||||
|
||||
async fn handle_database_execute(&self, request: JsonRpcRequest) -> JsonRpcResponse {
|
||||
#[derive(Deserialize)]
|
||||
struct Params {
|
||||
sql: String,
|
||||
#[serde(default)]
|
||||
params: Vec<serde_json::Value>,
|
||||
}
|
||||
|
||||
match serde_json::from_value::<Params>(request.params.clone()) {
|
||||
Ok(params) => {
|
||||
match self
|
||||
.context
|
||||
.database_execute(¶ms.sql, params.params)
|
||||
.await
|
||||
{
|
||||
Ok(affected) => JsonRpcResponse::success(
|
||||
request.id,
|
||||
serde_json::json!({ "affected": affected }),
|
||||
),
|
||||
Err(e) => JsonRpcResponse::error(request.id, -32000, e.to_string()),
|
||||
}
|
||||
}
|
||||
Err(e) => JsonRpcResponse::error(request.id, -32602, format!("Invalid params: {}", e)),
|
||||
}
|
||||
}
|
||||
|
||||
async fn handle_http_request(&self, request: JsonRpcRequest) -> JsonRpcResponse {
|
||||
#[derive(Deserialize)]
|
||||
struct Params {
|
||||
url: String,
|
||||
#[serde(default)]
|
||||
options: HttpRequestOptions,
|
||||
}
|
||||
|
||||
match serde_json::from_value::<Params>(request.params.clone()) {
|
||||
Ok(params) => match self.context.http_request(¶ms.url, params.options).await {
|
||||
Ok(response) => {
|
||||
JsonRpcResponse::success(request.id, serde_json::to_value(response).unwrap())
|
||||
}
|
||||
Err(e) => JsonRpcResponse::error(request.id, -32000, e.to_string()),
|
||||
},
|
||||
Err(e) => JsonRpcResponse::error(request.id, -32602, format!("Invalid params: {}", e)),
|
||||
}
|
||||
}
|
||||
|
||||
async fn handle_crypto_encrypt(&self, request: JsonRpcRequest) -> JsonRpcResponse {
|
||||
#[derive(Deserialize)]
|
||||
struct Params {
|
||||
data: String,
|
||||
}
|
||||
|
||||
match serde_json::from_value::<Params>(request.params.clone()) {
|
||||
Ok(params) => match self.context.crypto_encrypt(¶ms.data).await {
|
||||
Ok(encrypted) => JsonRpcResponse::success(
|
||||
request.id,
|
||||
serde_json::json!({ "encrypted": encrypted }),
|
||||
),
|
||||
Err(e) => JsonRpcResponse::error(request.id, -32000, e.to_string()),
|
||||
},
|
||||
Err(e) => JsonRpcResponse::error(request.id, -32602, format!("Invalid params: {}", e)),
|
||||
}
|
||||
}
|
||||
|
||||
async fn handle_crypto_decrypt(&self, request: JsonRpcRequest) -> JsonRpcResponse {
|
||||
#[derive(Deserialize)]
|
||||
struct Params {
|
||||
data: String,
|
||||
}
|
||||
|
||||
match serde_json::from_value::<Params>(request.params.clone()) {
|
||||
Ok(params) => match self.context.crypto_decrypt(¶ms.data).await {
|
||||
Ok(decrypted) => JsonRpcResponse::success(
|
||||
request.id,
|
||||
serde_json::json!({ "decrypted": decrypted }),
|
||||
),
|
||||
Err(e) => JsonRpcResponse::error(request.id, -32000, e.to_string()),
|
||||
},
|
||||
Err(e) => JsonRpcResponse::error(request.id, -32602, format!("Invalid params: {}", e)),
|
||||
}
|
||||
}
|
||||
|
||||
fn handle_notification(&self, request: JsonRpcRequest, level: &str) -> JsonRpcResponse {
|
||||
#[derive(Deserialize)]
|
||||
struct Params {
|
||||
message: String,
|
||||
}
|
||||
|
||||
match serde_json::from_value::<Params>(request.params.clone()) {
|
||||
Ok(params) => {
|
||||
let result = match level {
|
||||
"success" => self.context.notification_success(¶ms.message),
|
||||
"error" => self.context.notification_error(¶ms.message),
|
||||
"info" => self.context.notification_info(¶ms.message),
|
||||
_ => Ok(()),
|
||||
};
|
||||
match result {
|
||||
Ok(()) => JsonRpcResponse::success(request.id, serde_json::json!({})),
|
||||
Err(e) => JsonRpcResponse::error(request.id, -32000, e.to_string()),
|
||||
}
|
||||
}
|
||||
Err(e) => JsonRpcResponse::error(request.id, -32602, format!("Invalid params: {}", e)),
|
||||
}
|
||||
}
|
||||
|
||||
fn handle_event_emit(&self, request: JsonRpcRequest) -> JsonRpcResponse {
|
||||
#[derive(Deserialize)]
|
||||
struct Params {
|
||||
event: String,
|
||||
data: serde_json::Value,
|
||||
}
|
||||
|
||||
match serde_json::from_value::<Params>(request.params.clone()) {
|
||||
Ok(params) => match self.context.event_emit(¶ms.event, params.data) {
|
||||
Ok(()) => JsonRpcResponse::success(request.id, serde_json::json!({})),
|
||||
Err(e) => JsonRpcResponse::error(request.id, -32000, e.to_string()),
|
||||
},
|
||||
Err(e) => JsonRpcResponse::error(request.id, -32602, format!("Invalid params: {}", e)),
|
||||
}
|
||||
}
|
||||
|
||||
async fn handle_storage_get(&self, request: JsonRpcRequest) -> JsonRpcResponse {
|
||||
#[derive(Deserialize)]
|
||||
struct Params {
|
||||
key: String,
|
||||
}
|
||||
|
||||
match serde_json::from_value::<Params>(request.params.clone()) {
|
||||
Ok(params) => match self.context.storage_get(¶ms.key).await {
|
||||
Ok(value) => {
|
||||
JsonRpcResponse::success(request.id, serde_json::json!({ "value": value }))
|
||||
}
|
||||
Err(e) => JsonRpcResponse::error(request.id, -32000, e.to_string()),
|
||||
},
|
||||
Err(e) => JsonRpcResponse::error(request.id, -32602, format!("Invalid params: {}", e)),
|
||||
}
|
||||
}
|
||||
|
||||
async fn handle_storage_set(&self, request: JsonRpcRequest) -> JsonRpcResponse {
|
||||
#[derive(Deserialize)]
|
||||
struct Params {
|
||||
key: String,
|
||||
value: String,
|
||||
}
|
||||
|
||||
match serde_json::from_value::<Params>(request.params.clone()) {
|
||||
Ok(params) => match self.context.storage_set(¶ms.key, ¶ms.value).await {
|
||||
Ok(()) => JsonRpcResponse::success(request.id, serde_json::json!({})),
|
||||
Err(e) => JsonRpcResponse::error(request.id, -32000, e.to_string()),
|
||||
},
|
||||
Err(e) => JsonRpcResponse::error(request.id, -32602, format!("Invalid params: {}", e)),
|
||||
}
|
||||
}
|
||||
|
||||
async fn handle_storage_delete(&self, request: JsonRpcRequest) -> JsonRpcResponse {
|
||||
#[derive(Deserialize)]
|
||||
struct Params {
|
||||
key: String,
|
||||
}
|
||||
|
||||
match serde_json::from_value::<Params>(request.params.clone()) {
|
||||
Ok(params) => match self.context.storage_delete(¶ms.key).await {
|
||||
Ok(()) => JsonRpcResponse::success(request.id, serde_json::json!({})),
|
||||
Err(e) => JsonRpcResponse::error(request.id, -32000, e.to_string()),
|
||||
},
|
||||
Err(e) => JsonRpcResponse::error(request.id, -32602, format!("Invalid params: {}", e)),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
|
||||
#[test]
|
||||
fn test_sdk_context_permission_check() {
|
||||
let context = PluginSdkContext::new(
|
||||
"test-plugin".to_string(),
|
||||
vec![PluginPermission::DatabaseRead],
|
||||
);
|
||||
|
||||
assert!(context
|
||||
.check_permission(PluginPermission::DatabaseRead)
|
||||
.is_ok());
|
||||
assert!(context
|
||||
.check_permission(PluginPermission::DatabaseWrite)
|
||||
.is_err());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_json_rpc_response() {
|
||||
let success =
|
||||
JsonRpcResponse::success(serde_json::json!(1), serde_json::json!({"result": "ok"}));
|
||||
assert!(success.result.is_some());
|
||||
assert!(success.error.is_none());
|
||||
|
||||
let error = JsonRpcResponse::error(serde_json::json!(1), -32000, "Error".to_string());
|
||||
assert!(error.result.is_none());
|
||||
assert!(error.error.is_some());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_is_allowed_table() {
|
||||
let context = PluginSdkContext::new("kiro-provider".to_string(), vec![]);
|
||||
|
||||
// 公共表应该允许
|
||||
assert!(context.is_allowed_table("SELECT * FROM credential_provider_plugins"));
|
||||
assert!(context.is_allowed_table("SELECT * FROM plugin_credentials"));
|
||||
|
||||
// 插件自己的表应该允许
|
||||
assert!(context.is_allowed_table("SELECT * FROM plugin_kiro_provider.accounts"));
|
||||
|
||||
// 其他表应该禁止
|
||||
assert!(!context.is_allowed_table("SELECT * FROM api_keys"));
|
||||
assert!(!context.is_allowed_table("SELECT * FROM plugin_other.data"));
|
||||
}
|
||||
}
|
||||
@@ -237,6 +237,19 @@ impl CredentialSyncService {
|
||||
"iFlow Cookie 凭证暂不支持同步到配置".to_string(),
|
||||
));
|
||||
}
|
||||
CredentialData::AnthropicKey { api_key, base_url } => {
|
||||
// Anthropic API Key 保存到 claude 配置(使用相同的 API 格式)
|
||||
let entry = ApiKeyEntry {
|
||||
id: credential.uuid.clone(),
|
||||
api_key: api_key.clone(),
|
||||
base_url: base_url.clone(),
|
||||
disabled: credential.is_disabled,
|
||||
proxy_url: None,
|
||||
};
|
||||
// 注意:Anthropic 凭证保存到单独的 anthropic 配置(如果有的话)
|
||||
// 目前暂时保存到 claude 配置中
|
||||
config.credential_pool.claude.push(entry);
|
||||
}
|
||||
}
|
||||
|
||||
self.update_config(config)
|
||||
@@ -396,6 +409,15 @@ impl CredentialSyncService {
|
||||
"iFlow 凭证暂不支持同步到配置".to_string(),
|
||||
));
|
||||
}
|
||||
// API Key Provider 类型 - 不支持同步到配置
|
||||
PoolProviderType::Anthropic
|
||||
| PoolProviderType::AzureOpenai
|
||||
| PoolProviderType::AwsBedrock
|
||||
| PoolProviderType::Ollama => {
|
||||
return Err(SyncError::InvalidCredentialType(
|
||||
"API Key Provider 凭证不支持同步到配置".to_string(),
|
||||
));
|
||||
}
|
||||
}
|
||||
|
||||
if !found {
|
||||
@@ -571,6 +593,20 @@ impl CredentialSyncService {
|
||||
"iFlow Cookie 凭证暂不支持同步到配置".to_string(),
|
||||
));
|
||||
}
|
||||
CredentialData::AnthropicKey { api_key, base_url } => {
|
||||
// Anthropic API Key 更新到 claude 配置
|
||||
if let Some(entry) = config
|
||||
.credential_pool
|
||||
.claude
|
||||
.iter_mut()
|
||||
.find(|e| e.id == credential.uuid)
|
||||
{
|
||||
entry.api_key = api_key.clone();
|
||||
entry.base_url = base_url.clone();
|
||||
entry.disabled = credential.is_disabled;
|
||||
found = true;
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
if !found {
|
||||
|
||||
@@ -0,0 +1,328 @@
|
||||
//! 统一凭证管理器
|
||||
//!
|
||||
//! 整合 orchestrator 和 credential 模块,提供统一的凭证管理接口。
|
||||
//!
|
||||
//! ## 功能
|
||||
//!
|
||||
//! - 统一的凭证获取接口
|
||||
//! - 自动风控和冷却管理
|
||||
//! - 与 orchestrator 的模型选择集成
|
||||
|
||||
use super::balancer::{CredentialSelection, LoadBalancer};
|
||||
use super::pool::{CredentialPool, PoolError};
|
||||
use super::risk::{CooldownConfig, RateLimitEvent, RiskController, RiskLevel};
|
||||
use super::types::{Credential, CredentialData};
|
||||
use crate::orchestrator::get_global_orchestrator;
|
||||
use crate::ProviderType;
|
||||
use chrono::Duration;
|
||||
use std::sync::Arc;
|
||||
use tokio::sync::RwLock;
|
||||
use tracing::{debug, info, warn};
|
||||
|
||||
/// 统一凭证管理器
|
||||
///
|
||||
/// 整合 orchestrator 的模型选择和 credential 的凭证管理
|
||||
pub struct UnifiedCredentialManager {
|
||||
/// 负载均衡器
|
||||
load_balancer: LoadBalancer,
|
||||
/// 风控控制器
|
||||
risk_controller: RiskController,
|
||||
/// 是否启用风控
|
||||
risk_control_enabled: RwLock<bool>,
|
||||
}
|
||||
|
||||
impl UnifiedCredentialManager {
|
||||
/// 创建新的统一凭证管理器
|
||||
pub fn new() -> Self {
|
||||
Self {
|
||||
load_balancer: LoadBalancer::round_robin(),
|
||||
risk_controller: RiskController::with_defaults(),
|
||||
risk_control_enabled: RwLock::new(true),
|
||||
}
|
||||
}
|
||||
|
||||
/// 使用自定义配置创建
|
||||
pub fn with_config(cooldown_config: CooldownConfig) -> Self {
|
||||
Self {
|
||||
load_balancer: LoadBalancer::round_robin(),
|
||||
risk_controller: RiskController::new(cooldown_config),
|
||||
risk_control_enabled: RwLock::new(true),
|
||||
}
|
||||
}
|
||||
|
||||
/// 获取负载均衡器
|
||||
pub fn load_balancer(&self) -> &LoadBalancer {
|
||||
&self.load_balancer
|
||||
}
|
||||
|
||||
/// 获取风控控制器
|
||||
pub fn risk_controller(&self) -> &RiskController {
|
||||
&self.risk_controller
|
||||
}
|
||||
|
||||
/// 设置是否启用风控
|
||||
pub async fn set_risk_control_enabled(&self, enabled: bool) {
|
||||
let mut flag = self.risk_control_enabled.write().await;
|
||||
*flag = enabled;
|
||||
}
|
||||
|
||||
/// 检查风控是否启用
|
||||
pub async fn is_risk_control_enabled(&self) -> bool {
|
||||
*self.risk_control_enabled.read().await
|
||||
}
|
||||
|
||||
/// 注册凭证池
|
||||
pub fn register_pool(&self, pool: Arc<CredentialPool>) {
|
||||
self.load_balancer.register_pool(pool);
|
||||
}
|
||||
|
||||
/// 选择凭证(带风控检查)
|
||||
///
|
||||
/// # 参数
|
||||
/// - `provider`: Provider 类型
|
||||
///
|
||||
/// # 返回
|
||||
/// - `Ok(CredentialSelection)`: 选中的凭证和 HTTP 客户端
|
||||
/// - `Err(PoolError)`: 选择失败
|
||||
pub async fn select_credential(
|
||||
&self,
|
||||
provider: ProviderType,
|
||||
) -> Result<CredentialSelection, PoolError> {
|
||||
let risk_enabled = self.is_risk_control_enabled().await;
|
||||
|
||||
// 如果启用风控,先检查是否有凭证在冷却中
|
||||
if risk_enabled {
|
||||
let cooling = self.risk_controller.get_cooling_credentials();
|
||||
if !cooling.is_empty() {
|
||||
debug!("有 {} 个凭证在冷却中", cooling.len());
|
||||
}
|
||||
}
|
||||
|
||||
// 使用负载均衡器选择凭证
|
||||
let selection = self.load_balancer.select_with_client(provider)?;
|
||||
|
||||
// 检查选中的凭证是否在冷却中
|
||||
if risk_enabled
|
||||
&& self
|
||||
.risk_controller
|
||||
.is_in_cooldown(&selection.credential.id)
|
||||
{
|
||||
warn!(
|
||||
"凭证 {} 在冷却中,尝试选择其他凭证",
|
||||
selection.credential.id
|
||||
);
|
||||
// 尝试故障转移
|
||||
return self.load_balancer.select_with_failover(provider, None);
|
||||
}
|
||||
|
||||
Ok(selection)
|
||||
}
|
||||
|
||||
/// 报告请求成功
|
||||
pub fn report_success(&self, provider: ProviderType, credential_id: &str, latency_ms: u64) {
|
||||
// 更新负载均衡器统计
|
||||
let _ = self
|
||||
.load_balancer
|
||||
.report(provider, credential_id, true, latency_ms);
|
||||
|
||||
// 更新风控状态
|
||||
self.risk_controller.record_success(credential_id);
|
||||
}
|
||||
|
||||
/// 报告请求失败
|
||||
///
|
||||
/// # 参数
|
||||
/// - `provider`: Provider 类型
|
||||
/// - `credential_id`: 凭证 ID
|
||||
/// - `status_code`: HTTP 状态码
|
||||
/// - `error_body`: 错误响应体
|
||||
/// - `retry_after`: Retry-After 头的值
|
||||
///
|
||||
/// # 返回
|
||||
/// 如果是限流错误,返回建议的冷却时间(秒)
|
||||
pub async fn report_failure(
|
||||
&self,
|
||||
provider: ProviderType,
|
||||
credential_id: &str,
|
||||
status_code: Option<u16>,
|
||||
error_body: Option<&str>,
|
||||
retry_after: Option<&str>,
|
||||
) -> Option<u64> {
|
||||
// 更新负载均衡器统计
|
||||
let _ = self.load_balancer.report(provider, credential_id, false, 0);
|
||||
|
||||
// 检查是否为限流错误
|
||||
let is_rate_limit = status_code
|
||||
.map(|code| RiskController::is_rate_limit_error(code, error_body))
|
||||
.unwrap_or(false);
|
||||
|
||||
if !is_rate_limit {
|
||||
return None;
|
||||
}
|
||||
|
||||
// 解析 Retry-After
|
||||
let retry_after_secs = retry_after.and_then(RiskController::parse_retry_after);
|
||||
|
||||
// 记录限流事件
|
||||
let mut event = RateLimitEvent::new(credential_id.to_string());
|
||||
if let Some(code) = status_code {
|
||||
event = event.with_status_code(code);
|
||||
}
|
||||
if let Some(body) = error_body {
|
||||
event = event.with_error_message(body.to_string());
|
||||
}
|
||||
if let Some(secs) = retry_after_secs {
|
||||
event = event.with_retry_after(secs);
|
||||
}
|
||||
|
||||
let cooldown_secs = self.risk_controller.record_rate_limit(event);
|
||||
|
||||
// 在负载均衡器中标记冷却
|
||||
let _ = self.load_balancer.mark_cooldown(
|
||||
provider,
|
||||
credential_id,
|
||||
Duration::seconds(cooldown_secs as i64),
|
||||
);
|
||||
|
||||
info!("凭证 {} 触发限流,冷却 {} 秒", credential_id, cooldown_secs);
|
||||
|
||||
Some(cooldown_secs)
|
||||
}
|
||||
|
||||
/// 获取凭证的风险等级
|
||||
pub fn get_risk_level(&self, credential_id: &str) -> RiskLevel {
|
||||
self.risk_controller.get_risk_level(credential_id)
|
||||
}
|
||||
|
||||
/// 手动清除凭证的冷却状态
|
||||
pub fn clear_cooldown(&self, provider: ProviderType, credential_id: &str) {
|
||||
self.risk_controller.clear_cooldown(credential_id);
|
||||
let _ = self.load_balancer.mark_active(provider, credential_id);
|
||||
}
|
||||
|
||||
/// 从 orchestrator 同步凭证到凭证池
|
||||
///
|
||||
/// 将 orchestrator 的 CredentialInfo 转换为 credential 模块的 Credential
|
||||
pub async fn sync_from_orchestrator(&self) -> Result<usize, String> {
|
||||
let orchestrator = get_global_orchestrator().ok_or("编排器未初始化")?;
|
||||
|
||||
// 获取所有可用模型
|
||||
let models = orchestrator.get_all_models().await;
|
||||
|
||||
let mut synced_count = 0;
|
||||
|
||||
// 按 provider 分组
|
||||
for model in models {
|
||||
let provider_type = self.map_orchestrator_provider(&model.provider_type);
|
||||
|
||||
// 获取或创建凭证池
|
||||
let pool = self
|
||||
.load_balancer
|
||||
.get_pool(provider_type)
|
||||
.unwrap_or_else(|| {
|
||||
let new_pool = Arc::new(CredentialPool::new(provider_type));
|
||||
self.load_balancer.register_pool(new_pool.clone());
|
||||
new_pool
|
||||
});
|
||||
|
||||
// 检查凭证是否已存在
|
||||
if pool.get(&model.credential_id).is_none() {
|
||||
// 创建新凭证
|
||||
let credential = Credential::new(
|
||||
model.credential_id.clone(),
|
||||
provider_type,
|
||||
CredentialData::ApiKey {
|
||||
key: format!("synced-{}", model.credential_id),
|
||||
base_url: None,
|
||||
},
|
||||
);
|
||||
|
||||
if pool.add(credential).is_ok() {
|
||||
synced_count += 1;
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
info!("从 orchestrator 同步了 {} 个凭证", synced_count);
|
||||
Ok(synced_count)
|
||||
}
|
||||
|
||||
/// 映射 orchestrator 的 ProviderType 到 credential 的 ProviderType
|
||||
fn map_orchestrator_provider(&self, provider: &str) -> ProviderType {
|
||||
match provider.to_lowercase().as_str() {
|
||||
"anthropic" => ProviderType::ClaudeOAuth,
|
||||
"openai" => ProviderType::Codex,
|
||||
"google" | "gemini" => ProviderType::Gemini,
|
||||
"kiro" => ProviderType::Kiro,
|
||||
_ => ProviderType::Kiro, // 默认
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
impl Default for UnifiedCredentialManager {
|
||||
fn default() -> Self {
|
||||
Self::new()
|
||||
}
|
||||
}
|
||||
|
||||
/// 全局统一凭证管理器
|
||||
static GLOBAL_UNIFIED_MANAGER: once_cell::sync::OnceCell<Arc<UnifiedCredentialManager>> =
|
||||
once_cell::sync::OnceCell::new();
|
||||
|
||||
/// 初始化全局统一凭证管理器
|
||||
pub fn init_global_unified_manager() -> Arc<UnifiedCredentialManager> {
|
||||
GLOBAL_UNIFIED_MANAGER
|
||||
.get_or_init(|| Arc::new(UnifiedCredentialManager::new()))
|
||||
.clone()
|
||||
}
|
||||
|
||||
/// 获取全局统一凭证管理器
|
||||
pub fn get_global_unified_manager() -> Option<Arc<UnifiedCredentialManager>> {
|
||||
GLOBAL_UNIFIED_MANAGER.get().cloned()
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
|
||||
#[test]
|
||||
fn test_unified_manager_new() {
|
||||
let manager = UnifiedCredentialManager::new();
|
||||
assert!(manager.load_balancer().providers().is_empty());
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_risk_control_toggle() {
|
||||
let manager = UnifiedCredentialManager::new();
|
||||
|
||||
assert!(manager.is_risk_control_enabled().await);
|
||||
|
||||
manager.set_risk_control_enabled(false).await;
|
||||
assert!(!manager.is_risk_control_enabled().await);
|
||||
|
||||
manager.set_risk_control_enabled(true).await;
|
||||
assert!(manager.is_risk_control_enabled().await);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_map_orchestrator_provider() {
|
||||
let manager = UnifiedCredentialManager::new();
|
||||
|
||||
assert_eq!(
|
||||
manager.map_orchestrator_provider("anthropic"),
|
||||
ProviderType::ClaudeOAuth
|
||||
);
|
||||
assert_eq!(
|
||||
manager.map_orchestrator_provider("openai"),
|
||||
ProviderType::Codex
|
||||
);
|
||||
assert_eq!(
|
||||
manager.map_orchestrator_provider("google"),
|
||||
ProviderType::Gemini
|
||||
);
|
||||
assert_eq!(
|
||||
manager.map_orchestrator_provider("kiro"),
|
||||
ProviderType::Kiro
|
||||
);
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,55 @@
|
||||
# 数据库模块
|
||||
|
||||
本模块负责 SQLite 数据库的初始化、表结构定义和数据迁移。
|
||||
|
||||
## 文件索引
|
||||
|
||||
| 文件 | 说明 |
|
||||
|------|------|
|
||||
| `mod.rs` | 模块入口,数据库初始化 |
|
||||
| `schema.rs` | 表结构定义和创建 |
|
||||
| `migration.rs` | 数据迁移逻辑 |
|
||||
| `system_providers.rs` | 系统预设 Provider 配置 |
|
||||
| `dao/` | 数据访问对象层 |
|
||||
|
||||
## 数据库表
|
||||
|
||||
### 核心表
|
||||
|
||||
- `api_key_providers` - API Key Provider 配置
|
||||
- `api_keys` - API Key 条目(已迁移到 provider_pool_credentials)
|
||||
- `provider_pool_credentials` - 凭证池(统一管理所有凭证)
|
||||
- `providers` - Provider 配置
|
||||
- `settings` - 应用设置
|
||||
|
||||
### 功能表
|
||||
|
||||
- `mcp_servers` - MCP 服务器配置
|
||||
- `prompts` - 提示词模板
|
||||
- `skills` - 技能配置
|
||||
- `skill_repos` - 技能仓库
|
||||
- `installed_plugins` - 已安装插件
|
||||
|
||||
## 数据迁移
|
||||
|
||||
### API Keys 迁移
|
||||
|
||||
`migrate_api_keys_to_pool()` 函数将 `api_keys` 表中的数据迁移到 `provider_pool_credentials` 表:
|
||||
|
||||
- 根据 provider_type 自动转换为对应的 CredentialData 类型
|
||||
- 保留使用统计和错误计数
|
||||
- 标记来源为 `imported`
|
||||
- 迁移完成后设置 `migrated_api_keys_to_pool` 标记,避免重复迁移
|
||||
|
||||
## 使用示例
|
||||
|
||||
```rust
|
||||
use crate::database::{init_database, DbConnection};
|
||||
|
||||
// 初始化数据库
|
||||
let db: DbConnection = init_database()?;
|
||||
|
||||
// 使用 DAO 操作数据
|
||||
let conn = db.lock().unwrap();
|
||||
let providers = ApiKeyProviderDao::get_all_providers(&conn)?;
|
||||
```
|
||||
@@ -386,6 +386,93 @@ impl ApiKeyProviderDao {
|
||||
Ok(keys)
|
||||
}
|
||||
|
||||
/// 获取指定类型的所有启用的 API Keys(包括自定义 Provider)
|
||||
/// 返回 (ApiKeyEntry, ApiKeyProvider) 元组列表
|
||||
pub fn get_enabled_api_keys_by_type(
|
||||
conn: &Connection,
|
||||
provider_type: ApiProviderType,
|
||||
) -> Result<Vec<(ApiKeyEntry, ApiKeyProvider)>, rusqlite::Error> {
|
||||
let type_str = provider_type.to_string();
|
||||
let mut stmt = conn.prepare(
|
||||
"SELECT k.id, k.provider_id, k.api_key_encrypted, k.alias, k.enabled,
|
||||
k.usage_count, k.error_count, k.last_used_at, k.created_at,
|
||||
p.id, p.name, p.type, p.api_host, p.is_system, p.group_name, p.enabled,
|
||||
p.sort_order, p.api_version, p.project, p.location, p.region,
|
||||
p.created_at, p.updated_at
|
||||
FROM api_keys k
|
||||
JOIN api_key_providers p ON k.provider_id = p.id
|
||||
WHERE p.type = ?1 AND k.enabled = 1 AND p.enabled = 1
|
||||
ORDER BY p.sort_order ASC, k.created_at ASC",
|
||||
)?;
|
||||
|
||||
let rows = stmt.query_map([type_str], |row| {
|
||||
// 解析 API Key
|
||||
let last_used_at_str: Option<String> = row.get(7)?;
|
||||
let created_at_str: String = row.get(8)?;
|
||||
let last_used_at = last_used_at_str.and_then(|s| {
|
||||
DateTime::parse_from_rfc3339(&s)
|
||||
.ok()
|
||||
.map(|dt| dt.with_timezone(&Utc))
|
||||
});
|
||||
let key_created_at = DateTime::parse_from_rfc3339(&created_at_str)
|
||||
.map(|dt| dt.with_timezone(&Utc))
|
||||
.unwrap_or_else(|_| Utc::now());
|
||||
|
||||
let key = ApiKeyEntry {
|
||||
id: row.get(0)?,
|
||||
provider_id: row.get(1)?,
|
||||
api_key_encrypted: row.get(2)?,
|
||||
alias: row.get(3)?,
|
||||
enabled: row.get(4)?,
|
||||
usage_count: row.get(5)?,
|
||||
error_count: row.get(6)?,
|
||||
last_used_at,
|
||||
created_at: key_created_at,
|
||||
};
|
||||
|
||||
// 解析 Provider
|
||||
let provider_created_at_str: String = row.get(21)?;
|
||||
let provider_updated_at_str: String = row.get(22)?;
|
||||
let provider_created_at = DateTime::parse_from_rfc3339(&provider_created_at_str)
|
||||
.map(|dt| dt.with_timezone(&Utc))
|
||||
.unwrap_or_else(|_| Utc::now());
|
||||
let provider_updated_at = DateTime::parse_from_rfc3339(&provider_updated_at_str)
|
||||
.map(|dt| dt.with_timezone(&Utc))
|
||||
.unwrap_or_else(|_| Utc::now());
|
||||
|
||||
let provider = ApiKeyProvider {
|
||||
id: row.get(9)?,
|
||||
name: row.get(10)?,
|
||||
provider_type: row
|
||||
.get::<_, String>(11)?
|
||||
.parse()
|
||||
.unwrap_or(ApiProviderType::Openai),
|
||||
api_host: row.get(12)?,
|
||||
is_system: row.get(13)?,
|
||||
group: row
|
||||
.get::<_, String>(14)?
|
||||
.parse()
|
||||
.unwrap_or(ProviderGroup::Custom),
|
||||
enabled: row.get(15)?,
|
||||
sort_order: row.get(16)?,
|
||||
api_version: row.get(17)?,
|
||||
project: row.get(18)?,
|
||||
location: row.get(19)?,
|
||||
region: row.get(20)?,
|
||||
created_at: provider_created_at,
|
||||
updated_at: provider_updated_at,
|
||||
};
|
||||
|
||||
Ok((key, provider))
|
||||
})?;
|
||||
|
||||
let mut result = Vec::new();
|
||||
for item in rows.flatten() {
|
||||
result.push(item);
|
||||
}
|
||||
Ok(result)
|
||||
}
|
||||
|
||||
/// 根据 ID 获取 API Key
|
||||
pub fn get_api_key_by_id(
|
||||
conn: &Connection,
|
||||
|
||||
@@ -1,6 +1,8 @@
|
||||
pub mod api_key_provider;
|
||||
pub mod installed_plugins;
|
||||
pub mod mcp;
|
||||
pub mod orchestrator;
|
||||
pub mod plugin_credential;
|
||||
pub mod prompts;
|
||||
pub mod provider_pool;
|
||||
pub mod providers;
|
||||
|
||||
@@ -0,0 +1,647 @@
|
||||
//! Orchestrator DAO 模块
|
||||
//!
|
||||
//! 提供模型元数据和用户偏好的数据库操作
|
||||
|
||||
use rusqlite::{params, Connection, OptionalExtension};
|
||||
use serde::{Deserialize, Serialize};
|
||||
|
||||
/// 模型元数据
|
||||
#[derive(Debug, Clone, Serialize, Deserialize)]
|
||||
pub struct ModelMetadataRow {
|
||||
pub model_id: String,
|
||||
pub provider_type: String,
|
||||
pub display_name: String,
|
||||
pub family: Option<String>,
|
||||
pub tier: String,
|
||||
pub context_length: Option<i64>,
|
||||
pub max_output_tokens: Option<i64>,
|
||||
pub cost_input_per_million: Option<f64>,
|
||||
pub cost_output_per_million: Option<f64>,
|
||||
pub supports_vision: bool,
|
||||
pub supports_tools: bool,
|
||||
pub supports_streaming: bool,
|
||||
pub is_deprecated: bool,
|
||||
pub release_date: Option<String>,
|
||||
pub description: Option<String>,
|
||||
pub created_at: i64,
|
||||
pub updated_at: i64,
|
||||
}
|
||||
|
||||
/// 用户等级偏好
|
||||
#[derive(Debug, Clone, Serialize, Deserialize)]
|
||||
pub struct UserTierPreference {
|
||||
pub tier_id: String,
|
||||
pub strategy_id: String,
|
||||
pub preferred_provider: Option<String>,
|
||||
pub fallback_enabled: bool,
|
||||
pub max_retries: i32,
|
||||
pub created_at: i64,
|
||||
pub updated_at: i64,
|
||||
}
|
||||
|
||||
/// 模型使用统计
|
||||
#[derive(Debug, Clone, Serialize, Deserialize)]
|
||||
pub struct ModelUsageStats {
|
||||
pub model_id: String,
|
||||
pub credential_id: String,
|
||||
pub date: String,
|
||||
pub request_count: i64,
|
||||
pub success_count: i64,
|
||||
pub error_count: i64,
|
||||
pub total_tokens: i64,
|
||||
pub total_latency_ms: i64,
|
||||
pub avg_latency_ms: Option<f64>,
|
||||
}
|
||||
|
||||
/// Orchestrator DAO
|
||||
pub struct OrchestratorDao;
|
||||
|
||||
impl OrchestratorDao {
|
||||
// ========================================================================
|
||||
// 模型元数据操作
|
||||
// ========================================================================
|
||||
|
||||
/// 获取所有模型元数据
|
||||
pub fn get_all_model_metadata(conn: &Connection) -> Result<Vec<ModelMetadataRow>, String> {
|
||||
let mut stmt = conn
|
||||
.prepare(
|
||||
"SELECT model_id, provider_type, display_name, family, tier,
|
||||
context_length, max_output_tokens, cost_input_per_million,
|
||||
cost_output_per_million, supports_vision, supports_tools,
|
||||
supports_streaming, is_deprecated, release_date, description,
|
||||
created_at, updated_at
|
||||
FROM model_metadata
|
||||
WHERE is_deprecated = 0
|
||||
ORDER BY provider_type, tier, display_name",
|
||||
)
|
||||
.map_err(|e| e.to_string())?;
|
||||
|
||||
let rows = stmt
|
||||
.query_map([], |row| {
|
||||
Ok(ModelMetadataRow {
|
||||
model_id: row.get(0)?,
|
||||
provider_type: row.get(1)?,
|
||||
display_name: row.get(2)?,
|
||||
family: row.get(3)?,
|
||||
tier: row.get(4)?,
|
||||
context_length: row.get(5)?,
|
||||
max_output_tokens: row.get(6)?,
|
||||
cost_input_per_million: row.get(7)?,
|
||||
cost_output_per_million: row.get(8)?,
|
||||
supports_vision: row.get::<_, i32>(9)? != 0,
|
||||
supports_tools: row.get::<_, i32>(10)? != 0,
|
||||
supports_streaming: row.get::<_, i32>(11)? != 0,
|
||||
is_deprecated: row.get::<_, i32>(12)? != 0,
|
||||
release_date: row.get(13)?,
|
||||
description: row.get(14)?,
|
||||
created_at: row.get(15)?,
|
||||
updated_at: row.get(16)?,
|
||||
})
|
||||
})
|
||||
.map_err(|e| e.to_string())?;
|
||||
|
||||
rows.collect::<Result<Vec<_>, _>>()
|
||||
.map_err(|e| e.to_string())
|
||||
}
|
||||
|
||||
/// 按 Provider 获取模型元数据
|
||||
pub fn get_model_metadata_by_provider(
|
||||
conn: &Connection,
|
||||
provider_type: &str,
|
||||
) -> Result<Vec<ModelMetadataRow>, String> {
|
||||
let mut stmt = conn
|
||||
.prepare(
|
||||
"SELECT model_id, provider_type, display_name, family, tier,
|
||||
context_length, max_output_tokens, cost_input_per_million,
|
||||
cost_output_per_million, supports_vision, supports_tools,
|
||||
supports_streaming, is_deprecated, release_date, description,
|
||||
created_at, updated_at
|
||||
FROM model_metadata
|
||||
WHERE provider_type = ?1 AND is_deprecated = 0
|
||||
ORDER BY tier, display_name",
|
||||
)
|
||||
.map_err(|e| e.to_string())?;
|
||||
|
||||
let rows = stmt
|
||||
.query_map([provider_type], |row| {
|
||||
Ok(ModelMetadataRow {
|
||||
model_id: row.get(0)?,
|
||||
provider_type: row.get(1)?,
|
||||
display_name: row.get(2)?,
|
||||
family: row.get(3)?,
|
||||
tier: row.get(4)?,
|
||||
context_length: row.get(5)?,
|
||||
max_output_tokens: row.get(6)?,
|
||||
cost_input_per_million: row.get(7)?,
|
||||
cost_output_per_million: row.get(8)?,
|
||||
supports_vision: row.get::<_, i32>(9)? != 0,
|
||||
supports_tools: row.get::<_, i32>(10)? != 0,
|
||||
supports_streaming: row.get::<_, i32>(11)? != 0,
|
||||
is_deprecated: row.get::<_, i32>(12)? != 0,
|
||||
release_date: row.get(13)?,
|
||||
description: row.get(14)?,
|
||||
created_at: row.get(15)?,
|
||||
updated_at: row.get(16)?,
|
||||
})
|
||||
})
|
||||
.map_err(|e| e.to_string())?;
|
||||
|
||||
rows.collect::<Result<Vec<_>, _>>()
|
||||
.map_err(|e| e.to_string())
|
||||
}
|
||||
|
||||
/// 按等级获取模型元数据
|
||||
pub fn get_model_metadata_by_tier(
|
||||
conn: &Connection,
|
||||
tier: &str,
|
||||
) -> Result<Vec<ModelMetadataRow>, String> {
|
||||
let mut stmt = conn
|
||||
.prepare(
|
||||
"SELECT model_id, provider_type, display_name, family, tier,
|
||||
context_length, max_output_tokens, cost_input_per_million,
|
||||
cost_output_per_million, supports_vision, supports_tools,
|
||||
supports_streaming, is_deprecated, release_date, description,
|
||||
created_at, updated_at
|
||||
FROM model_metadata
|
||||
WHERE tier = ?1 AND is_deprecated = 0
|
||||
ORDER BY provider_type, display_name",
|
||||
)
|
||||
.map_err(|e| e.to_string())?;
|
||||
|
||||
let rows = stmt
|
||||
.query_map([tier], |row| {
|
||||
Ok(ModelMetadataRow {
|
||||
model_id: row.get(0)?,
|
||||
provider_type: row.get(1)?,
|
||||
display_name: row.get(2)?,
|
||||
family: row.get(3)?,
|
||||
tier: row.get(4)?,
|
||||
context_length: row.get(5)?,
|
||||
max_output_tokens: row.get(6)?,
|
||||
cost_input_per_million: row.get(7)?,
|
||||
cost_output_per_million: row.get(8)?,
|
||||
supports_vision: row.get::<_, i32>(9)? != 0,
|
||||
supports_tools: row.get::<_, i32>(10)? != 0,
|
||||
supports_streaming: row.get::<_, i32>(11)? != 0,
|
||||
is_deprecated: row.get::<_, i32>(12)? != 0,
|
||||
release_date: row.get(13)?,
|
||||
description: row.get(14)?,
|
||||
created_at: row.get(15)?,
|
||||
updated_at: row.get(16)?,
|
||||
})
|
||||
})
|
||||
.map_err(|e| e.to_string())?;
|
||||
|
||||
rows.collect::<Result<Vec<_>, _>>()
|
||||
.map_err(|e| e.to_string())
|
||||
}
|
||||
|
||||
/// 获取单个模型元数据
|
||||
pub fn get_model_metadata(
|
||||
conn: &Connection,
|
||||
model_id: &str,
|
||||
) -> Result<Option<ModelMetadataRow>, String> {
|
||||
conn.query_row(
|
||||
"SELECT model_id, provider_type, display_name, family, tier,
|
||||
context_length, max_output_tokens, cost_input_per_million,
|
||||
cost_output_per_million, supports_vision, supports_tools,
|
||||
supports_streaming, is_deprecated, release_date, description,
|
||||
created_at, updated_at
|
||||
FROM model_metadata
|
||||
WHERE model_id = ?1",
|
||||
[model_id],
|
||||
|row| {
|
||||
Ok(ModelMetadataRow {
|
||||
model_id: row.get(0)?,
|
||||
provider_type: row.get(1)?,
|
||||
display_name: row.get(2)?,
|
||||
family: row.get(3)?,
|
||||
tier: row.get(4)?,
|
||||
context_length: row.get(5)?,
|
||||
max_output_tokens: row.get(6)?,
|
||||
cost_input_per_million: row.get(7)?,
|
||||
cost_output_per_million: row.get(8)?,
|
||||
supports_vision: row.get::<_, i32>(9)? != 0,
|
||||
supports_tools: row.get::<_, i32>(10)? != 0,
|
||||
supports_streaming: row.get::<_, i32>(11)? != 0,
|
||||
is_deprecated: row.get::<_, i32>(12)? != 0,
|
||||
release_date: row.get(13)?,
|
||||
description: row.get(14)?,
|
||||
created_at: row.get(15)?,
|
||||
updated_at: row.get(16)?,
|
||||
})
|
||||
},
|
||||
)
|
||||
.optional()
|
||||
.map_err(|e| e.to_string())
|
||||
}
|
||||
|
||||
/// 插入或更新模型元数据
|
||||
pub fn upsert_model_metadata(
|
||||
conn: &Connection,
|
||||
metadata: &ModelMetadataRow,
|
||||
) -> Result<(), String> {
|
||||
conn.execute(
|
||||
"INSERT INTO model_metadata (
|
||||
model_id, provider_type, display_name, family, tier,
|
||||
context_length, max_output_tokens, cost_input_per_million,
|
||||
cost_output_per_million, supports_vision, supports_tools,
|
||||
supports_streaming, is_deprecated, release_date, description,
|
||||
created_at, updated_at
|
||||
) VALUES (?1, ?2, ?3, ?4, ?5, ?6, ?7, ?8, ?9, ?10, ?11, ?12, ?13, ?14, ?15, ?16, ?17)
|
||||
ON CONFLICT(model_id) DO UPDATE SET
|
||||
provider_type = excluded.provider_type,
|
||||
display_name = excluded.display_name,
|
||||
family = excluded.family,
|
||||
tier = excluded.tier,
|
||||
context_length = excluded.context_length,
|
||||
max_output_tokens = excluded.max_output_tokens,
|
||||
cost_input_per_million = excluded.cost_input_per_million,
|
||||
cost_output_per_million = excluded.cost_output_per_million,
|
||||
supports_vision = excluded.supports_vision,
|
||||
supports_tools = excluded.supports_tools,
|
||||
supports_streaming = excluded.supports_streaming,
|
||||
is_deprecated = excluded.is_deprecated,
|
||||
release_date = excluded.release_date,
|
||||
description = excluded.description,
|
||||
updated_at = excluded.updated_at",
|
||||
params![
|
||||
metadata.model_id,
|
||||
metadata.provider_type,
|
||||
metadata.display_name,
|
||||
metadata.family,
|
||||
metadata.tier,
|
||||
metadata.context_length,
|
||||
metadata.max_output_tokens,
|
||||
metadata.cost_input_per_million,
|
||||
metadata.cost_output_per_million,
|
||||
metadata.supports_vision as i32,
|
||||
metadata.supports_tools as i32,
|
||||
metadata.supports_streaming as i32,
|
||||
metadata.is_deprecated as i32,
|
||||
metadata.release_date,
|
||||
metadata.description,
|
||||
metadata.created_at,
|
||||
metadata.updated_at,
|
||||
],
|
||||
)
|
||||
.map_err(|e| e.to_string())?;
|
||||
|
||||
Ok(())
|
||||
}
|
||||
|
||||
/// 批量插入模型元数据
|
||||
pub fn bulk_upsert_model_metadata(
|
||||
conn: &Connection,
|
||||
metadata_list: &[ModelMetadataRow],
|
||||
) -> Result<usize, String> {
|
||||
let mut count = 0;
|
||||
for metadata in metadata_list {
|
||||
Self::upsert_model_metadata(conn, metadata)?;
|
||||
count += 1;
|
||||
}
|
||||
Ok(count)
|
||||
}
|
||||
|
||||
// ========================================================================
|
||||
// 用户等级偏好操作
|
||||
// ========================================================================
|
||||
|
||||
/// 获取所有用户等级偏好
|
||||
pub fn get_all_tier_preferences(conn: &Connection) -> Result<Vec<UserTierPreference>, String> {
|
||||
let mut stmt = conn
|
||||
.prepare(
|
||||
"SELECT tier_id, strategy_id, preferred_provider, fallback_enabled,
|
||||
max_retries, created_at, updated_at
|
||||
FROM user_tier_preferences
|
||||
ORDER BY tier_id",
|
||||
)
|
||||
.map_err(|e| e.to_string())?;
|
||||
|
||||
let rows = stmt
|
||||
.query_map([], |row| {
|
||||
Ok(UserTierPreference {
|
||||
tier_id: row.get(0)?,
|
||||
strategy_id: row.get(1)?,
|
||||
preferred_provider: row.get(2)?,
|
||||
fallback_enabled: row.get::<_, i32>(3)? != 0,
|
||||
max_retries: row.get(4)?,
|
||||
created_at: row.get(5)?,
|
||||
updated_at: row.get(6)?,
|
||||
})
|
||||
})
|
||||
.map_err(|e| e.to_string())?;
|
||||
|
||||
rows.collect::<Result<Vec<_>, _>>()
|
||||
.map_err(|e| e.to_string())
|
||||
}
|
||||
|
||||
/// 获取单个等级偏好
|
||||
pub fn get_tier_preference(
|
||||
conn: &Connection,
|
||||
tier_id: &str,
|
||||
) -> Result<Option<UserTierPreference>, String> {
|
||||
conn.query_row(
|
||||
"SELECT tier_id, strategy_id, preferred_provider, fallback_enabled,
|
||||
max_retries, created_at, updated_at
|
||||
FROM user_tier_preferences
|
||||
WHERE tier_id = ?1",
|
||||
[tier_id],
|
||||
|row| {
|
||||
Ok(UserTierPreference {
|
||||
tier_id: row.get(0)?,
|
||||
strategy_id: row.get(1)?,
|
||||
preferred_provider: row.get(2)?,
|
||||
fallback_enabled: row.get::<_, i32>(3)? != 0,
|
||||
max_retries: row.get(4)?,
|
||||
created_at: row.get(5)?,
|
||||
updated_at: row.get(6)?,
|
||||
})
|
||||
},
|
||||
)
|
||||
.optional()
|
||||
.map_err(|e| e.to_string())
|
||||
}
|
||||
|
||||
/// 插入或更新等级偏好
|
||||
pub fn upsert_tier_preference(
|
||||
conn: &Connection,
|
||||
pref: &UserTierPreference,
|
||||
) -> Result<(), String> {
|
||||
conn.execute(
|
||||
"INSERT INTO user_tier_preferences (
|
||||
tier_id, strategy_id, preferred_provider, fallback_enabled,
|
||||
max_retries, created_at, updated_at
|
||||
) VALUES (?1, ?2, ?3, ?4, ?5, ?6, ?7)
|
||||
ON CONFLICT(tier_id) DO UPDATE SET
|
||||
strategy_id = excluded.strategy_id,
|
||||
preferred_provider = excluded.preferred_provider,
|
||||
fallback_enabled = excluded.fallback_enabled,
|
||||
max_retries = excluded.max_retries,
|
||||
updated_at = excluded.updated_at",
|
||||
params![
|
||||
pref.tier_id,
|
||||
pref.strategy_id,
|
||||
pref.preferred_provider,
|
||||
pref.fallback_enabled as i32,
|
||||
pref.max_retries,
|
||||
pref.created_at,
|
||||
pref.updated_at,
|
||||
],
|
||||
)
|
||||
.map_err(|e| e.to_string())?;
|
||||
|
||||
Ok(())
|
||||
}
|
||||
|
||||
/// 初始化默认等级偏好
|
||||
pub fn init_default_tier_preferences(conn: &Connection) -> Result<(), String> {
|
||||
let now = chrono::Utc::now().timestamp();
|
||||
|
||||
let defaults = vec![
|
||||
UserTierPreference {
|
||||
tier_id: "mini".to_string(),
|
||||
strategy_id: "speed_optimized".to_string(),
|
||||
preferred_provider: None,
|
||||
fallback_enabled: true,
|
||||
max_retries: 3,
|
||||
created_at: now,
|
||||
updated_at: now,
|
||||
},
|
||||
UserTierPreference {
|
||||
tier_id: "pro".to_string(),
|
||||
strategy_id: "task_based".to_string(),
|
||||
preferred_provider: None,
|
||||
fallback_enabled: true,
|
||||
max_retries: 3,
|
||||
created_at: now,
|
||||
updated_at: now,
|
||||
},
|
||||
UserTierPreference {
|
||||
tier_id: "max".to_string(),
|
||||
strategy_id: "quality_first".to_string(),
|
||||
preferred_provider: None,
|
||||
fallback_enabled: true,
|
||||
max_retries: 3,
|
||||
created_at: now,
|
||||
updated_at: now,
|
||||
},
|
||||
];
|
||||
|
||||
for pref in defaults {
|
||||
// 只在不存在时插入
|
||||
let exists = Self::get_tier_preference(conn, &pref.tier_id)?.is_some();
|
||||
if !exists {
|
||||
Self::upsert_tier_preference(conn, &pref)?;
|
||||
}
|
||||
}
|
||||
|
||||
Ok(())
|
||||
}
|
||||
|
||||
// ========================================================================
|
||||
// 模型使用统计操作
|
||||
// ========================================================================
|
||||
|
||||
/// 记录模型使用
|
||||
pub fn record_model_usage(
|
||||
conn: &Connection,
|
||||
model_id: &str,
|
||||
credential_id: &str,
|
||||
success: bool,
|
||||
tokens: i64,
|
||||
latency_ms: i64,
|
||||
) -> Result<(), String> {
|
||||
let today = chrono::Utc::now().format("%Y-%m-%d").to_string();
|
||||
|
||||
// 尝试更新现有记录
|
||||
let updated = conn
|
||||
.execute(
|
||||
"UPDATE model_usage_stats SET
|
||||
request_count = request_count + 1,
|
||||
success_count = success_count + ?1,
|
||||
error_count = error_count + ?2,
|
||||
total_tokens = total_tokens + ?3,
|
||||
total_latency_ms = total_latency_ms + ?4,
|
||||
avg_latency_ms = CAST((total_latency_ms + ?4) AS REAL) / (request_count + 1)
|
||||
WHERE model_id = ?5 AND credential_id = ?6 AND date = ?7",
|
||||
params![
|
||||
if success { 1 } else { 0 },
|
||||
if success { 0 } else { 1 },
|
||||
tokens,
|
||||
latency_ms,
|
||||
model_id,
|
||||
credential_id,
|
||||
today,
|
||||
],
|
||||
)
|
||||
.map_err(|e| e.to_string())?;
|
||||
|
||||
// 如果没有更新到记录,插入新记录
|
||||
if updated == 0 {
|
||||
conn.execute(
|
||||
"INSERT INTO model_usage_stats (
|
||||
model_id, credential_id, date, request_count, success_count,
|
||||
error_count, total_tokens, total_latency_ms, avg_latency_ms
|
||||
) VALUES (?1, ?2, ?3, 1, ?4, ?5, ?6, ?7, ?8)",
|
||||
params![
|
||||
model_id,
|
||||
credential_id,
|
||||
today,
|
||||
if success { 1 } else { 0 },
|
||||
if success { 0 } else { 1 },
|
||||
tokens,
|
||||
latency_ms,
|
||||
latency_ms as f64,
|
||||
],
|
||||
)
|
||||
.map_err(|e| e.to_string())?;
|
||||
}
|
||||
|
||||
Ok(())
|
||||
}
|
||||
|
||||
/// 获取模型使用统计
|
||||
pub fn get_model_usage_stats(
|
||||
conn: &Connection,
|
||||
model_id: &str,
|
||||
days: i32,
|
||||
) -> Result<Vec<ModelUsageStats>, String> {
|
||||
let mut stmt = conn
|
||||
.prepare(
|
||||
"SELECT model_id, credential_id, date, request_count, success_count,
|
||||
error_count, total_tokens, total_latency_ms, avg_latency_ms
|
||||
FROM model_usage_stats
|
||||
WHERE model_id = ?1 AND date >= date('now', ?2)
|
||||
ORDER BY date DESC",
|
||||
)
|
||||
.map_err(|e| e.to_string())?;
|
||||
|
||||
let days_param = format!("-{} days", days);
|
||||
let rows = stmt
|
||||
.query_map(params![model_id, days_param], |row| {
|
||||
Ok(ModelUsageStats {
|
||||
model_id: row.get(0)?,
|
||||
credential_id: row.get(1)?,
|
||||
date: row.get(2)?,
|
||||
request_count: row.get(3)?,
|
||||
success_count: row.get(4)?,
|
||||
error_count: row.get(5)?,
|
||||
total_tokens: row.get(6)?,
|
||||
total_latency_ms: row.get(7)?,
|
||||
avg_latency_ms: row.get(8)?,
|
||||
})
|
||||
})
|
||||
.map_err(|e| e.to_string())?;
|
||||
|
||||
rows.collect::<Result<Vec<_>, _>>()
|
||||
.map_err(|e| e.to_string())
|
||||
}
|
||||
|
||||
/// 清理旧的使用统计
|
||||
pub fn cleanup_old_usage_stats(conn: &Connection, days: i32) -> Result<usize, String> {
|
||||
let days_param = format!("-{} days", days);
|
||||
conn.execute(
|
||||
"DELETE FROM model_usage_stats WHERE date < date('now', ?1)",
|
||||
[days_param],
|
||||
)
|
||||
.map_err(|e| e.to_string())
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
use rusqlite::Connection;
|
||||
|
||||
fn setup_test_db() -> Connection {
|
||||
let conn = Connection::open_in_memory().unwrap();
|
||||
crate::database::schema::create_tables(&conn).unwrap();
|
||||
conn
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_model_metadata_crud() {
|
||||
let conn = setup_test_db();
|
||||
let now = chrono::Utc::now().timestamp();
|
||||
|
||||
let metadata = ModelMetadataRow {
|
||||
model_id: "claude-3-opus".to_string(),
|
||||
provider_type: "anthropic".to_string(),
|
||||
display_name: "Claude 3 Opus".to_string(),
|
||||
family: Some("opus".to_string()),
|
||||
tier: "max".to_string(),
|
||||
context_length: Some(200000),
|
||||
max_output_tokens: Some(4096),
|
||||
cost_input_per_million: Some(15.0),
|
||||
cost_output_per_million: Some(75.0),
|
||||
supports_vision: true,
|
||||
supports_tools: true,
|
||||
supports_streaming: true,
|
||||
is_deprecated: false,
|
||||
release_date: Some("2024-03-04".to_string()),
|
||||
description: Some("Most capable Claude model".to_string()),
|
||||
created_at: now,
|
||||
updated_at: now,
|
||||
};
|
||||
|
||||
// Insert
|
||||
OrchestratorDao::upsert_model_metadata(&conn, &metadata).unwrap();
|
||||
|
||||
// Read
|
||||
let result = OrchestratorDao::get_model_metadata(&conn, "claude-3-opus")
|
||||
.unwrap()
|
||||
.unwrap();
|
||||
assert_eq!(result.display_name, "Claude 3 Opus");
|
||||
assert!(result.supports_vision);
|
||||
|
||||
// Update
|
||||
let mut updated = metadata.clone();
|
||||
updated.display_name = "Claude 3 Opus (Updated)".to_string();
|
||||
OrchestratorDao::upsert_model_metadata(&conn, &updated).unwrap();
|
||||
|
||||
let result = OrchestratorDao::get_model_metadata(&conn, "claude-3-opus")
|
||||
.unwrap()
|
||||
.unwrap();
|
||||
assert_eq!(result.display_name, "Claude 3 Opus (Updated)");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_tier_preferences() {
|
||||
let conn = setup_test_db();
|
||||
|
||||
// Init defaults
|
||||
OrchestratorDao::init_default_tier_preferences(&conn).unwrap();
|
||||
|
||||
// Check defaults exist
|
||||
let prefs = OrchestratorDao::get_all_tier_preferences(&conn).unwrap();
|
||||
assert_eq!(prefs.len(), 3);
|
||||
|
||||
// Get specific
|
||||
let pro = OrchestratorDao::get_tier_preference(&conn, "pro")
|
||||
.unwrap()
|
||||
.unwrap();
|
||||
assert_eq!(pro.strategy_id, "task_based");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_usage_stats() {
|
||||
let conn = setup_test_db();
|
||||
|
||||
// Record usage
|
||||
OrchestratorDao::record_model_usage(&conn, "claude-3-opus", "cred-1", true, 1000, 500)
|
||||
.unwrap();
|
||||
OrchestratorDao::record_model_usage(&conn, "claude-3-opus", "cred-1", true, 2000, 600)
|
||||
.unwrap();
|
||||
OrchestratorDao::record_model_usage(&conn, "claude-3-opus", "cred-1", false, 0, 100)
|
||||
.unwrap();
|
||||
|
||||
// Get stats
|
||||
let stats = OrchestratorDao::get_model_usage_stats(&conn, "claude-3-opus", 7).unwrap();
|
||||
assert_eq!(stats.len(), 1);
|
||||
assert_eq!(stats[0].request_count, 3);
|
||||
assert_eq!(stats[0].success_count, 2);
|
||||
assert_eq!(stats[0].error_count, 1);
|
||||
assert_eq!(stats[0].total_tokens, 3000);
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,530 @@
|
||||
//! 插件凭证数据访问对象
|
||||
//!
|
||||
//! 提供 OAuth Provider 插件凭证的 CRUD 操作。
|
||||
|
||||
#![allow(dead_code)]
|
||||
|
||||
use chrono::{DateTime, Utc};
|
||||
use rusqlite::{params, Connection, OptionalExtension};
|
||||
use serde::{Deserialize, Serialize};
|
||||
|
||||
/// 凭证状态
|
||||
#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
|
||||
#[serde(rename_all = "lowercase")]
|
||||
pub enum CredentialStatus {
|
||||
/// 活跃可用
|
||||
Active,
|
||||
/// 已禁用
|
||||
Disabled,
|
||||
/// 已过期
|
||||
Expired,
|
||||
/// 错误状态
|
||||
Error,
|
||||
}
|
||||
|
||||
impl CredentialStatus {
|
||||
pub fn as_str(&self) -> &'static str {
|
||||
match self {
|
||||
CredentialStatus::Active => "active",
|
||||
CredentialStatus::Disabled => "disabled",
|
||||
CredentialStatus::Expired => "expired",
|
||||
CredentialStatus::Error => "error",
|
||||
}
|
||||
}
|
||||
|
||||
pub fn from_str(s: &str) -> Self {
|
||||
match s {
|
||||
"active" => CredentialStatus::Active,
|
||||
"disabled" => CredentialStatus::Disabled,
|
||||
"expired" => CredentialStatus::Expired,
|
||||
"error" => CredentialStatus::Error,
|
||||
_ => CredentialStatus::Active,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
/// 插件凭证记录
|
||||
#[derive(Debug, Clone, Serialize, Deserialize)]
|
||||
pub struct PluginCredentialRecord {
|
||||
/// 凭证 ID
|
||||
pub id: String,
|
||||
/// 插件 ID
|
||||
pub plugin_id: String,
|
||||
/// 认证类型 (oauth, api_key, cookie, etc.)
|
||||
pub auth_type: String,
|
||||
/// 显示名称
|
||||
pub display_name: Option<String>,
|
||||
/// 状态
|
||||
pub status: CredentialStatus,
|
||||
/// 加密配置 (JSON)
|
||||
pub config_encrypted: String,
|
||||
/// 使用次数
|
||||
pub usage_count: u32,
|
||||
/// 错误次数
|
||||
pub error_count: u32,
|
||||
/// 最后使用时间
|
||||
pub last_used_at: Option<DateTime<Utc>>,
|
||||
/// 最后错误时间
|
||||
pub last_error_at: Option<DateTime<Utc>>,
|
||||
/// 最后错误消息
|
||||
pub last_error_message: Option<String>,
|
||||
/// 创建时间
|
||||
pub created_at: DateTime<Utc>,
|
||||
/// 更新时间
|
||||
pub updated_at: DateTime<Utc>,
|
||||
}
|
||||
|
||||
/// 新建凭证参数
|
||||
#[derive(Debug, Clone)]
|
||||
pub struct NewPluginCredential {
|
||||
pub id: String,
|
||||
pub plugin_id: String,
|
||||
pub auth_type: String,
|
||||
pub display_name: Option<String>,
|
||||
pub config_encrypted: String,
|
||||
}
|
||||
|
||||
/// 数据库行结构
|
||||
struct CredentialRow {
|
||||
id: String,
|
||||
plugin_id: String,
|
||||
auth_type: String,
|
||||
display_name: Option<String>,
|
||||
status: String,
|
||||
config_encrypted: String,
|
||||
usage_count: i32,
|
||||
error_count: i32,
|
||||
last_used_at: Option<String>,
|
||||
last_error_at: Option<String>,
|
||||
last_error_message: Option<String>,
|
||||
created_at: String,
|
||||
updated_at: String,
|
||||
}
|
||||
|
||||
impl CredentialRow {
|
||||
fn into_record(self) -> Result<PluginCredentialRecord, String> {
|
||||
let created_at = DateTime::parse_from_rfc3339(&self.created_at)
|
||||
.map_err(|e| format!("无效的创建时间格式: {}", e))?
|
||||
.with_timezone(&Utc);
|
||||
|
||||
let updated_at = DateTime::parse_from_rfc3339(&self.updated_at)
|
||||
.map_err(|e| format!("无效的更新时间格式: {}", e))?
|
||||
.with_timezone(&Utc);
|
||||
|
||||
let last_used_at = self
|
||||
.last_used_at
|
||||
.map(|s| DateTime::parse_from_rfc3339(&s).ok())
|
||||
.flatten()
|
||||
.map(|dt| dt.with_timezone(&Utc));
|
||||
|
||||
let last_error_at = self
|
||||
.last_error_at
|
||||
.map(|s| DateTime::parse_from_rfc3339(&s).ok())
|
||||
.flatten()
|
||||
.map(|dt| dt.with_timezone(&Utc));
|
||||
|
||||
Ok(PluginCredentialRecord {
|
||||
id: self.id,
|
||||
plugin_id: self.plugin_id,
|
||||
auth_type: self.auth_type,
|
||||
display_name: self.display_name,
|
||||
status: CredentialStatus::from_str(&self.status),
|
||||
config_encrypted: self.config_encrypted,
|
||||
usage_count: self.usage_count as u32,
|
||||
error_count: self.error_count as u32,
|
||||
last_used_at,
|
||||
last_error_at,
|
||||
last_error_message: self.last_error_message,
|
||||
created_at,
|
||||
updated_at,
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
pub struct PluginCredentialDao;
|
||||
|
||||
impl PluginCredentialDao {
|
||||
/// 创建凭证
|
||||
pub fn create(
|
||||
conn: &Connection,
|
||||
credential: &NewPluginCredential,
|
||||
) -> Result<(), rusqlite::Error> {
|
||||
let now = Utc::now().to_rfc3339();
|
||||
|
||||
conn.execute(
|
||||
"INSERT INTO plugin_credentials
|
||||
(id, plugin_id, auth_type, display_name, status, config_encrypted,
|
||||
usage_count, error_count, created_at, updated_at)
|
||||
VALUES (?1, ?2, ?3, ?4, 'active', ?5, 0, 0, ?6, ?7)",
|
||||
params![
|
||||
credential.id,
|
||||
credential.plugin_id,
|
||||
credential.auth_type,
|
||||
credential.display_name,
|
||||
credential.config_encrypted,
|
||||
now,
|
||||
now,
|
||||
],
|
||||
)?;
|
||||
|
||||
Ok(())
|
||||
}
|
||||
|
||||
/// 获取单个凭证
|
||||
pub fn get(
|
||||
conn: &Connection,
|
||||
credential_id: &str,
|
||||
) -> Result<Option<PluginCredentialRecord>, String> {
|
||||
let result = conn
|
||||
.query_row(
|
||||
"SELECT id, plugin_id, auth_type, display_name, status, config_encrypted,
|
||||
usage_count, error_count, last_used_at, last_error_at,
|
||||
last_error_message, created_at, updated_at
|
||||
FROM plugin_credentials WHERE id = ?1",
|
||||
params![credential_id],
|
||||
|row| {
|
||||
Ok(CredentialRow {
|
||||
id: row.get(0)?,
|
||||
plugin_id: row.get(1)?,
|
||||
auth_type: row.get(2)?,
|
||||
display_name: row.get(3)?,
|
||||
status: row.get(4)?,
|
||||
config_encrypted: row.get(5)?,
|
||||
usage_count: row.get(6)?,
|
||||
error_count: row.get(7)?,
|
||||
last_used_at: row.get(8)?,
|
||||
last_error_at: row.get(9)?,
|
||||
last_error_message: row.get(10)?,
|
||||
created_at: row.get(11)?,
|
||||
updated_at: row.get(12)?,
|
||||
})
|
||||
},
|
||||
)
|
||||
.optional()
|
||||
.map_err(|e| format!("数据库错误: {}", e))?;
|
||||
|
||||
match result {
|
||||
Some(row) => Ok(Some(row.into_record()?)),
|
||||
None => Ok(None),
|
||||
}
|
||||
}
|
||||
|
||||
/// 列出插件的所有凭证
|
||||
pub fn list_by_plugin(
|
||||
conn: &Connection,
|
||||
plugin_id: &str,
|
||||
) -> Result<Vec<PluginCredentialRecord>, String> {
|
||||
let mut stmt = conn
|
||||
.prepare(
|
||||
"SELECT id, plugin_id, auth_type, display_name, status, config_encrypted,
|
||||
usage_count, error_count, last_used_at, last_error_at,
|
||||
last_error_message, created_at, updated_at
|
||||
FROM plugin_credentials WHERE plugin_id = ?1 ORDER BY created_at DESC",
|
||||
)
|
||||
.map_err(|e| format!("数据库错误: {}", e))?;
|
||||
|
||||
let rows = stmt
|
||||
.query_map(params![plugin_id], |row| {
|
||||
Ok(CredentialRow {
|
||||
id: row.get(0)?,
|
||||
plugin_id: row.get(1)?,
|
||||
auth_type: row.get(2)?,
|
||||
display_name: row.get(3)?,
|
||||
status: row.get(4)?,
|
||||
config_encrypted: row.get(5)?,
|
||||
usage_count: row.get(6)?,
|
||||
error_count: row.get(7)?,
|
||||
last_used_at: row.get(8)?,
|
||||
last_error_at: row.get(9)?,
|
||||
last_error_message: row.get(10)?,
|
||||
created_at: row.get(11)?,
|
||||
updated_at: row.get(12)?,
|
||||
})
|
||||
})
|
||||
.map_err(|e| format!("数据库错误: {}", e))?;
|
||||
|
||||
let mut credentials = Vec::new();
|
||||
for row in rows {
|
||||
let row = row.map_err(|e| format!("数据库错误: {}", e))?;
|
||||
credentials.push(row.into_record()?);
|
||||
}
|
||||
|
||||
Ok(credentials)
|
||||
}
|
||||
|
||||
/// 列出所有活跃凭证
|
||||
pub fn list_active(conn: &Connection) -> Result<Vec<PluginCredentialRecord>, String> {
|
||||
let mut stmt = conn
|
||||
.prepare(
|
||||
"SELECT id, plugin_id, auth_type, display_name, status, config_encrypted,
|
||||
usage_count, error_count, last_used_at, last_error_at,
|
||||
last_error_message, created_at, updated_at
|
||||
FROM plugin_credentials WHERE status = 'active' ORDER BY usage_count DESC",
|
||||
)
|
||||
.map_err(|e| format!("数据库错误: {}", e))?;
|
||||
|
||||
let rows = stmt
|
||||
.query_map([], |row| {
|
||||
Ok(CredentialRow {
|
||||
id: row.get(0)?,
|
||||
plugin_id: row.get(1)?,
|
||||
auth_type: row.get(2)?,
|
||||
display_name: row.get(3)?,
|
||||
status: row.get(4)?,
|
||||
config_encrypted: row.get(5)?,
|
||||
usage_count: row.get(6)?,
|
||||
error_count: row.get(7)?,
|
||||
last_used_at: row.get(8)?,
|
||||
last_error_at: row.get(9)?,
|
||||
last_error_message: row.get(10)?,
|
||||
created_at: row.get(11)?,
|
||||
updated_at: row.get(12)?,
|
||||
})
|
||||
})
|
||||
.map_err(|e| format!("数据库错误: {}", e))?;
|
||||
|
||||
let mut credentials = Vec::new();
|
||||
for row in rows {
|
||||
let row = row.map_err(|e| format!("数据库错误: {}", e))?;
|
||||
credentials.push(row.into_record()?);
|
||||
}
|
||||
|
||||
Ok(credentials)
|
||||
}
|
||||
|
||||
/// 更新凭证配置
|
||||
pub fn update_config(
|
||||
conn: &Connection,
|
||||
credential_id: &str,
|
||||
config_encrypted: &str,
|
||||
) -> Result<bool, rusqlite::Error> {
|
||||
let now = Utc::now().to_rfc3339();
|
||||
let rows_affected = conn.execute(
|
||||
"UPDATE plugin_credentials SET config_encrypted = ?1, updated_at = ?2 WHERE id = ?3",
|
||||
params![config_encrypted, now, credential_id],
|
||||
)?;
|
||||
|
||||
Ok(rows_affected > 0)
|
||||
}
|
||||
|
||||
/// 更新凭证状态
|
||||
pub fn update_status(
|
||||
conn: &Connection,
|
||||
credential_id: &str,
|
||||
status: CredentialStatus,
|
||||
) -> Result<bool, rusqlite::Error> {
|
||||
let now = Utc::now().to_rfc3339();
|
||||
let rows_affected = conn.execute(
|
||||
"UPDATE plugin_credentials SET status = ?1, updated_at = ?2 WHERE id = ?3",
|
||||
params![status.as_str(), now, credential_id],
|
||||
)?;
|
||||
|
||||
Ok(rows_affected > 0)
|
||||
}
|
||||
|
||||
/// 记录使用
|
||||
pub fn record_usage(conn: &Connection, credential_id: &str) -> Result<bool, rusqlite::Error> {
|
||||
let now = Utc::now().to_rfc3339();
|
||||
let rows_affected = conn.execute(
|
||||
"UPDATE plugin_credentials
|
||||
SET usage_count = usage_count + 1, last_used_at = ?1, updated_at = ?2
|
||||
WHERE id = ?3",
|
||||
params![now, now, credential_id],
|
||||
)?;
|
||||
|
||||
Ok(rows_affected > 0)
|
||||
}
|
||||
|
||||
/// 记录错误
|
||||
pub fn record_error(
|
||||
conn: &Connection,
|
||||
credential_id: &str,
|
||||
error_message: &str,
|
||||
) -> Result<bool, rusqlite::Error> {
|
||||
let now = Utc::now().to_rfc3339();
|
||||
let rows_affected = conn.execute(
|
||||
"UPDATE plugin_credentials
|
||||
SET error_count = error_count + 1, last_error_at = ?1,
|
||||
last_error_message = ?2, updated_at = ?3
|
||||
WHERE id = ?4",
|
||||
params![now, error_message, now, credential_id],
|
||||
)?;
|
||||
|
||||
Ok(rows_affected > 0)
|
||||
}
|
||||
|
||||
/// 重置错误计数
|
||||
pub fn reset_errors(conn: &Connection, credential_id: &str) -> Result<bool, rusqlite::Error> {
|
||||
let now = Utc::now().to_rfc3339();
|
||||
let rows_affected = conn.execute(
|
||||
"UPDATE plugin_credentials
|
||||
SET error_count = 0, last_error_at = NULL, last_error_message = NULL,
|
||||
status = 'active', updated_at = ?1
|
||||
WHERE id = ?2",
|
||||
params![now, credential_id],
|
||||
)?;
|
||||
|
||||
Ok(rows_affected > 0)
|
||||
}
|
||||
|
||||
/// 删除凭证
|
||||
pub fn delete(conn: &Connection, credential_id: &str) -> Result<bool, rusqlite::Error> {
|
||||
let rows_affected = conn.execute(
|
||||
"DELETE FROM plugin_credentials WHERE id = ?1",
|
||||
params![credential_id],
|
||||
)?;
|
||||
|
||||
Ok(rows_affected > 0)
|
||||
}
|
||||
|
||||
/// 删除插件的所有凭证
|
||||
pub fn delete_by_plugin(conn: &Connection, plugin_id: &str) -> Result<u32, rusqlite::Error> {
|
||||
let rows_affected = conn.execute(
|
||||
"DELETE FROM plugin_credentials WHERE plugin_id = ?1",
|
||||
params![plugin_id],
|
||||
)?;
|
||||
|
||||
Ok(rows_affected as u32)
|
||||
}
|
||||
|
||||
/// 统计插件凭证数量
|
||||
pub fn count_by_plugin(conn: &Connection, plugin_id: &str) -> Result<u32, rusqlite::Error> {
|
||||
let count: i32 = conn.query_row(
|
||||
"SELECT COUNT(*) FROM plugin_credentials WHERE plugin_id = ?1",
|
||||
params![plugin_id],
|
||||
|row| row.get(0),
|
||||
)?;
|
||||
|
||||
Ok(count as u32)
|
||||
}
|
||||
|
||||
/// 统计活跃凭证数量
|
||||
pub fn count_active_by_plugin(
|
||||
conn: &Connection,
|
||||
plugin_id: &str,
|
||||
) -> Result<u32, rusqlite::Error> {
|
||||
let count: i32 = conn.query_row(
|
||||
"SELECT COUNT(*) FROM plugin_credentials WHERE plugin_id = ?1 AND status = 'active'",
|
||||
params![plugin_id],
|
||||
|row| row.get(0),
|
||||
)?;
|
||||
|
||||
Ok(count as u32)
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
|
||||
fn create_test_connection() -> Connection {
|
||||
let conn = Connection::open_in_memory().unwrap();
|
||||
conn.execute(
|
||||
"CREATE TABLE IF NOT EXISTS plugin_credentials (
|
||||
id TEXT PRIMARY KEY,
|
||||
plugin_id TEXT NOT NULL,
|
||||
auth_type TEXT NOT NULL,
|
||||
display_name TEXT,
|
||||
status TEXT NOT NULL DEFAULT 'active',
|
||||
config_encrypted TEXT NOT NULL,
|
||||
usage_count INTEGER DEFAULT 0,
|
||||
error_count INTEGER DEFAULT 0,
|
||||
last_used_at TEXT,
|
||||
last_error_at TEXT,
|
||||
last_error_message TEXT,
|
||||
created_at TEXT NOT NULL,
|
||||
updated_at TEXT NOT NULL
|
||||
)",
|
||||
[],
|
||||
)
|
||||
.unwrap();
|
||||
conn
|
||||
}
|
||||
|
||||
fn create_test_credential(id: &str, plugin_id: &str) -> NewPluginCredential {
|
||||
NewPluginCredential {
|
||||
id: id.to_string(),
|
||||
plugin_id: plugin_id.to_string(),
|
||||
auth_type: "oauth".to_string(),
|
||||
display_name: Some("Test Credential".to_string()),
|
||||
config_encrypted: r#"{"token":"test"}"#.to_string(),
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_create_and_get() {
|
||||
let conn = create_test_connection();
|
||||
let credential = create_test_credential("cred-1", "plugin-1");
|
||||
|
||||
PluginCredentialDao::create(&conn, &credential).unwrap();
|
||||
|
||||
let retrieved = PluginCredentialDao::get(&conn, "cred-1").unwrap().unwrap();
|
||||
assert_eq!(retrieved.id, "cred-1");
|
||||
assert_eq!(retrieved.plugin_id, "plugin-1");
|
||||
assert_eq!(retrieved.auth_type, "oauth");
|
||||
assert_eq!(retrieved.status, CredentialStatus::Active);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_list_by_plugin() {
|
||||
let conn = create_test_connection();
|
||||
|
||||
PluginCredentialDao::create(&conn, &create_test_credential("cred-1", "plugin-1")).unwrap();
|
||||
PluginCredentialDao::create(&conn, &create_test_credential("cred-2", "plugin-1")).unwrap();
|
||||
PluginCredentialDao::create(&conn, &create_test_credential("cred-3", "plugin-2")).unwrap();
|
||||
|
||||
let credentials = PluginCredentialDao::list_by_plugin(&conn, "plugin-1").unwrap();
|
||||
assert_eq!(credentials.len(), 2);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_update_status() {
|
||||
let conn = create_test_connection();
|
||||
PluginCredentialDao::create(&conn, &create_test_credential("cred-1", "plugin-1")).unwrap();
|
||||
|
||||
PluginCredentialDao::update_status(&conn, "cred-1", CredentialStatus::Disabled).unwrap();
|
||||
|
||||
let retrieved = PluginCredentialDao::get(&conn, "cred-1").unwrap().unwrap();
|
||||
assert_eq!(retrieved.status, CredentialStatus::Disabled);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_record_usage() {
|
||||
let conn = create_test_connection();
|
||||
PluginCredentialDao::create(&conn, &create_test_credential("cred-1", "plugin-1")).unwrap();
|
||||
|
||||
PluginCredentialDao::record_usage(&conn, "cred-1").unwrap();
|
||||
PluginCredentialDao::record_usage(&conn, "cred-1").unwrap();
|
||||
|
||||
let retrieved = PluginCredentialDao::get(&conn, "cred-1").unwrap().unwrap();
|
||||
assert_eq!(retrieved.usage_count, 2);
|
||||
assert!(retrieved.last_used_at.is_some());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_record_error() {
|
||||
let conn = create_test_connection();
|
||||
PluginCredentialDao::create(&conn, &create_test_credential("cred-1", "plugin-1")).unwrap();
|
||||
|
||||
PluginCredentialDao::record_error(&conn, "cred-1", "Token expired").unwrap();
|
||||
|
||||
let retrieved = PluginCredentialDao::get(&conn, "cred-1").unwrap().unwrap();
|
||||
assert_eq!(retrieved.error_count, 1);
|
||||
assert_eq!(
|
||||
retrieved.last_error_message,
|
||||
Some("Token expired".to_string())
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_delete() {
|
||||
let conn = create_test_connection();
|
||||
PluginCredentialDao::create(&conn, &create_test_credential("cred-1", "plugin-1")).unwrap();
|
||||
|
||||
let deleted = PluginCredentialDao::delete(&conn, "cred-1").unwrap();
|
||||
assert!(deleted);
|
||||
|
||||
let retrieved = PluginCredentialDao::get(&conn, "cred-1").unwrap();
|
||||
assert!(retrieved.is_none());
|
||||
}
|
||||
}
|
||||
@@ -267,3 +267,100 @@ struct ApiKeyMigrationRow {
|
||||
api_host: String,
|
||||
provider_name: String,
|
||||
}
|
||||
|
||||
/// 清理旧的 API Key 凭证(OpenAIKey 和 ClaudeKey 类型)
|
||||
///
|
||||
/// 这些凭证是通过旧的 UI 添加的,现在已经被新的 API Key Provider 系统取代。
|
||||
/// 此函数会删除 provider_pool_credentials 表中的 openai_key 和 claude_key 类型凭证。
|
||||
pub fn cleanup_legacy_api_key_credentials(conn: &Connection) -> Result<usize, String> {
|
||||
// 检查是否已经清理过
|
||||
let cleaned: bool = conn
|
||||
.query_row(
|
||||
"SELECT value FROM settings WHERE key = 'cleaned_legacy_api_key_credentials'",
|
||||
[],
|
||||
|row| row.get::<_, String>(0),
|
||||
)
|
||||
.map(|v| v == "true")
|
||||
.unwrap_or(false);
|
||||
|
||||
if cleaned {
|
||||
tracing::debug!("[清理] 旧 API Key 凭证已清理过,跳过");
|
||||
return Ok(0);
|
||||
}
|
||||
|
||||
tracing::info!("[清理] 开始清理旧的 API Key 凭证(openai_key, claude_key 类型)");
|
||||
|
||||
// 查询需要清理的凭证数量
|
||||
let count: i64 = conn
|
||||
.query_row(
|
||||
"SELECT COUNT(*) FROM provider_pool_credentials
|
||||
WHERE credential_data LIKE '%\"type\":\"openai_key\"%'
|
||||
OR credential_data LIKE '%\"type\":\"claude_key\"%'",
|
||||
[],
|
||||
|row| row.get(0),
|
||||
)
|
||||
.unwrap_or(0);
|
||||
|
||||
if count == 0 {
|
||||
tracing::info!("[清理] 没有需要清理的旧 API Key 凭证");
|
||||
// 标记清理完成
|
||||
conn.execute(
|
||||
"INSERT OR REPLACE INTO settings (key, value) VALUES ('cleaned_legacy_api_key_credentials', 'true')",
|
||||
[],
|
||||
)
|
||||
.map_err(|e| format!("标记清理完成失败: {}", e))?;
|
||||
return Ok(0);
|
||||
}
|
||||
|
||||
// 记录将要删除的凭证信息
|
||||
let mut stmt = conn
|
||||
.prepare(
|
||||
"SELECT uuid, name, provider_type, credential_data
|
||||
FROM provider_pool_credentials
|
||||
WHERE credential_data LIKE '%\"type\":\"openai_key\"%'
|
||||
OR credential_data LIKE '%\"type\":\"claude_key\"%'",
|
||||
)
|
||||
.map_err(|e| format!("准备查询语句失败: {}", e))?;
|
||||
|
||||
let rows = stmt
|
||||
.query_map([], |row| {
|
||||
Ok((
|
||||
row.get::<_, String>(0)?,
|
||||
row.get::<_, Option<String>>(1)?,
|
||||
row.get::<_, String>(2)?,
|
||||
))
|
||||
})
|
||||
.map_err(|e| format!("查询旧凭证失败: {}", e))?;
|
||||
|
||||
for row_result in rows {
|
||||
if let Ok((uuid, name, provider_type)) = row_result {
|
||||
tracing::info!(
|
||||
"[清理] 将删除旧凭证: {} (name: {}, type: {})",
|
||||
uuid,
|
||||
name.as_deref().unwrap_or("未命名"),
|
||||
provider_type
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
// 删除旧的 API Key 凭证
|
||||
let deleted = conn
|
||||
.execute(
|
||||
"DELETE FROM provider_pool_credentials
|
||||
WHERE credential_data LIKE '%\"type\":\"openai_key\"%'
|
||||
OR credential_data LIKE '%\"type\":\"claude_key\"%'",
|
||||
[],
|
||||
)
|
||||
.map_err(|e| format!("删除旧凭证失败: {}", e))?;
|
||||
|
||||
// 标记清理完成
|
||||
conn.execute(
|
||||
"INSERT OR REPLACE INTO settings (key, value) VALUES ('cleaned_legacy_api_key_credentials', 'true')",
|
||||
[],
|
||||
)
|
||||
.map_err(|e| format!("标记清理完成失败: {}", e))?;
|
||||
|
||||
tracing::info!("[清理] 旧 API Key 凭证清理完成,共删除 {} 条记录", deleted);
|
||||
|
||||
Ok(deleted)
|
||||
}
|
||||
|
||||
@@ -39,5 +39,17 @@ pub fn init_database() -> Result<DbConnection, String> {
|
||||
}
|
||||
}
|
||||
|
||||
// 清理旧的 API Key 凭证(openai_key, claude_key 类型)
|
||||
match migration::cleanup_legacy_api_key_credentials(&conn) {
|
||||
Ok(count) => {
|
||||
if count > 0 {
|
||||
tracing::info!("[数据库] 已清理 {} 条旧 API Key 凭证", count);
|
||||
}
|
||||
}
|
||||
Err(e) => {
|
||||
tracing::warn!("[数据库] 旧 API Key 凭证清理失败(非致命): {}", e);
|
||||
}
|
||||
}
|
||||
|
||||
Ok(Arc::new(Mutex::new(conn)))
|
||||
}
|
||||
|
||||
@@ -246,6 +246,210 @@ pub fn create_tables(conn: &Connection) -> Result<(), rusqlite::Error> {
|
||||
[],
|
||||
)?;
|
||||
|
||||
// OAuth Provider 插件表
|
||||
// 存储已安装的 OAuth Provider 插件信息
|
||||
conn.execute(
|
||||
"CREATE TABLE IF NOT EXISTS credential_provider_plugins (
|
||||
id TEXT PRIMARY KEY,
|
||||
display_name TEXT NOT NULL,
|
||||
version TEXT NOT NULL,
|
||||
description TEXT,
|
||||
author TEXT,
|
||||
homepage TEXT,
|
||||
license TEXT,
|
||||
target_protocol TEXT NOT NULL,
|
||||
install_path TEXT NOT NULL,
|
||||
binary_path TEXT,
|
||||
ui_entry TEXT,
|
||||
enabled INTEGER DEFAULT 1,
|
||||
config TEXT DEFAULT '{}',
|
||||
installed_at TEXT NOT NULL,
|
||||
updated_at TEXT NOT NULL,
|
||||
last_used_at TEXT,
|
||||
source_type TEXT NOT NULL DEFAULT 'local',
|
||||
source_data TEXT
|
||||
)",
|
||||
[],
|
||||
)?;
|
||||
|
||||
// 创建 credential_provider_plugins 索引
|
||||
conn.execute(
|
||||
"CREATE INDEX IF NOT EXISTS idx_credential_provider_plugins_protocol
|
||||
ON credential_provider_plugins(target_protocol)",
|
||||
[],
|
||||
)?;
|
||||
|
||||
// 插件凭证表
|
||||
// 存储每个插件管理的凭证
|
||||
conn.execute(
|
||||
"CREATE TABLE IF NOT EXISTS plugin_credentials (
|
||||
id TEXT PRIMARY KEY,
|
||||
plugin_id TEXT NOT NULL,
|
||||
auth_type TEXT NOT NULL,
|
||||
display_name TEXT,
|
||||
status TEXT NOT NULL DEFAULT 'active',
|
||||
config_encrypted TEXT NOT NULL,
|
||||
usage_count INTEGER DEFAULT 0,
|
||||
error_count INTEGER DEFAULT 0,
|
||||
last_used_at TEXT,
|
||||
last_error_at TEXT,
|
||||
last_error_message TEXT,
|
||||
created_at TEXT NOT NULL,
|
||||
updated_at TEXT NOT NULL,
|
||||
FOREIGN KEY (plugin_id) REFERENCES credential_provider_plugins(id) ON DELETE CASCADE
|
||||
)",
|
||||
[],
|
||||
)?;
|
||||
|
||||
// 创建 plugin_credentials 索引
|
||||
conn.execute(
|
||||
"CREATE INDEX IF NOT EXISTS idx_plugin_credentials_plugin
|
||||
ON plugin_credentials(plugin_id)",
|
||||
[],
|
||||
)?;
|
||||
conn.execute(
|
||||
"CREATE INDEX IF NOT EXISTS idx_plugin_credentials_status
|
||||
ON plugin_credentials(status)",
|
||||
[],
|
||||
)?;
|
||||
|
||||
// 插件存储表
|
||||
// 提供给插件的键值存储
|
||||
conn.execute(
|
||||
"CREATE TABLE IF NOT EXISTS plugin_storage (
|
||||
plugin_id TEXT NOT NULL,
|
||||
key TEXT NOT NULL,
|
||||
value TEXT NOT NULL,
|
||||
created_at TEXT NOT NULL,
|
||||
updated_at TEXT NOT NULL,
|
||||
PRIMARY KEY (plugin_id, key),
|
||||
FOREIGN KEY (plugin_id) REFERENCES credential_provider_plugins(id) ON DELETE CASCADE
|
||||
)",
|
||||
[],
|
||||
)?;
|
||||
|
||||
// 插件事件日志表
|
||||
// 记录插件的重要事件
|
||||
conn.execute(
|
||||
"CREATE TABLE IF NOT EXISTS plugin_event_logs (
|
||||
id INTEGER PRIMARY KEY AUTOINCREMENT,
|
||||
plugin_id TEXT NOT NULL,
|
||||
event_type TEXT NOT NULL,
|
||||
event_data TEXT,
|
||||
created_at TEXT NOT NULL,
|
||||
FOREIGN KEY (plugin_id) REFERENCES credential_provider_plugins(id) ON DELETE CASCADE
|
||||
)",
|
||||
[],
|
||||
)?;
|
||||
|
||||
// 创建 plugin_event_logs 索引
|
||||
conn.execute(
|
||||
"CREATE INDEX IF NOT EXISTS idx_plugin_event_logs_plugin
|
||||
ON plugin_event_logs(plugin_id)",
|
||||
[],
|
||||
)?;
|
||||
conn.execute(
|
||||
"CREATE INDEX IF NOT EXISTS idx_plugin_event_logs_type
|
||||
ON plugin_event_logs(event_type)",
|
||||
[],
|
||||
)?;
|
||||
conn.execute(
|
||||
"CREATE INDEX IF NOT EXISTS idx_plugin_event_logs_created
|
||||
ON plugin_event_logs(created_at)",
|
||||
[],
|
||||
)?;
|
||||
|
||||
// ============================================================================
|
||||
// Orchestrator 相关表
|
||||
// ============================================================================
|
||||
|
||||
// 模型元数据表
|
||||
// 存储模型的静态信息,用于智能选择
|
||||
conn.execute(
|
||||
"CREATE TABLE IF NOT EXISTS model_metadata (
|
||||
model_id TEXT PRIMARY KEY,
|
||||
provider_type TEXT NOT NULL,
|
||||
display_name TEXT NOT NULL,
|
||||
family TEXT,
|
||||
tier TEXT NOT NULL DEFAULT 'pro',
|
||||
context_length INTEGER,
|
||||
max_output_tokens INTEGER,
|
||||
cost_input_per_million REAL,
|
||||
cost_output_per_million REAL,
|
||||
supports_vision INTEGER DEFAULT 0,
|
||||
supports_tools INTEGER DEFAULT 0,
|
||||
supports_streaming INTEGER DEFAULT 1,
|
||||
is_deprecated INTEGER DEFAULT 0,
|
||||
release_date TEXT,
|
||||
description TEXT,
|
||||
created_at INTEGER NOT NULL,
|
||||
updated_at INTEGER NOT NULL
|
||||
)",
|
||||
[],
|
||||
)?;
|
||||
|
||||
// 创建 model_metadata 索引
|
||||
conn.execute(
|
||||
"CREATE INDEX IF NOT EXISTS idx_model_metadata_provider
|
||||
ON model_metadata(provider_type)",
|
||||
[],
|
||||
)?;
|
||||
conn.execute(
|
||||
"CREATE INDEX IF NOT EXISTS idx_model_metadata_tier
|
||||
ON model_metadata(tier)",
|
||||
[],
|
||||
)?;
|
||||
conn.execute(
|
||||
"CREATE INDEX IF NOT EXISTS idx_model_metadata_family
|
||||
ON model_metadata(family)",
|
||||
[],
|
||||
)?;
|
||||
|
||||
// 用户等级偏好表
|
||||
// 存储用户对每个服务等级的策略偏好
|
||||
conn.execute(
|
||||
"CREATE TABLE IF NOT EXISTS user_tier_preferences (
|
||||
tier_id TEXT PRIMARY KEY,
|
||||
strategy_id TEXT NOT NULL DEFAULT 'task_based',
|
||||
preferred_provider TEXT,
|
||||
fallback_enabled INTEGER DEFAULT 1,
|
||||
max_retries INTEGER DEFAULT 3,
|
||||
created_at INTEGER NOT NULL,
|
||||
updated_at INTEGER NOT NULL
|
||||
)",
|
||||
[],
|
||||
)?;
|
||||
|
||||
// 模型使用统计表
|
||||
// 记录每个模型的使用情况,用于智能选择
|
||||
conn.execute(
|
||||
"CREATE TABLE IF NOT EXISTS model_usage_stats (
|
||||
model_id TEXT NOT NULL,
|
||||
credential_id TEXT NOT NULL,
|
||||
date TEXT NOT NULL,
|
||||
request_count INTEGER DEFAULT 0,
|
||||
success_count INTEGER DEFAULT 0,
|
||||
error_count INTEGER DEFAULT 0,
|
||||
total_tokens INTEGER DEFAULT 0,
|
||||
total_latency_ms INTEGER DEFAULT 0,
|
||||
avg_latency_ms REAL,
|
||||
PRIMARY KEY (model_id, credential_id, date)
|
||||
)",
|
||||
[],
|
||||
)?;
|
||||
|
||||
// 创建 model_usage_stats 索引
|
||||
conn.execute(
|
||||
"CREATE INDEX IF NOT EXISTS idx_model_usage_stats_date
|
||||
ON model_usage_stats(date)",
|
||||
[],
|
||||
)?;
|
||||
conn.execute(
|
||||
"CREATE INDEX IF NOT EXISTS idx_model_usage_stats_model
|
||||
ON model_usage_stats(model_id)",
|
||||
[],
|
||||
)?;
|
||||
|
||||
Ok(())
|
||||
}
|
||||
|
||||
|
||||
+23
-2453
File diff suppressed because it is too large
Load Diff
@@ -84,6 +84,11 @@ pub enum CredentialData {
|
||||
IFlowOAuth { creds_file_path: String },
|
||||
/// iFlow Cookie 凭证
|
||||
IFlowCookie { creds_file_path: String },
|
||||
/// Anthropic API Key 凭证(直接使用 Anthropic API)
|
||||
AnthropicKey {
|
||||
api_key: String,
|
||||
base_url: Option<String>,
|
||||
},
|
||||
}
|
||||
|
||||
impl CredentialData {
|
||||
@@ -132,6 +137,9 @@ impl CredentialData {
|
||||
CredentialData::IFlowCookie { creds_file_path } => {
|
||||
format!("iFlow Cookie: {}", mask_path(creds_file_path))
|
||||
}
|
||||
CredentialData::AnthropicKey { api_key, .. } => {
|
||||
format!("Anthropic: {}", mask_key(api_key))
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
@@ -150,6 +158,7 @@ impl CredentialData {
|
||||
CredentialData::ClaudeOAuth { .. } => PoolProviderType::ClaudeOAuth,
|
||||
CredentialData::IFlowOAuth { .. } => PoolProviderType::IFlow,
|
||||
CredentialData::IFlowCookie { .. } => PoolProviderType::IFlow,
|
||||
CredentialData::AnthropicKey { .. } => PoolProviderType::Anthropic,
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -511,6 +520,11 @@ pub fn get_default_check_model(provider_type: PoolProviderType) -> &'static str
|
||||
PoolProviderType::Codex => "gpt-4o-mini",
|
||||
PoolProviderType::ClaudeOAuth => "claude-sonnet-4-5-20250929",
|
||||
PoolProviderType::IFlow => "deepseek-chat",
|
||||
// API Key Provider 类型
|
||||
PoolProviderType::Anthropic => "claude-sonnet-4-5-20250929",
|
||||
PoolProviderType::AzureOpenai => "gpt-4o-mini",
|
||||
PoolProviderType::AwsBedrock => "claude-sonnet-4-5-20250929",
|
||||
PoolProviderType::Ollama => "llama3.2",
|
||||
}
|
||||
}
|
||||
|
||||
@@ -563,6 +577,7 @@ fn get_credential_type(cred: &CredentialData) -> String {
|
||||
CredentialData::ClaudeOAuth { .. } => "claude_oauth".to_string(),
|
||||
CredentialData::IFlowOAuth { .. } => "iflow_oauth".to_string(),
|
||||
CredentialData::IFlowCookie { .. } => "iflow_cookie".to_string(),
|
||||
CredentialData::AnthropicKey { .. } => "anthropic_key".to_string(),
|
||||
}
|
||||
}
|
||||
|
||||
@@ -592,6 +607,7 @@ fn get_base_url(cred: &CredentialData) -> Option<String> {
|
||||
match cred {
|
||||
CredentialData::OpenAIKey { base_url, .. } => base_url.clone(),
|
||||
CredentialData::ClaudeKey { base_url, .. } => base_url.clone(),
|
||||
CredentialData::AnthropicKey { base_url, .. } => base_url.clone(),
|
||||
_ => None,
|
||||
}
|
||||
}
|
||||
@@ -601,6 +617,7 @@ fn get_api_key(cred: &CredentialData) -> Option<String> {
|
||||
match cred {
|
||||
CredentialData::OpenAIKey { api_key, .. } => Some(api_key.clone()),
|
||||
CredentialData::ClaudeKey { api_key, .. } => Some(api_key.clone()),
|
||||
CredentialData::AnthropicKey { api_key, .. } => Some(api_key.clone()),
|
||||
_ => None,
|
||||
}
|
||||
}
|
||||
|
||||
@@ -0,0 +1,322 @@
|
||||
//! 降级处理器
|
||||
//!
|
||||
//! 处理模型选择失败时的降级逻辑。
|
||||
|
||||
use super::tier::{AvailableModel, ServiceTier};
|
||||
use serde::{Deserialize, Serialize};
|
||||
|
||||
/// 降级策略
|
||||
#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
|
||||
#[serde(rename_all = "snake_case")]
|
||||
pub enum FallbackPolicy {
|
||||
/// 不降级,直接失败
|
||||
None,
|
||||
/// 降级到下一个等级
|
||||
NextTier,
|
||||
/// 降级到任意可用模型
|
||||
AnyAvailable,
|
||||
/// 使用指定的备用模型
|
||||
Specific,
|
||||
}
|
||||
|
||||
impl Default for FallbackPolicy {
|
||||
fn default() -> Self {
|
||||
FallbackPolicy::NextTier
|
||||
}
|
||||
}
|
||||
|
||||
/// 降级结果
|
||||
#[derive(Debug, Clone, Serialize, Deserialize)]
|
||||
pub struct FallbackResult {
|
||||
/// 是否成功降级
|
||||
pub success: bool,
|
||||
/// 降级后的模型
|
||||
pub model: Option<AvailableModel>,
|
||||
/// 原始等级
|
||||
pub original_tier: ServiceTier,
|
||||
/// 降级后的等级
|
||||
pub fallback_tier: Option<ServiceTier>,
|
||||
/// 降级原因
|
||||
pub reason: String,
|
||||
/// 尝试次数
|
||||
pub attempts: u32,
|
||||
}
|
||||
|
||||
/// 降级处理器
|
||||
pub struct FallbackHandler {
|
||||
/// 降级策略
|
||||
policy: FallbackPolicy,
|
||||
/// 最大尝试次数
|
||||
max_attempts: u32,
|
||||
/// 备用模型 ID(用于 Specific 策略)
|
||||
fallback_model_id: Option<String>,
|
||||
}
|
||||
|
||||
impl FallbackHandler {
|
||||
/// 创建新的降级处理器
|
||||
pub fn new(policy: FallbackPolicy) -> Self {
|
||||
Self {
|
||||
policy,
|
||||
max_attempts: 3,
|
||||
fallback_model_id: None,
|
||||
}
|
||||
}
|
||||
|
||||
/// 设置最大尝试次数
|
||||
pub fn with_max_attempts(mut self, max: u32) -> Self {
|
||||
self.max_attempts = max;
|
||||
self
|
||||
}
|
||||
|
||||
/// 设置备用模型 ID
|
||||
pub fn with_fallback_model(mut self, model_id: &str) -> Self {
|
||||
self.fallback_model_id = Some(model_id.to_string());
|
||||
self
|
||||
}
|
||||
|
||||
/// 获取降级策略
|
||||
pub fn policy(&self) -> FallbackPolicy {
|
||||
self.policy
|
||||
}
|
||||
|
||||
/// 获取下一个降级等级
|
||||
pub fn next_tier(tier: ServiceTier) -> Option<ServiceTier> {
|
||||
match tier {
|
||||
ServiceTier::Max => Some(ServiceTier::Pro),
|
||||
ServiceTier::Pro => Some(ServiceTier::Mini),
|
||||
ServiceTier::Mini => None,
|
||||
}
|
||||
}
|
||||
|
||||
/// 获取所有降级等级(按优先级排序)
|
||||
pub fn fallback_tiers(tier: ServiceTier) -> Vec<ServiceTier> {
|
||||
match tier {
|
||||
ServiceTier::Max => vec![ServiceTier::Pro, ServiceTier::Mini],
|
||||
ServiceTier::Pro => vec![ServiceTier::Mini],
|
||||
ServiceTier::Mini => vec![],
|
||||
}
|
||||
}
|
||||
|
||||
/// 处理降级
|
||||
pub fn handle(
|
||||
&self,
|
||||
original_tier: ServiceTier,
|
||||
available_models: &[(ServiceTier, Vec<AvailableModel>)],
|
||||
reason: &str,
|
||||
) -> FallbackResult {
|
||||
match self.policy {
|
||||
FallbackPolicy::None => FallbackResult {
|
||||
success: false,
|
||||
model: None,
|
||||
original_tier,
|
||||
fallback_tier: None,
|
||||
reason: format!("降级策略为 None,不进行降级: {}", reason),
|
||||
attempts: 0,
|
||||
},
|
||||
|
||||
FallbackPolicy::NextTier => {
|
||||
let fallback_tiers = Self::fallback_tiers(original_tier);
|
||||
let mut attempts = 0;
|
||||
|
||||
for tier in fallback_tiers {
|
||||
attempts += 1;
|
||||
if attempts > self.max_attempts {
|
||||
break;
|
||||
}
|
||||
|
||||
if let Some((_, models)) = available_models.iter().find(|(t, _)| *t == tier) {
|
||||
if let Some(model) = models.iter().find(|m| m.is_healthy).cloned() {
|
||||
return FallbackResult {
|
||||
success: true,
|
||||
model: Some(model),
|
||||
original_tier,
|
||||
fallback_tier: Some(tier),
|
||||
reason: format!("从 {} 降级到 {}: {}", original_tier, tier, reason),
|
||||
attempts,
|
||||
};
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
FallbackResult {
|
||||
success: false,
|
||||
model: None,
|
||||
original_tier,
|
||||
fallback_tier: None,
|
||||
reason: format!("所有降级等级都没有可用模型: {}", reason),
|
||||
attempts,
|
||||
}
|
||||
}
|
||||
|
||||
FallbackPolicy::AnyAvailable => {
|
||||
let mut attempts = 0;
|
||||
|
||||
// 按等级优先级遍历所有模型
|
||||
for tier in [ServiceTier::Max, ServiceTier::Pro, ServiceTier::Mini] {
|
||||
attempts += 1;
|
||||
if attempts > self.max_attempts {
|
||||
break;
|
||||
}
|
||||
|
||||
if let Some((_, models)) = available_models.iter().find(|(t, _)| *t == tier) {
|
||||
if let Some(model) = models.iter().find(|m| m.is_healthy).cloned() {
|
||||
return FallbackResult {
|
||||
success: true,
|
||||
model: Some(model),
|
||||
original_tier,
|
||||
fallback_tier: Some(tier),
|
||||
reason: format!("选择任意可用模型 (等级 {}): {}", tier, reason),
|
||||
attempts,
|
||||
};
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
FallbackResult {
|
||||
success: false,
|
||||
model: None,
|
||||
original_tier,
|
||||
fallback_tier: None,
|
||||
reason: format!("没有任何可用模型: {}", reason),
|
||||
attempts,
|
||||
}
|
||||
}
|
||||
|
||||
FallbackPolicy::Specific => {
|
||||
if let Some(fallback_id) = &self.fallback_model_id {
|
||||
for (tier, models) in available_models {
|
||||
if let Some(model) = models
|
||||
.iter()
|
||||
.find(|m| m.id == *fallback_id && m.is_healthy)
|
||||
.cloned()
|
||||
{
|
||||
return FallbackResult {
|
||||
success: true,
|
||||
model: Some(model),
|
||||
original_tier,
|
||||
fallback_tier: Some(*tier),
|
||||
reason: format!("使用指定备用模型 {}: {}", fallback_id, reason),
|
||||
attempts: 1,
|
||||
};
|
||||
}
|
||||
}
|
||||
|
||||
FallbackResult {
|
||||
success: false,
|
||||
model: None,
|
||||
original_tier,
|
||||
fallback_tier: None,
|
||||
reason: format!("指定的备用模型 {} 不可用: {}", fallback_id, reason),
|
||||
attempts: 1,
|
||||
}
|
||||
} else {
|
||||
FallbackResult {
|
||||
success: false,
|
||||
model: None,
|
||||
original_tier,
|
||||
fallback_tier: None,
|
||||
reason: format!("未配置备用模型: {}", reason),
|
||||
attempts: 0,
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
impl Default for FallbackHandler {
|
||||
fn default() -> Self {
|
||||
Self::new(FallbackPolicy::NextTier)
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
|
||||
fn create_test_models() -> Vec<(ServiceTier, Vec<AvailableModel>)> {
|
||||
vec![
|
||||
(
|
||||
ServiceTier::Mini,
|
||||
vec![AvailableModel {
|
||||
id: "haiku".to_string(),
|
||||
display_name: "Claude Haiku".to_string(),
|
||||
provider_type: "anthropic".to_string(),
|
||||
family: Some("haiku".to_string()),
|
||||
credential_id: "cred-1".to_string(),
|
||||
context_length: None,
|
||||
supports_vision: false,
|
||||
supports_tools: false,
|
||||
input_cost_per_million: None,
|
||||
output_cost_per_million: None,
|
||||
is_healthy: true,
|
||||
current_load: None,
|
||||
}],
|
||||
),
|
||||
(
|
||||
ServiceTier::Pro,
|
||||
vec![AvailableModel {
|
||||
id: "sonnet".to_string(),
|
||||
display_name: "Claude Sonnet".to_string(),
|
||||
provider_type: "anthropic".to_string(),
|
||||
family: Some("sonnet".to_string()),
|
||||
credential_id: "cred-2".to_string(),
|
||||
context_length: None,
|
||||
supports_vision: false,
|
||||
supports_tools: false,
|
||||
input_cost_per_million: None,
|
||||
output_cost_per_million: None,
|
||||
is_healthy: true,
|
||||
current_load: None,
|
||||
}],
|
||||
),
|
||||
]
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_fallback_next_tier() {
|
||||
let handler = FallbackHandler::new(FallbackPolicy::NextTier);
|
||||
let models = create_test_models();
|
||||
|
||||
let result = handler.handle(ServiceTier::Max, &models, "测试降级");
|
||||
|
||||
assert!(result.success);
|
||||
assert_eq!(result.fallback_tier, Some(ServiceTier::Pro));
|
||||
assert_eq!(result.model.unwrap().id, "sonnet");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_fallback_none() {
|
||||
let handler = FallbackHandler::new(FallbackPolicy::None);
|
||||
let models = create_test_models();
|
||||
|
||||
let result = handler.handle(ServiceTier::Max, &models, "测试降级");
|
||||
|
||||
assert!(!result.success);
|
||||
assert!(result.model.is_none());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_fallback_specific() {
|
||||
let handler = FallbackHandler::new(FallbackPolicy::Specific).with_fallback_model("haiku");
|
||||
let models = create_test_models();
|
||||
|
||||
let result = handler.handle(ServiceTier::Max, &models, "测试降级");
|
||||
|
||||
assert!(result.success);
|
||||
assert_eq!(result.model.unwrap().id, "haiku");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_next_tier() {
|
||||
assert_eq!(
|
||||
FallbackHandler::next_tier(ServiceTier::Max),
|
||||
Some(ServiceTier::Pro)
|
||||
);
|
||||
assert_eq!(
|
||||
FallbackHandler::next_tier(ServiceTier::Pro),
|
||||
Some(ServiceTier::Mini)
|
||||
);
|
||||
assert_eq!(FallbackHandler::next_tier(ServiceTier::Mini), None);
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,43 @@
|
||||
//! 模型编排器模块
|
||||
//!
|
||||
//! 提供 Mini/Pro/Max 服务等级的智能路由系统。
|
||||
//!
|
||||
//! ## 模块结构
|
||||
//!
|
||||
//! - `tier` - 服务等级定义 (Mini/Pro/Max)
|
||||
//! - `strategy` - 选择策略 trait 和注册表
|
||||
//! - `strategies` - 内置策略实现
|
||||
//! - `selector` - 模型选择器
|
||||
//! - `fallback` - 降级处理器
|
||||
//! - `pool_builder` - 动态模型池构建
|
||||
//! - `orchestrator` - 统一编排接口
|
||||
//!
|
||||
//! ## 使用模式
|
||||
//!
|
||||
//! 1. **简单模式(默认)**: Mini/Pro/Max 三档,动态根据用户凭证组合模型池
|
||||
//! 2. **专家模式**: 直接选择具体模型
|
||||
|
||||
mod fallback;
|
||||
mod orchestrator;
|
||||
mod pool_builder;
|
||||
mod selector;
|
||||
pub mod strategies;
|
||||
mod strategy;
|
||||
mod tier;
|
||||
|
||||
pub use fallback::{FallbackHandler, FallbackPolicy, FallbackResult};
|
||||
pub use orchestrator::{
|
||||
get_global_orchestrator, init_global_orchestrator, ModelOrchestrator, OrchestratorConfig,
|
||||
PoolStats,
|
||||
};
|
||||
pub use pool_builder::{
|
||||
builtin_model_metadata, builtin_provider_definitions, CredentialInfo, DynamicPoolBuilder,
|
||||
ModelFamily, ModelMetadata, ProviderDefinition, ProviderType,
|
||||
};
|
||||
pub use selector::{ModelSelector, SelectionResult};
|
||||
pub use strategies::*;
|
||||
pub use strategy::{
|
||||
ModelSelection, SelectionContext, SelectionStrategy, StrategyError, StrategyInfo,
|
||||
StrategyRegistry, StrategyResult, TaskHint,
|
||||
};
|
||||
pub use tier::{AvailableModel, ServiceTier, TierConfig, TierPool};
|
||||
@@ -0,0 +1,369 @@
|
||||
//! 模型编排器
|
||||
//!
|
||||
//! 统一的模型编排接口,整合模型池构建、策略选择和降级处理。
|
||||
|
||||
use super::fallback::{FallbackHandler, FallbackPolicy, FallbackResult};
|
||||
use super::pool_builder::{CredentialInfo, DynamicPoolBuilder, ProviderType};
|
||||
use super::selector::{ModelSelector, SelectionResult};
|
||||
use super::strategies::create_default_registry;
|
||||
use super::strategy::{SelectionContext, StrategyError, StrategyInfo, StrategyResult, TaskHint};
|
||||
use super::tier::{AvailableModel, ServiceTier, TierConfig, TierPool};
|
||||
use serde::{Deserialize, Serialize};
|
||||
use std::sync::Arc;
|
||||
use tokio::sync::RwLock;
|
||||
use tracing::{debug, info, warn};
|
||||
|
||||
/// 编排器配置
|
||||
#[derive(Debug, Clone, Serialize, Deserialize)]
|
||||
pub struct OrchestratorConfig {
|
||||
/// 默认服务等级
|
||||
pub default_tier: ServiceTier,
|
||||
/// 是否启用自动降级
|
||||
pub auto_fallback: bool,
|
||||
/// 降级策略
|
||||
pub fallback_policy: FallbackPolicy,
|
||||
/// 是否启用负载均衡
|
||||
pub load_balancing: bool,
|
||||
/// 模型池刷新间隔(秒)
|
||||
pub pool_refresh_interval: u64,
|
||||
}
|
||||
|
||||
impl Default for OrchestratorConfig {
|
||||
fn default() -> Self {
|
||||
Self {
|
||||
default_tier: ServiceTier::Pro,
|
||||
auto_fallback: true,
|
||||
fallback_policy: FallbackPolicy::NextTier,
|
||||
load_balancing: true,
|
||||
pool_refresh_interval: 60,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
/// 模型编排器
|
||||
///
|
||||
/// 提供统一的模型选择和管理接口
|
||||
pub struct ModelOrchestrator {
|
||||
/// 配置
|
||||
config: RwLock<OrchestratorConfig>,
|
||||
/// 模型选择器
|
||||
selector: ModelSelector,
|
||||
/// 模型池构建器
|
||||
pool_builder: DynamicPoolBuilder,
|
||||
/// 降级处理器
|
||||
fallback_handler: FallbackHandler,
|
||||
/// 当前凭证列表
|
||||
credentials: RwLock<Vec<CredentialInfo>>,
|
||||
}
|
||||
|
||||
impl ModelOrchestrator {
|
||||
/// 创建新的编排器
|
||||
pub fn new() -> Self {
|
||||
let registry = create_default_registry();
|
||||
let config = OrchestratorConfig::default();
|
||||
|
||||
Self {
|
||||
fallback_handler: FallbackHandler::new(config.fallback_policy),
|
||||
config: RwLock::new(config),
|
||||
selector: ModelSelector::new(registry),
|
||||
pool_builder: DynamicPoolBuilder::new(),
|
||||
credentials: RwLock::new(Vec::new()),
|
||||
}
|
||||
}
|
||||
|
||||
/// 使用自定义配置创建
|
||||
pub fn with_config(config: OrchestratorConfig) -> Self {
|
||||
let registry = create_default_registry();
|
||||
|
||||
Self {
|
||||
fallback_handler: FallbackHandler::new(config.fallback_policy),
|
||||
config: RwLock::new(config),
|
||||
selector: ModelSelector::new(registry),
|
||||
pool_builder: DynamicPoolBuilder::new(),
|
||||
credentials: RwLock::new(Vec::new()),
|
||||
}
|
||||
}
|
||||
|
||||
/// 更新配置
|
||||
pub async fn update_config(&self, config: OrchestratorConfig) {
|
||||
let mut current = self.config.write().await;
|
||||
*current = config;
|
||||
info!("编排器配置已更新");
|
||||
}
|
||||
|
||||
/// 获取配置
|
||||
pub async fn get_config(&self) -> OrchestratorConfig {
|
||||
self.config.read().await.clone()
|
||||
}
|
||||
|
||||
/// 更新凭证列表
|
||||
pub async fn update_credentials(&self, credentials: Vec<CredentialInfo>) {
|
||||
info!("更新凭证列表: {} 个凭证", credentials.len());
|
||||
|
||||
// 构建新的模型池
|
||||
let pool = self.pool_builder.build_pool(&credentials);
|
||||
|
||||
info!(
|
||||
"模型池已构建: Mini={}, Pro={}, Max={}",
|
||||
pool.get(ServiceTier::Mini).len(),
|
||||
pool.get(ServiceTier::Pro).len(),
|
||||
pool.get(ServiceTier::Max).len()
|
||||
);
|
||||
|
||||
// 更新选择器的模型池
|
||||
self.selector.update_pool(pool).await;
|
||||
|
||||
// 保存凭证列表
|
||||
let mut creds = self.credentials.write().await;
|
||||
*creds = credentials;
|
||||
}
|
||||
|
||||
/// 添加凭证
|
||||
pub async fn add_credential(&self, credential: CredentialInfo) {
|
||||
let mut creds = self.credentials.write().await;
|
||||
creds.push(credential);
|
||||
|
||||
// 重新构建模型池
|
||||
let pool = self.pool_builder.build_pool(&creds);
|
||||
drop(creds);
|
||||
|
||||
self.selector.update_pool(pool).await;
|
||||
}
|
||||
|
||||
/// 移除凭证
|
||||
pub async fn remove_credential(&self, credential_id: &str) {
|
||||
let mut creds = self.credentials.write().await;
|
||||
creds.retain(|c| c.id != credential_id);
|
||||
|
||||
// 重新构建模型池
|
||||
let pool = self.pool_builder.build_pool(&creds);
|
||||
drop(creds);
|
||||
|
||||
self.selector.update_pool(pool).await;
|
||||
}
|
||||
|
||||
/// 选择模型
|
||||
pub async fn select(&self, ctx: &SelectionContext) -> StrategyResult<SelectionResult> {
|
||||
debug!("选择模型: 等级={}, 任务={:?}", ctx.tier, ctx.task_hint);
|
||||
|
||||
self.selector.select(ctx).await
|
||||
}
|
||||
|
||||
/// 使用指定策略选择模型
|
||||
pub async fn select_with_strategy(
|
||||
&self,
|
||||
strategy_id: &str,
|
||||
ctx: &SelectionContext,
|
||||
) -> StrategyResult<SelectionResult> {
|
||||
self.selector.select_with_strategy(strategy_id, ctx).await
|
||||
}
|
||||
|
||||
/// 快速选择(使用默认等级和策略)
|
||||
pub async fn quick_select(&self) -> StrategyResult<SelectionResult> {
|
||||
let config = self.config.read().await;
|
||||
let ctx = SelectionContext::new(config.default_tier);
|
||||
drop(config);
|
||||
|
||||
self.select(&ctx).await
|
||||
}
|
||||
|
||||
/// 为特定任务选择模型
|
||||
pub async fn select_for_task(
|
||||
&self,
|
||||
tier: ServiceTier,
|
||||
task: TaskHint,
|
||||
) -> StrategyResult<SelectionResult> {
|
||||
let ctx = SelectionContext::new(tier).with_task_hint(task);
|
||||
self.select(&ctx).await
|
||||
}
|
||||
|
||||
/// 获取当前模型池
|
||||
pub async fn get_pool(&self) -> TierPool {
|
||||
self.selector.get_pool().await
|
||||
}
|
||||
|
||||
/// 获取指定等级的可用模型
|
||||
pub async fn get_models(&self, tier: ServiceTier) -> Vec<AvailableModel> {
|
||||
let pool = self.selector.get_pool().await;
|
||||
pool.get(tier).to_vec()
|
||||
}
|
||||
|
||||
/// 获取所有可用模型
|
||||
pub async fn get_all_models(&self) -> Vec<AvailableModel> {
|
||||
let pool = self.selector.get_pool().await;
|
||||
let mut all = Vec::new();
|
||||
all.extend(pool.get(ServiceTier::Mini).iter().cloned());
|
||||
all.extend(pool.get(ServiceTier::Pro).iter().cloned());
|
||||
all.extend(pool.get(ServiceTier::Max).iter().cloned());
|
||||
all
|
||||
}
|
||||
|
||||
/// 列出所有可用策略
|
||||
pub async fn list_strategies(&self) -> Vec<StrategyInfo> {
|
||||
self.selector.list_strategies().await
|
||||
}
|
||||
|
||||
/// 设置等级的默认策略
|
||||
pub fn set_tier_strategy(&mut self, tier: ServiceTier, strategy_id: &str) {
|
||||
self.selector.set_tier_strategy(tier, strategy_id);
|
||||
}
|
||||
|
||||
/// 获取模型池统计
|
||||
pub async fn get_pool_stats(&self) -> PoolStats {
|
||||
let pool = self.selector.get_pool().await;
|
||||
|
||||
PoolStats {
|
||||
mini_count: pool.get(ServiceTier::Mini).len(),
|
||||
pro_count: pool.get(ServiceTier::Pro).len(),
|
||||
max_count: pool.get(ServiceTier::Max).len(),
|
||||
total_count: pool.total_count(),
|
||||
healthy_count: pool
|
||||
.get(ServiceTier::Mini)
|
||||
.iter()
|
||||
.chain(pool.get(ServiceTier::Pro).iter())
|
||||
.chain(pool.get(ServiceTier::Max).iter())
|
||||
.filter(|m| m.is_healthy)
|
||||
.count(),
|
||||
}
|
||||
}
|
||||
|
||||
/// 标记模型为不健康
|
||||
pub async fn mark_unhealthy(&self, model_id: &str, credential_id: &str) {
|
||||
warn!("标记模型为不健康: {} (凭证: {})", model_id, credential_id);
|
||||
|
||||
let mut creds = self.credentials.write().await;
|
||||
if let Some(cred) = creds.iter_mut().find(|c| c.id == credential_id) {
|
||||
cred.is_healthy = false;
|
||||
}
|
||||
|
||||
// 重新构建模型池
|
||||
let pool = self.pool_builder.build_pool(&creds);
|
||||
drop(creds);
|
||||
|
||||
self.selector.update_pool(pool).await;
|
||||
}
|
||||
|
||||
/// 标记模型为健康
|
||||
pub async fn mark_healthy(&self, credential_id: &str) {
|
||||
info!("标记凭证为健康: {}", credential_id);
|
||||
|
||||
let mut creds = self.credentials.write().await;
|
||||
if let Some(cred) = creds.iter_mut().find(|c| c.id == credential_id) {
|
||||
cred.is_healthy = true;
|
||||
}
|
||||
|
||||
// 重新构建模型池
|
||||
let pool = self.pool_builder.build_pool(&creds);
|
||||
drop(creds);
|
||||
|
||||
self.selector.update_pool(pool).await;
|
||||
}
|
||||
|
||||
/// 更新凭证负载
|
||||
pub async fn update_load(&self, credential_id: &str, load: u8) {
|
||||
let mut creds = self.credentials.write().await;
|
||||
if let Some(cred) = creds.iter_mut().find(|c| c.id == credential_id) {
|
||||
cred.current_load = Some(load);
|
||||
}
|
||||
|
||||
// 重新构建模型池
|
||||
let pool = self.pool_builder.build_pool(&creds);
|
||||
drop(creds);
|
||||
|
||||
self.selector.update_pool(pool).await;
|
||||
}
|
||||
}
|
||||
|
||||
impl Default for ModelOrchestrator {
|
||||
fn default() -> Self {
|
||||
Self::new()
|
||||
}
|
||||
}
|
||||
|
||||
/// 模型池统计
|
||||
#[derive(Debug, Clone, Serialize, Deserialize)]
|
||||
pub struct PoolStats {
|
||||
/// Mini 等级模型数
|
||||
pub mini_count: usize,
|
||||
/// Pro 等级模型数
|
||||
pub pro_count: usize,
|
||||
/// Max 等级模型数
|
||||
pub max_count: usize,
|
||||
/// 总模型数
|
||||
pub total_count: usize,
|
||||
/// 健康模型数
|
||||
pub healthy_count: usize,
|
||||
}
|
||||
|
||||
/// 全局编排器实例
|
||||
static GLOBAL_ORCHESTRATOR: once_cell::sync::OnceCell<Arc<ModelOrchestrator>> =
|
||||
once_cell::sync::OnceCell::new();
|
||||
|
||||
/// 初始化全局编排器
|
||||
pub fn init_global_orchestrator() -> Arc<ModelOrchestrator> {
|
||||
GLOBAL_ORCHESTRATOR
|
||||
.get_or_init(|| Arc::new(ModelOrchestrator::new()))
|
||||
.clone()
|
||||
}
|
||||
|
||||
/// 获取全局编排器
|
||||
pub fn get_global_orchestrator() -> Option<Arc<ModelOrchestrator>> {
|
||||
GLOBAL_ORCHESTRATOR.get().cloned()
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_orchestrator_basic() {
|
||||
let orchestrator = ModelOrchestrator::new();
|
||||
|
||||
// 添加凭证
|
||||
orchestrator
|
||||
.update_credentials(vec![CredentialInfo {
|
||||
id: "cred-1".to_string(),
|
||||
provider_type: ProviderType::Anthropic,
|
||||
supported_models: vec![
|
||||
"claude-sonnet-4-5-20250514".to_string(),
|
||||
"claude-3-5-haiku-20241022".to_string(),
|
||||
],
|
||||
is_healthy: true,
|
||||
current_load: Some(30),
|
||||
}])
|
||||
.await;
|
||||
|
||||
// 获取统计
|
||||
let stats = orchestrator.get_pool_stats().await;
|
||||
assert!(stats.total_count > 0);
|
||||
|
||||
// 选择模型
|
||||
let ctx = SelectionContext::new(ServiceTier::Pro);
|
||||
let result = orchestrator.select(&ctx).await;
|
||||
assert!(result.is_ok());
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_orchestrator_task_selection() {
|
||||
let orchestrator = ModelOrchestrator::new();
|
||||
|
||||
orchestrator
|
||||
.update_credentials(vec![CredentialInfo {
|
||||
id: "cred-1".to_string(),
|
||||
provider_type: ProviderType::Anthropic,
|
||||
supported_models: vec![
|
||||
"claude-sonnet-4-5-20250514".to_string(),
|
||||
"claude-3-5-haiku-20241022".to_string(),
|
||||
],
|
||||
is_healthy: true,
|
||||
current_load: Some(30),
|
||||
}])
|
||||
.await;
|
||||
|
||||
// 为代码任务选择
|
||||
let result = orchestrator
|
||||
.select_for_task(ServiceTier::Pro, TaskHint::Coding)
|
||||
.await;
|
||||
assert!(result.is_ok());
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,643 @@
|
||||
//! 模型编排器 - 动态模型池构建
|
||||
//!
|
||||
//! 根据用户凭证动态构建各等级的模型池。
|
||||
|
||||
use super::tier::{AvailableModel, ServiceTier, TierPool};
|
||||
use serde::{Deserialize, Serialize};
|
||||
use std::collections::HashMap;
|
||||
|
||||
/// Provider 类型
|
||||
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash, Serialize, Deserialize)]
|
||||
#[serde(rename_all = "lowercase")]
|
||||
pub enum ProviderType {
|
||||
Anthropic,
|
||||
OpenAI,
|
||||
Google,
|
||||
Kiro,
|
||||
Azure,
|
||||
Bedrock,
|
||||
Custom,
|
||||
}
|
||||
|
||||
impl ProviderType {
|
||||
/// 从字符串解析
|
||||
pub fn from_str(s: &str) -> Option<Self> {
|
||||
match s.to_lowercase().as_str() {
|
||||
"anthropic" => Some(ProviderType::Anthropic),
|
||||
"openai" => Some(ProviderType::OpenAI),
|
||||
"google" | "gemini" => Some(ProviderType::Google),
|
||||
"kiro" | "codewhisperer" => Some(ProviderType::Kiro),
|
||||
"azure" => Some(ProviderType::Azure),
|
||||
"bedrock" => Some(ProviderType::Bedrock),
|
||||
_ => Some(ProviderType::Custom),
|
||||
}
|
||||
}
|
||||
|
||||
/// 获取显示名称
|
||||
pub fn display_name(&self) -> &'static str {
|
||||
match self {
|
||||
ProviderType::Anthropic => "Anthropic",
|
||||
ProviderType::OpenAI => "OpenAI",
|
||||
ProviderType::Google => "Google",
|
||||
ProviderType::Kiro => "Kiro",
|
||||
ProviderType::Azure => "Azure",
|
||||
ProviderType::Bedrock => "Bedrock",
|
||||
ProviderType::Custom => "Custom",
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
/// 模型家族定义
|
||||
#[derive(Debug, Clone, Serialize, Deserialize)]
|
||||
pub struct ModelFamily {
|
||||
/// 家族名称
|
||||
pub name: String,
|
||||
/// 匹配模式(glob 风格)
|
||||
pub pattern: String,
|
||||
/// 对应的服务等级 (1=Mini, 2=Pro, 3=Max)
|
||||
pub tier: u8,
|
||||
/// 描述
|
||||
pub description: Option<String>,
|
||||
}
|
||||
|
||||
/// Provider 定义
|
||||
#[derive(Debug, Clone, Serialize, Deserialize)]
|
||||
pub struct ProviderDefinition {
|
||||
/// Provider 类型
|
||||
pub provider_type: ProviderType,
|
||||
/// 显示名称
|
||||
pub display_name: String,
|
||||
/// 模型家族列表(按优先级排序)
|
||||
pub families: Vec<ModelFamily>,
|
||||
/// 默认 base URL
|
||||
pub default_base_url: Option<String>,
|
||||
}
|
||||
|
||||
impl ProviderDefinition {
|
||||
/// 获取模型的家族
|
||||
pub fn get_family(&self, model_id: &str) -> Option<&ModelFamily> {
|
||||
let model_lower = model_id.to_lowercase();
|
||||
self.families.iter().find(|f| {
|
||||
let pattern_lower = f.pattern.to_lowercase();
|
||||
if pattern_lower.contains('*') {
|
||||
// 简单的 glob 匹配
|
||||
let parts: Vec<&str> = pattern_lower.split('*').collect();
|
||||
if parts.len() == 2 {
|
||||
let prefix = parts[0];
|
||||
let suffix = parts[1];
|
||||
model_lower.starts_with(prefix) && model_lower.ends_with(suffix)
|
||||
} else if parts.len() == 1 {
|
||||
model_lower.starts_with(parts[0])
|
||||
} else {
|
||||
false
|
||||
}
|
||||
} else {
|
||||
model_lower.contains(&pattern_lower)
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
/// 获取模型对应的服务等级
|
||||
pub fn get_tier(&self, model_id: &str) -> Option<ServiceTier> {
|
||||
self.get_family(model_id).map(|f| match f.tier {
|
||||
1 => ServiceTier::Mini,
|
||||
2 => ServiceTier::Pro,
|
||||
3 => ServiceTier::Max,
|
||||
_ => ServiceTier::Pro,
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
/// 内置 Provider 定义
|
||||
pub fn builtin_provider_definitions() -> Vec<ProviderDefinition> {
|
||||
vec![
|
||||
// Anthropic
|
||||
ProviderDefinition {
|
||||
provider_type: ProviderType::Anthropic,
|
||||
display_name: "Anthropic".to_string(),
|
||||
families: vec![
|
||||
ModelFamily {
|
||||
name: "opus".to_string(),
|
||||
pattern: "claude-*opus*".to_string(),
|
||||
tier: 3,
|
||||
description: Some("Claude Opus - 最强能力".to_string()),
|
||||
},
|
||||
ModelFamily {
|
||||
name: "sonnet".to_string(),
|
||||
pattern: "claude-*sonnet*".to_string(),
|
||||
tier: 2,
|
||||
description: Some("Claude Sonnet - 均衡选择".to_string()),
|
||||
},
|
||||
ModelFamily {
|
||||
name: "haiku".to_string(),
|
||||
pattern: "claude-*haiku*".to_string(),
|
||||
tier: 1,
|
||||
description: Some("Claude Haiku - 快速响应".to_string()),
|
||||
},
|
||||
],
|
||||
default_base_url: Some("https://api.anthropic.com".to_string()),
|
||||
},
|
||||
// OpenAI
|
||||
ProviderDefinition {
|
||||
provider_type: ProviderType::OpenAI,
|
||||
display_name: "OpenAI".to_string(),
|
||||
families: vec![
|
||||
ModelFamily {
|
||||
name: "o1".to_string(),
|
||||
pattern: "o1*".to_string(),
|
||||
tier: 3,
|
||||
description: Some("O1 - 推理能力最强".to_string()),
|
||||
},
|
||||
ModelFamily {
|
||||
name: "gpt-4o".to_string(),
|
||||
pattern: "gpt-4o*".to_string(),
|
||||
tier: 2,
|
||||
description: Some("GPT-4o - 多模态均衡".to_string()),
|
||||
},
|
||||
ModelFamily {
|
||||
name: "gpt-4".to_string(),
|
||||
pattern: "gpt-4*".to_string(),
|
||||
tier: 2,
|
||||
description: Some("GPT-4 - 强大能力".to_string()),
|
||||
},
|
||||
ModelFamily {
|
||||
name: "gpt-3.5".to_string(),
|
||||
pattern: "gpt-3.5*".to_string(),
|
||||
tier: 1,
|
||||
description: Some("GPT-3.5 - 快速响应".to_string()),
|
||||
},
|
||||
],
|
||||
default_base_url: Some("https://api.openai.com".to_string()),
|
||||
},
|
||||
// Google
|
||||
ProviderDefinition {
|
||||
provider_type: ProviderType::Google,
|
||||
display_name: "Google".to_string(),
|
||||
families: vec![
|
||||
ModelFamily {
|
||||
name: "ultra".to_string(),
|
||||
pattern: "gemini-*ultra*".to_string(),
|
||||
tier: 3,
|
||||
description: Some("Gemini Ultra - 最强能力".to_string()),
|
||||
},
|
||||
ModelFamily {
|
||||
name: "pro".to_string(),
|
||||
pattern: "gemini-*pro*".to_string(),
|
||||
tier: 2,
|
||||
description: Some("Gemini Pro - 均衡选择".to_string()),
|
||||
},
|
||||
ModelFamily {
|
||||
name: "flash".to_string(),
|
||||
pattern: "gemini-*flash*".to_string(),
|
||||
tier: 1,
|
||||
description: Some("Gemini Flash - 快速响应".to_string()),
|
||||
},
|
||||
],
|
||||
default_base_url: Some("https://generativelanguage.googleapis.com".to_string()),
|
||||
},
|
||||
// Kiro (CodeWhisperer)
|
||||
ProviderDefinition {
|
||||
provider_type: ProviderType::Kiro,
|
||||
display_name: "Kiro".to_string(),
|
||||
families: vec![
|
||||
ModelFamily {
|
||||
name: "opus".to_string(),
|
||||
pattern: "claude-*opus*".to_string(),
|
||||
tier: 3,
|
||||
description: Some("Claude Opus via Kiro".to_string()),
|
||||
},
|
||||
ModelFamily {
|
||||
name: "sonnet".to_string(),
|
||||
pattern: "claude-*sonnet*".to_string(),
|
||||
tier: 2,
|
||||
description: Some("Claude Sonnet via Kiro".to_string()),
|
||||
},
|
||||
ModelFamily {
|
||||
name: "haiku".to_string(),
|
||||
pattern: "claude-*haiku*".to_string(),
|
||||
tier: 1,
|
||||
description: Some("Claude Haiku via Kiro".to_string()),
|
||||
},
|
||||
],
|
||||
default_base_url: None,
|
||||
},
|
||||
]
|
||||
}
|
||||
|
||||
/// 模型元数据
|
||||
#[derive(Debug, Clone, Serialize, Deserialize)]
|
||||
pub struct ModelMetadata {
|
||||
/// 模型 ID
|
||||
pub id: String,
|
||||
/// 显示名称
|
||||
pub display_name: String,
|
||||
/// Provider 类型
|
||||
pub provider_type: ProviderType,
|
||||
/// 模型家族
|
||||
pub family: Option<String>,
|
||||
/// 上下文长度
|
||||
pub context_length: Option<u32>,
|
||||
/// 是否支持视觉
|
||||
pub supports_vision: bool,
|
||||
/// 是否支持工具调用
|
||||
pub supports_tools: bool,
|
||||
/// 输入价格(每 1M tokens)
|
||||
pub input_cost_per_million: Option<f64>,
|
||||
/// 输出价格(每 1M tokens)
|
||||
pub output_cost_per_million: Option<f64>,
|
||||
/// 发布日期
|
||||
pub release_date: Option<String>,
|
||||
/// 是否是最新版本
|
||||
pub is_latest: bool,
|
||||
}
|
||||
|
||||
/// 内置模型元数据
|
||||
pub fn builtin_model_metadata() -> Vec<ModelMetadata> {
|
||||
vec![
|
||||
// Anthropic Models
|
||||
ModelMetadata {
|
||||
id: "claude-opus-4-5-20251101".to_string(),
|
||||
display_name: "Claude Opus 4.5".to_string(),
|
||||
provider_type: ProviderType::Anthropic,
|
||||
family: Some("opus".to_string()),
|
||||
context_length: Some(200000),
|
||||
supports_vision: true,
|
||||
supports_tools: true,
|
||||
input_cost_per_million: Some(15.0),
|
||||
output_cost_per_million: Some(75.0),
|
||||
release_date: Some("2025-11-01".to_string()),
|
||||
is_latest: true,
|
||||
},
|
||||
ModelMetadata {
|
||||
id: "claude-sonnet-4-5-20250514".to_string(),
|
||||
display_name: "Claude Sonnet 4.5".to_string(),
|
||||
provider_type: ProviderType::Anthropic,
|
||||
family: Some("sonnet".to_string()),
|
||||
context_length: Some(200000),
|
||||
supports_vision: true,
|
||||
supports_tools: true,
|
||||
input_cost_per_million: Some(3.0),
|
||||
output_cost_per_million: Some(15.0),
|
||||
release_date: Some("2025-05-14".to_string()),
|
||||
is_latest: true,
|
||||
},
|
||||
ModelMetadata {
|
||||
id: "claude-3-5-sonnet-20241022".to_string(),
|
||||
display_name: "Claude 3.5 Sonnet".to_string(),
|
||||
provider_type: ProviderType::Anthropic,
|
||||
family: Some("sonnet".to_string()),
|
||||
context_length: Some(200000),
|
||||
supports_vision: true,
|
||||
supports_tools: true,
|
||||
input_cost_per_million: Some(3.0),
|
||||
output_cost_per_million: Some(15.0),
|
||||
release_date: Some("2024-10-22".to_string()),
|
||||
is_latest: false,
|
||||
},
|
||||
ModelMetadata {
|
||||
id: "claude-3-5-haiku-20241022".to_string(),
|
||||
display_name: "Claude 3.5 Haiku".to_string(),
|
||||
provider_type: ProviderType::Anthropic,
|
||||
family: Some("haiku".to_string()),
|
||||
context_length: Some(200000),
|
||||
supports_vision: true,
|
||||
supports_tools: true,
|
||||
input_cost_per_million: Some(0.25),
|
||||
output_cost_per_million: Some(1.25),
|
||||
release_date: Some("2024-10-22".to_string()),
|
||||
is_latest: true,
|
||||
},
|
||||
// OpenAI Models
|
||||
ModelMetadata {
|
||||
id: "o1".to_string(),
|
||||
display_name: "O1".to_string(),
|
||||
provider_type: ProviderType::OpenAI,
|
||||
family: Some("o1".to_string()),
|
||||
context_length: Some(200000),
|
||||
supports_vision: true,
|
||||
supports_tools: true,
|
||||
input_cost_per_million: Some(15.0),
|
||||
output_cost_per_million: Some(60.0),
|
||||
release_date: Some("2024-12-01".to_string()),
|
||||
is_latest: true,
|
||||
},
|
||||
ModelMetadata {
|
||||
id: "gpt-4o".to_string(),
|
||||
display_name: "GPT-4o".to_string(),
|
||||
provider_type: ProviderType::OpenAI,
|
||||
family: Some("gpt-4o".to_string()),
|
||||
context_length: Some(128000),
|
||||
supports_vision: true,
|
||||
supports_tools: true,
|
||||
input_cost_per_million: Some(2.5),
|
||||
output_cost_per_million: Some(10.0),
|
||||
release_date: Some("2024-05-13".to_string()),
|
||||
is_latest: true,
|
||||
},
|
||||
ModelMetadata {
|
||||
id: "gpt-4-turbo".to_string(),
|
||||
display_name: "GPT-4 Turbo".to_string(),
|
||||
provider_type: ProviderType::OpenAI,
|
||||
family: Some("gpt-4".to_string()),
|
||||
context_length: Some(128000),
|
||||
supports_vision: true,
|
||||
supports_tools: true,
|
||||
input_cost_per_million: Some(10.0),
|
||||
output_cost_per_million: Some(30.0),
|
||||
release_date: Some("2024-04-09".to_string()),
|
||||
is_latest: false,
|
||||
},
|
||||
ModelMetadata {
|
||||
id: "gpt-3.5-turbo".to_string(),
|
||||
display_name: "GPT-3.5 Turbo".to_string(),
|
||||
provider_type: ProviderType::OpenAI,
|
||||
family: Some("gpt-3.5".to_string()),
|
||||
context_length: Some(16385),
|
||||
supports_vision: false,
|
||||
supports_tools: true,
|
||||
input_cost_per_million: Some(0.5),
|
||||
output_cost_per_million: Some(1.5),
|
||||
release_date: Some("2023-11-06".to_string()),
|
||||
is_latest: true,
|
||||
},
|
||||
// Google Models
|
||||
ModelMetadata {
|
||||
id: "gemini-2.0-flash".to_string(),
|
||||
display_name: "Gemini 2.0 Flash".to_string(),
|
||||
provider_type: ProviderType::Google,
|
||||
family: Some("flash".to_string()),
|
||||
context_length: Some(1000000),
|
||||
supports_vision: true,
|
||||
supports_tools: true,
|
||||
input_cost_per_million: Some(0.075),
|
||||
output_cost_per_million: Some(0.3),
|
||||
release_date: Some("2024-12-11".to_string()),
|
||||
is_latest: true,
|
||||
},
|
||||
ModelMetadata {
|
||||
id: "gemini-1.5-pro".to_string(),
|
||||
display_name: "Gemini 1.5 Pro".to_string(),
|
||||
provider_type: ProviderType::Google,
|
||||
family: Some("pro".to_string()),
|
||||
context_length: Some(2000000),
|
||||
supports_vision: true,
|
||||
supports_tools: true,
|
||||
input_cost_per_million: Some(1.25),
|
||||
output_cost_per_million: Some(5.0),
|
||||
release_date: Some("2024-05-14".to_string()),
|
||||
is_latest: true,
|
||||
},
|
||||
]
|
||||
}
|
||||
|
||||
/// 凭证信息(用于构建模型池)
|
||||
#[derive(Debug, Clone)]
|
||||
pub struct CredentialInfo {
|
||||
/// 凭证 ID
|
||||
pub id: String,
|
||||
/// Provider 类型
|
||||
pub provider_type: ProviderType,
|
||||
/// 支持的模型列表
|
||||
pub supported_models: Vec<String>,
|
||||
/// 是否健康
|
||||
pub is_healthy: bool,
|
||||
/// 当前负载
|
||||
pub current_load: Option<u8>,
|
||||
}
|
||||
|
||||
/// 动态模型池构建器
|
||||
pub struct DynamicPoolBuilder {
|
||||
/// Provider 定义
|
||||
provider_definitions: Vec<ProviderDefinition>,
|
||||
/// 模型元数据
|
||||
model_metadata: HashMap<String, ModelMetadata>,
|
||||
}
|
||||
|
||||
impl DynamicPoolBuilder {
|
||||
/// 创建新的构建器
|
||||
pub fn new() -> Self {
|
||||
let definitions = builtin_provider_definitions();
|
||||
let metadata: HashMap<_, _> = builtin_model_metadata()
|
||||
.into_iter()
|
||||
.map(|m| (m.id.clone(), m))
|
||||
.collect();
|
||||
|
||||
Self {
|
||||
provider_definitions: definitions,
|
||||
model_metadata: metadata,
|
||||
}
|
||||
}
|
||||
|
||||
/// 添加自定义 Provider 定义
|
||||
pub fn add_provider_definition(&mut self, definition: ProviderDefinition) {
|
||||
self.provider_definitions.push(definition);
|
||||
}
|
||||
|
||||
/// 添加模型元数据
|
||||
pub fn add_model_metadata(&mut self, metadata: ModelMetadata) {
|
||||
self.model_metadata.insert(metadata.id.clone(), metadata);
|
||||
}
|
||||
|
||||
/// 获取 Provider 定义
|
||||
pub fn get_provider_definition(
|
||||
&self,
|
||||
provider_type: ProviderType,
|
||||
) -> Option<&ProviderDefinition> {
|
||||
self.provider_definitions
|
||||
.iter()
|
||||
.find(|d| d.provider_type == provider_type)
|
||||
}
|
||||
|
||||
/// 根据凭证构建模型池
|
||||
pub fn build_pool(&self, credentials: &[CredentialInfo]) -> TierPool {
|
||||
let mut pool = TierPool::new();
|
||||
|
||||
for credential in credentials {
|
||||
if !credential.is_healthy {
|
||||
continue;
|
||||
}
|
||||
|
||||
let provider_def = match self.get_provider_definition(credential.provider_type) {
|
||||
Some(def) => def,
|
||||
None => continue,
|
||||
};
|
||||
|
||||
for model_id in &credential.supported_models {
|
||||
// 获取模型元数据
|
||||
let metadata = self.model_metadata.get(model_id);
|
||||
|
||||
// 确定服务等级
|
||||
let tier = provider_def.get_tier(model_id).unwrap_or(ServiceTier::Pro);
|
||||
|
||||
// 获取家族名称
|
||||
let family = provider_def
|
||||
.get_family(model_id)
|
||||
.map(|f| f.name.clone())
|
||||
.or_else(|| metadata.as_ref().and_then(|m| m.family.clone()));
|
||||
|
||||
// 构建 AvailableModel
|
||||
let available_model = AvailableModel {
|
||||
id: model_id.clone(),
|
||||
display_name: metadata
|
||||
.as_ref()
|
||||
.map(|m| m.display_name.clone())
|
||||
.unwrap_or_else(|| model_id.clone()),
|
||||
provider_type: format!("{:?}", credential.provider_type).to_lowercase(),
|
||||
family,
|
||||
credential_id: credential.id.clone(),
|
||||
context_length: metadata.as_ref().and_then(|m| m.context_length),
|
||||
supports_vision: metadata
|
||||
.as_ref()
|
||||
.map(|m| m.supports_vision)
|
||||
.unwrap_or(false),
|
||||
supports_tools: metadata.as_ref().map(|m| m.supports_tools).unwrap_or(false),
|
||||
input_cost_per_million: metadata
|
||||
.as_ref()
|
||||
.and_then(|m| m.input_cost_per_million),
|
||||
output_cost_per_million: metadata
|
||||
.as_ref()
|
||||
.and_then(|m| m.output_cost_per_million),
|
||||
is_healthy: credential.is_healthy,
|
||||
current_load: credential.current_load,
|
||||
};
|
||||
|
||||
pool.add(tier, available_model);
|
||||
}
|
||||
}
|
||||
|
||||
// 按评分排序
|
||||
pool.sort_by_score();
|
||||
|
||||
pool
|
||||
}
|
||||
|
||||
/// 为每个等级选择最佳模型(每个 Provider 一个)
|
||||
pub fn build_best_pool(&self, credentials: &[CredentialInfo]) -> TierPool {
|
||||
let full_pool = self.build_pool(credentials);
|
||||
let mut best_pool = TierPool::new();
|
||||
|
||||
for tier in ServiceTier::all() {
|
||||
let models = full_pool.get(*tier);
|
||||
let mut seen_providers: HashMap<String, bool> = HashMap::new();
|
||||
|
||||
for model in models {
|
||||
// 每个 Provider 只选择一个最佳模型
|
||||
if !seen_providers.contains_key(&model.provider_type) {
|
||||
seen_providers.insert(model.provider_type.clone(), true);
|
||||
best_pool.add(*tier, model.clone());
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
best_pool
|
||||
}
|
||||
}
|
||||
|
||||
impl Default for DynamicPoolBuilder {
|
||||
fn default() -> Self {
|
||||
Self::new()
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
|
||||
#[test]
|
||||
fn test_provider_definition_get_family() {
|
||||
let definitions = builtin_provider_definitions();
|
||||
let anthropic = definitions
|
||||
.iter()
|
||||
.find(|d| d.provider_type == ProviderType::Anthropic)
|
||||
.unwrap();
|
||||
|
||||
let family = anthropic.get_family("claude-3-5-sonnet-20241022");
|
||||
assert!(family.is_some());
|
||||
assert_eq!(family.unwrap().name, "sonnet");
|
||||
|
||||
let family = anthropic.get_family("claude-opus-4-5-20251101");
|
||||
assert!(family.is_some());
|
||||
assert_eq!(family.unwrap().name, "opus");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_provider_definition_get_tier() {
|
||||
let definitions = builtin_provider_definitions();
|
||||
let anthropic = definitions
|
||||
.iter()
|
||||
.find(|d| d.provider_type == ProviderType::Anthropic)
|
||||
.unwrap();
|
||||
|
||||
assert_eq!(
|
||||
anthropic.get_tier("claude-3-5-haiku-20241022"),
|
||||
Some(ServiceTier::Mini)
|
||||
);
|
||||
assert_eq!(
|
||||
anthropic.get_tier("claude-3-5-sonnet-20241022"),
|
||||
Some(ServiceTier::Pro)
|
||||
);
|
||||
assert_eq!(
|
||||
anthropic.get_tier("claude-opus-4-5-20251101"),
|
||||
Some(ServiceTier::Max)
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_dynamic_pool_builder() {
|
||||
let builder = DynamicPoolBuilder::new();
|
||||
|
||||
let credentials = vec![
|
||||
CredentialInfo {
|
||||
id: "cred-1".to_string(),
|
||||
provider_type: ProviderType::Anthropic,
|
||||
supported_models: vec![
|
||||
"claude-opus-4-5-20251101".to_string(),
|
||||
"claude-sonnet-4-5-20250514".to_string(),
|
||||
"claude-3-5-haiku-20241022".to_string(),
|
||||
],
|
||||
is_healthy: true,
|
||||
current_load: Some(30),
|
||||
},
|
||||
CredentialInfo {
|
||||
id: "cred-2".to_string(),
|
||||
provider_type: ProviderType::OpenAI,
|
||||
supported_models: vec!["gpt-4o".to_string(), "gpt-3.5-turbo".to_string()],
|
||||
is_healthy: true,
|
||||
current_load: Some(20),
|
||||
},
|
||||
];
|
||||
|
||||
let pool = builder.build_pool(&credentials);
|
||||
|
||||
assert!(!pool.is_empty());
|
||||
assert!(!pool.get(ServiceTier::Mini).is_empty());
|
||||
assert!(!pool.get(ServiceTier::Pro).is_empty());
|
||||
assert!(!pool.get(ServiceTier::Max).is_empty());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_build_best_pool() {
|
||||
let builder = DynamicPoolBuilder::new();
|
||||
|
||||
let credentials = vec![CredentialInfo {
|
||||
id: "cred-1".to_string(),
|
||||
provider_type: ProviderType::Anthropic,
|
||||
supported_models: vec![
|
||||
"claude-sonnet-4-5-20250514".to_string(),
|
||||
"claude-3-5-sonnet-20241022".to_string(),
|
||||
],
|
||||
is_healthy: true,
|
||||
current_load: Some(30),
|
||||
}];
|
||||
|
||||
let pool = builder.build_best_pool(&credentials);
|
||||
|
||||
// Pro 等级应该只有一个 Anthropic 模型
|
||||
let pro_models = pool.get(ServiceTier::Pro);
|
||||
let anthropic_count = pro_models
|
||||
.iter()
|
||||
.filter(|m| m.provider_type == "anthropic")
|
||||
.count();
|
||||
assert_eq!(anthropic_count, 1);
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,325 @@
|
||||
//! 模型选择器
|
||||
//!
|
||||
//! 提供统一的模型选择接口,整合策略和模型池。
|
||||
|
||||
use super::strategy::{SelectionContext, StrategyError, StrategyRegistry, StrategyResult};
|
||||
use super::tier::{AvailableModel, ServiceTier, TierConfig, TierPool};
|
||||
use serde::{Deserialize, Serialize};
|
||||
use std::collections::HashMap;
|
||||
use std::sync::Arc;
|
||||
use tokio::sync::RwLock;
|
||||
use tracing::{debug, info, warn};
|
||||
|
||||
/// 选择结果
|
||||
#[derive(Debug, Clone, Serialize, Deserialize)]
|
||||
pub struct SelectionResult {
|
||||
/// 选中的模型
|
||||
pub model: AvailableModel,
|
||||
/// 使用的策略 ID
|
||||
pub strategy_id: String,
|
||||
/// 选择原因
|
||||
pub reason: String,
|
||||
/// 置信度 (0-100)
|
||||
pub confidence: u8,
|
||||
/// 服务等级
|
||||
pub tier: ServiceTier,
|
||||
/// 是否是降级选择
|
||||
pub is_fallback: bool,
|
||||
/// 降级原因(如果是降级)
|
||||
pub fallback_reason: Option<String>,
|
||||
}
|
||||
|
||||
/// 模型选择器
|
||||
pub struct ModelSelector {
|
||||
/// 策略注册表
|
||||
registry: Arc<RwLock<StrategyRegistry>>,
|
||||
/// 等级配置
|
||||
tier_configs: HashMap<ServiceTier, TierConfig>,
|
||||
/// 模型池
|
||||
pool: Arc<RwLock<TierPool>>,
|
||||
}
|
||||
|
||||
impl ModelSelector {
|
||||
/// 创建新的模型选择器
|
||||
pub fn new(registry: StrategyRegistry) -> Self {
|
||||
Self {
|
||||
registry: Arc::new(RwLock::new(registry)),
|
||||
tier_configs: TierConfig::defaults(),
|
||||
pool: Arc::new(RwLock::new(TierPool::new())),
|
||||
}
|
||||
}
|
||||
|
||||
/// 使用自定义配置创建
|
||||
pub fn with_configs(
|
||||
registry: StrategyRegistry,
|
||||
configs: HashMap<ServiceTier, TierConfig>,
|
||||
) -> Self {
|
||||
Self {
|
||||
registry: Arc::new(RwLock::new(registry)),
|
||||
tier_configs: configs,
|
||||
pool: Arc::new(RwLock::new(TierPool::new())),
|
||||
}
|
||||
}
|
||||
|
||||
/// 更新模型池
|
||||
pub async fn update_pool(&self, pool: TierPool) {
|
||||
let mut current = self.pool.write().await;
|
||||
*current = pool;
|
||||
info!(
|
||||
"模型池已更新: Mini={}, Pro={}, Max={}",
|
||||
current.mini.len(),
|
||||
current.pro.len(),
|
||||
current.max.len()
|
||||
);
|
||||
}
|
||||
|
||||
/// 获取模型池
|
||||
pub async fn get_pool(&self) -> TierPool {
|
||||
self.pool.read().await.clone()
|
||||
}
|
||||
|
||||
/// 选择模型
|
||||
pub async fn select(&self, ctx: &SelectionContext) -> StrategyResult<SelectionResult> {
|
||||
let pool = self.pool.read().await;
|
||||
let models = pool.get(ctx.tier);
|
||||
|
||||
if models.is_empty() {
|
||||
warn!("等级 {} 没有可用模型,尝试降级", ctx.tier);
|
||||
return self.select_with_fallback(ctx).await;
|
||||
}
|
||||
|
||||
// 获取等级配置
|
||||
let config = self
|
||||
.tier_configs
|
||||
.get(&ctx.tier)
|
||||
.cloned()
|
||||
.unwrap_or_else(|| TierConfig::pro());
|
||||
|
||||
// 获取策略
|
||||
let registry = self.registry.read().await;
|
||||
let strategy = registry
|
||||
.get(&config.default_strategy)
|
||||
.or_else(|| registry.get_default())
|
||||
.ok_or_else(|| StrategyError::StrategyNotFound(config.default_strategy.clone()))?;
|
||||
|
||||
debug!("使用策略 {} 选择模型 (等级: {})", strategy.id(), ctx.tier);
|
||||
|
||||
// 执行选择
|
||||
let selection = strategy.select(models, ctx).await?;
|
||||
|
||||
Ok(SelectionResult {
|
||||
model: selection.model,
|
||||
strategy_id: strategy.id().to_string(),
|
||||
reason: selection.reason,
|
||||
confidence: selection.confidence,
|
||||
tier: ctx.tier,
|
||||
is_fallback: false,
|
||||
fallback_reason: None,
|
||||
})
|
||||
}
|
||||
|
||||
/// 使用指定策略选择模型
|
||||
pub async fn select_with_strategy(
|
||||
&self,
|
||||
strategy_id: &str,
|
||||
ctx: &SelectionContext,
|
||||
) -> StrategyResult<SelectionResult> {
|
||||
let pool = self.pool.read().await;
|
||||
let models = pool.get(ctx.tier);
|
||||
|
||||
if models.is_empty() {
|
||||
return Err(StrategyError::NoAvailableModels);
|
||||
}
|
||||
|
||||
let registry = self.registry.read().await;
|
||||
let strategy = registry
|
||||
.get(strategy_id)
|
||||
.ok_or_else(|| StrategyError::StrategyNotFound(strategy_id.to_string()))?;
|
||||
|
||||
let selection = strategy.select(models, ctx).await?;
|
||||
|
||||
Ok(SelectionResult {
|
||||
model: selection.model,
|
||||
strategy_id: strategy.id().to_string(),
|
||||
reason: selection.reason,
|
||||
confidence: selection.confidence,
|
||||
tier: ctx.tier,
|
||||
is_fallback: false,
|
||||
fallback_reason: None,
|
||||
})
|
||||
}
|
||||
|
||||
/// 带降级的选择
|
||||
async fn select_with_fallback(
|
||||
&self,
|
||||
ctx: &SelectionContext,
|
||||
) -> StrategyResult<SelectionResult> {
|
||||
let pool = self.pool.read().await;
|
||||
|
||||
// 尝试降级到更低等级
|
||||
let fallback_tiers = match ctx.tier {
|
||||
ServiceTier::Max => vec![ServiceTier::Pro, ServiceTier::Mini],
|
||||
ServiceTier::Pro => vec![ServiceTier::Mini],
|
||||
ServiceTier::Mini => vec![],
|
||||
};
|
||||
|
||||
for fallback_tier in fallback_tiers {
|
||||
let models = pool.get(fallback_tier);
|
||||
if !models.is_empty() {
|
||||
let mut fallback_ctx = ctx.clone();
|
||||
fallback_ctx.tier = fallback_tier;
|
||||
|
||||
let config = self
|
||||
.tier_configs
|
||||
.get(&fallback_tier)
|
||||
.cloned()
|
||||
.unwrap_or_else(|| TierConfig::pro());
|
||||
|
||||
let registry = self.registry.read().await;
|
||||
let strategy = registry
|
||||
.get(&config.default_strategy)
|
||||
.or_else(|| registry.get_default())
|
||||
.ok_or_else(|| {
|
||||
StrategyError::StrategyNotFound(config.default_strategy.clone())
|
||||
})?;
|
||||
|
||||
let selection = strategy.select(models, &fallback_ctx).await?;
|
||||
|
||||
info!(
|
||||
"降级选择: {} -> {} (模型: {})",
|
||||
ctx.tier, fallback_tier, selection.model.id
|
||||
);
|
||||
|
||||
return Ok(SelectionResult {
|
||||
model: selection.model,
|
||||
strategy_id: strategy.id().to_string(),
|
||||
reason: selection.reason,
|
||||
confidence: selection.confidence.saturating_sub(20), // 降级降低置信度
|
||||
tier: fallback_tier,
|
||||
is_fallback: true,
|
||||
fallback_reason: Some(format!(
|
||||
"等级 {} 无可用模型,降级到 {}",
|
||||
ctx.tier, fallback_tier
|
||||
)),
|
||||
});
|
||||
}
|
||||
}
|
||||
|
||||
Err(StrategyError::NoAvailableModels)
|
||||
}
|
||||
|
||||
/// 获取策略注册表
|
||||
pub async fn get_registry(&self) -> Arc<RwLock<StrategyRegistry>> {
|
||||
self.registry.clone()
|
||||
}
|
||||
|
||||
/// 列出所有可用策略
|
||||
pub async fn list_strategies(&self) -> Vec<super::strategy::StrategyInfo> {
|
||||
let registry = self.registry.read().await;
|
||||
registry.list_all()
|
||||
}
|
||||
|
||||
/// 设置等级的默认策略
|
||||
pub fn set_tier_strategy(&mut self, tier: ServiceTier, strategy_id: &str) {
|
||||
if let Some(config) = self.tier_configs.get_mut(&tier) {
|
||||
config.default_strategy = strategy_id.to_string();
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
use crate::orchestrator::strategies::create_default_registry;
|
||||
|
||||
fn create_test_pool() -> TierPool {
|
||||
let mut pool = TierPool::new();
|
||||
|
||||
pool.add(
|
||||
ServiceTier::Mini,
|
||||
AvailableModel {
|
||||
id: "haiku".to_string(),
|
||||
display_name: "Claude Haiku".to_string(),
|
||||
provider_type: "anthropic".to_string(),
|
||||
family: Some("haiku".to_string()),
|
||||
credential_id: "cred-1".to_string(),
|
||||
context_length: Some(200000),
|
||||
supports_vision: true,
|
||||
supports_tools: true,
|
||||
input_cost_per_million: None,
|
||||
output_cost_per_million: None,
|
||||
is_healthy: true,
|
||||
current_load: Some(20),
|
||||
},
|
||||
);
|
||||
|
||||
pool.add(
|
||||
ServiceTier::Pro,
|
||||
AvailableModel {
|
||||
id: "sonnet".to_string(),
|
||||
display_name: "Claude Sonnet".to_string(),
|
||||
provider_type: "anthropic".to_string(),
|
||||
family: Some("sonnet".to_string()),
|
||||
credential_id: "cred-2".to_string(),
|
||||
context_length: Some(200000),
|
||||
supports_vision: true,
|
||||
supports_tools: true,
|
||||
input_cost_per_million: None,
|
||||
output_cost_per_million: None,
|
||||
is_healthy: true,
|
||||
current_load: Some(30),
|
||||
},
|
||||
);
|
||||
|
||||
pool
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_model_selector() {
|
||||
let registry = create_default_registry();
|
||||
let selector = ModelSelector::new(registry);
|
||||
|
||||
selector.update_pool(create_test_pool()).await;
|
||||
|
||||
let ctx = SelectionContext::new(ServiceTier::Pro);
|
||||
let result = selector.select(&ctx).await.unwrap();
|
||||
|
||||
assert_eq!(result.tier, ServiceTier::Pro);
|
||||
assert!(!result.is_fallback);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_fallback_selection() {
|
||||
let registry = create_default_registry();
|
||||
let selector = ModelSelector::new(registry);
|
||||
|
||||
// 只有 Mini 等级有模型
|
||||
let mut pool = TierPool::new();
|
||||
pool.add(
|
||||
ServiceTier::Mini,
|
||||
AvailableModel {
|
||||
id: "haiku".to_string(),
|
||||
display_name: "Claude Haiku".to_string(),
|
||||
provider_type: "anthropic".to_string(),
|
||||
family: Some("haiku".to_string()),
|
||||
credential_id: "cred-1".to_string(),
|
||||
context_length: None,
|
||||
supports_vision: false,
|
||||
supports_tools: false,
|
||||
input_cost_per_million: None,
|
||||
output_cost_per_million: None,
|
||||
is_healthy: true,
|
||||
current_load: None,
|
||||
},
|
||||
);
|
||||
selector.update_pool(pool).await;
|
||||
|
||||
// 请求 Max 等级,应该降级到 Mini
|
||||
let ctx = SelectionContext::new(ServiceTier::Max);
|
||||
let result = selector.select(&ctx).await.unwrap();
|
||||
|
||||
assert_eq!(result.tier, ServiceTier::Mini);
|
||||
assert!(result.is_fallback);
|
||||
assert!(result.fallback_reason.is_some());
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,157 @@
|
||||
//! 成本优化策略
|
||||
//!
|
||||
//! 选择成本最低的模型。
|
||||
|
||||
use crate::orchestrator::strategy::{
|
||||
ModelSelection, SelectionContext, SelectionStrategy, StrategyError, StrategyResult,
|
||||
};
|
||||
use crate::orchestrator::tier::AvailableModel;
|
||||
use async_trait::async_trait;
|
||||
|
||||
/// 成本优化策略
|
||||
pub struct CostOptimizedStrategy;
|
||||
|
||||
impl CostOptimizedStrategy {
|
||||
/// 创建新的成本优化策略
|
||||
pub fn new() -> Self {
|
||||
Self
|
||||
}
|
||||
|
||||
/// 计算模型的成本得分(越低越好)
|
||||
fn cost_score(model: &AvailableModel) -> f64 {
|
||||
// 如果有价格信息,使用价格
|
||||
if let (Some(input), Some(output)) =
|
||||
(model.input_cost_per_million, model.output_cost_per_million)
|
||||
{
|
||||
// 假设输入输出比例为 1:1
|
||||
return input + output;
|
||||
}
|
||||
|
||||
// 否则根据家族估算成本
|
||||
let family = model.family.as_deref().unwrap_or("").to_lowercase();
|
||||
|
||||
if family.contains("haiku") || family.contains("flash") || family.contains("gpt-3.5") {
|
||||
1.0 // 最便宜
|
||||
} else if family.contains("sonnet") || family.contains("pro") {
|
||||
5.0 // 中等
|
||||
} else if family.contains("opus") || family.contains("ultra") || family.contains("o1") {
|
||||
15.0 // 最贵
|
||||
} else if family.contains("gpt-4") {
|
||||
10.0 // 较贵
|
||||
} else {
|
||||
5.0 // 默认中等
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
impl Default for CostOptimizedStrategy {
|
||||
fn default() -> Self {
|
||||
Self::new()
|
||||
}
|
||||
}
|
||||
|
||||
#[async_trait]
|
||||
impl SelectionStrategy for CostOptimizedStrategy {
|
||||
fn id(&self) -> &str {
|
||||
"cost_optimized"
|
||||
}
|
||||
|
||||
fn display_name(&self) -> &str {
|
||||
"成本优先"
|
||||
}
|
||||
|
||||
fn description(&self) -> &str {
|
||||
"选择成本最低的模型"
|
||||
}
|
||||
|
||||
async fn select(
|
||||
&self,
|
||||
pool: &[AvailableModel],
|
||||
ctx: &SelectionContext,
|
||||
) -> StrategyResult<ModelSelection> {
|
||||
// 过滤可用模型
|
||||
let mut available: Vec<_> = pool
|
||||
.iter()
|
||||
.filter(|m| {
|
||||
m.is_healthy
|
||||
&& !ctx.excluded_models.contains(&m.id)
|
||||
&& (!ctx.requires_vision || m.supports_vision)
|
||||
&& (!ctx.requires_tools || m.supports_tools)
|
||||
})
|
||||
.cloned()
|
||||
.collect();
|
||||
|
||||
if available.is_empty() {
|
||||
return Err(StrategyError::NoAvailableModels);
|
||||
}
|
||||
|
||||
// 按成本排序(从低到高)
|
||||
available.sort_by(|a, b| {
|
||||
let cost_a = Self::cost_score(a);
|
||||
let cost_b = Self::cost_score(b);
|
||||
cost_a
|
||||
.partial_cmp(&cost_b)
|
||||
.unwrap_or(std::cmp::Ordering::Equal)
|
||||
});
|
||||
|
||||
let selected = available.remove(0);
|
||||
let cost = Self::cost_score(&selected);
|
||||
|
||||
Ok(ModelSelection {
|
||||
model: selected,
|
||||
reason: format!("成本优先选择 (估算成本: {:.2})", cost),
|
||||
confidence: 90,
|
||||
alternatives: available,
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
use crate::orchestrator::tier::ServiceTier;
|
||||
|
||||
fn create_test_models() -> Vec<AvailableModel> {
|
||||
vec![
|
||||
AvailableModel {
|
||||
id: "claude-opus".to_string(),
|
||||
display_name: "Claude Opus".to_string(),
|
||||
provider_type: "anthropic".to_string(),
|
||||
family: Some("opus".to_string()),
|
||||
credential_id: "cred-1".to_string(),
|
||||
context_length: Some(200000),
|
||||
supports_vision: true,
|
||||
supports_tools: true,
|
||||
input_cost_per_million: Some(15.0),
|
||||
output_cost_per_million: Some(75.0),
|
||||
is_healthy: true,
|
||||
current_load: None,
|
||||
},
|
||||
AvailableModel {
|
||||
id: "claude-haiku".to_string(),
|
||||
display_name: "Claude Haiku".to_string(),
|
||||
provider_type: "anthropic".to_string(),
|
||||
family: Some("haiku".to_string()),
|
||||
credential_id: "cred-2".to_string(),
|
||||
context_length: Some(200000),
|
||||
supports_vision: true,
|
||||
supports_tools: true,
|
||||
input_cost_per_million: Some(0.25),
|
||||
output_cost_per_million: Some(1.25),
|
||||
is_healthy: true,
|
||||
current_load: None,
|
||||
},
|
||||
]
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_cost_optimized_selection() {
|
||||
let strategy = CostOptimizedStrategy::new();
|
||||
let models = create_test_models();
|
||||
let ctx = SelectionContext::new(ServiceTier::Pro);
|
||||
|
||||
let result = strategy.select(&models, &ctx).await.unwrap();
|
||||
// 应该选择最便宜的 Haiku
|
||||
assert_eq!(result.model.id, "claude-haiku");
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,158 @@
|
||||
//! 负载均衡策略
|
||||
//!
|
||||
//! 根据当前负载选择模型,实现负载均衡。
|
||||
|
||||
use crate::orchestrator::strategy::{
|
||||
ModelSelection, SelectionContext, SelectionStrategy, StrategyError, StrategyResult,
|
||||
};
|
||||
use crate::orchestrator::tier::AvailableModel;
|
||||
use async_trait::async_trait;
|
||||
|
||||
/// 负载均衡策略
|
||||
pub struct LoadBalancedStrategy;
|
||||
|
||||
impl LoadBalancedStrategy {
|
||||
/// 创建新的负载均衡策略
|
||||
pub fn new() -> Self {
|
||||
Self
|
||||
}
|
||||
|
||||
/// 计算模型的负载得分(越低越好)
|
||||
fn load_score(model: &AvailableModel) -> f64 {
|
||||
// 基础负载
|
||||
let load = model.current_load.unwrap_or(50) as f64;
|
||||
|
||||
// 如果不健康,给予最高负载
|
||||
if !model.is_healthy {
|
||||
return 1000.0;
|
||||
}
|
||||
|
||||
load
|
||||
}
|
||||
}
|
||||
|
||||
impl Default for LoadBalancedStrategy {
|
||||
fn default() -> Self {
|
||||
Self::new()
|
||||
}
|
||||
}
|
||||
|
||||
#[async_trait]
|
||||
impl SelectionStrategy for LoadBalancedStrategy {
|
||||
fn id(&self) -> &str {
|
||||
"load_balanced"
|
||||
}
|
||||
|
||||
fn display_name(&self) -> &str {
|
||||
"负载均衡"
|
||||
}
|
||||
|
||||
fn description(&self) -> &str {
|
||||
"根据当前负载选择模型,实现负载均衡"
|
||||
}
|
||||
|
||||
async fn select(
|
||||
&self,
|
||||
pool: &[AvailableModel],
|
||||
ctx: &SelectionContext,
|
||||
) -> StrategyResult<ModelSelection> {
|
||||
// 过滤可用模型
|
||||
let mut available: Vec<_> = pool
|
||||
.iter()
|
||||
.filter(|m| {
|
||||
m.is_healthy
|
||||
&& !ctx.excluded_models.contains(&m.id)
|
||||
&& (!ctx.requires_vision || m.supports_vision)
|
||||
&& (!ctx.requires_tools || m.supports_tools)
|
||||
})
|
||||
.cloned()
|
||||
.collect();
|
||||
|
||||
if available.is_empty() {
|
||||
return Err(StrategyError::NoAvailableModels);
|
||||
}
|
||||
|
||||
// 按负载排序(从低到高)
|
||||
available.sort_by(|a, b| {
|
||||
let load_a = Self::load_score(a);
|
||||
let load_b = Self::load_score(b);
|
||||
load_a
|
||||
.partial_cmp(&load_b)
|
||||
.unwrap_or(std::cmp::Ordering::Equal)
|
||||
});
|
||||
|
||||
let selected = available.remove(0);
|
||||
let load = selected.current_load.unwrap_or(50);
|
||||
|
||||
Ok(ModelSelection {
|
||||
model: selected,
|
||||
reason: format!("负载均衡选择 (当前负载: {}%)", load),
|
||||
confidence: 80,
|
||||
alternatives: available,
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
use crate::orchestrator::tier::ServiceTier;
|
||||
|
||||
fn create_test_models() -> Vec<AvailableModel> {
|
||||
vec![
|
||||
AvailableModel {
|
||||
id: "model-high-load".to_string(),
|
||||
display_name: "High Load Model".to_string(),
|
||||
provider_type: "test".to_string(),
|
||||
family: None,
|
||||
credential_id: "cred-1".to_string(),
|
||||
context_length: None,
|
||||
supports_vision: false,
|
||||
supports_tools: false,
|
||||
input_cost_per_million: None,
|
||||
output_cost_per_million: None,
|
||||
is_healthy: true,
|
||||
current_load: Some(80),
|
||||
},
|
||||
AvailableModel {
|
||||
id: "model-low-load".to_string(),
|
||||
display_name: "Low Load Model".to_string(),
|
||||
provider_type: "test".to_string(),
|
||||
family: None,
|
||||
credential_id: "cred-2".to_string(),
|
||||
context_length: None,
|
||||
supports_vision: false,
|
||||
supports_tools: false,
|
||||
input_cost_per_million: None,
|
||||
output_cost_per_million: None,
|
||||
is_healthy: true,
|
||||
current_load: Some(20),
|
||||
},
|
||||
AvailableModel {
|
||||
id: "model-medium-load".to_string(),
|
||||
display_name: "Medium Load Model".to_string(),
|
||||
provider_type: "test".to_string(),
|
||||
family: None,
|
||||
credential_id: "cred-3".to_string(),
|
||||
context_length: None,
|
||||
supports_vision: false,
|
||||
supports_tools: false,
|
||||
input_cost_per_million: None,
|
||||
output_cost_per_million: None,
|
||||
is_healthy: true,
|
||||
current_load: Some(50),
|
||||
},
|
||||
]
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_load_balanced_selection() {
|
||||
let strategy = LoadBalancedStrategy::new();
|
||||
let models = create_test_models();
|
||||
let ctx = SelectionContext::new(ServiceTier::Pro);
|
||||
|
||||
let result = strategy.select(&models, &ctx).await.unwrap();
|
||||
// 应该选择负载最低的模型
|
||||
assert_eq!(result.model.id, "model-low-load");
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,34 @@
|
||||
//! 内置选择策略
|
||||
//!
|
||||
//! 提供多种模型选择策略实现。
|
||||
|
||||
mod cost_optimized;
|
||||
mod load_balanced;
|
||||
mod round_robin;
|
||||
mod speed_optimized;
|
||||
mod task_based;
|
||||
|
||||
pub use cost_optimized::CostOptimizedStrategy;
|
||||
pub use load_balanced::LoadBalancedStrategy;
|
||||
pub use round_robin::RoundRobinStrategy;
|
||||
pub use speed_optimized::SpeedOptimizedStrategy;
|
||||
pub use task_based::TaskBasedStrategy;
|
||||
|
||||
use super::strategy::StrategyRegistry;
|
||||
use std::sync::Arc;
|
||||
|
||||
/// 注册所有内置策略
|
||||
pub fn register_builtin_strategies(registry: &mut StrategyRegistry) {
|
||||
registry.register(Arc::new(RoundRobinStrategy::new()));
|
||||
registry.register(Arc::new(TaskBasedStrategy::new()));
|
||||
registry.register(Arc::new(CostOptimizedStrategy::new()));
|
||||
registry.register(Arc::new(SpeedOptimizedStrategy::new()));
|
||||
registry.register(Arc::new(LoadBalancedStrategy::new()));
|
||||
}
|
||||
|
||||
/// 创建带有内置策略的注册表
|
||||
pub fn create_default_registry() -> StrategyRegistry {
|
||||
let mut registry = StrategyRegistry::new();
|
||||
register_builtin_strategies(&mut registry);
|
||||
registry
|
||||
}
|
||||
@@ -0,0 +1,172 @@
|
||||
//! 轮询策略
|
||||
//!
|
||||
//! 按顺序轮询选择模型,实现简单的负载分散。
|
||||
|
||||
use crate::orchestrator::strategy::{
|
||||
ModelSelection, SelectionContext, SelectionStrategy, StrategyError, StrategyResult,
|
||||
};
|
||||
use crate::orchestrator::tier::AvailableModel;
|
||||
use async_trait::async_trait;
|
||||
use std::sync::atomic::{AtomicUsize, Ordering};
|
||||
|
||||
/// 轮询策略
|
||||
pub struct RoundRobinStrategy {
|
||||
/// 当前索引
|
||||
index: AtomicUsize,
|
||||
}
|
||||
|
||||
impl RoundRobinStrategy {
|
||||
/// 创建新的轮询策略
|
||||
pub fn new() -> Self {
|
||||
Self {
|
||||
index: AtomicUsize::new(0),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
impl Default for RoundRobinStrategy {
|
||||
fn default() -> Self {
|
||||
Self::new()
|
||||
}
|
||||
}
|
||||
|
||||
#[async_trait]
|
||||
impl SelectionStrategy for RoundRobinStrategy {
|
||||
fn id(&self) -> &str {
|
||||
"round_robin"
|
||||
}
|
||||
|
||||
fn display_name(&self) -> &str {
|
||||
"轮询"
|
||||
}
|
||||
|
||||
fn description(&self) -> &str {
|
||||
"按顺序轮询选择模型,实现简单的负载分散"
|
||||
}
|
||||
|
||||
async fn select(
|
||||
&self,
|
||||
pool: &[AvailableModel],
|
||||
ctx: &SelectionContext,
|
||||
) -> StrategyResult<ModelSelection> {
|
||||
// 过滤可用模型
|
||||
let available: Vec<_> = pool
|
||||
.iter()
|
||||
.filter(|m| {
|
||||
m.is_healthy
|
||||
&& !ctx.excluded_models.contains(&m.id)
|
||||
&& (!ctx.requires_vision || m.supports_vision)
|
||||
&& (!ctx.requires_tools || m.supports_tools)
|
||||
})
|
||||
.collect();
|
||||
|
||||
if available.is_empty() {
|
||||
return Err(StrategyError::NoAvailableModels);
|
||||
}
|
||||
|
||||
// 获取下一个索引
|
||||
let idx = self.index.fetch_add(1, Ordering::Relaxed) % available.len();
|
||||
let selected = available[idx].clone();
|
||||
|
||||
// 构建备选列表
|
||||
let alternatives: Vec<_> = available
|
||||
.iter()
|
||||
.enumerate()
|
||||
.filter(|(i, _)| *i != idx)
|
||||
.map(|(_, m)| (*m).clone())
|
||||
.collect();
|
||||
|
||||
Ok(ModelSelection {
|
||||
model: selected,
|
||||
reason: format!("轮询选择 (索引 {})", idx),
|
||||
confidence: 80,
|
||||
alternatives,
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
use crate::orchestrator::tier::ServiceTier;
|
||||
|
||||
fn create_test_models() -> Vec<AvailableModel> {
|
||||
vec![
|
||||
AvailableModel {
|
||||
id: "model-1".to_string(),
|
||||
display_name: "Model 1".to_string(),
|
||||
provider_type: "test".to_string(),
|
||||
family: None,
|
||||
credential_id: "cred-1".to_string(),
|
||||
context_length: None,
|
||||
supports_vision: false,
|
||||
supports_tools: false,
|
||||
input_cost_per_million: None,
|
||||
output_cost_per_million: None,
|
||||
is_healthy: true,
|
||||
current_load: None,
|
||||
},
|
||||
AvailableModel {
|
||||
id: "model-2".to_string(),
|
||||
display_name: "Model 2".to_string(),
|
||||
provider_type: "test".to_string(),
|
||||
family: None,
|
||||
credential_id: "cred-2".to_string(),
|
||||
context_length: None,
|
||||
supports_vision: false,
|
||||
supports_tools: false,
|
||||
input_cost_per_million: None,
|
||||
output_cost_per_million: None,
|
||||
is_healthy: true,
|
||||
current_load: None,
|
||||
},
|
||||
AvailableModel {
|
||||
id: "model-3".to_string(),
|
||||
display_name: "Model 3".to_string(),
|
||||
provider_type: "test".to_string(),
|
||||
family: None,
|
||||
credential_id: "cred-3".to_string(),
|
||||
context_length: None,
|
||||
supports_vision: false,
|
||||
supports_tools: false,
|
||||
input_cost_per_million: None,
|
||||
output_cost_per_million: None,
|
||||
is_healthy: true,
|
||||
current_load: None,
|
||||
},
|
||||
]
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_round_robin_selection() {
|
||||
let strategy = RoundRobinStrategy::new();
|
||||
let models = create_test_models();
|
||||
let ctx = SelectionContext::new(ServiceTier::Pro);
|
||||
|
||||
// 第一次选择
|
||||
let result1 = strategy.select(&models, &ctx).await.unwrap();
|
||||
assert_eq!(result1.model.id, "model-1");
|
||||
|
||||
// 第二次选择
|
||||
let result2 = strategy.select(&models, &ctx).await.unwrap();
|
||||
assert_eq!(result2.model.id, "model-2");
|
||||
|
||||
// 第三次选择
|
||||
let result3 = strategy.select(&models, &ctx).await.unwrap();
|
||||
assert_eq!(result3.model.id, "model-3");
|
||||
|
||||
// 第四次选择(回到第一个)
|
||||
let result4 = strategy.select(&models, &ctx).await.unwrap();
|
||||
assert_eq!(result4.model.id, "model-1");
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_round_robin_empty_pool() {
|
||||
let strategy = RoundRobinStrategy::new();
|
||||
let models: Vec<AvailableModel> = vec![];
|
||||
let ctx = SelectionContext::new(ServiceTier::Pro);
|
||||
|
||||
let result = strategy.select(&models, &ctx).await;
|
||||
assert!(result.is_err());
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,157 @@
|
||||
//! 速度优化策略
|
||||
//!
|
||||
//! 选择响应速度最快的模型。
|
||||
|
||||
use crate::orchestrator::strategy::{
|
||||
ModelSelection, SelectionContext, SelectionStrategy, StrategyError, StrategyResult,
|
||||
};
|
||||
use crate::orchestrator::tier::AvailableModel;
|
||||
use async_trait::async_trait;
|
||||
|
||||
/// 速度优化策略
|
||||
pub struct SpeedOptimizedStrategy;
|
||||
|
||||
impl SpeedOptimizedStrategy {
|
||||
/// 创建新的速度优化策略
|
||||
pub fn new() -> Self {
|
||||
Self
|
||||
}
|
||||
|
||||
/// 计算模型的速度得分(越高越好)
|
||||
fn speed_score(model: &AvailableModel) -> f64 {
|
||||
let mut score = 100.0;
|
||||
|
||||
// 根据家族估算速度
|
||||
let family = model.family.as_deref().unwrap_or("").to_lowercase();
|
||||
|
||||
if family.contains("haiku") || family.contains("flash") {
|
||||
score += 50.0; // 最快
|
||||
} else if family.contains("gpt-3.5") {
|
||||
score += 40.0;
|
||||
} else if family.contains("sonnet") || family.contains("pro") {
|
||||
score += 20.0; // 中等
|
||||
} else if family.contains("gpt-4") {
|
||||
score += 10.0;
|
||||
} else if family.contains("opus") || family.contains("ultra") || family.contains("o1") {
|
||||
score += 0.0; // 最慢
|
||||
}
|
||||
|
||||
// 负载惩罚(负载越高,速度越慢)
|
||||
if let Some(load) = model.current_load {
|
||||
score -= load as f64 * 0.5;
|
||||
}
|
||||
|
||||
score
|
||||
}
|
||||
}
|
||||
|
||||
impl Default for SpeedOptimizedStrategy {
|
||||
fn default() -> Self {
|
||||
Self::new()
|
||||
}
|
||||
}
|
||||
|
||||
#[async_trait]
|
||||
impl SelectionStrategy for SpeedOptimizedStrategy {
|
||||
fn id(&self) -> &str {
|
||||
"speed_optimized"
|
||||
}
|
||||
|
||||
fn display_name(&self) -> &str {
|
||||
"速度优先"
|
||||
}
|
||||
|
||||
fn description(&self) -> &str {
|
||||
"选择响应速度最快的模型"
|
||||
}
|
||||
|
||||
async fn select(
|
||||
&self,
|
||||
pool: &[AvailableModel],
|
||||
ctx: &SelectionContext,
|
||||
) -> StrategyResult<ModelSelection> {
|
||||
// 过滤可用模型
|
||||
let mut available: Vec<_> = pool
|
||||
.iter()
|
||||
.filter(|m| {
|
||||
m.is_healthy
|
||||
&& !ctx.excluded_models.contains(&m.id)
|
||||
&& (!ctx.requires_vision || m.supports_vision)
|
||||
&& (!ctx.requires_tools || m.supports_tools)
|
||||
})
|
||||
.cloned()
|
||||
.collect();
|
||||
|
||||
if available.is_empty() {
|
||||
return Err(StrategyError::NoAvailableModels);
|
||||
}
|
||||
|
||||
// 按速度排序(从高到低)
|
||||
available.sort_by(|a, b| {
|
||||
let speed_a = Self::speed_score(a);
|
||||
let speed_b = Self::speed_score(b);
|
||||
speed_b
|
||||
.partial_cmp(&speed_a)
|
||||
.unwrap_or(std::cmp::Ordering::Equal)
|
||||
});
|
||||
|
||||
let selected = available.remove(0);
|
||||
|
||||
Ok(ModelSelection {
|
||||
model: selected,
|
||||
reason: "速度优先选择".to_string(),
|
||||
confidence: 85,
|
||||
alternatives: available,
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
use crate::orchestrator::tier::ServiceTier;
|
||||
|
||||
fn create_test_models() -> Vec<AvailableModel> {
|
||||
vec![
|
||||
AvailableModel {
|
||||
id: "claude-opus".to_string(),
|
||||
display_name: "Claude Opus".to_string(),
|
||||
provider_type: "anthropic".to_string(),
|
||||
family: Some("opus".to_string()),
|
||||
credential_id: "cred-1".to_string(),
|
||||
context_length: Some(200000),
|
||||
supports_vision: true,
|
||||
supports_tools: true,
|
||||
input_cost_per_million: None,
|
||||
output_cost_per_million: None,
|
||||
is_healthy: true,
|
||||
current_load: Some(20),
|
||||
},
|
||||
AvailableModel {
|
||||
id: "claude-haiku".to_string(),
|
||||
display_name: "Claude Haiku".to_string(),
|
||||
provider_type: "anthropic".to_string(),
|
||||
family: Some("haiku".to_string()),
|
||||
credential_id: "cred-2".to_string(),
|
||||
context_length: Some(200000),
|
||||
supports_vision: true,
|
||||
supports_tools: true,
|
||||
input_cost_per_million: None,
|
||||
output_cost_per_million: None,
|
||||
is_healthy: true,
|
||||
current_load: Some(10),
|
||||
},
|
||||
]
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_speed_optimized_selection() {
|
||||
let strategy = SpeedOptimizedStrategy::new();
|
||||
let models = create_test_models();
|
||||
let ctx = SelectionContext::new(ServiceTier::Mini);
|
||||
|
||||
let result = strategy.select(&models, &ctx).await.unwrap();
|
||||
// 应该选择最快的 Haiku
|
||||
assert_eq!(result.model.id, "claude-haiku");
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,264 @@
|
||||
//! 任务匹配策略
|
||||
//!
|
||||
//! 根据任务类型选择最适合的模型。
|
||||
|
||||
use crate::orchestrator::strategy::{
|
||||
ModelSelection, SelectionContext, SelectionStrategy, StrategyError, StrategyResult, TaskHint,
|
||||
};
|
||||
use crate::orchestrator::tier::AvailableModel;
|
||||
use async_trait::async_trait;
|
||||
|
||||
/// 任务匹配策略
|
||||
pub struct TaskBasedStrategy;
|
||||
|
||||
impl TaskBasedStrategy {
|
||||
/// 创建新的任务匹配策略
|
||||
pub fn new() -> Self {
|
||||
Self
|
||||
}
|
||||
|
||||
/// 根据任务类型计算模型得分
|
||||
fn score_for_task(model: &AvailableModel, task: Option<TaskHint>) -> f64 {
|
||||
let mut score = 0.0;
|
||||
|
||||
// 基础分:健康状态
|
||||
if !model.is_healthy {
|
||||
return 0.0;
|
||||
}
|
||||
|
||||
let family = model.family.as_deref().unwrap_or("").to_lowercase();
|
||||
|
||||
match task {
|
||||
Some(TaskHint::Coding) => {
|
||||
// 代码任务偏好 Sonnet/GPT-4 级别
|
||||
if family.contains("sonnet") || family.contains("gpt-4") {
|
||||
score += 100.0;
|
||||
} else if family.contains("opus") || family.contains("o1") {
|
||||
score += 90.0;
|
||||
} else if family.contains("haiku") || family.contains("flash") {
|
||||
score += 60.0;
|
||||
}
|
||||
// 工具调用对代码任务很重要
|
||||
if model.supports_tools {
|
||||
score += 20.0;
|
||||
}
|
||||
}
|
||||
Some(TaskHint::Writing) | Some(TaskHint::Analysis) => {
|
||||
// 写作/分析任务偏好 Opus/O1 级别
|
||||
if family.contains("opus") || family.contains("o1") {
|
||||
score += 100.0;
|
||||
} else if family.contains("sonnet") || family.contains("gpt-4") {
|
||||
score += 80.0;
|
||||
} else {
|
||||
score += 50.0;
|
||||
}
|
||||
}
|
||||
Some(TaskHint::Chat) => {
|
||||
// 对话任务偏好快速响应
|
||||
if family.contains("haiku") || family.contains("flash") {
|
||||
score += 100.0;
|
||||
} else if family.contains("sonnet") {
|
||||
score += 80.0;
|
||||
} else {
|
||||
score += 60.0;
|
||||
}
|
||||
}
|
||||
Some(TaskHint::Math) => {
|
||||
// 数学任务偏好推理能力强的模型
|
||||
if family.contains("o1") {
|
||||
score += 100.0;
|
||||
} else if family.contains("opus") {
|
||||
score += 90.0;
|
||||
} else if family.contains("sonnet") || family.contains("gpt-4") {
|
||||
score += 70.0;
|
||||
} else {
|
||||
score += 50.0;
|
||||
}
|
||||
}
|
||||
Some(TaskHint::Translation) | Some(TaskHint::Summarization) => {
|
||||
// 翻译/摘要任务偏好均衡模型
|
||||
if family.contains("sonnet") || family.contains("gpt-4") {
|
||||
score += 100.0;
|
||||
} else if family.contains("opus") {
|
||||
score += 80.0;
|
||||
} else {
|
||||
score += 60.0;
|
||||
}
|
||||
}
|
||||
_ => {
|
||||
// 默认:按家族等级评分
|
||||
if family.contains("opus") || family.contains("o1") {
|
||||
score += 90.0;
|
||||
} else if family.contains("sonnet") || family.contains("gpt-4") {
|
||||
score += 80.0;
|
||||
} else {
|
||||
score += 70.0;
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// 上下文长度加分
|
||||
if let Some(ctx_len) = model.context_length {
|
||||
score += (ctx_len as f64 / 50000.0).min(10.0);
|
||||
}
|
||||
|
||||
// 负载惩罚
|
||||
if let Some(load) = model.current_load {
|
||||
score -= load as f64 * 0.3;
|
||||
}
|
||||
|
||||
score
|
||||
}
|
||||
}
|
||||
|
||||
impl Default for TaskBasedStrategy {
|
||||
fn default() -> Self {
|
||||
Self::new()
|
||||
}
|
||||
}
|
||||
|
||||
#[async_trait]
|
||||
impl SelectionStrategy for TaskBasedStrategy {
|
||||
fn id(&self) -> &str {
|
||||
"task_based"
|
||||
}
|
||||
|
||||
fn display_name(&self) -> &str {
|
||||
"任务匹配"
|
||||
}
|
||||
|
||||
fn description(&self) -> &str {
|
||||
"根据任务类型选择最适合的模型"
|
||||
}
|
||||
|
||||
fn supports_task(&self, _task: TaskHint) -> bool {
|
||||
true
|
||||
}
|
||||
|
||||
async fn select(
|
||||
&self,
|
||||
pool: &[AvailableModel],
|
||||
ctx: &SelectionContext,
|
||||
) -> StrategyResult<ModelSelection> {
|
||||
// 过滤可用模型
|
||||
let mut available: Vec<_> = pool
|
||||
.iter()
|
||||
.filter(|m| {
|
||||
m.is_healthy
|
||||
&& !ctx.excluded_models.contains(&m.id)
|
||||
&& (!ctx.requires_vision || m.supports_vision)
|
||||
&& (!ctx.requires_tools || m.supports_tools)
|
||||
})
|
||||
.cloned()
|
||||
.collect();
|
||||
|
||||
if available.is_empty() {
|
||||
return Err(StrategyError::NoAvailableModels);
|
||||
}
|
||||
|
||||
// 按任务类型评分排序
|
||||
available.sort_by(|a, b| {
|
||||
let score_a = Self::score_for_task(a, ctx.task_hint);
|
||||
let score_b = Self::score_for_task(b, ctx.task_hint);
|
||||
score_b
|
||||
.partial_cmp(&score_a)
|
||||
.unwrap_or(std::cmp::Ordering::Equal)
|
||||
});
|
||||
|
||||
let selected = available.remove(0);
|
||||
let task_name = ctx.task_hint.map(|t| t.display_name()).unwrap_or("通用");
|
||||
|
||||
Ok(ModelSelection {
|
||||
model: selected,
|
||||
reason: format!("任务匹配选择 (任务类型: {})", task_name),
|
||||
confidence: 85,
|
||||
alternatives: available,
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
use crate::orchestrator::tier::ServiceTier;
|
||||
|
||||
fn create_test_models() -> Vec<AvailableModel> {
|
||||
vec![
|
||||
AvailableModel {
|
||||
id: "claude-opus".to_string(),
|
||||
display_name: "Claude Opus".to_string(),
|
||||
provider_type: "anthropic".to_string(),
|
||||
family: Some("opus".to_string()),
|
||||
credential_id: "cred-1".to_string(),
|
||||
context_length: Some(200000),
|
||||
supports_vision: true,
|
||||
supports_tools: true,
|
||||
input_cost_per_million: None,
|
||||
output_cost_per_million: None,
|
||||
is_healthy: true,
|
||||
current_load: Some(20),
|
||||
},
|
||||
AvailableModel {
|
||||
id: "claude-sonnet".to_string(),
|
||||
display_name: "Claude Sonnet".to_string(),
|
||||
provider_type: "anthropic".to_string(),
|
||||
family: Some("sonnet".to_string()),
|
||||
credential_id: "cred-2".to_string(),
|
||||
context_length: Some(200000),
|
||||
supports_vision: true,
|
||||
supports_tools: true,
|
||||
input_cost_per_million: None,
|
||||
output_cost_per_million: None,
|
||||
is_healthy: true,
|
||||
current_load: Some(30),
|
||||
},
|
||||
AvailableModel {
|
||||
id: "claude-haiku".to_string(),
|
||||
display_name: "Claude Haiku".to_string(),
|
||||
provider_type: "anthropic".to_string(),
|
||||
family: Some("haiku".to_string()),
|
||||
credential_id: "cred-3".to_string(),
|
||||
context_length: Some(200000),
|
||||
supports_vision: true,
|
||||
supports_tools: true,
|
||||
input_cost_per_million: None,
|
||||
output_cost_per_million: None,
|
||||
is_healthy: true,
|
||||
current_load: Some(10),
|
||||
},
|
||||
]
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_task_based_coding() {
|
||||
let strategy = TaskBasedStrategy::new();
|
||||
let models = create_test_models();
|
||||
let ctx = SelectionContext::new(ServiceTier::Pro).with_task_hint(TaskHint::Coding);
|
||||
|
||||
let result = strategy.select(&models, &ctx).await.unwrap();
|
||||
// 代码任务应该选择 Sonnet
|
||||
assert_eq!(result.model.family, Some("sonnet".to_string()));
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_task_based_chat() {
|
||||
let strategy = TaskBasedStrategy::new();
|
||||
let models = create_test_models();
|
||||
let ctx = SelectionContext::new(ServiceTier::Mini).with_task_hint(TaskHint::Chat);
|
||||
|
||||
let result = strategy.select(&models, &ctx).await.unwrap();
|
||||
// 对话任务应该选择 Haiku
|
||||
assert_eq!(result.model.family, Some("haiku".to_string()));
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_task_based_analysis() {
|
||||
let strategy = TaskBasedStrategy::new();
|
||||
let models = create_test_models();
|
||||
let ctx = SelectionContext::new(ServiceTier::Max).with_task_hint(TaskHint::Analysis);
|
||||
|
||||
let result = strategy.select(&models, &ctx).await.unwrap();
|
||||
// 分析任务应该选择 Opus
|
||||
assert_eq!(result.model.family, Some("opus".to_string()));
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,383 @@
|
||||
//! 选择策略 trait 和注册表
|
||||
//!
|
||||
//! 定义模型选择策略的接口和策略注册表。
|
||||
|
||||
use super::tier::{AvailableModel, ServiceTier};
|
||||
use async_trait::async_trait;
|
||||
use serde::{Deserialize, Serialize};
|
||||
use std::collections::HashMap;
|
||||
use std::sync::Arc;
|
||||
use thiserror::Error;
|
||||
|
||||
/// 策略错误
|
||||
#[derive(Error, Debug)]
|
||||
pub enum StrategyError {
|
||||
#[error("没有可用的模型")]
|
||||
NoAvailableModels,
|
||||
|
||||
#[error("策略不存在: {0}")]
|
||||
StrategyNotFound(String),
|
||||
|
||||
#[error("选择失败: {0}")]
|
||||
SelectionFailed(String),
|
||||
|
||||
#[error("配置错误: {0}")]
|
||||
ConfigError(String),
|
||||
}
|
||||
|
||||
pub type StrategyResult<T> = Result<T, StrategyError>;
|
||||
|
||||
/// 选择上下文
|
||||
///
|
||||
/// 包含选择模型时需要的所有上下文信息
|
||||
#[derive(Debug, Clone, Serialize, Deserialize)]
|
||||
pub struct SelectionContext {
|
||||
/// 服务等级
|
||||
pub tier: ServiceTier,
|
||||
/// 请求的模型名称(如果有)
|
||||
pub requested_model: Option<String>,
|
||||
/// 任务类型提示
|
||||
pub task_hint: Option<TaskHint>,
|
||||
/// 是否需要视觉能力
|
||||
pub requires_vision: bool,
|
||||
/// 是否需要工具调用
|
||||
pub requires_tools: bool,
|
||||
/// 预估输入 tokens
|
||||
pub estimated_input_tokens: Option<u32>,
|
||||
/// 预估输出 tokens
|
||||
pub estimated_output_tokens: Option<u32>,
|
||||
/// 用户偏好的 Provider
|
||||
pub preferred_provider: Option<String>,
|
||||
/// 排除的模型 ID 列表
|
||||
pub excluded_models: Vec<String>,
|
||||
/// 额外元数据
|
||||
pub metadata: HashMap<String, serde_json::Value>,
|
||||
}
|
||||
|
||||
impl Default for SelectionContext {
|
||||
fn default() -> Self {
|
||||
Self {
|
||||
tier: ServiceTier::Pro,
|
||||
requested_model: None,
|
||||
task_hint: None,
|
||||
requires_vision: false,
|
||||
requires_tools: false,
|
||||
estimated_input_tokens: None,
|
||||
estimated_output_tokens: None,
|
||||
preferred_provider: None,
|
||||
excluded_models: Vec::new(),
|
||||
metadata: HashMap::new(),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
impl SelectionContext {
|
||||
/// 创建新的选择上下文
|
||||
pub fn new(tier: ServiceTier) -> Self {
|
||||
Self {
|
||||
tier,
|
||||
..Default::default()
|
||||
}
|
||||
}
|
||||
|
||||
/// 设置任务提示
|
||||
pub fn with_task_hint(mut self, hint: TaskHint) -> Self {
|
||||
self.task_hint = Some(hint);
|
||||
self
|
||||
}
|
||||
|
||||
/// 设置视觉需求
|
||||
pub fn with_vision(mut self, requires: bool) -> Self {
|
||||
self.requires_vision = requires;
|
||||
self
|
||||
}
|
||||
|
||||
/// 设置工具调用需求
|
||||
pub fn with_tools(mut self, requires: bool) -> Self {
|
||||
self.requires_tools = requires;
|
||||
self
|
||||
}
|
||||
|
||||
/// 设置偏好的 Provider
|
||||
pub fn with_preferred_provider(mut self, provider: &str) -> Self {
|
||||
self.preferred_provider = Some(provider.to_string());
|
||||
self
|
||||
}
|
||||
|
||||
/// 添加排除的模型
|
||||
pub fn exclude_model(mut self, model_id: &str) -> Self {
|
||||
self.excluded_models.push(model_id.to_string());
|
||||
self
|
||||
}
|
||||
}
|
||||
|
||||
/// 任务类型提示
|
||||
#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
|
||||
#[serde(rename_all = "snake_case")]
|
||||
pub enum TaskHint {
|
||||
/// 代码生成/编辑
|
||||
Coding,
|
||||
/// 写作/创意
|
||||
Writing,
|
||||
/// 分析/推理
|
||||
Analysis,
|
||||
/// 对话/聊天
|
||||
Chat,
|
||||
/// 翻译
|
||||
Translation,
|
||||
/// 摘要
|
||||
Summarization,
|
||||
/// 数学/计算
|
||||
Math,
|
||||
/// 其他
|
||||
Other,
|
||||
}
|
||||
|
||||
impl TaskHint {
|
||||
/// 获取任务提示的显示名称
|
||||
pub fn display_name(&self) -> &'static str {
|
||||
match self {
|
||||
TaskHint::Coding => "代码",
|
||||
TaskHint::Writing => "写作",
|
||||
TaskHint::Analysis => "分析",
|
||||
TaskHint::Chat => "对话",
|
||||
TaskHint::Translation => "翻译",
|
||||
TaskHint::Summarization => "摘要",
|
||||
TaskHint::Math => "数学",
|
||||
TaskHint::Other => "其他",
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
/// 模型选择结果
|
||||
#[derive(Debug, Clone, Serialize, Deserialize)]
|
||||
pub struct ModelSelection {
|
||||
/// 选中的模型
|
||||
pub model: AvailableModel,
|
||||
/// 选择原因
|
||||
pub reason: String,
|
||||
/// 置信度 (0-100)
|
||||
pub confidence: u8,
|
||||
/// 备选模型列表
|
||||
pub alternatives: Vec<AvailableModel>,
|
||||
}
|
||||
|
||||
/// 选择策略 trait
|
||||
///
|
||||
/// 所有模型选择策略必须实现此 trait
|
||||
#[async_trait]
|
||||
pub trait SelectionStrategy: Send + Sync {
|
||||
/// 策略 ID
|
||||
fn id(&self) -> &str;
|
||||
|
||||
/// 策略显示名称
|
||||
fn display_name(&self) -> &str;
|
||||
|
||||
/// 策略描述
|
||||
fn description(&self) -> &str {
|
||||
""
|
||||
}
|
||||
|
||||
/// 选择模型
|
||||
async fn select(
|
||||
&self,
|
||||
pool: &[AvailableModel],
|
||||
ctx: &SelectionContext,
|
||||
) -> StrategyResult<ModelSelection>;
|
||||
|
||||
/// 是否支持指定的任务类型
|
||||
fn supports_task(&self, _task: TaskHint) -> bool {
|
||||
true
|
||||
}
|
||||
|
||||
/// 获取策略配置 Schema
|
||||
fn config_schema(&self) -> serde_json::Value {
|
||||
serde_json::json!({})
|
||||
}
|
||||
|
||||
/// 更新策略配置
|
||||
fn update_config(&mut self, _config: serde_json::Value) -> StrategyResult<()> {
|
||||
Ok(())
|
||||
}
|
||||
}
|
||||
|
||||
/// 策略注册表
|
||||
///
|
||||
/// 管理所有可用的选择策略
|
||||
pub struct StrategyRegistry {
|
||||
/// 已注册的策略
|
||||
strategies: HashMap<String, Arc<dyn SelectionStrategy>>,
|
||||
/// 默认策略 ID
|
||||
default_strategy: String,
|
||||
}
|
||||
|
||||
impl StrategyRegistry {
|
||||
/// 创建新的策略注册表
|
||||
pub fn new() -> Self {
|
||||
Self {
|
||||
strategies: HashMap::new(),
|
||||
default_strategy: "round_robin".to_string(),
|
||||
}
|
||||
}
|
||||
|
||||
/// 注册策略
|
||||
pub fn register(&mut self, strategy: Arc<dyn SelectionStrategy>) {
|
||||
let id = strategy.id().to_string();
|
||||
tracing::info!("注册选择策略: {} ({})", id, strategy.display_name());
|
||||
self.strategies.insert(id, strategy);
|
||||
}
|
||||
|
||||
/// 获取策略
|
||||
pub fn get(&self, id: &str) -> Option<Arc<dyn SelectionStrategy>> {
|
||||
self.strategies.get(id).cloned()
|
||||
}
|
||||
|
||||
/// 获取默认策略
|
||||
pub fn get_default(&self) -> Option<Arc<dyn SelectionStrategy>> {
|
||||
self.get(&self.default_strategy)
|
||||
}
|
||||
|
||||
/// 设置默认策略
|
||||
pub fn set_default(&mut self, id: &str) -> StrategyResult<()> {
|
||||
if self.strategies.contains_key(id) {
|
||||
self.default_strategy = id.to_string();
|
||||
Ok(())
|
||||
} else {
|
||||
Err(StrategyError::StrategyNotFound(id.to_string()))
|
||||
}
|
||||
}
|
||||
|
||||
/// 获取所有策略 ID
|
||||
pub fn list_ids(&self) -> Vec<&str> {
|
||||
self.strategies.keys().map(|s| s.as_str()).collect()
|
||||
}
|
||||
|
||||
/// 获取所有策略信息
|
||||
pub fn list_all(&self) -> Vec<StrategyInfo> {
|
||||
self.strategies
|
||||
.values()
|
||||
.map(|s| StrategyInfo {
|
||||
id: s.id().to_string(),
|
||||
display_name: s.display_name().to_string(),
|
||||
description: s.description().to_string(),
|
||||
is_default: s.id() == self.default_strategy,
|
||||
})
|
||||
.collect()
|
||||
}
|
||||
|
||||
/// 使用指定策略选择模型
|
||||
pub async fn select_with(
|
||||
&self,
|
||||
strategy_id: &str,
|
||||
pool: &[AvailableModel],
|
||||
ctx: &SelectionContext,
|
||||
) -> StrategyResult<ModelSelection> {
|
||||
let strategy = self
|
||||
.get(strategy_id)
|
||||
.ok_or_else(|| StrategyError::StrategyNotFound(strategy_id.to_string()))?;
|
||||
|
||||
strategy.select(pool, ctx).await
|
||||
}
|
||||
|
||||
/// 使用默认策略选择模型
|
||||
pub async fn select(
|
||||
&self,
|
||||
pool: &[AvailableModel],
|
||||
ctx: &SelectionContext,
|
||||
) -> StrategyResult<ModelSelection> {
|
||||
let strategy = self
|
||||
.get_default()
|
||||
.ok_or_else(|| StrategyError::StrategyNotFound(self.default_strategy.clone()))?;
|
||||
|
||||
strategy.select(pool, ctx).await
|
||||
}
|
||||
}
|
||||
|
||||
impl Default for StrategyRegistry {
|
||||
fn default() -> Self {
|
||||
Self::new()
|
||||
}
|
||||
}
|
||||
|
||||
/// 策略信息
|
||||
#[derive(Debug, Clone, Serialize, Deserialize)]
|
||||
pub struct StrategyInfo {
|
||||
/// 策略 ID
|
||||
pub id: String,
|
||||
/// 显示名称
|
||||
pub display_name: String,
|
||||
/// 描述
|
||||
pub description: String,
|
||||
/// 是否是默认策略
|
||||
pub is_default: bool,
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
|
||||
struct MockStrategy {
|
||||
id: String,
|
||||
}
|
||||
|
||||
#[async_trait]
|
||||
impl SelectionStrategy for MockStrategy {
|
||||
fn id(&self) -> &str {
|
||||
&self.id
|
||||
}
|
||||
|
||||
fn display_name(&self) -> &str {
|
||||
"Mock Strategy"
|
||||
}
|
||||
|
||||
async fn select(
|
||||
&self,
|
||||
pool: &[AvailableModel],
|
||||
_ctx: &SelectionContext,
|
||||
) -> StrategyResult<ModelSelection> {
|
||||
if pool.is_empty() {
|
||||
return Err(StrategyError::NoAvailableModels);
|
||||
}
|
||||
|
||||
Ok(ModelSelection {
|
||||
model: pool[0].clone(),
|
||||
reason: "Mock selection".to_string(),
|
||||
confidence: 100,
|
||||
alternatives: pool[1..].to_vec(),
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_selection_context() {
|
||||
let ctx = SelectionContext::new(ServiceTier::Pro)
|
||||
.with_task_hint(TaskHint::Coding)
|
||||
.with_vision(true)
|
||||
.with_tools(true)
|
||||
.with_preferred_provider("anthropic")
|
||||
.exclude_model("model-1");
|
||||
|
||||
assert_eq!(ctx.tier, ServiceTier::Pro);
|
||||
assert_eq!(ctx.task_hint, Some(TaskHint::Coding));
|
||||
assert!(ctx.requires_vision);
|
||||
assert!(ctx.requires_tools);
|
||||
assert_eq!(ctx.preferred_provider, Some("anthropic".to_string()));
|
||||
assert!(ctx.excluded_models.contains(&"model-1".to_string()));
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_strategy_registry() {
|
||||
let mut registry = StrategyRegistry::new();
|
||||
|
||||
let strategy = Arc::new(MockStrategy {
|
||||
id: "mock".to_string(),
|
||||
});
|
||||
registry.register(strategy);
|
||||
|
||||
assert!(registry.get("mock").is_some());
|
||||
assert!(registry.get("nonexistent").is_none());
|
||||
|
||||
registry.set_default("mock").unwrap();
|
||||
assert!(registry.get_default().is_some());
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,363 @@
|
||||
//! 服务等级定义
|
||||
//!
|
||||
//! 定义 Mini/Pro/Max 三个服务等级及其配置。
|
||||
|
||||
use serde::{Deserialize, Serialize};
|
||||
use std::collections::HashMap;
|
||||
|
||||
/// 服务等级
|
||||
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash, Serialize, Deserialize)]
|
||||
#[serde(rename_all = "lowercase")]
|
||||
pub enum ServiceTier {
|
||||
/// Mini - 快速响应,适合简单任务
|
||||
Mini,
|
||||
/// Pro - 均衡选择,适合大多数任务
|
||||
Pro,
|
||||
/// Max - 最强能力,适合复杂任务
|
||||
Max,
|
||||
}
|
||||
|
||||
impl ServiceTier {
|
||||
/// 获取等级的显示名称
|
||||
pub fn display_name(&self) -> &'static str {
|
||||
match self {
|
||||
ServiceTier::Mini => "Mini",
|
||||
ServiceTier::Pro => "Pro",
|
||||
ServiceTier::Max => "Max",
|
||||
}
|
||||
}
|
||||
|
||||
/// 获取等级的描述
|
||||
pub fn description(&self) -> &'static str {
|
||||
match self {
|
||||
ServiceTier::Mini => "快速响应,适合简单任务",
|
||||
ServiceTier::Pro => "均衡选择,适合大多数任务",
|
||||
ServiceTier::Max => "最强能力,适合复杂任务",
|
||||
}
|
||||
}
|
||||
|
||||
/// 获取等级的数值(用于排序)
|
||||
pub fn level(&self) -> u8 {
|
||||
match self {
|
||||
ServiceTier::Mini => 1,
|
||||
ServiceTier::Pro => 2,
|
||||
ServiceTier::Max => 3,
|
||||
}
|
||||
}
|
||||
|
||||
/// 从字符串解析
|
||||
pub fn from_str(s: &str) -> Option<Self> {
|
||||
match s.to_lowercase().as_str() {
|
||||
"mini" => Some(ServiceTier::Mini),
|
||||
"pro" => Some(ServiceTier::Pro),
|
||||
"max" => Some(ServiceTier::Max),
|
||||
_ => None,
|
||||
}
|
||||
}
|
||||
|
||||
/// 获取所有等级
|
||||
pub fn all() -> &'static [ServiceTier] {
|
||||
&[ServiceTier::Mini, ServiceTier::Pro, ServiceTier::Max]
|
||||
}
|
||||
}
|
||||
|
||||
impl Default for ServiceTier {
|
||||
fn default() -> Self {
|
||||
ServiceTier::Pro
|
||||
}
|
||||
}
|
||||
|
||||
impl std::fmt::Display for ServiceTier {
|
||||
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
|
||||
write!(f, "{}", self.display_name())
|
||||
}
|
||||
}
|
||||
|
||||
/// 等级配置
|
||||
#[derive(Debug, Clone, Serialize, Deserialize)]
|
||||
pub struct TierConfig {
|
||||
/// 等级
|
||||
pub tier: ServiceTier,
|
||||
/// 默认策略 ID
|
||||
pub default_strategy: String,
|
||||
/// 模型家族优先级(按优先级排序)
|
||||
pub family_priorities: Vec<String>,
|
||||
/// 最大并发请求数
|
||||
pub max_concurrent: Option<u32>,
|
||||
/// 超时时间(毫秒)
|
||||
pub timeout_ms: Option<u64>,
|
||||
}
|
||||
|
||||
impl TierConfig {
|
||||
/// 创建 Mini 等级的默认配置
|
||||
pub fn mini() -> Self {
|
||||
Self {
|
||||
tier: ServiceTier::Mini,
|
||||
default_strategy: "speed_optimized".to_string(),
|
||||
family_priorities: vec![
|
||||
"haiku".to_string(),
|
||||
"flash".to_string(),
|
||||
"gpt-3.5".to_string(),
|
||||
],
|
||||
max_concurrent: Some(10),
|
||||
timeout_ms: Some(30000),
|
||||
}
|
||||
}
|
||||
|
||||
/// 创建 Pro 等级的默认配置
|
||||
pub fn pro() -> Self {
|
||||
Self {
|
||||
tier: ServiceTier::Pro,
|
||||
default_strategy: "load_balanced".to_string(),
|
||||
family_priorities: vec!["sonnet".to_string(), "pro".to_string(), "gpt-4".to_string()],
|
||||
max_concurrent: Some(5),
|
||||
timeout_ms: Some(120000),
|
||||
}
|
||||
}
|
||||
|
||||
/// 创建 Max 等级的默认配置
|
||||
pub fn max() -> Self {
|
||||
Self {
|
||||
tier: ServiceTier::Max,
|
||||
default_strategy: "task_based".to_string(),
|
||||
family_priorities: vec!["opus".to_string(), "ultra".to_string(), "o1".to_string()],
|
||||
max_concurrent: Some(3),
|
||||
timeout_ms: Some(300000),
|
||||
}
|
||||
}
|
||||
|
||||
/// 获取默认配置映射
|
||||
pub fn defaults() -> HashMap<ServiceTier, TierConfig> {
|
||||
let mut map = HashMap::new();
|
||||
map.insert(ServiceTier::Mini, TierConfig::mini());
|
||||
map.insert(ServiceTier::Pro, TierConfig::pro());
|
||||
map.insert(ServiceTier::Max, TierConfig::max());
|
||||
map
|
||||
}
|
||||
}
|
||||
|
||||
/// 可用模型信息
|
||||
#[derive(Debug, Clone, Serialize, Deserialize)]
|
||||
pub struct AvailableModel {
|
||||
/// 模型 ID
|
||||
pub id: String,
|
||||
/// 显示名称
|
||||
pub display_name: String,
|
||||
/// Provider 类型
|
||||
pub provider_type: String,
|
||||
/// 模型家族
|
||||
pub family: Option<String>,
|
||||
/// 凭证 ID
|
||||
pub credential_id: String,
|
||||
/// 上下文长度
|
||||
pub context_length: Option<u32>,
|
||||
/// 是否支持视觉
|
||||
pub supports_vision: bool,
|
||||
/// 是否支持工具调用
|
||||
pub supports_tools: bool,
|
||||
/// 输入价格(每 1M tokens)
|
||||
pub input_cost_per_million: Option<f64>,
|
||||
/// 输出价格(每 1M tokens)
|
||||
pub output_cost_per_million: Option<f64>,
|
||||
/// 健康状态
|
||||
pub is_healthy: bool,
|
||||
/// 当前负载(0-100)
|
||||
pub current_load: Option<u8>,
|
||||
}
|
||||
|
||||
impl AvailableModel {
|
||||
/// 计算模型的综合评分
|
||||
pub fn score(&self, tier: ServiceTier) -> f64 {
|
||||
let mut score = 0.0;
|
||||
|
||||
// 基础分:健康状态
|
||||
if !self.is_healthy {
|
||||
return 0.0;
|
||||
}
|
||||
|
||||
// 家族匹配分
|
||||
if let Some(family) = &self.family {
|
||||
let tier_families = match tier {
|
||||
ServiceTier::Mini => vec!["haiku", "flash", "gpt-3.5"],
|
||||
ServiceTier::Pro => vec!["sonnet", "pro", "gpt-4"],
|
||||
ServiceTier::Max => vec!["opus", "ultra", "o1"],
|
||||
};
|
||||
|
||||
for (i, f) in tier_families.iter().enumerate() {
|
||||
if family.to_lowercase().contains(f) {
|
||||
score += 100.0 - (i as f64 * 10.0);
|
||||
break;
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// 负载分(负载越低越好)
|
||||
if let Some(load) = self.current_load {
|
||||
score += (100 - load) as f64 * 0.5;
|
||||
} else {
|
||||
score += 50.0; // 默认中等负载
|
||||
}
|
||||
|
||||
// 能力分
|
||||
if self.supports_vision {
|
||||
score += 10.0;
|
||||
}
|
||||
if self.supports_tools {
|
||||
score += 10.0;
|
||||
}
|
||||
|
||||
// 上下文长度分
|
||||
if let Some(ctx) = self.context_length {
|
||||
score += (ctx as f64 / 10000.0).min(20.0);
|
||||
}
|
||||
|
||||
score
|
||||
}
|
||||
}
|
||||
|
||||
/// 等级模型池
|
||||
#[derive(Debug, Clone, Default)]
|
||||
pub struct TierPool {
|
||||
/// Mini 等级可用模型
|
||||
pub mini: Vec<AvailableModel>,
|
||||
/// Pro 等级可用模型
|
||||
pub pro: Vec<AvailableModel>,
|
||||
/// Max 等级可用模型
|
||||
pub max: Vec<AvailableModel>,
|
||||
}
|
||||
|
||||
impl TierPool {
|
||||
/// 创建新的模型池
|
||||
pub fn new() -> Self {
|
||||
Self::default()
|
||||
}
|
||||
|
||||
/// 获取指定等级的模型列表
|
||||
pub fn get(&self, tier: ServiceTier) -> &[AvailableModel] {
|
||||
match tier {
|
||||
ServiceTier::Mini => &self.mini,
|
||||
ServiceTier::Pro => &self.pro,
|
||||
ServiceTier::Max => &self.max,
|
||||
}
|
||||
}
|
||||
|
||||
/// 获取指定等级的可变模型列表
|
||||
pub fn get_mut(&mut self, tier: ServiceTier) -> &mut Vec<AvailableModel> {
|
||||
match tier {
|
||||
ServiceTier::Mini => &mut self.mini,
|
||||
ServiceTier::Pro => &mut self.pro,
|
||||
ServiceTier::Max => &mut self.max,
|
||||
}
|
||||
}
|
||||
|
||||
/// 添加模型到指定等级
|
||||
pub fn add(&mut self, tier: ServiceTier, model: AvailableModel) {
|
||||
self.get_mut(tier).push(model);
|
||||
}
|
||||
|
||||
/// 获取所有等级的模型总数
|
||||
pub fn total_count(&self) -> usize {
|
||||
self.mini.len() + self.pro.len() + self.max.len()
|
||||
}
|
||||
|
||||
/// 检查是否为空
|
||||
pub fn is_empty(&self) -> bool {
|
||||
self.mini.is_empty() && self.pro.is_empty() && self.max.is_empty()
|
||||
}
|
||||
|
||||
/// 按评分排序所有等级的模型
|
||||
pub fn sort_by_score(&mut self) {
|
||||
self.mini.sort_by(|a, b| {
|
||||
b.score(ServiceTier::Mini)
|
||||
.partial_cmp(&a.score(ServiceTier::Mini))
|
||||
.unwrap_or(std::cmp::Ordering::Equal)
|
||||
});
|
||||
self.pro.sort_by(|a, b| {
|
||||
b.score(ServiceTier::Pro)
|
||||
.partial_cmp(&a.score(ServiceTier::Pro))
|
||||
.unwrap_or(std::cmp::Ordering::Equal)
|
||||
});
|
||||
self.max.sort_by(|a, b| {
|
||||
b.score(ServiceTier::Max)
|
||||
.partial_cmp(&a.score(ServiceTier::Max))
|
||||
.unwrap_or(std::cmp::Ordering::Equal)
|
||||
});
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
|
||||
#[test]
|
||||
fn test_service_tier_basics() {
|
||||
assert_eq!(ServiceTier::Mini.level(), 1);
|
||||
assert_eq!(ServiceTier::Pro.level(), 2);
|
||||
assert_eq!(ServiceTier::Max.level(), 3);
|
||||
|
||||
assert_eq!(ServiceTier::from_str("mini"), Some(ServiceTier::Mini));
|
||||
assert_eq!(ServiceTier::from_str("PRO"), Some(ServiceTier::Pro));
|
||||
assert_eq!(ServiceTier::from_str("invalid"), None);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_tier_config_defaults() {
|
||||
let defaults = TierConfig::defaults();
|
||||
assert_eq!(defaults.len(), 3);
|
||||
assert!(defaults.contains_key(&ServiceTier::Mini));
|
||||
assert!(defaults.contains_key(&ServiceTier::Pro));
|
||||
assert!(defaults.contains_key(&ServiceTier::Max));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_available_model_score() {
|
||||
let model = AvailableModel {
|
||||
id: "claude-3-5-haiku".to_string(),
|
||||
display_name: "Claude 3.5 Haiku".to_string(),
|
||||
provider_type: "anthropic".to_string(),
|
||||
family: Some("haiku".to_string()),
|
||||
credential_id: "cred-1".to_string(),
|
||||
context_length: Some(200000),
|
||||
supports_vision: true,
|
||||
supports_tools: true,
|
||||
input_cost_per_million: None,
|
||||
output_cost_per_million: None,
|
||||
is_healthy: true,
|
||||
current_load: Some(30),
|
||||
};
|
||||
|
||||
// Haiku 模型在 Mini 等级应该得分最高
|
||||
let mini_score = model.score(ServiceTier::Mini);
|
||||
let pro_score = model.score(ServiceTier::Pro);
|
||||
|
||||
assert!(mini_score > pro_score);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_tier_pool() {
|
||||
let mut pool = TierPool::new();
|
||||
assert!(pool.is_empty());
|
||||
|
||||
pool.add(
|
||||
ServiceTier::Mini,
|
||||
AvailableModel {
|
||||
id: "test".to_string(),
|
||||
display_name: "Test".to_string(),
|
||||
provider_type: "test".to_string(),
|
||||
family: None,
|
||||
credential_id: "cred".to_string(),
|
||||
context_length: None,
|
||||
supports_vision: false,
|
||||
supports_tools: false,
|
||||
input_cost_per_million: None,
|
||||
output_cost_per_million: None,
|
||||
is_healthy: true,
|
||||
current_load: None,
|
||||
},
|
||||
);
|
||||
|
||||
assert!(!pool.is_empty());
|
||||
assert_eq!(pool.total_count(), 1);
|
||||
assert_eq!(pool.get(ServiceTier::Mini).len(), 1);
|
||||
}
|
||||
}
|
||||
@@ -783,6 +783,109 @@ pub async fn chat_completions(
|
||||
}
|
||||
};
|
||||
|
||||
// 如果 Provider Pool 中没有找到凭证,尝试从 API Key Provider 获取
|
||||
let credential = if credential.is_none() {
|
||||
eprintln!("[CHAT_COMPLETIONS] Provider Pool 中未找到凭证,尝试 API Key Provider...");
|
||||
|
||||
// 根据 selected_provider 映射到 ApiProviderType
|
||||
use crate::database::dao::api_key_provider::ApiProviderType;
|
||||
let api_provider_type = match selected_provider.to_lowercase().as_str() {
|
||||
"anthropic" | "claude" => Some(ApiProviderType::Anthropic),
|
||||
"openai" => Some(ApiProviderType::Openai),
|
||||
"gemini" => Some(ApiProviderType::Gemini),
|
||||
// 以下都是 OpenAI 兼容的 Provider
|
||||
"deepseek" | "moonshot" | "groq" | "grok" | "mistral" | "perplexity" | "cohere"
|
||||
| "openrouter" | "silicon" => Some(ApiProviderType::Openai),
|
||||
_ => None,
|
||||
};
|
||||
|
||||
if let (Some(db), Some(api_type)) = (&state.db, api_provider_type) {
|
||||
eprintln!(
|
||||
"[CHAT_COMPLETIONS] 尝试从 API Key Provider 类型 '{:?}' 获取凭证",
|
||||
api_type
|
||||
);
|
||||
|
||||
// 使用按类型获取的方法(包括自定义 Provider)
|
||||
match state.api_key_service.get_next_api_key_by_type(db, api_type) {
|
||||
Ok(Some((_key_id, api_key, provider_info))) => {
|
||||
eprintln!(
|
||||
"[CHAT_COMPLETIONS] 从 API Key Provider 获取到凭证: provider={}, api_host={}",
|
||||
provider_info.name,
|
||||
provider_info.api_host
|
||||
);
|
||||
|
||||
let base_url = if provider_info.api_host.is_empty() {
|
||||
None
|
||||
} else {
|
||||
Some(provider_info.api_host.clone())
|
||||
};
|
||||
|
||||
let provider_type = match provider_info.provider_type {
|
||||
ApiProviderType::Anthropic => crate::ProviderType::Anthropic,
|
||||
ApiProviderType::Openai | ApiProviderType::OpenaiResponse => {
|
||||
crate::ProviderType::OpenAI
|
||||
}
|
||||
ApiProviderType::Gemini => crate::ProviderType::GeminiApiKey,
|
||||
_ => crate::ProviderType::OpenAI,
|
||||
};
|
||||
|
||||
// 根据 provider_type 创建对应的 CredentialData
|
||||
let credential_data = match provider_type {
|
||||
crate::ProviderType::Anthropic => {
|
||||
crate::models::provider_pool_model::CredentialData::AnthropicKey {
|
||||
api_key: api_key.clone(),
|
||||
base_url,
|
||||
}
|
||||
}
|
||||
crate::ProviderType::GeminiApiKey => {
|
||||
crate::models::provider_pool_model::CredentialData::GeminiApiKey {
|
||||
api_key: api_key.clone(),
|
||||
base_url,
|
||||
excluded_models: vec![],
|
||||
}
|
||||
}
|
||||
_ => crate::models::provider_pool_model::CredentialData::OpenAIKey {
|
||||
api_key: api_key.clone(),
|
||||
base_url,
|
||||
},
|
||||
};
|
||||
|
||||
// 构建 ProviderCredential
|
||||
let mut cred = crate::models::provider_pool_model::ProviderCredential::new(
|
||||
provider_type,
|
||||
credential_data,
|
||||
);
|
||||
cred.name = Some(provider_info.name.clone());
|
||||
|
||||
state.logs.write().await.add(
|
||||
"info",
|
||||
&format!(
|
||||
"[ROUTE] Using API Key Provider credential: provider={}, type={:?}",
|
||||
provider_info.name, provider_info.provider_type
|
||||
),
|
||||
);
|
||||
|
||||
Some(cred)
|
||||
}
|
||||
Ok(None) => {
|
||||
eprintln!(
|
||||
"[CHAT_COMPLETIONS] API Key Provider 类型 '{:?}' 没有可用的 API Key",
|
||||
api_type
|
||||
);
|
||||
None
|
||||
}
|
||||
Err(e) => {
|
||||
eprintln!("[CHAT_COMPLETIONS] 从 API Key Provider 获取凭证失败: {}", e);
|
||||
None
|
||||
}
|
||||
}
|
||||
} else {
|
||||
None
|
||||
}
|
||||
} else {
|
||||
credential
|
||||
};
|
||||
|
||||
// 如果找到凭证池中的凭证,使用它
|
||||
if let Some(cred) = credential {
|
||||
eprintln!(
|
||||
@@ -1624,6 +1727,84 @@ pub async fn anthropic_messages(
|
||||
None => None,
|
||||
};
|
||||
|
||||
// 如果 Provider Pool 中没有找到凭证,尝试从 API Key Provider 获取
|
||||
let credential = if credential.is_none() {
|
||||
// 根据 selected_provider 映射到 ApiProviderType
|
||||
use crate::database::dao::api_key_provider::ApiProviderType;
|
||||
let api_provider_type = match selected_provider.to_lowercase().as_str() {
|
||||
"anthropic" | "claude" => Some(ApiProviderType::Anthropic),
|
||||
"openai" => Some(ApiProviderType::Openai),
|
||||
"gemini" => Some(ApiProviderType::Gemini),
|
||||
_ => None,
|
||||
};
|
||||
|
||||
if let (Some(db), Some(api_type)) = (&state.db, api_provider_type) {
|
||||
// 使用按类型获取的方法(包括自定义 Provider)
|
||||
match state.api_key_service.get_next_api_key_by_type(db, api_type) {
|
||||
Ok(Some((_key_id, api_key, provider_info))) => {
|
||||
state.logs.write().await.add(
|
||||
"info",
|
||||
&format!(
|
||||
"[ROUTE] Using API Key Provider credential: provider={}, type={:?}, api_host={}",
|
||||
provider_info.name, provider_info.provider_type, provider_info.api_host
|
||||
),
|
||||
);
|
||||
|
||||
let base_url = if provider_info.api_host.is_empty() {
|
||||
None
|
||||
} else {
|
||||
Some(provider_info.api_host.clone())
|
||||
};
|
||||
|
||||
let provider_type = match provider_info.provider_type {
|
||||
ApiProviderType::Anthropic => crate::ProviderType::Anthropic,
|
||||
ApiProviderType::Openai | ApiProviderType::OpenaiResponse => {
|
||||
crate::ProviderType::OpenAI
|
||||
}
|
||||
ApiProviderType::Gemini => crate::ProviderType::GeminiApiKey,
|
||||
_ => crate::ProviderType::OpenAI,
|
||||
};
|
||||
|
||||
// 根据 provider_type 创建对应的 CredentialData
|
||||
let credential_data = match provider_type {
|
||||
crate::ProviderType::Anthropic => {
|
||||
crate::models::provider_pool_model::CredentialData::AnthropicKey {
|
||||
api_key: api_key.clone(),
|
||||
base_url,
|
||||
}
|
||||
}
|
||||
crate::ProviderType::GeminiApiKey => {
|
||||
crate::models::provider_pool_model::CredentialData::GeminiApiKey {
|
||||
api_key: api_key.clone(),
|
||||
base_url,
|
||||
excluded_models: vec![],
|
||||
}
|
||||
}
|
||||
_ => crate::models::provider_pool_model::CredentialData::OpenAIKey {
|
||||
api_key: api_key.clone(),
|
||||
base_url,
|
||||
},
|
||||
};
|
||||
|
||||
// 构建 ProviderCredential
|
||||
let mut cred = crate::models::provider_pool_model::ProviderCredential::new(
|
||||
provider_type,
|
||||
credential_data,
|
||||
);
|
||||
cred.name = Some(provider_info.name.clone());
|
||||
|
||||
Some(cred)
|
||||
}
|
||||
Ok(None) => None,
|
||||
Err(_) => None,
|
||||
}
|
||||
} else {
|
||||
None
|
||||
}
|
||||
} else {
|
||||
credential
|
||||
};
|
||||
|
||||
// 如果找到凭证池中的凭证,使用它
|
||||
if let Some(cred) = credential {
|
||||
state.logs.write().await.add(
|
||||
@@ -2548,3 +2729,176 @@ fn build_stream_error_response(
|
||||
.into_response()
|
||||
})
|
||||
}
|
||||
|
||||
// ============================================================================
|
||||
// API Key Provider 辅助函数
|
||||
// ============================================================================
|
||||
|
||||
/// 将 provider_type 映射到 API Key Provider ID
|
||||
fn map_to_api_key_provider_id(provider_type: &str) -> String {
|
||||
match provider_type.to_lowercase().as_str() {
|
||||
"openai" | "gpt" => "openai".to_string(),
|
||||
"anthropic" | "claude" => "anthropic".to_string(),
|
||||
"gemini" | "google" => "gemini".to_string(),
|
||||
"azure" | "azure-openai" | "azure_openai" => "azure-openai".to_string(),
|
||||
"vertexai" | "vertex" => "vertexai".to_string(),
|
||||
"bedrock" | "aws-bedrock" | "aws_bedrock" => "aws-bedrock".to_string(),
|
||||
"ollama" => "ollama".to_string(),
|
||||
_ => provider_type.to_string(),
|
||||
}
|
||||
}
|
||||
|
||||
/// 根据 API Provider 类型构建额外的请求头
|
||||
fn build_api_key_headers(
|
||||
provider_type: &crate::database::dao::api_key_provider::ApiProviderType,
|
||||
api_key: &str,
|
||||
) -> HashMap<String, String> {
|
||||
use crate::database::dao::api_key_provider::ApiProviderType;
|
||||
|
||||
let mut headers = HashMap::new();
|
||||
|
||||
match provider_type {
|
||||
ApiProviderType::Anthropic => {
|
||||
headers.insert("x-api-key".to_string(), api_key.to_string());
|
||||
headers.insert("anthropic-version".to_string(), "2023-06-01".to_string());
|
||||
}
|
||||
ApiProviderType::Gemini => {
|
||||
headers.insert("x-goog-api-key".to_string(), api_key.to_string());
|
||||
}
|
||||
ApiProviderType::AzureOpenai => {
|
||||
headers.insert("api-key".to_string(), api_key.to_string());
|
||||
}
|
||||
_ => {
|
||||
headers.insert("Authorization".to_string(), format!("Bearer {}", api_key));
|
||||
}
|
||||
}
|
||||
|
||||
headers
|
||||
}
|
||||
|
||||
/// 获取默认的 API Host
|
||||
fn get_default_api_host(
|
||||
provider_type: &crate::database::dao::api_key_provider::ApiProviderType,
|
||||
) -> String {
|
||||
use crate::database::dao::api_key_provider::ApiProviderType;
|
||||
|
||||
match provider_type {
|
||||
ApiProviderType::Openai | ApiProviderType::OpenaiResponse => {
|
||||
"https://api.openai.com".to_string()
|
||||
}
|
||||
ApiProviderType::Anthropic => "https://api.anthropic.com".to_string(),
|
||||
ApiProviderType::Gemini => "https://generativelanguage.googleapis.com".to_string(),
|
||||
ApiProviderType::Ollama => "http://localhost:11434".to_string(),
|
||||
_ => "https://api.openai.com".to_string(),
|
||||
}
|
||||
}
|
||||
|
||||
/// 将 OpenAI 格式请求转换为 Anthropic 格式
|
||||
fn convert_openai_to_anthropic(request: &ChatCompletionRequest) -> serde_json::Value {
|
||||
let mut messages = Vec::new();
|
||||
let mut system_prompt = None;
|
||||
|
||||
for msg in &request.messages {
|
||||
if msg.role == "system" {
|
||||
// 提取 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
|
||||
.iter()
|
||||
.filter_map(|p| {
|
||||
if let crate::models::openai::ContentPart::Text { text } = p {
|
||||
Some(text.clone())
|
||||
} else {
|
||||
None
|
||||
}
|
||||
})
|
||||
.collect::<Vec<_>>()
|
||||
.join("\n"),
|
||||
});
|
||||
}
|
||||
} else {
|
||||
// 转换其他消息
|
||||
let content = match &msg.content {
|
||||
Some(c) => match c {
|
||||
crate::models::openai::MessageContent::Text(s) => s.clone(),
|
||||
crate::models::openai::MessageContent::Parts(parts) => parts
|
||||
.iter()
|
||||
.filter_map(|p| {
|
||||
if let crate::models::openai::ContentPart::Text { text } = p {
|
||||
Some(text.clone())
|
||||
} else {
|
||||
None
|
||||
}
|
||||
})
|
||||
.collect::<Vec<_>>()
|
||||
.join("\n"),
|
||||
},
|
||||
None => String::new(),
|
||||
};
|
||||
|
||||
messages.push(serde_json::json!({
|
||||
"role": msg.role,
|
||||
"content": content
|
||||
}));
|
||||
}
|
||||
}
|
||||
|
||||
let mut result = serde_json::json!({
|
||||
"model": request.model,
|
||||
"messages": messages,
|
||||
"max_tokens": request.max_tokens.unwrap_or(4096),
|
||||
"stream": request.stream
|
||||
});
|
||||
|
||||
if let Some(system) = system_prompt {
|
||||
result["system"] = serde_json::Value::String(system);
|
||||
}
|
||||
|
||||
if let Some(temp) = request.temperature {
|
||||
result["temperature"] = serde_json::Value::Number(
|
||||
serde_json::Number::from_f64(temp as f64).unwrap_or(serde_json::Number::from(1)),
|
||||
);
|
||||
}
|
||||
|
||||
result
|
||||
}
|
||||
|
||||
/// 将 Anthropic 响应转换为 OpenAI 格式
|
||||
fn convert_anthropic_response_to_openai(anthropic_resp: &serde_json::Value, model: &str) -> String {
|
||||
let content = anthropic_resp["content"]
|
||||
.as_array()
|
||||
.and_then(|arr| arr.first())
|
||||
.and_then(|c| c["text"].as_str())
|
||||
.unwrap_or("");
|
||||
|
||||
let usage = serde_json::json!({
|
||||
"prompt_tokens": anthropic_resp["usage"]["input_tokens"].as_u64().unwrap_or(0),
|
||||
"completion_tokens": anthropic_resp["usage"]["output_tokens"].as_u64().unwrap_or(0),
|
||||
"total_tokens": anthropic_resp["usage"]["input_tokens"].as_u64().unwrap_or(0)
|
||||
+ anthropic_resp["usage"]["output_tokens"].as_u64().unwrap_or(0)
|
||||
});
|
||||
|
||||
let openai_resp = serde_json::json!({
|
||||
"id": anthropic_resp["id"].as_str().unwrap_or("chatcmpl-unknown"),
|
||||
"object": "chat.completion",
|
||||
"created": chrono::Utc::now().timestamp(),
|
||||
"model": model,
|
||||
"choices": [{
|
||||
"index": 0,
|
||||
"message": {
|
||||
"role": "assistant",
|
||||
"content": content
|
||||
},
|
||||
"finish_reason": match anthropic_resp["stop_reason"].as_str() {
|
||||
Some("end_turn") => "stop",
|
||||
Some("max_tokens") => "length",
|
||||
Some("tool_use") => "tool_calls",
|
||||
_ => "stop"
|
||||
}
|
||||
}],
|
||||
"usage": usage
|
||||
});
|
||||
|
||||
serde_json::to_string(&openai_resp).unwrap_or_default()
|
||||
}
|
||||
|
||||
@@ -1,6 +1,10 @@
|
||||
//! 凭证 API 端点(用于 aster Agent 集成)
|
||||
//!
|
||||
//! 为 aster 子进程提供凭证查询接口,支持所有 11 种 Provider 类型。
|
||||
//! 为 aster 子进程提供凭证查询接口,支持多种凭证类型:
|
||||
//! - OAuth 凭证(Kiro, Gemini, Qwen, Antigravity 等)
|
||||
//! - API Key Provider(OpenAI, Anthropic, Gemini API Key 等)
|
||||
//! - OAuth 插件凭证(动态加载的第三方插件)
|
||||
//!
|
||||
//! 此 API 仅供内部使用,返回完整的凭证信息(包括未脱敏的 access_token)。
|
||||
|
||||
use axum::{
|
||||
@@ -12,6 +16,7 @@ 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;
|
||||
@@ -20,10 +25,27 @@ use crate::server::AppState;
|
||||
#[derive(Debug, Deserialize)]
|
||||
pub struct SelectCredentialRequest {
|
||||
/// Provider 类型(kiro, gemini, qwen, openai, claude, etc.)
|
||||
/// 支持 OAuth 凭证类型和 API Key Provider 类型
|
||||
pub provider_type: String,
|
||||
/// 指定模型(可选)
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub model: Option<String>,
|
||||
/// 凭证来源偏好(可选):oauth, api_key, plugin
|
||||
/// 如果不指定,会按优先级自动选择
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub source_preference: Option<String>,
|
||||
}
|
||||
|
||||
/// 凭证类型
|
||||
#[derive(Debug, Clone, Serialize, Deserialize, PartialEq)]
|
||||
#[serde(rename_all = "snake_case")]
|
||||
pub enum CredentialType {
|
||||
/// OAuth 凭证(凭证池)
|
||||
OAuth,
|
||||
/// API Key(API Key Provider)
|
||||
ApiKey,
|
||||
/// OAuth 插件凭证
|
||||
Plugin,
|
||||
}
|
||||
|
||||
/// 凭证信息响应
|
||||
@@ -33,6 +55,8 @@ pub struct CredentialResponse {
|
||||
pub uuid: String,
|
||||
/// Provider 类型
|
||||
pub provider_type: String,
|
||||
/// 凭证类型
|
||||
pub credential_type: CredentialType,
|
||||
/// Access Token(完整,未脱敏)
|
||||
pub access_token: String,
|
||||
/// Base URL
|
||||
@@ -43,6 +67,9 @@ pub struct CredentialResponse {
|
||||
/// 凭证名称
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub name: Option<String>,
|
||||
/// 额外的请求头(用于某些 Provider)
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub extra_headers: Option<std::collections::HashMap<String, String>>,
|
||||
}
|
||||
|
||||
/// API 错误响应
|
||||
@@ -62,15 +89,26 @@ impl IntoResponse for CredentialApiError {
|
||||
}
|
||||
|
||||
/// POST /v1/credentials/select - 选择可用凭证
|
||||
///
|
||||
/// 支持多种凭证来源:
|
||||
/// 1. OAuth 凭证池(Kiro, Gemini, Qwen, Antigravity 等)
|
||||
/// 2. API Key Provider(OpenAI, Anthropic, Gemini API Key 等)
|
||||
/// 3. OAuth 插件凭证(动态加载的第三方插件)
|
||||
///
|
||||
/// 选择优先级(如果未指定 source_preference):
|
||||
/// 1. 首先尝试 OAuth 凭证池
|
||||
/// 2. 然后尝试 API Key Provider
|
||||
/// 3. 最后尝试 OAuth 插件
|
||||
pub async fn credentials_select(
|
||||
State(state): State<AppState>,
|
||||
_headers: HeaderMap,
|
||||
Json(request): Json<SelectCredentialRequest>,
|
||||
) -> Result<Json<CredentialResponse>, CredentialApiError> {
|
||||
tracing::info!(
|
||||
"[CREDENTIALS_API] 选择凭证请求: provider_type={}, model={:?}",
|
||||
"[CREDENTIALS_API] 选择凭证请求: provider_type={}, model={:?}, source_preference={:?}",
|
||||
request.provider_type,
|
||||
request.model
|
||||
request.model,
|
||||
request.source_preference
|
||||
);
|
||||
|
||||
let db = state.db.as_ref().ok_or_else(|| CredentialApiError {
|
||||
@@ -79,34 +117,199 @@ pub async fn credentials_select(
|
||||
status_code: 503,
|
||||
})?;
|
||||
|
||||
// 根据 source_preference 决定选择策略
|
||||
let source_pref = request.source_preference.as_deref();
|
||||
|
||||
// 尝试从 OAuth 凭证池选择
|
||||
if source_pref.is_none() || source_pref == Some("oauth") {
|
||||
if let Some(response) = try_select_oauth_credential(&state, db, &request).await? {
|
||||
return Ok(Json(response));
|
||||
}
|
||||
}
|
||||
|
||||
// 尝试从 API Key Provider 选择
|
||||
if source_pref.is_none() || source_pref == Some("api_key") {
|
||||
if let Some(response) = try_select_api_key_credential(&state, db, &request).await? {
|
||||
return Ok(Json(response));
|
||||
}
|
||||
}
|
||||
|
||||
// 尝试从 OAuth 插件选择
|
||||
if source_pref.is_none() || source_pref == Some("plugin") {
|
||||
if let Some(response) = try_select_plugin_credential(&state, &request).await? {
|
||||
return Ok(Json(response));
|
||||
}
|
||||
}
|
||||
|
||||
// 没有找到可用凭证
|
||||
Err(CredentialApiError {
|
||||
error: "no_available_credentials".to_string(),
|
||||
message: format!("没有可用的 {} 凭证", request.provider_type),
|
||||
status_code: 503,
|
||||
})
|
||||
}
|
||||
|
||||
/// 尝试从 OAuth 凭证池选择凭证
|
||||
async fn try_select_oauth_credential(
|
||||
state: &AppState,
|
||||
db: &crate::database::DbConnection,
|
||||
request: &SelectCredentialRequest,
|
||||
) -> Result<Option<CredentialResponse>, CredentialApiError> {
|
||||
// 使用 ProviderPoolService 智能选择凭证
|
||||
let credential = state
|
||||
.pool_service
|
||||
.select_credential(db, &request.provider_type, request.model.as_deref())
|
||||
.map_err(|e| CredentialApiError {
|
||||
error: "selection_error".to_string(),
|
||||
message: format!("凭证选择失败: {}", e),
|
||||
status_code: 500,
|
||||
})?
|
||||
.ok_or_else(|| CredentialApiError {
|
||||
error: "no_available_credentials".to_string(),
|
||||
message: format!("没有可用的 {} 凭证", request.provider_type),
|
||||
status_code: 503,
|
||||
})?;
|
||||
let credential = match state.pool_service.select_credential(
|
||||
db,
|
||||
&request.provider_type,
|
||||
request.model.as_deref(),
|
||||
) {
|
||||
Ok(Some(cred)) => cred,
|
||||
Ok(None) => return Ok(None),
|
||||
Err(_) => return Ok(None),
|
||||
};
|
||||
|
||||
// 获取 access_token
|
||||
let access_token = credential
|
||||
let access_token = match credential
|
||||
.cached_token
|
||||
.as_ref()
|
||||
.and_then(|cache| cache.access_token.clone())
|
||||
.ok_or_else(|| CredentialApiError {
|
||||
error: "no_cached_token".to_string(),
|
||||
message: "凭证没有缓存的 Token".to_string(),
|
||||
status_code: 503,
|
||||
})?;
|
||||
{
|
||||
Some(token) => token,
|
||||
None => return Ok(None),
|
||||
};
|
||||
|
||||
// 根据 Provider 类型确定 base_url
|
||||
let base_url = match credential.provider_type {
|
||||
let base_url = get_oauth_base_url(&credential.provider_type);
|
||||
|
||||
let response = CredentialResponse {
|
||||
uuid: credential.uuid.clone(),
|
||||
provider_type: credential.provider_type.to_string(),
|
||||
credential_type: CredentialType::OAuth,
|
||||
access_token,
|
||||
base_url,
|
||||
expires_at: credential
|
||||
.cached_token
|
||||
.as_ref()
|
||||
.and_then(|cache| cache.expiry_time),
|
||||
name: credential.name.clone(),
|
||||
extra_headers: None,
|
||||
};
|
||||
|
||||
tracing::info!(
|
||||
"[CREDENTIALS_API] OAuth 凭证选择成功: {} ({})",
|
||||
response.name.as_deref().unwrap_or("未命名"),
|
||||
response.uuid
|
||||
);
|
||||
|
||||
Ok(Some(response))
|
||||
}
|
||||
|
||||
/// 尝试从 API Key Provider 选择凭证
|
||||
async fn try_select_api_key_credential(
|
||||
state: &AppState,
|
||||
db: &crate::database::DbConnection,
|
||||
request: &SelectCredentialRequest,
|
||||
) -> Result<Option<CredentialResponse>, CredentialApiError> {
|
||||
// 将 provider_type 映射到 API Key Provider ID
|
||||
let provider_id = map_to_api_key_provider_id(&request.provider_type);
|
||||
|
||||
// 获取 API Key Provider Service
|
||||
let api_key_service = &state.api_key_service;
|
||||
|
||||
// 尝试获取下一个可用的 API Key
|
||||
let (key_id, api_key) = match api_key_service.get_next_api_key_entry(db, &provider_id) {
|
||||
Ok(Some((id, key))) => (id, key),
|
||||
Ok(None) => return Ok(None),
|
||||
Err(_) => return Ok(None),
|
||||
};
|
||||
|
||||
// 获取 Provider 信息以确定 base_url
|
||||
let conn = db.lock().map_err(|e| CredentialApiError {
|
||||
error: "database_lock_error".to_string(),
|
||||
message: format!("数据库锁定失败: {}", e),
|
||||
status_code: 500,
|
||||
})?;
|
||||
|
||||
let provider = match ApiKeyProviderDao::get_provider_by_id(&conn, &provider_id) {
|
||||
Ok(Some(p)) => p,
|
||||
Ok(None) => return Ok(None),
|
||||
Err(_) => return Ok(None),
|
||||
};
|
||||
drop(conn);
|
||||
|
||||
// 构建额外的请求头
|
||||
let extra_headers = build_api_key_headers(&provider.provider_type, &api_key);
|
||||
|
||||
let response = CredentialResponse {
|
||||
uuid: key_id,
|
||||
provider_type: request.provider_type.clone(),
|
||||
credential_type: CredentialType::ApiKey,
|
||||
access_token: api_key,
|
||||
base_url: provider.api_host,
|
||||
expires_at: None, // API Key 通常没有过期时间
|
||||
name: Some(provider.name),
|
||||
extra_headers: Some(extra_headers),
|
||||
};
|
||||
|
||||
tracing::info!(
|
||||
"[CREDENTIALS_API] API Key 凭证选择成功: {} ({})",
|
||||
response.name.as_deref().unwrap_or("未命名"),
|
||||
response.uuid
|
||||
);
|
||||
|
||||
Ok(Some(response))
|
||||
}
|
||||
|
||||
/// 尝试从 OAuth 插件选择凭证
|
||||
async fn try_select_plugin_credential(
|
||||
_state: &AppState,
|
||||
request: &SelectCredentialRequest,
|
||||
) -> Result<Option<CredentialResponse>, CredentialApiError> {
|
||||
// 获取 OAuth 插件注册表
|
||||
let registry = match crate::credential::registry::get_global_registry() {
|
||||
Some(r) => r,
|
||||
None => return Ok(None),
|
||||
};
|
||||
|
||||
// 根据模型查找插件
|
||||
let model = request.model.as_deref().unwrap_or("");
|
||||
let plugin = match registry.find_by_model(model).await {
|
||||
Some(p) => p,
|
||||
None => return Ok(None),
|
||||
};
|
||||
|
||||
// 获取凭证
|
||||
let acquired = match plugin.acquire_credential(model).await {
|
||||
Ok(cred) => cred,
|
||||
Err(_) => return Ok(None),
|
||||
};
|
||||
|
||||
// 构建响应
|
||||
let response = CredentialResponse {
|
||||
uuid: acquired.id.clone(),
|
||||
provider_type: plugin.id().to_string(),
|
||||
credential_type: CredentialType::Plugin,
|
||||
access_token: acquired
|
||||
.headers
|
||||
.get("Authorization")
|
||||
.map(|h| h.trim_start_matches("Bearer ").to_string())
|
||||
.unwrap_or_default(),
|
||||
base_url: acquired.base_url.unwrap_or_default(),
|
||||
expires_at: None,
|
||||
name: acquired.name,
|
||||
extra_headers: Some(acquired.headers),
|
||||
};
|
||||
|
||||
tracing::info!(
|
||||
"[CREDENTIALS_API] OAuth 插件凭证选择成功: {} ({})",
|
||||
response.name.as_deref().unwrap_or("未命名"),
|
||||
response.uuid
|
||||
);
|
||||
|
||||
Ok(Some(response))
|
||||
}
|
||||
|
||||
/// 根据 OAuth Provider 类型获取 base_url
|
||||
fn get_oauth_base_url(provider_type: &PoolProviderType) -> String {
|
||||
match provider_type {
|
||||
PoolProviderType::Kiro => "https://api.anthropic.com".to_string(),
|
||||
PoolProviderType::Gemini => "https://generativelanguage.googleapis.com".to_string(),
|
||||
PoolProviderType::Qwen => "https://dashscope.aliyuncs.com/compatible-mode/v1".to_string(),
|
||||
@@ -116,37 +319,55 @@ pub async fn credentials_select(
|
||||
PoolProviderType::Codex => "https://api.openai.com/v1".to_string(),
|
||||
PoolProviderType::ClaudeOAuth => "https://api.anthropic.com".to_string(),
|
||||
PoolProviderType::IFlow => "https://chat.iflyrec.com".to_string(),
|
||||
_ => {
|
||||
return Err(CredentialApiError {
|
||||
error: "unsupported_provider".to_string(),
|
||||
message: format!("不支持的 Provider 类型: {:?}", credential.provider_type),
|
||||
status_code: 400,
|
||||
})
|
||||
_ => "https://api.openai.com/v1".to_string(),
|
||||
}
|
||||
}
|
||||
|
||||
/// 将 provider_type 映射到 API Key Provider ID
|
||||
fn map_to_api_key_provider_id(provider_type: &str) -> String {
|
||||
match provider_type.to_lowercase().as_str() {
|
||||
"openai" | "gpt" => "openai".to_string(),
|
||||
"anthropic" | "claude" => "anthropic".to_string(),
|
||||
"gemini" | "google" => "gemini".to_string(),
|
||||
"azure" | "azure-openai" => "azure-openai".to_string(),
|
||||
"vertexai" | "vertex" => "vertexai".to_string(),
|
||||
"bedrock" | "aws-bedrock" => "aws-bedrock".to_string(),
|
||||
"ollama" => "ollama".to_string(),
|
||||
_ => provider_type.to_string(),
|
||||
}
|
||||
}
|
||||
|
||||
/// 根据 API Provider 类型构建额外的请求头
|
||||
fn build_api_key_headers(
|
||||
provider_type: &ApiProviderType,
|
||||
api_key: &str,
|
||||
) -> std::collections::HashMap<String, String> {
|
||||
let mut headers = std::collections::HashMap::new();
|
||||
|
||||
match provider_type {
|
||||
ApiProviderType::Anthropic => {
|
||||
headers.insert("x-api-key".to_string(), api_key.to_string());
|
||||
headers.insert("anthropic-version".to_string(), "2023-06-01".to_string());
|
||||
}
|
||||
};
|
||||
ApiProviderType::Gemini => {
|
||||
headers.insert("x-goog-api-key".to_string(), api_key.to_string());
|
||||
}
|
||||
ApiProviderType::AzureOpenai => {
|
||||
headers.insert("api-key".to_string(), api_key.to_string());
|
||||
}
|
||||
_ => {
|
||||
headers.insert("Authorization".to_string(), format!("Bearer {}", api_key));
|
||||
}
|
||||
}
|
||||
|
||||
let response = CredentialResponse {
|
||||
uuid: credential.uuid.clone(),
|
||||
provider_type: credential.provider_type.to_string(),
|
||||
access_token,
|
||||
base_url,
|
||||
expires_at: credential
|
||||
.cached_token
|
||||
.as_ref()
|
||||
.and_then(|cache| cache.expiry_time),
|
||||
name: credential.name.clone(),
|
||||
};
|
||||
|
||||
tracing::info!(
|
||||
"[CREDENTIALS_API] 凭证选择成功: {} ({})",
|
||||
response.name.as_deref().unwrap_or("未命名"),
|
||||
response.uuid
|
||||
);
|
||||
|
||||
Ok(Json(response))
|
||||
headers
|
||||
}
|
||||
|
||||
/// GET /v1/credentials/{uuid}/token - 获取指定凭证的 Token
|
||||
///
|
||||
/// 支持多种凭证类型:
|
||||
/// - OAuth 凭证池中的凭证
|
||||
/// - API Key Provider 中的 API Key
|
||||
pub async fn credentials_get_token(
|
||||
State(state): State<AppState>,
|
||||
Path(uuid): Path<String>,
|
||||
@@ -160,6 +381,30 @@ pub async fn credentials_get_token(
|
||||
status_code: 503,
|
||||
})?;
|
||||
|
||||
// 首先尝试从 OAuth 凭证池查询
|
||||
if let Some(response) = try_get_oauth_token(&state, db, &uuid).await? {
|
||||
return Ok(Json(response));
|
||||
}
|
||||
|
||||
// 然后尝试从 API Key Provider 查询
|
||||
if let Some(response) = try_get_api_key_token(&state, db, &uuid).await? {
|
||||
return Ok(Json(response));
|
||||
}
|
||||
|
||||
// 未找到凭证
|
||||
Err(CredentialApiError {
|
||||
error: "credential_not_found".to_string(),
|
||||
message: format!("未找到 UUID 为 {} 的凭证", uuid),
|
||||
status_code: 404,
|
||||
})
|
||||
}
|
||||
|
||||
/// 尝试从 OAuth 凭证池获取 Token
|
||||
async fn try_get_oauth_token(
|
||||
state: &AppState,
|
||||
db: &crate::database::DbConnection,
|
||||
uuid: &str,
|
||||
) -> Result<Option<CredentialResponse>, CredentialApiError> {
|
||||
// 查询凭证
|
||||
let credential = {
|
||||
let conn = db.lock().map_err(|e| CredentialApiError {
|
||||
@@ -168,17 +413,11 @@ pub async fn credentials_get_token(
|
||||
status_code: 500,
|
||||
})?;
|
||||
|
||||
ProviderPoolDao::get_by_uuid(&conn, &uuid)
|
||||
.map_err(|e| CredentialApiError {
|
||||
error: "database_query_error".to_string(),
|
||||
message: format!("查询凭证失败: {}", e),
|
||||
status_code: 500,
|
||||
})?
|
||||
.ok_or_else(|| CredentialApiError {
|
||||
error: "credential_not_found".to_string(),
|
||||
message: format!("未找到 UUID 为 {} 的凭证", uuid),
|
||||
status_code: 404,
|
||||
})?
|
||||
match ProviderPoolDao::get_by_uuid(&conn, uuid) {
|
||||
Ok(Some(cred)) => cred,
|
||||
Ok(None) => return Ok(None),
|
||||
Err(_) => return Ok(None),
|
||||
}
|
||||
};
|
||||
|
||||
// 如果 Token 即将过期,尝试刷新
|
||||
@@ -194,7 +433,7 @@ pub async fn credentials_get_token(
|
||||
.token_cache
|
||||
.refresh_and_cache_with_events(
|
||||
db,
|
||||
&uuid,
|
||||
uuid,
|
||||
false,
|
||||
Some(state.kiro_event_service.clone()),
|
||||
)
|
||||
@@ -219,31 +458,13 @@ pub async fn credentials_get_token(
|
||||
None
|
||||
};
|
||||
|
||||
let access_token = cached_token.ok_or_else(|| CredentialApiError {
|
||||
error: "no_cached_token".to_string(),
|
||||
message: "凭证没有缓存的 Token".to_string(),
|
||||
status_code: 503,
|
||||
})?;
|
||||
let access_token = match cached_token {
|
||||
Some(token) => token,
|
||||
None => return Ok(None),
|
||||
};
|
||||
|
||||
// 根据 Provider 类型确定 base_url
|
||||
let base_url = match credential.provider_type {
|
||||
PoolProviderType::Kiro => "https://api.anthropic.com".to_string(),
|
||||
PoolProviderType::Gemini => "https://generativelanguage.googleapis.com".to_string(),
|
||||
PoolProviderType::Qwen => "https://dashscope.aliyuncs.com/compatible-mode/v1".to_string(),
|
||||
PoolProviderType::Antigravity => "https://api.anthropic.com".to_string(),
|
||||
PoolProviderType::Vertex => "https://vertex-ai.googleapis.com".to_string(),
|
||||
PoolProviderType::GeminiApiKey => "https://generativelanguage.googleapis.com".to_string(),
|
||||
PoolProviderType::Codex => "https://api.openai.com/v1".to_string(),
|
||||
PoolProviderType::ClaudeOAuth => "https://api.anthropic.com".to_string(),
|
||||
PoolProviderType::IFlow => "https://chat.iflyrec.com".to_string(),
|
||||
_ => {
|
||||
return Err(CredentialApiError {
|
||||
error: "unsupported_provider".to_string(),
|
||||
message: format!("不支持的 Provider 类型: {:?}", credential.provider_type),
|
||||
status_code: 400,
|
||||
})
|
||||
}
|
||||
};
|
||||
let base_url = get_oauth_base_url(&credential.provider_type);
|
||||
|
||||
// 重新查询凭证以获取更新后的 expires_at
|
||||
let updated_credential = {
|
||||
@@ -253,22 +474,17 @@ pub async fn credentials_get_token(
|
||||
status_code: 500,
|
||||
})?;
|
||||
|
||||
ProviderPoolDao::get_by_uuid(&conn, &uuid)
|
||||
.map_err(|e| CredentialApiError {
|
||||
error: "database_query_error".to_string(),
|
||||
message: format!("查询凭证失败: {}", e),
|
||||
status_code: 500,
|
||||
})?
|
||||
.ok_or_else(|| CredentialApiError {
|
||||
error: "credential_not_found".to_string(),
|
||||
message: format!("未找到 UUID 为 {} 的凭证", uuid),
|
||||
status_code: 404,
|
||||
})?
|
||||
match ProviderPoolDao::get_by_uuid(&conn, uuid) {
|
||||
Ok(Some(cred)) => cred,
|
||||
Ok(None) => return Ok(None),
|
||||
Err(_) => return Ok(None),
|
||||
}
|
||||
};
|
||||
|
||||
let response = CredentialResponse {
|
||||
uuid: updated_credential.uuid.clone(),
|
||||
provider_type: updated_credential.provider_type.to_string(),
|
||||
credential_type: CredentialType::OAuth,
|
||||
access_token,
|
||||
base_url,
|
||||
expires_at: updated_credential
|
||||
@@ -276,13 +492,74 @@ pub async fn credentials_get_token(
|
||||
.as_ref()
|
||||
.and_then(|cache| cache.expiry_time),
|
||||
name: updated_credential.name.clone(),
|
||||
extra_headers: None,
|
||||
};
|
||||
|
||||
tracing::info!(
|
||||
"[CREDENTIALS_API] 返回凭证 Token: {} ({})",
|
||||
"[CREDENTIALS_API] 返回 OAuth 凭证 Token: {} ({})",
|
||||
response.name.as_deref().unwrap_or("未命名"),
|
||||
response.uuid
|
||||
);
|
||||
|
||||
Ok(Json(response))
|
||||
Ok(Some(response))
|
||||
}
|
||||
|
||||
/// 尝试从 API Key Provider 获取 Token
|
||||
async fn try_get_api_key_token(
|
||||
state: &AppState,
|
||||
db: &crate::database::DbConnection,
|
||||
uuid: &str,
|
||||
) -> Result<Option<CredentialResponse>, CredentialApiError> {
|
||||
let conn = db.lock().map_err(|e| CredentialApiError {
|
||||
error: "database_lock_error".to_string(),
|
||||
message: format!("数据库锁定失败: {}", e),
|
||||
status_code: 500,
|
||||
})?;
|
||||
|
||||
// 查询 API Key
|
||||
let api_key_entry = match ApiKeyProviderDao::get_api_key_by_id(&conn, uuid) {
|
||||
Ok(Some(key)) => key,
|
||||
Ok(None) => return Ok(None),
|
||||
Err(_) => return Ok(None),
|
||||
};
|
||||
|
||||
// 获取 Provider 信息
|
||||
let provider = match ApiKeyProviderDao::get_provider_by_id(&conn, &api_key_entry.provider_id) {
|
||||
Ok(Some(p)) => p,
|
||||
Ok(None) => return Ok(None),
|
||||
Err(_) => return Ok(None),
|
||||
};
|
||||
drop(conn);
|
||||
|
||||
// 解密 API Key
|
||||
let api_key = state
|
||||
.api_key_service
|
||||
.decrypt_api_key(&api_key_entry.api_key_encrypted)
|
||||
.map_err(|e| CredentialApiError {
|
||||
error: "decryption_error".to_string(),
|
||||
message: format!("API Key 解密失败: {}", e),
|
||||
status_code: 500,
|
||||
})?;
|
||||
|
||||
// 构建额外的请求头
|
||||
let extra_headers = build_api_key_headers(&provider.provider_type, &api_key);
|
||||
|
||||
let response = CredentialResponse {
|
||||
uuid: api_key_entry.id.clone(),
|
||||
provider_type: provider.provider_type.to_string(),
|
||||
credential_type: CredentialType::ApiKey,
|
||||
access_token: api_key,
|
||||
base_url: provider.api_host,
|
||||
expires_at: None,
|
||||
name: api_key_entry.alias.or(Some(provider.name)),
|
||||
extra_headers: Some(extra_headers),
|
||||
};
|
||||
|
||||
tracing::info!(
|
||||
"[CREDENTIALS_API] 返回 API Key 凭证: {} ({})",
|
||||
response.name.as_deref().unwrap_or("未命名"),
|
||||
response.uuid
|
||||
);
|
||||
|
||||
Ok(Some(response))
|
||||
}
|
||||
|
||||
@@ -431,6 +431,37 @@ pub async fn management_add_credential(
|
||||
);
|
||||
}
|
||||
}
|
||||
// Anthropic API Key Provider
|
||||
PoolProviderType::Anthropic => {
|
||||
if let Some(api_key) = request.api_key {
|
||||
CredentialData::AnthropicKey {
|
||||
api_key,
|
||||
base_url: request.base_url,
|
||||
}
|
||||
} else {
|
||||
return (
|
||||
StatusCode::BAD_REQUEST,
|
||||
Json(AddCredentialResponse {
|
||||
success: false,
|
||||
message: "API key is required for Anthropic provider".to_string(),
|
||||
id: None,
|
||||
}),
|
||||
);
|
||||
}
|
||||
}
|
||||
// API Key Provider 类型 - 不支持通过此接口添加凭证
|
||||
PoolProviderType::AzureOpenai | PoolProviderType::AwsBedrock | PoolProviderType::Ollama => {
|
||||
return (
|
||||
StatusCode::BAD_REQUEST,
|
||||
Json(AddCredentialResponse {
|
||||
success: false,
|
||||
message:
|
||||
"This provider type should be configured via API Key Provider settings"
|
||||
.to_string(),
|
||||
id: None,
|
||||
}),
|
||||
);
|
||||
}
|
||||
};
|
||||
|
||||
// 创建凭证
|
||||
|
||||
@@ -778,6 +778,301 @@ pub async fn call_provider_anthropic(
|
||||
)
|
||||
.into_response()
|
||||
}
|
||||
// Anthropic API Key - 根据 base_url 决定调用方式
|
||||
CredentialData::AnthropicKey { api_key, base_url } => {
|
||||
// 如果有自定义 base_url,假设是 OpenAI 兼容的代理服务器
|
||||
// 需要将 Anthropic 请求转换为 OpenAI 请求,然后将响应转换回来
|
||||
if let Some(custom_url) = base_url {
|
||||
state.logs.write().await.add(
|
||||
"info",
|
||||
&format!(
|
||||
"[ANTHROPIC_COMPAT] 使用 OpenAI 兼容 API: base_url={} credential_uuid={} stream={}",
|
||||
custom_url,
|
||||
&credential.uuid[..8],
|
||||
request.stream
|
||||
),
|
||||
);
|
||||
|
||||
// 将 Anthropic 请求转换为 OpenAI 请求
|
||||
let openai_request = crate::converter::anthropic_to_openai::convert_anthropic_to_openai(request);
|
||||
|
||||
// 使用 OpenAI 兼容 API 调用
|
||||
let openai = OpenAICustomProvider::with_config(api_key.clone(), Some(custom_url.clone()));
|
||||
match openai.call_api(&openai_request).await {
|
||||
Ok(resp) => {
|
||||
let status = resp.status();
|
||||
state.logs.write().await.add(
|
||||
"info",
|
||||
&format!(
|
||||
"[ANTHROPIC_COMPAT] 响应状态: status={} model={} stream={}",
|
||||
status,
|
||||
request.model,
|
||||
request.stream
|
||||
),
|
||||
);
|
||||
|
||||
// 流式请求暂不支持格式转换,直接透传 OpenAI SSE 格式
|
||||
// TODO: 实现 OpenAI SSE -> Anthropic SSE 的流式转换
|
||||
if request.stream && status.is_success() {
|
||||
state.logs.write().await.add(
|
||||
"info",
|
||||
"[ANTHROPIC_COMPAT] 流式请求,透传 OpenAI SSE 响应(暂不支持格式转换)",
|
||||
);
|
||||
if let Some(db) = &state.db {
|
||||
let _ = state.pool_service.mark_healthy(
|
||||
db,
|
||||
&credential.uuid,
|
||||
Some(&request.model),
|
||||
);
|
||||
let _ = state.pool_service.record_usage(db, &credential.uuid);
|
||||
}
|
||||
|
||||
// 直接透传 OpenAI SSE 流
|
||||
let stream = resp.bytes_stream();
|
||||
return Response::builder()
|
||||
.status(StatusCode::OK)
|
||||
.header(header::CONTENT_TYPE, "text/event-stream")
|
||||
.header(header::CACHE_CONTROL, "no-cache")
|
||||
.header("Connection", "keep-alive")
|
||||
.body(Body::from_stream(stream))
|
||||
.unwrap_or_else(|_| {
|
||||
(
|
||||
StatusCode::INTERNAL_SERVER_ERROR,
|
||||
Json(serde_json::json!({"error": {"message": "Failed to build stream response"}})),
|
||||
)
|
||||
.into_response()
|
||||
});
|
||||
}
|
||||
|
||||
// 非流式响应需要转换格式
|
||||
if status.is_success() {
|
||||
if let Some(db) = &state.db {
|
||||
let _ = state.pool_service.mark_healthy(
|
||||
db,
|
||||
&credential.uuid,
|
||||
Some(&request.model),
|
||||
);
|
||||
let _ = state.pool_service.record_usage(db, &credential.uuid);
|
||||
}
|
||||
} else {
|
||||
if let Some(db) = &state.db {
|
||||
let _ = state.pool_service.mark_unhealthy(
|
||||
db,
|
||||
&credential.uuid,
|
||||
Some(&format!("API error: {}", status)),
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
match resp.text().await {
|
||||
Ok(body) => {
|
||||
if status.is_success() {
|
||||
// 将 OpenAI 响应转换为 Anthropic 响应
|
||||
match serde_json::from_str::<crate::models::openai::ChatCompletionResponse>(&body) {
|
||||
Ok(openai_resp) => {
|
||||
let anthropic_resp = convert_openai_response_to_anthropic(&openai_resp, &request.model);
|
||||
Response::builder()
|
||||
.status(StatusCode::OK)
|
||||
.header(header::CONTENT_TYPE, "application/json")
|
||||
.body(Body::from(serde_json::to_string(&anthropic_resp).unwrap_or_default()))
|
||||
.unwrap_or_else(|_| {
|
||||
(
|
||||
StatusCode::INTERNAL_SERVER_ERROR,
|
||||
Json(serde_json::json!({"error": {"message": "Failed to build response"}})),
|
||||
)
|
||||
.into_response()
|
||||
})
|
||||
}
|
||||
Err(e) => {
|
||||
state.logs.write().await.add(
|
||||
"error",
|
||||
&format!("[ANTHROPIC_COMPAT] 解析 OpenAI 响应失败: {}", e),
|
||||
);
|
||||
// 返回原始响应
|
||||
Response::builder()
|
||||
.status(StatusCode::OK)
|
||||
.header(header::CONTENT_TYPE, "application/json")
|
||||
.body(Body::from(body))
|
||||
.unwrap_or_else(|_| {
|
||||
(
|
||||
StatusCode::INTERNAL_SERVER_ERROR,
|
||||
Json(serde_json::json!({"error": {"message": "Failed to build response"}})),
|
||||
)
|
||||
.into_response()
|
||||
})
|
||||
}
|
||||
}
|
||||
} else {
|
||||
(
|
||||
StatusCode::from_u16(status.as_u16())
|
||||
.unwrap_or(StatusCode::INTERNAL_SERVER_ERROR),
|
||||
Json(serde_json::json!({"error": {"message": body}})),
|
||||
)
|
||||
.into_response()
|
||||
}
|
||||
}
|
||||
Err(e) => (
|
||||
StatusCode::INTERNAL_SERVER_ERROR,
|
||||
Json(serde_json::json!({"error": {"message": format!("Failed to read response: {}", e)}})),
|
||||
)
|
||||
.into_response(),
|
||||
}
|
||||
}
|
||||
Err(e) => {
|
||||
if let Some(db) = &state.db {
|
||||
let _ = state.pool_service.mark_unhealthy(
|
||||
db,
|
||||
&credential.uuid,
|
||||
Some(&format!("API call failed: {}", e)),
|
||||
);
|
||||
}
|
||||
(
|
||||
StatusCode::INTERNAL_SERVER_ERROR,
|
||||
Json(serde_json::json!({"error": {"message": format!("OpenAI compatible API call failed: {}", e)}})),
|
||||
)
|
||||
.into_response()
|
||||
}
|
||||
}
|
||||
} else {
|
||||
// 没有自定义 base_url,使用原生 Anthropic API
|
||||
let claude = ClaudeCustomProvider::with_config(api_key.clone(), None);
|
||||
let request_url = claude.get_base_url();
|
||||
state.logs.write().await.add(
|
||||
"info",
|
||||
&format!(
|
||||
"[ANTHROPIC] 使用 Anthropic API: base_url=https://api.anthropic.com -> {}/v1/messages credential_uuid={} stream={}",
|
||||
request_url,
|
||||
&credential.uuid[..8],
|
||||
request.stream
|
||||
),
|
||||
);
|
||||
match claude.call_api(request).await {
|
||||
Ok(resp) => {
|
||||
let status = resp.status();
|
||||
state.logs.write().await.add(
|
||||
"info",
|
||||
&format!(
|
||||
"[ANTHROPIC] 响应状态: status={} model={} stream={}",
|
||||
status,
|
||||
request.model,
|
||||
request.stream
|
||||
),
|
||||
);
|
||||
|
||||
// 如果是流式请求,直接透传流式响应
|
||||
if request.stream && status.is_success() {
|
||||
state.logs.write().await.add(
|
||||
"info",
|
||||
"[ANTHROPIC] 流式请求,透传 SSE 响应",
|
||||
);
|
||||
if let Some(db) = &state.db {
|
||||
let _ = state.pool_service.mark_healthy(
|
||||
db,
|
||||
&credential.uuid,
|
||||
Some(&request.model),
|
||||
);
|
||||
let _ = state.pool_service.record_usage(db, &credential.uuid);
|
||||
}
|
||||
let stream = resp.bytes_stream();
|
||||
return Response::builder()
|
||||
.status(StatusCode::OK)
|
||||
.header(header::CONTENT_TYPE, "text/event-stream")
|
||||
.header(header::CACHE_CONTROL, "no-cache")
|
||||
.header("Connection", "keep-alive")
|
||||
.body(Body::from_stream(stream))
|
||||
.unwrap_or_else(|_| {
|
||||
(
|
||||
StatusCode::INTERNAL_SERVER_ERROR,
|
||||
Json(serde_json::json!({"error": {"message": "Failed to build stream response"}})),
|
||||
)
|
||||
.into_response()
|
||||
});
|
||||
}
|
||||
|
||||
// 非流式请求,读取完整响应
|
||||
match resp.text().await {
|
||||
Ok(body) => {
|
||||
if status.is_success() {
|
||||
if let Some(db) = &state.db {
|
||||
let _ = state.pool_service.mark_healthy(
|
||||
db,
|
||||
&credential.uuid,
|
||||
Some(&request.model),
|
||||
);
|
||||
let _ = state.pool_service.record_usage(db, &credential.uuid);
|
||||
}
|
||||
Response::builder()
|
||||
.status(StatusCode::OK)
|
||||
.header(header::CONTENT_TYPE, "application/json")
|
||||
.body(Body::from(body))
|
||||
.unwrap_or_else(|_| {
|
||||
(
|
||||
StatusCode::INTERNAL_SERVER_ERROR,
|
||||
Json(serde_json::json!({"error": {"message": "Failed to build response"}})),
|
||||
)
|
||||
.into_response()
|
||||
})
|
||||
} else {
|
||||
state.logs.write().await.add(
|
||||
"error",
|
||||
&format!(
|
||||
"[ANTHROPIC] 请求失败: status={} body={}",
|
||||
status,
|
||||
&body.chars().take(200).collect::<String>()
|
||||
),
|
||||
);
|
||||
if let Some(db) = &state.db {
|
||||
let _ = state.pool_service.mark_unhealthy(
|
||||
db,
|
||||
&credential.uuid,
|
||||
Some(&body),
|
||||
);
|
||||
}
|
||||
(
|
||||
StatusCode::from_u16(status.as_u16())
|
||||
.unwrap_or(StatusCode::INTERNAL_SERVER_ERROR),
|
||||
Json(serde_json::json!({"error": {"message": body}})),
|
||||
)
|
||||
.into_response()
|
||||
}
|
||||
}
|
||||
Err(e) => {
|
||||
state.logs.write().await.add(
|
||||
"error",
|
||||
&format!("[ANTHROPIC] 读取响应失败: {}", e),
|
||||
);
|
||||
if let Some(db) = &state.db {
|
||||
let _ = state.pool_service.mark_unhealthy(
|
||||
db,
|
||||
&credential.uuid,
|
||||
Some(&e.to_string()),
|
||||
);
|
||||
}
|
||||
(
|
||||
StatusCode::INTERNAL_SERVER_ERROR,
|
||||
Json(serde_json::json!({"error": {"message": e.to_string()}})),
|
||||
)
|
||||
.into_response()
|
||||
}
|
||||
}
|
||||
}
|
||||
Err(e) => {
|
||||
if let Some(db) = &state.db {
|
||||
let _ = state.pool_service.mark_unhealthy(
|
||||
db,
|
||||
&credential.uuid,
|
||||
Some(&e.to_string()),
|
||||
);
|
||||
}
|
||||
(
|
||||
StatusCode::INTERNAL_SERVER_ERROR,
|
||||
Json(serde_json::json!({"error": {"message": e.to_string()}})),
|
||||
)
|
||||
.into_response()
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
@@ -1663,6 +1958,125 @@ pub async fn call_provider_openai(
|
||||
)
|
||||
.into_response()
|
||||
}
|
||||
// AnthropicKey - 如果有自定义 base_url,使用 OpenAI 兼容格式调用
|
||||
CredentialData::AnthropicKey { api_key, base_url } => {
|
||||
// 如果有自定义 base_url,假设是 OpenAI 兼容的代理服务器
|
||||
if let Some(custom_url) = base_url {
|
||||
let openai = OpenAICustomProvider::with_config(api_key.clone(), Some(custom_url.clone()));
|
||||
state.logs.write().await.add(
|
||||
"info",
|
||||
&format!(
|
||||
"[OPENAI_COMPAT] 使用 OpenAI 兼容 API: base_url={} credential_uuid={} stream={}",
|
||||
custom_url,
|
||||
&credential.uuid[..8],
|
||||
request.stream
|
||||
),
|
||||
);
|
||||
match openai.call_api(request).await {
|
||||
Ok(resp) => {
|
||||
let status = resp.status();
|
||||
state.logs.write().await.add(
|
||||
"info",
|
||||
&format!(
|
||||
"[OPENAI_COMPAT] 响应状态: status={} model={} stream={}",
|
||||
status,
|
||||
request.model,
|
||||
request.stream
|
||||
),
|
||||
);
|
||||
|
||||
if request.stream && status.is_success() {
|
||||
state.logs.write().await.add(
|
||||
"info",
|
||||
"[OPENAI_COMPAT] 流式请求,透传 SSE 响应",
|
||||
);
|
||||
if let Some(db) = &state.db {
|
||||
let _ = state.pool_service.mark_healthy(
|
||||
db,
|
||||
&credential.uuid,
|
||||
Some(&request.model),
|
||||
);
|
||||
let _ = state.pool_service.record_usage(db, &credential.uuid);
|
||||
}
|
||||
let stream = resp.bytes_stream();
|
||||
return Response::builder()
|
||||
.status(StatusCode::OK)
|
||||
.header(header::CONTENT_TYPE, "text/event-stream")
|
||||
.header(header::CACHE_CONTROL, "no-cache")
|
||||
.header("Connection", "keep-alive")
|
||||
.body(Body::from_stream(stream))
|
||||
.unwrap_or_else(|_| {
|
||||
(
|
||||
StatusCode::INTERNAL_SERVER_ERROR,
|
||||
Json(serde_json::json!({"error": {"message": "Failed to build stream response"}})),
|
||||
)
|
||||
.into_response()
|
||||
});
|
||||
}
|
||||
|
||||
// 非流式响应
|
||||
if status.is_success() {
|
||||
if let Some(db) = &state.db {
|
||||
let _ = state.pool_service.mark_healthy(
|
||||
db,
|
||||
&credential.uuid,
|
||||
Some(&request.model),
|
||||
);
|
||||
let _ = state.pool_service.record_usage(db, &credential.uuid);
|
||||
}
|
||||
} else {
|
||||
if let Some(db) = &state.db {
|
||||
let _ = state.pool_service.mark_unhealthy(
|
||||
db,
|
||||
&credential.uuid,
|
||||
Some(&format!("API error: {}", status)),
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
match resp.bytes().await {
|
||||
Ok(body) => Response::builder()
|
||||
.status(status)
|
||||
.header(header::CONTENT_TYPE, "application/json")
|
||||
.body(Body::from(body))
|
||||
.unwrap_or_else(|_| {
|
||||
(
|
||||
StatusCode::INTERNAL_SERVER_ERROR,
|
||||
Json(serde_json::json!({"error": {"message": "Failed to build response"}})),
|
||||
)
|
||||
.into_response()
|
||||
}),
|
||||
Err(e) => (
|
||||
StatusCode::INTERNAL_SERVER_ERROR,
|
||||
Json(serde_json::json!({"error": {"message": format!("Failed to read response: {}", e)}})),
|
||||
)
|
||||
.into_response(),
|
||||
}
|
||||
}
|
||||
Err(e) => {
|
||||
if let Some(db) = &state.db {
|
||||
let _ = state.pool_service.mark_unhealthy(
|
||||
db,
|
||||
&credential.uuid,
|
||||
Some(&format!("API call failed: {}", e)),
|
||||
);
|
||||
}
|
||||
(
|
||||
StatusCode::INTERNAL_SERVER_ERROR,
|
||||
Json(serde_json::json!({"error": {"message": format!("OpenAI compatible API call failed: {}", e)}})),
|
||||
)
|
||||
.into_response()
|
||||
}
|
||||
}
|
||||
} else {
|
||||
// 没有自定义 base_url,不支持 OpenAI 格式
|
||||
(
|
||||
StatusCode::BAD_REQUEST,
|
||||
Json(serde_json::json!({"error": {"message": "AnthropicKey without custom base_url does not support OpenAI format. Use Anthropic format endpoint instead."}})),
|
||||
)
|
||||
.into_response()
|
||||
}
|
||||
}
|
||||
// 新增的凭证类型暂不支持 OpenAI 格式
|
||||
CredentialData::CodexOAuth { .. }
|
||||
| CredentialData::ClaudeOAuth { .. }
|
||||
@@ -2936,3 +3350,74 @@ fn convert_gemini_chunk_to_openai_sse(json: &serde_json::Value, model: &str) ->
|
||||
|
||||
Some(format!("data: {}\n\n", response.to_string()))
|
||||
}
|
||||
|
||||
/// 将 OpenAI ChatCompletionResponse 转换为 Anthropic MessagesResponse 格式
|
||||
fn convert_openai_response_to_anthropic(
|
||||
openai_resp: &crate::models::openai::ChatCompletionResponse,
|
||||
model: &str,
|
||||
) -> serde_json::Value {
|
||||
// 提取第一个 choice 的内容
|
||||
let content = openai_resp
|
||||
.choices
|
||||
.first()
|
||||
.and_then(|c| c.message.content.as_ref())
|
||||
.cloned()
|
||||
.unwrap_or_default();
|
||||
|
||||
// 提取 tool_calls
|
||||
let tool_use: Vec<serde_json::Value> = openai_resp
|
||||
.choices
|
||||
.first()
|
||||
.and_then(|c| c.message.tool_calls.as_ref())
|
||||
.map(|calls| {
|
||||
calls
|
||||
.iter()
|
||||
.map(|tc| {
|
||||
serde_json::json!({
|
||||
"type": "tool_use",
|
||||
"id": tc.id,
|
||||
"name": tc.function.name,
|
||||
"input": serde_json::from_str::<serde_json::Value>(&tc.function.arguments).unwrap_or_default()
|
||||
})
|
||||
})
|
||||
.collect()
|
||||
})
|
||||
.unwrap_or_default();
|
||||
|
||||
// 构建 content 数组
|
||||
let mut content_array: Vec<serde_json::Value> = Vec::new();
|
||||
if !content.is_empty() {
|
||||
content_array.push(serde_json::json!({
|
||||
"type": "text",
|
||||
"text": content
|
||||
}));
|
||||
}
|
||||
content_array.extend(tool_use);
|
||||
|
||||
// 转换 finish_reason
|
||||
let stop_reason = openai_resp
|
||||
.choices
|
||||
.first()
|
||||
.map(|c| match c.finish_reason.as_str() {
|
||||
"stop" => "end_turn",
|
||||
"length" => "max_tokens",
|
||||
"tool_calls" => "tool_use",
|
||||
_ => "end_turn",
|
||||
})
|
||||
.unwrap_or("end_turn");
|
||||
|
||||
// 构建 Anthropic 响应
|
||||
serde_json::json!({
|
||||
"id": format!("msg_{}", uuid::Uuid::new_v4()),
|
||||
"type": "message",
|
||||
"role": "assistant",
|
||||
"content": content_array,
|
||||
"model": model,
|
||||
"stop_reason": stop_reason,
|
||||
"stop_sequence": null,
|
||||
"usage": {
|
||||
"input_tokens": openai_resp.usage.prompt_tokens,
|
||||
"output_tokens": openai_resp.usage.completion_tokens
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
@@ -394,6 +394,8 @@ pub struct AppState {
|
||||
pub endpoint_providers: Arc<RwLock<EndpointProvidersConfig>>,
|
||||
/// Kiro 事件服务
|
||||
pub kiro_event_service: Arc<KiroEventService>,
|
||||
/// API Key Provider 服务
|
||||
pub api_key_service: Arc<crate::services::api_key_provider_service::ApiKeyProviderService>,
|
||||
}
|
||||
|
||||
/// 启动配置文件监控
|
||||
@@ -765,6 +767,10 @@ async fn run_server(
|
||||
// 创建 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,
|
||||
@@ -789,6 +795,7 @@ async fn run_server(
|
||||
flow_interceptor,
|
||||
endpoint_providers,
|
||||
kiro_event_service,
|
||||
api_key_service,
|
||||
};
|
||||
|
||||
// 启动配置文件监控
|
||||
|
||||
@@ -472,6 +472,41 @@ impl ApiKeyProviderService {
|
||||
.map_err(|e| e.to_string())
|
||||
}
|
||||
|
||||
/// 按 Provider 类型获取下一个可用的 API Key(轮询负载均衡)
|
||||
/// 这个方法会查找所有该类型的 Provider(包括自定义 Provider)
|
||||
pub fn get_next_api_key_by_type(
|
||||
&self,
|
||||
db: &DbConnection,
|
||||
provider_type: ApiProviderType,
|
||||
) -> Result<Option<(String, String, ApiKeyProvider)>, String> {
|
||||
let conn = db.lock().map_err(|e| e.to_string())?;
|
||||
|
||||
// 获取所有启用的 API Keys(按类型)
|
||||
let keys = ApiKeyProviderDao::get_enabled_api_keys_by_type(&conn, provider_type)
|
||||
.map_err(|e| e.to_string())?;
|
||||
|
||||
if keys.is_empty() {
|
||||
return Ok(None);
|
||||
}
|
||||
|
||||
// 使用类型名称作为轮询索引的 key
|
||||
let type_key = format!("type:{}", provider_type);
|
||||
let index = {
|
||||
let mut indices = self.round_robin_index.write().map_err(|e| e.to_string())?;
|
||||
indices
|
||||
.entry(type_key)
|
||||
.or_insert_with(|| AtomicUsize::new(0))
|
||||
.fetch_add(1, Ordering::SeqCst)
|
||||
};
|
||||
|
||||
// 选择 API Key
|
||||
let (selected_key, provider) = &keys[index % keys.len()];
|
||||
|
||||
// 解密并返回
|
||||
let decrypted = self.encryption.decrypt(&selected_key.api_key_encrypted)?;
|
||||
Ok(Some((selected_key.id.clone(), decrypted, provider.clone())))
|
||||
}
|
||||
|
||||
/// 记录 API Key 错误
|
||||
pub fn record_error(&self, db: &DbConnection, key_id: &str) -> Result<(), String> {
|
||||
let conn = db.lock().map_err(|e| e.to_string())?;
|
||||
|
||||
@@ -224,20 +224,83 @@ impl ProviderPoolService {
|
||||
) -> Result<Option<ProviderCredential>, String> {
|
||||
let pt: PoolProviderType = provider_type.parse().map_err(|e: String| e)?;
|
||||
let conn = db.lock().map_err(|e| e.to_string())?;
|
||||
let credentials = ProviderPoolDao::get_by_type(&conn, &pt).map_err(|e| e.to_string())?;
|
||||
|
||||
// 获取凭证,对于 Anthropic 类型,也查找 Claude 类型的凭证
|
||||
let mut credentials =
|
||||
ProviderPoolDao::get_by_type(&conn, &pt).map_err(|e| e.to_string())?;
|
||||
eprintln!(
|
||||
"[SELECT_CREDENTIAL] provider_type={}, pt={:?}, initial_count={}",
|
||||
provider_type,
|
||||
pt,
|
||||
credentials.len()
|
||||
);
|
||||
|
||||
// Anthropic 和 Claude 共享凭证(都使用 Anthropic API)
|
||||
if pt == PoolProviderType::Anthropic {
|
||||
let claude_creds = ProviderPoolDao::get_by_type(&conn, &PoolProviderType::Claude)
|
||||
.map_err(|e| e.to_string())?;
|
||||
eprintln!(
|
||||
"[SELECT_CREDENTIAL] Anthropic: adding {} Claude credentials",
|
||||
claude_creds.len()
|
||||
);
|
||||
credentials.extend(claude_creds);
|
||||
} else if pt == PoolProviderType::Claude {
|
||||
let anthropic_creds = ProviderPoolDao::get_by_type(&conn, &PoolProviderType::Anthropic)
|
||||
.map_err(|e| e.to_string())?;
|
||||
eprintln!(
|
||||
"[SELECT_CREDENTIAL] Claude: adding {} Anthropic credentials",
|
||||
anthropic_creds.len()
|
||||
);
|
||||
credentials.extend(anthropic_creds);
|
||||
}
|
||||
|
||||
drop(conn);
|
||||
|
||||
eprintln!(
|
||||
"[SELECT_CREDENTIAL] total_credentials={}, model={:?}",
|
||||
credentials.len(),
|
||||
model
|
||||
);
|
||||
|
||||
// 过滤可用的凭证
|
||||
let mut available: Vec<_> = credentials
|
||||
.into_iter()
|
||||
.filter(|c| c.is_available())
|
||||
.filter(|c| {
|
||||
let is_avail = c.is_available();
|
||||
eprintln!(
|
||||
"[SELECT_CREDENTIAL] credential {} (type={}) is_available={}",
|
||||
c.name.as_deref().unwrap_or("unnamed"),
|
||||
c.provider_type,
|
||||
is_avail
|
||||
);
|
||||
is_avail
|
||||
})
|
||||
.collect();
|
||||
|
||||
eprintln!(
|
||||
"[SELECT_CREDENTIAL] after is_available filter: {}",
|
||||
available.len()
|
||||
);
|
||||
|
||||
// 如果指定了模型,进一步过滤支持该模型的凭证
|
||||
if let Some(m) = model {
|
||||
available.retain(|c| c.supports_model(m));
|
||||
available.retain(|c| {
|
||||
let supports = c.supports_model(m);
|
||||
eprintln!(
|
||||
"[SELECT_CREDENTIAL] credential {} supports_model({})={}",
|
||||
c.name.as_deref().unwrap_or("unnamed"),
|
||||
m,
|
||||
supports
|
||||
);
|
||||
supports
|
||||
});
|
||||
}
|
||||
|
||||
eprintln!(
|
||||
"[SELECT_CREDENTIAL] final available count: {}",
|
||||
available.len()
|
||||
);
|
||||
|
||||
if available.is_empty() {
|
||||
return Ok(None);
|
||||
}
|
||||
@@ -821,6 +884,11 @@ impl ProviderPoolService {
|
||||
CredentialData::IFlowCookie { creds_file_path } => {
|
||||
self.check_iflow_cookie_health(creds_file_path, model).await
|
||||
}
|
||||
CredentialData::AnthropicKey { api_key, base_url } => {
|
||||
// Anthropic API Key 使用与 Claude API Key 相同的健康检查逻辑
|
||||
self.check_claude_health(api_key, base_url.as_deref(), model)
|
||||
.await
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
@@ -494,6 +494,17 @@ impl TokenCacheService {
|
||||
CredentialData::IFlowCookie { creds_file_path } => {
|
||||
self.refresh_iflow_cookie(creds_file_path).await
|
||||
}
|
||||
CredentialData::AnthropicKey { api_key, .. } => {
|
||||
// API Key 不需要刷新,直接返回
|
||||
Ok(CachedTokenInfo {
|
||||
access_token: Some(api_key.clone()),
|
||||
refresh_token: None,
|
||||
expiry_time: None, // 永不过期
|
||||
last_refresh: Some(Utc::now()),
|
||||
refresh_error_count: 0,
|
||||
last_refresh_error: None,
|
||||
})
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
@@ -1057,6 +1068,14 @@ impl TokenCacheService {
|
||||
last_refresh_error: None,
|
||||
})
|
||||
}
|
||||
CredentialData::AnthropicKey { api_key, .. } => Ok(CachedTokenInfo {
|
||||
access_token: Some(api_key.clone()),
|
||||
refresh_token: None,
|
||||
expiry_time: None,
|
||||
last_refresh: None,
|
||||
refresh_error_count: 0,
|
||||
last_refresh_error: None,
|
||||
}),
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
@@ -1,7 +1,7 @@
|
||||
{
|
||||
"$schema": "https://schema.tauri.app/config/2",
|
||||
"productName": "ProxyCast",
|
||||
"version": "0.27.0",
|
||||
"version": "0.28.0",
|
||||
"identifier": "com.proxycast.app",
|
||||
"build": {
|
||||
"beforeDevCommand": "npm run dev",
|
||||
|
||||
Reference in New Issue
Block a user