diff --git a/src-tauri/src/commands/model_test.rs b/src-tauri/src/commands/model_test.rs deleted file mode 100644 index a751875f2..000000000 --- a/src-tauri/src/commands/model_test.rs +++ /dev/null @@ -1,128 +0,0 @@ -//! 模型测试相关命令 - -use crate::app_config::AppType; -use crate::error::AppError; -use crate::services::model_test::{ - ModelTestConfig, ModelTestLog, ModelTestResult, ModelTestService, -}; -use crate::store::AppState; -use tauri::State; - -/// 测试单个供应商的模型可用性 -#[tauri::command] -pub async fn test_provider_model( - state: State<'_, AppState>, - app_type: AppType, - provider_id: String, -) -> Result { - // 获取测试配置 - let config = state.db.get_model_test_config()?; - - // 获取供应商 - let providers = state.db.get_all_providers(app_type.as_str())?; - let provider = providers - .get(&provider_id) - .ok_or_else(|| AppError::Message(format!("供应商 {provider_id} 不存在")))?; - - // 执行测试 - let result = ModelTestService::test_provider(&app_type, provider, &config).await?; - - // 记录日志 - let _ = state.db.save_model_test_log( - &provider_id, - &provider.name, - app_type.as_str(), - &result.model_used, - &config.test_prompt, - &result, - ); - - Ok(result) -} - -/// 批量测试所有供应商 -#[tauri::command] -pub async fn test_all_providers_model( - state: State<'_, AppState>, - app_type: AppType, - proxy_targets_only: bool, -) -> Result, AppError> { - let config = state.db.get_model_test_config()?; - let providers = state.db.get_all_providers(app_type.as_str())?; - - let mut results = Vec::new(); - - for (id, provider) in providers { - // 如果只测试代理目标,跳过非代理目标 - if proxy_targets_only && !provider.is_proxy_target.unwrap_or(false) { - continue; - } - - match ModelTestService::test_provider(&app_type, &provider, &config).await { - Ok(result) => { - // 记录日志 - let _ = state.db.save_model_test_log( - &id, - &provider.name, - app_type.as_str(), - &result.model_used, - &config.test_prompt, - &result, - ); - results.push((id, result)); - } - Err(e) => { - let error_result = ModelTestResult { - success: false, - message: e.to_string(), - response_time_ms: None, - http_status: None, - model_used: String::new(), - tested_at: chrono::Utc::now().timestamp(), - }; - results.push((id, error_result)); - } - } - } - - Ok(results) -} - -/// 获取模型测试配置 -#[tauri::command] -pub fn get_model_test_config(state: State<'_, AppState>) -> Result { - state.db.get_model_test_config() -} - -/// 保存模型测试配置 -#[tauri::command] -pub fn save_model_test_config( - state: State<'_, AppState>, - config: ModelTestConfig, -) -> Result<(), AppError> { - state.db.save_model_test_config(&config) -} - -/// 获取模型测试日志 -#[tauri::command] -pub fn get_model_test_logs( - state: State<'_, AppState>, - app_type: Option, - provider_id: Option, - limit: Option, -) -> Result, AppError> { - state.db.get_model_test_logs( - app_type.as_deref(), - provider_id.as_deref(), - limit.unwrap_or(50), - ) -} - -/// 清理旧的测试日志 -#[tauri::command] -pub fn cleanup_model_test_logs( - state: State<'_, AppState>, - keep_count: Option, -) -> Result { - state.db.cleanup_model_test_logs(keep_count.unwrap_or(100)) -} diff --git a/src-tauri/src/services/model_test.rs b/src-tauri/src/services/model_test.rs deleted file mode 100644 index a62dacf7b..000000000 --- a/src-tauri/src/services/model_test.rs +++ /dev/null @@ -1,510 +0,0 @@ -//! 模型测试服务 -//! -//! 提供独立的模型可用性测试功能,复用现有 Provider 适配器逻辑, -//! 但不影响正常代理数据流程。测试结果记录到独立的日志表。 - -use crate::app_config::AppType; -use crate::database::Database; -use crate::error::AppError; -use crate::provider::Provider; -use crate::proxy::providers::{get_adapter, AuthInfo, ProviderAdapter}; -use reqwest::Client; -use serde::{Deserialize, Serialize}; -use serde_json::{json, Value}; -use std::time::{Duration, Instant}; - -/// 模型测试配置 -#[derive(Debug, Clone, Serialize, Deserialize)] -#[serde(rename_all = "camelCase")] -pub struct ModelTestConfig { - /// 默认测试模型(Claude) - pub claude_model: String, - /// 默认测试模型(Codex/OpenAI) - pub codex_model: String, - /// 默认测试模型(Gemini) - pub gemini_model: String, - /// 测试提示词 - pub test_prompt: String, - /// 超时时间(秒) - pub timeout_secs: u64, -} - -impl Default for ModelTestConfig { - fn default() -> Self { - Self { - claude_model: "claude-haiku-4-5-20251001".to_string(), - codex_model: "gpt-5.1-low".to_string(), - gemini_model: "gemini-3-pro-low".to_string(), - test_prompt: "ping".to_string(), - timeout_secs: 15, - } - } -} - -/// 模型测试结果 -#[derive(Debug, Clone, Serialize, Deserialize)] -#[serde(rename_all = "camelCase")] -pub struct ModelTestResult { - pub success: bool, - pub message: String, - pub response_time_ms: Option, - pub http_status: Option, - pub model_used: String, - pub tested_at: i64, -} - -/// 模型测试日志记录 -#[derive(Debug, Clone, Serialize, Deserialize)] -#[serde(rename_all = "camelCase")] -pub struct ModelTestLog { - pub id: i64, - pub provider_id: String, - pub provider_name: String, - pub app_type: String, - pub model: String, - pub prompt: String, - pub success: bool, - pub message: String, - pub response_time_ms: Option, - pub http_status: Option, - pub tested_at: i64, -} - -/// 模型测试服务 -pub struct ModelTestService; - -impl ModelTestService { - /// 测试单个供应商的模型可用性 - pub async fn test_provider( - app_type: &AppType, - provider: &Provider, - config: &ModelTestConfig, - ) -> Result { - let start = Instant::now(); - let adapter = get_adapter(app_type); - - // 构建 HTTP 客户端(独立于代理服务) - let client = Client::builder() - .timeout(Duration::from_secs(config.timeout_secs)) - .build() - .map_err(|e| AppError::Message(format!("创建 HTTP 客户端失败: {e}")))?; - - // 根据 AppType 选择测试模型 - let model = match app_type { - AppType::Claude => &config.claude_model, - AppType::Codex => &config.codex_model, - AppType::Gemini => &config.gemini_model, - }; - - let result = match app_type { - AppType::Claude => { - Self::test_claude( - &client, - provider, - adapter.as_ref(), - model, - &config.test_prompt, - ) - .await - } - AppType::Codex => { - Self::test_codex( - &client, - provider, - adapter.as_ref(), - model, - &config.test_prompt, - ) - .await - } - AppType::Gemini => { - Self::test_gemini( - &client, - provider, - adapter.as_ref(), - model, - &config.test_prompt, - ) - .await - } - }; - - let response_time = start.elapsed().as_millis() as u64; - let tested_at = chrono::Utc::now().timestamp(); - - match result { - Ok((status, msg)) => Ok(ModelTestResult { - success: true, - message: msg, - response_time_ms: Some(response_time), - http_status: Some(status), - model_used: model.clone(), - tested_at, - }), - Err(e) => Ok(ModelTestResult { - success: false, - message: e.to_string(), - response_time_ms: Some(response_time), - http_status: None, - model_used: model.clone(), - tested_at, - }), - } - } - - /// 测试 Claude (Anthropic Messages API) - async fn test_claude( - client: &Client, - provider: &Provider, - adapter: &dyn ProviderAdapter, - model: &str, - prompt: &str, - ) -> Result<(u16, String), AppError> { - let base_url = adapter - .extract_base_url(provider) - .map_err(|e| AppError::Message(format!("提取 base_url 失败: {e}")))?; - - let auth = adapter - .extract_auth(provider) - .ok_or_else(|| AppError::Message("未找到 API Key".to_string()))?; - - // 智能拼接 URL,避免重复 /v1 - let base = base_url.trim_end_matches('/'); - let url = if base.ends_with("/v1") { - format!("{base}/messages") - } else { - format!("{base}/v1/messages") - }; - - let body = json!({ - "model": model, - "max_tokens": 1, - "messages": [{ - "role": "user", - "content": prompt - }] - }); - - let mut request = client.post(&url).json(&body); - request = Self::add_claude_auth(request, &auth); - - let response = request.send().await.map_err(|e| { - if e.is_timeout() { - AppError::Message("请求超时".to_string()) - } else if e.is_connect() { - AppError::Message(format!("连接失败: {e}")) - } else { - AppError::Message(e.to_string()) - } - })?; - - let status = response.status().as_u16(); - - if response.status().is_success() { - // 先获取文本,再尝试解析 JSON(兼容流式响应) - let text = response.text().await.unwrap_or_default(); - - // 尝试解析 JSON - if let Ok(data) = serde_json::from_str::(&text) { - if data.get("type").is_some() - || data.get("content").is_some() - || data.get("id").is_some() - { - return Ok((status, "模型测试成功".to_string())); - } - } - - // 即使无法解析 JSON,只要状态码是 200 就认为成功 - Ok((status, "模型测试成功".to_string())) - } else { - let error_text = response.text().await.unwrap_or_default(); - Err(AppError::Message(format!("HTTP {status}: {error_text}"))) - } - } - - /// 测试 Codex (OpenAI Chat Completions API) - async fn test_codex( - client: &Client, - provider: &Provider, - adapter: &dyn ProviderAdapter, - model: &str, - prompt: &str, - ) -> Result<(u16, String), AppError> { - let base_url = adapter - .extract_base_url(provider) - .map_err(|e| AppError::Message(format!("提取 base_url 失败: {e}")))?; - - let auth = adapter - .extract_auth(provider) - .ok_or_else(|| AppError::Message("未找到 API Key".to_string()))?; - - // 智能拼接 URL,避免重复 /v1 - let base = base_url.trim_end_matches('/'); - let url = if base.ends_with("/v1") { - format!("{base}/chat/completions") - } else { - format!("{base}/v1/chat/completions") - }; - - let body = json!({ - "model": model, - "messages": [{ - "role": "user", - "content": prompt - }], - "max_tokens": 1, - "stream": false - }); - - let request = client - .post(&url) - .header("Authorization", format!("Bearer {}", auth.api_key)) - .header("Content-Type", "application/json") - .json(&body); - - let response = request.send().await.map_err(|e| { - if e.is_timeout() { - AppError::Message("请求超时".to_string()) - } else if e.is_connect() { - AppError::Message(format!("连接失败: {e}")) - } else { - AppError::Message(e.to_string()) - } - })?; - - let status = response.status().as_u16(); - - if response.status().is_success() { - // 先获取文本,再尝试解析 JSON - let text = response.text().await.unwrap_or_default(); - - if let Ok(data) = serde_json::from_str::(&text) { - if data.get("choices").is_some() || data.get("id").is_some() { - return Ok((status, "模型测试成功".to_string())); - } - } - - // 即使无法解析 JSON,只要状态码是 200 就认为成功 - Ok((status, "模型测试成功".to_string())) - } else { - let error_text = response.text().await.unwrap_or_default(); - Err(AppError::Message(format!("HTTP {status}: {error_text}"))) - } - } - - /// 测试 Gemini (Google Generative AI API) - async fn test_gemini( - client: &Client, - provider: &Provider, - adapter: &dyn ProviderAdapter, - model: &str, - prompt: &str, - ) -> Result<(u16, String), AppError> { - let base_url = adapter - .extract_base_url(provider) - .map_err(|e| AppError::Message(format!("提取 base_url 失败: {e}")))?; - - let auth = adapter - .extract_auth(provider) - .ok_or_else(|| AppError::Message("未找到 API Key".to_string()))?; - - let url = format!( - "{}/v1beta/models/{}:generateContent?key={}", - base_url.trim_end_matches('/'), - model, - auth.api_key - ); - - let body = json!({ - "contents": [{ - "parts": [{ - "text": prompt - }] - }], - "generationConfig": { - "maxOutputTokens": 1 - } - }); - - let request = client - .post(&url) - .header("Content-Type", "application/json") - .json(&body); - - let response = request.send().await.map_err(|e| { - if e.is_timeout() { - AppError::Message("请求超时".to_string()) - } else if e.is_connect() { - AppError::Message(format!("连接失败: {e}")) - } else { - AppError::Message(e.to_string()) - } - })?; - - let status = response.status().as_u16(); - - if response.status().is_success() { - let data: Value = response - .json() - .await - .map_err(|e| AppError::Message(format!("解析响应失败: {e}")))?; - - if data.get("candidates").is_some() { - Ok((status, "模型测试成功".to_string())) - } else { - Err(AppError::Message("响应格式异常".to_string())) - } - } else { - let error_text = response.text().await.unwrap_or_default(); - Err(AppError::Message(format!("HTTP {status}: {error_text}"))) - } - } - - /// 添加 Claude 认证头 - fn add_claude_auth( - request: reqwest::RequestBuilder, - auth: &AuthInfo, - ) -> reqwest::RequestBuilder { - request - .header("x-api-key", &auth.api_key) - .header("anthropic-version", "2023-06-01") - .header("Content-Type", "application/json") - } -} - -// ===== 数据库操作 ===== - -impl Database { - /// 保存模型测试日志 - pub fn save_model_test_log( - &self, - provider_id: &str, - provider_name: &str, - app_type: &str, - model: &str, - prompt: &str, - result: &ModelTestResult, - ) -> Result { - let conn = self - .conn - .lock() - .map_err(|e| AppError::Database(format!("获取数据库连接失败: {e}")))?; - - conn.execute( - "INSERT INTO model_test_logs - (provider_id, provider_name, app_type, model, prompt, success, message, response_time_ms, http_status, tested_at) - VALUES (?1, ?2, ?3, ?4, ?5, ?6, ?7, ?8, ?9, ?10)", - rusqlite::params![ - provider_id, - provider_name, - app_type, - model, - prompt, - result.success, - result.message, - result.response_time_ms.map(|t| t as i64), - result.http_status.map(|s| s as i64), - result.tested_at, - ], - ) - .map_err(|e| AppError::Database(e.to_string()))?; - - Ok(conn.last_insert_rowid()) - } - - /// 获取模型测试日志 - pub fn get_model_test_logs( - &self, - app_type: Option<&str>, - provider_id: Option<&str>, - limit: u32, - ) -> Result, AppError> { - let conn = self - .conn - .lock() - .map_err(|e| AppError::Database(format!("获取数据库连接失败: {e}")))?; - - let mut sql = String::from( - "SELECT id, provider_id, provider_name, app_type, model, prompt, success, message, response_time_ms, http_status, tested_at - FROM model_test_logs WHERE 1=1" - ); - - let mut params: Vec> = Vec::new(); - - if let Some(at) = app_type { - sql.push_str(" AND app_type = ?"); - params.push(Box::new(at.to_string())); - } - - if let Some(pid) = provider_id { - sql.push_str(" AND provider_id = ?"); - params.push(Box::new(pid.to_string())); - } - - sql.push_str(" ORDER BY tested_at DESC LIMIT ?"); - params.push(Box::new(limit as i64)); - - let params_refs: Vec<&dyn rusqlite::ToSql> = params.iter().map(|p| p.as_ref()).collect(); - - let mut stmt = conn - .prepare(&sql) - .map_err(|e| AppError::Database(e.to_string()))?; - - let logs = stmt - .query_map(params_refs.as_slice(), |row| { - Ok(ModelTestLog { - id: row.get(0)?, - provider_id: row.get(1)?, - provider_name: row.get(2)?, - app_type: row.get(3)?, - model: row.get(4)?, - prompt: row.get(5)?, - success: row.get(6)?, - message: row.get(7)?, - response_time_ms: row.get(8)?, - http_status: row.get(9)?, - tested_at: row.get(10)?, - }) - }) - .map_err(|e| AppError::Database(e.to_string()))? - .collect::, _>>() - .map_err(|e| AppError::Database(e.to_string()))?; - - Ok(logs) - } - - /// 获取模型测试配置 - pub fn get_model_test_config(&self) -> Result { - match self.get_setting("model_test_config")? { - Some(json) => serde_json::from_str(&json) - .map_err(|e| AppError::Message(format!("解析模型测试配置失败: {e}"))), - None => Ok(ModelTestConfig::default()), - } - } - - /// 保存模型测试配置 - pub fn save_model_test_config(&self, config: &ModelTestConfig) -> Result<(), AppError> { - let json = serde_json::to_string(config) - .map_err(|e| AppError::Message(format!("序列化模型测试配置失败: {e}")))?; - self.set_setting("model_test_config", &json) - } - - /// 清理旧的测试日志(保留最近 N 条) - pub fn cleanup_model_test_logs(&self, keep_count: u32) -> Result { - let conn = self - .conn - .lock() - .map_err(|e| AppError::Database(format!("获取数据库连接失败: {e}")))?; - - let deleted = conn - .execute( - "DELETE FROM model_test_logs WHERE id NOT IN ( - SELECT id FROM model_test_logs ORDER BY tested_at DESC LIMIT ? - )", - rusqlite::params![keep_count as i64], - ) - .map_err(|e| AppError::Database(e.to_string()))?; - - Ok(deleted as u64) - } -} diff --git a/src-tauri/src/services/stream_check.rs b/src-tauri/src/services/stream_check.rs index eaf2517b8..36004344f 100644 --- a/src-tauri/src/services/stream_check.rs +++ b/src-tauri/src/services/stream_check.rs @@ -29,6 +29,12 @@ pub struct StreamCheckConfig { pub timeout_secs: u64, pub max_retries: u32, pub degraded_threshold_ms: u64, + /// Claude 测试模型 + pub claude_model: String, + /// Codex 测试模型 + pub codex_model: String, + /// Gemini 测试模型 + pub gemini_model: String, } impl Default for StreamCheckConfig { @@ -37,6 +43,9 @@ impl Default for StreamCheckConfig { timeout_secs: 45, max_retries: 2, degraded_threshold_ms: 6000, + claude_model: "claude-haiku-4-5-20251001".to_string(), + codex_model: "gpt-5.1-codex@low".to_string(), + gemini_model: "gemini-3-pro-preview".to_string(), } } } @@ -133,9 +142,15 @@ impl StreamCheckService { .map_err(|e| AppError::Message(format!("创建客户端失败: {e}")))?; let result = match app_type { - AppType::Claude => Self::check_claude_stream(&client, &base_url, &auth).await, - AppType::Codex => Self::check_codex_stream(&client, &base_url, &auth).await, - AppType::Gemini => Self::check_gemini_stream(&client, &base_url, &auth).await, + AppType::Claude => { + Self::check_claude_stream(&client, &base_url, &auth, &config.claude_model).await + } + AppType::Codex => { + Self::check_codex_stream(&client, &base_url, &auth, &config.codex_model).await + } + AppType::Gemini => { + Self::check_gemini_stream(&client, &base_url, &auth, &config.gemini_model).await + } }; let response_time = start.elapsed().as_millis() as u64; @@ -174,6 +189,7 @@ impl StreamCheckService { client: &Client, base_url: &str, auth: &AuthInfo, + model: &str, ) -> Result<(u16, String), AppError> { let base = base_url.trim_end_matches('/'); let url = if base.ends_with("/v1") { @@ -182,8 +198,6 @@ impl StreamCheckService { format!("{base}/v1/messages") }; - let model = "claude-3-5-haiku-latest"; - let body = json!({ "model": model, "max_tokens": 1, @@ -225,6 +239,7 @@ impl StreamCheckService { client: &Client, base_url: &str, auth: &AuthInfo, + model: &str, ) -> Result<(u16, String), AppError> { let base = base_url.trim_end_matches('/'); let url = if base.ends_with("/v1") { @@ -233,10 +248,11 @@ impl StreamCheckService { format!("{base}/v1/chat/completions") }; - let model = "gpt-4o-mini"; + // 解析模型名和推理等级 (支持 model@level 或 model#level 格式) + let (actual_model, reasoning_effort) = Self::parse_model_with_effort(model); - let body = json!({ - "model": model, + let mut body = json!({ + "model": actual_model, "messages": [ { "role": "system", "content": "" }, { "role": "assistant", "content": "" }, @@ -247,6 +263,11 @@ impl StreamCheckService { "stream": true }); + // 如果是推理模型,添加 reasoning_effort + if let Some(effort) = reasoning_effort { + body["reasoning_effort"] = json!(effort); + } + let response = client .post(&url) .header("Authorization", format!("Bearer {}", auth.api_key)) @@ -279,12 +300,11 @@ impl StreamCheckService { client: &Client, base_url: &str, auth: &AuthInfo, + model: &str, ) -> Result<(u16, String), AppError> { let base = base_url.trim_end_matches('/'); let url = format!("{base}/v1/chat/completions"); - let model = "gemini-1.5-flash"; - let body = json!({ "model": model, "messages": [{ "role": "user", "content": "hi" }], @@ -328,6 +348,20 @@ impl StreamCheckService { } } + /// 解析模型名和推理等级 (支持 model@level 或 model#level 格式) + /// 返回 (实际模型名, Option<推理等级>) + fn parse_model_with_effort(model: &str) -> (String, Option) { + // 查找 @ 或 # 分隔符 + if let Some(pos) = model.find('@').or_else(|| model.find('#')) { + let actual_model = model[..pos].to_string(); + let effort = model[pos + 1..].to_string(); + if !effort.is_empty() { + return (actual_model, Some(effort)); + } + } + (model.to_string(), None) + } + fn should_retry(msg: &str) -> bool { let lower = msg.to_lowercase(); lower.contains("timeout") @@ -381,4 +415,22 @@ mod tests { assert_eq!(config.max_retries, 2); assert_eq!(config.degraded_threshold_ms, 6000); } + + #[test] + fn test_parse_model_with_effort() { + // 带 @ 分隔符 + let (model, effort) = StreamCheckService::parse_model_with_effort("gpt-5.1-codex@low"); + assert_eq!(model, "gpt-5.1-codex"); + assert_eq!(effort, Some("low".to_string())); + + // 带 # 分隔符 + let (model, effort) = StreamCheckService::parse_model_with_effort("o1-preview#high"); + assert_eq!(model, "o1-preview"); + assert_eq!(effort, Some("high".to_string())); + + // 无分隔符 + let (model, effort) = StreamCheckService::parse_model_with_effort("gpt-4o-mini"); + assert_eq!(model, "gpt-4o-mini"); + assert_eq!(effort, None); + } }