diff --git a/src-tauri/src/converter/openai_to_antigravity.rs b/src-tauri/src/converter/openai_to_antigravity.rs index 5cfce66ac..fdf122c96 100644 --- a/src-tauri/src/converter/openai_to_antigravity.rs +++ b/src-tauri/src/converter/openai_to_antigravity.rs @@ -15,6 +15,7 @@ //! - 2025-12-28: 修复请求格式,对齐 CLIProxyAPI 实现 use crate::models::openai::*; +use crate::session::{get_thought_signature, SessionManager}; use serde::{Deserialize, Serialize}; use uuid::Uuid; @@ -177,8 +178,10 @@ fn generate_request_id() -> String { format!("agent-{}", Uuid::new_v4()) } -/// 生成随机会话 ID -fn generate_session_id() -> String { +/// 生成随机会话 ID(兜底方案) +/// +/// 当无法从请求中提取稳定的会话 ID 时使用 +fn generate_random_session_id() -> String { let uuid = Uuid::new_v4(); let bytes = uuid.as_bytes(); let n: u64 = u64::from_le_bytes([ @@ -216,14 +219,28 @@ fn default_safety_settings() -> Vec { /// 模型名称映射 fn model_mapping(model: &str) -> &str { match model { + // Claude 模型映射 "claude-sonnet-4-5-thinking" => "claude-sonnet-4-5", "claude-opus-4-5" => "claude-opus-4-5-thinking", + + // Gemini 模型映射 "gemini-2.5-flash-thinking" => "gemini-2.5-flash", "gemini-2.5-computer-use-preview-10-2025" => "rev19-uic3-1p", + + // Gemini 3 preview 模型映射到正式名称 "gemini-3-pro-image-preview" => "gemini-3-pro-image", + "gemini-3-flash-preview" => "gemini-3-flash", "gemini-3-pro-preview" => "gemini-3-pro-high", + + // Gemini 2.5 preview 模型映射 + "gemini-2.5-flash-preview" => "gemini-2.5-flash", + + // Claude via Antigravity 映射 "gemini-claude-sonnet-4-5" => "claude-sonnet-4-5", "gemini-claude-sonnet-4-5-thinking" => "claude-sonnet-4-5-thinking", + "gemini-claude-opus-4-5-thinking" => "claude-opus-4-5-thinking", + + // 其他模型直接透传 _ => model, } } @@ -392,6 +409,14 @@ pub fn convert_openai_to_antigravity_with_context( if let Some(tool_calls) = &msg.tool_calls { let mut function_ids: Vec = Vec::new(); + // 获取全局存储的 thoughtSignature(如果有) + let global_sig = get_thought_signature(); + let thought_sig = global_sig.unwrap_or_else(|| { + // 如果没有缓存的签名,使用跳过验证的标记 + // 注意:Vertex AI 不接受此标记,但 Cloud Code API 接受 + GEMINI_CLI_FUNCTION_THOUGHT_SIGNATURE.to_string() + }); + for tc in tool_calls { let args: serde_json::Value = serde_json::from_str(&tc.function.arguments) .unwrap_or(serde_json::json!({})); @@ -405,9 +430,7 @@ pub fn convert_openai_to_antigravity_with_context( args, // 直接使用 args,不要包装 }), function_response: None, - thought_signature: Some( - GEMINI_CLI_FUNCTION_THOUGHT_SIGNATURE.to_string(), - ), + thought_signature: Some(thought_sig.clone()), }); function_ids.push(tc.id.clone()); @@ -654,13 +677,28 @@ pub fn convert_openai_to_antigravity_with_context( } }); + // 构建 toolConfig(如果有工具定义) + let tool_config: Option = if tools.is_some() { + Some(serde_json::json!({ + "functionCallingConfig": { + "mode": "AUTO" + } + })) + } else { + None + }; + + // 使用 SessionManager 生成稳定的会话 ID + let session_id = SessionManager::extract_session_id(request); + eprintln!("[CONVERT] 生成的稳定 SessionId: {}", session_id); + let inner = AntigravityRequestInner { contents, system_instruction, generation_config: Some(generation_config), tools, - tool_config: None, - session_id: Some(generate_session_id()), + tool_config, + session_id: Some(session_id), safety_settings: Some(default_safety_settings()), }; diff --git a/src-tauri/src/lib.rs b/src-tauri/src/lib.rs index 96cc44021..faa9edd98 100644 --- a/src-tauri/src/lib.rs +++ b/src-tauri/src/lib.rs @@ -25,6 +25,7 @@ pub mod resilience; pub mod router; pub mod screenshot; pub mod services; +pub mod session; pub mod session_files; pub mod stream; pub mod streaming; diff --git a/src-tauri/src/providers/antigravity.rs b/src-tauri/src/providers/antigravity.rs index bf2886e26..5c36a2a0b 100644 --- a/src-tauri/src/providers/antigravity.rs +++ b/src-tauri/src/providers/antigravity.rs @@ -214,19 +214,31 @@ pub const ANTIGRAVITY_MODELS_FALLBACK: &[&str] = &[ "gemini-3-pro-preview", "gemini-3-flash-preview", "gemini-2.5-flash-preview", + "gemini-2.5-flash", + "gemini-2.5-pro", + "gemini-3-flash", + "gemini-3-pro-high", + "gemini-3-pro-low", "gemini-claude-sonnet-4-5", "gemini-claude-sonnet-4-5-thinking", "gemini-claude-opus-4-5-thinking", + "claude-sonnet-4-5", + "claude-sonnet-4-5-thinking", + "claude-opus-4-5-thinking", ]; /// 模型别名映射(fallback,当无法从 models 仓库获取时使用) /// 格式:用户友好名称 -> 内部 API 名称 pub const ANTIGRAVITY_ALIAS_FALLBACK: &[(&str, &str)] = &[ + // 需要映射的模型 ("gemini-2.5-computer-use-preview-10-2025", "rev19-uic3-1p"), ("gemini-3-pro-image-preview", "gemini-3-pro-image"), - ("gemini-3-pro-preview", "gemini-3-pro-high"), + // Gemini 3 preview 模型映射到正式名称 ("gemini-3-flash-preview", "gemini-3-flash"), + ("gemini-3-pro-preview", "gemini-3-pro-high"), + // Gemini 2.5 preview 模型映射 ("gemini-2.5-flash-preview", "gemini-2.5-flash"), + // Claude via Antigravity ("gemini-claude-sonnet-4-5", "claude-sonnet-4-5"), ( "gemini-claude-sonnet-4-5-thinking", diff --git a/src-tauri/src/server/handlers/provider_calls.rs b/src-tauri/src/server/handlers/provider_calls.rs index ad258988c..1c2c24b97 100644 --- a/src-tauri/src/server/handlers/provider_calls.rs +++ b/src-tauri/src/server/handlers/provider_calls.rs @@ -67,6 +67,7 @@ use crate::server_utils::{ build_anthropic_response, build_anthropic_stream_response, build_error_response, build_error_response_with_status, parse_cw_response, safe_truncate, CWParsedResponse, }; +use crate::session::store_thought_signature; use crate::stream::{PipelineConfig, StreamPipeline}; use crate::streaming::traits::StreamingProvider; use crate::streaming::{ @@ -3112,6 +3113,21 @@ fn extract_content_from_json(json: &serde_json::Value) -> Option<(String, Vec<(S .and_then(|t| t.as_bool()) .unwrap_or(false); + // 捕获 thoughtSignature 到全局存储(用于后续请求) + if let Some(sig) = part + .get("thoughtSignature") + .or_else(|| part.get("thought_signature")) + .and_then(|s| s.as_str()) + { + if !sig.is_empty() { + eprintln!( + "[ANTIGRAVITY_PARSE] 捕获 thoughtSignature (长度: {})", + sig.len() + ); + store_thought_signature(sig); + } + } + // 跳过纯 thoughtSignature 部分 let has_thought_signature = part .get("thoughtSignature") diff --git a/src-tauri/src/session/mod.rs b/src-tauri/src/session/mod.rs new file mode 100644 index 000000000..652852a98 --- /dev/null +++ b/src-tauri/src/session/mod.rs @@ -0,0 +1,25 @@ +//! 会话管理模块 +//! +//! 提供以下功能: +//! - 稳定的 SessionId 生成(基于请求内容哈希) +//! - thoughtSignature 全局缓存 +//! - 会话粘性管理(会话与账号映射) +//! - 调度模式配置 +//! - 增强的限流处理(Duration 解析、指数退避) + +mod rate_limit; +mod session_manager; +mod signature_store; +mod sticky_config; +mod sticky_manager; + +pub use rate_limit::{ + extract_retry_delay, parse_duration_string, RateLimitReason, RateLimitRecord, RateLimitTracker, +}; +pub use session_manager::SessionManager; +pub use signature_store::{ + clear_thought_signature, get_thought_signature, has_valid_signature, store_thought_signature, + take_thought_signature, +}; +pub use sticky_config::{SchedulingMode, StickySessionConfig}; +pub use sticky_manager::{AccountInfo, StickySessionManager}; diff --git a/src-tauri/src/session/rate_limit.rs b/src-tauri/src/session/rate_limit.rs new file mode 100644 index 000000000..75c76fd74 --- /dev/null +++ b/src-tauri/src/session/rate_limit.rs @@ -0,0 +1,450 @@ +//! 增强的限流处理模块 +//! +//! 提供以下功能: +//! - Duration 字符串解析(如 "1.5s", "1h16m0.667s") +//! - 指数退避策略 +//! - 账号级别和模型级别限流 +//! - 连续失败计数 + +use chrono::{DateTime, Duration, Utc}; +use dashmap::DashMap; +use serde::{Deserialize, Serialize}; +use std::sync::atomic::{AtomicU32, Ordering}; + +/// 限流原因类型 +#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)] +pub enum RateLimitReason { + /// 配额耗尽 + QuotaExhausted, + /// 速率限制 + RateLimitExceeded, + /// 模型容量耗尽 + ModelCapacityExhausted, + /// 服务器错误 (5xx) + ServerError, + /// 未知原因 + Unknown, +} + +impl std::fmt::Display for RateLimitReason { + fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + match self { + Self::QuotaExhausted => write!(f, "QuotaExhausted"), + Self::RateLimitExceeded => write!(f, "RateLimitExceeded"), + Self::ModelCapacityExhausted => write!(f, "ModelCapacityExhausted"), + Self::ServerError => write!(f, "ServerError"), + Self::Unknown => write!(f, "Unknown"), + } + } +} + +/// 限流记录 +#[derive(Debug, Clone, Serialize, Deserialize)] +pub struct RateLimitRecord { + /// 账号/凭证 ID + pub account_id: String, + /// 限流原因 + pub reason: RateLimitReason, + /// 限流开始时间 + pub started_at: DateTime, + /// 限流结束时间(预计) + pub reset_at: DateTime, + /// 连续失败次数 + pub consecutive_failures: u32, + /// 模型名称(如果是模型级别限流) + pub model: Option, +} + +/// 限流追踪器 +#[derive(Debug)] +pub struct RateLimitTracker { + /// 账号级别限流记录 + account_limits: DashMap, + /// 模型级别限流记录 (account_id:model -> record) + model_limits: DashMap, + /// 连续失败计数 (account_id -> count) + failure_counts: DashMap, + /// 基础退避时间(秒) + base_backoff_seconds: u64, + /// 最大退避时间(秒) + max_backoff_seconds: u64, +} + +impl Default for RateLimitTracker { + fn default() -> Self { + Self::new(5, 300) // 默认 5 秒基础退避,最大 5 分钟 + } +} + +impl RateLimitTracker { + /// 创建新的限流追踪器 + pub fn new(base_backoff_seconds: u64, max_backoff_seconds: u64) -> Self { + Self { + account_limits: DashMap::new(), + model_limits: DashMap::new(), + failure_counts: DashMap::new(), + base_backoff_seconds, + max_backoff_seconds, + } + } + + /// 标记账号限流 + pub fn mark_rate_limited( + &self, + account_id: &str, + reason: RateLimitReason, + retry_after: Option, + model: Option<&str>, + ) -> RateLimitRecord { + let now = Utc::now(); + + // 增加连续失败计数 + let failures = self + .failure_counts + .entry(account_id.to_string()) + .or_insert_with(|| AtomicU32::new(0)); + let failure_count = failures.fetch_add(1, Ordering::SeqCst) + 1; + + // 计算退避时间 + let backoff = if let Some(retry) = retry_after { + retry + } else { + self.calculate_exponential_backoff(failure_count) + }; + + let reset_at = now + backoff; + + let record = RateLimitRecord { + account_id: account_id.to_string(), + reason, + started_at: now, + reset_at, + consecutive_failures: failure_count, + model: model.map(|s| s.to_string()), + }; + + // 根据是否有模型信息决定存储位置 + if let Some(m) = model { + let key = format!("{}:{}", account_id, m); + self.model_limits.insert(key, record.clone()); + } else { + self.account_limits + .insert(account_id.to_string(), record.clone()); + } + + tracing::warn!( + account_id = %account_id, + reason = %reason, + reset_at = %reset_at, + failures = failure_count, + model = ?model, + "账号被限流" + ); + + record + } + + /// 计算指数退避时间 + fn calculate_exponential_backoff(&self, failure_count: u32) -> Duration { + // 指数退避: base * 2^(failures-1),但不超过最大值 + let exponent = (failure_count - 1).min(10); // 防止溢出 + let backoff_secs = self.base_backoff_seconds * (1 << exponent); + let capped_secs = backoff_secs.min(self.max_backoff_seconds); + Duration::seconds(capped_secs as i64) + } + + /// 检查账号是否被限流 + pub fn is_rate_limited(&self, account_id: &str) -> bool { + self.get_remaining_wait(account_id) > 0 + } + + /// 检查特定模型是否被限流 + pub fn is_model_rate_limited(&self, account_id: &str, model: &str) -> bool { + let key = format!("{}:{}", account_id, model); + if let Some(record) = self.model_limits.get(&key) { + return Utc::now() < record.reset_at; + } + false + } + + /// 获取剩余等待时间(秒) + pub fn get_remaining_wait(&self, account_id: &str) -> i64 { + if let Some(record) = self.account_limits.get(account_id) { + let remaining = (record.reset_at - Utc::now()).num_seconds(); + if remaining > 0 { + return remaining; + } + } + 0 + } + + /// 获取模型的剩余等待时间(秒) + pub fn get_model_remaining_wait(&self, account_id: &str, model: &str) -> i64 { + let key = format!("{}:{}", account_id, model); + if let Some(record) = self.model_limits.get(&key) { + let remaining = (record.reset_at - Utc::now()).num_seconds(); + if remaining > 0 { + return remaining; + } + } + 0 + } + + /// 清除账号的限流状态(成功请求后调用) + pub fn clear_rate_limit(&self, account_id: &str) { + self.account_limits.remove(account_id); + // 重置连续失败计数 + if let Some(counter) = self.failure_counts.get(account_id) { + counter.store(0, Ordering::SeqCst); + } + } + + /// 清除模型的限流状态 + pub fn clear_model_rate_limit(&self, account_id: &str, model: &str) { + let key = format!("{}:{}", account_id, model); + self.model_limits.remove(&key); + } + + /// 清理过期的限流记录 + pub fn cleanup_expired(&self) { + let now = Utc::now(); + + // 清理账号级别限流 + self.account_limits + .retain(|_, record| record.reset_at > now); + + // 清理模型级别限流 + self.model_limits.retain(|_, record| record.reset_at > now); + } + + /// 获取所有被限流的账号 + pub fn get_rate_limited_accounts(&self) -> Vec { + let now = Utc::now(); + self.account_limits + .iter() + .filter(|entry| entry.value().reset_at > now) + .map(|entry| entry.key().clone()) + .collect() + } +} + +/// 解析 Duration 字符串 +/// +/// 支持格式: +/// - "1.5s" -> 1.5 秒 +/// - "1h16m0.667s" -> 1 小时 16 分钟 0.667 秒 +/// - "30m" -> 30 分钟 +/// - "2h" -> 2 小时 +/// +/// # 参数 +/// - `s`: Duration 字符串 +/// +/// # 返回 +/// 解析后的 Duration,如果解析失败返回 None +pub fn parse_duration_string(s: &str) -> Option { + let s = s.trim(); + if s.is_empty() { + return None; + } + + let mut total_millis: i64 = 0; + let mut current_num = String::new(); + let mut chars = s.chars().peekable(); + + while let Some(c) = chars.next() { + if c.is_ascii_digit() || c == '.' { + current_num.push(c); + } else { + if current_num.is_empty() { + continue; + } + + let num: f64 = current_num.parse().ok()?; + current_num.clear(); + + match c { + 'h' => total_millis += (num * 3600.0 * 1000.0) as i64, + 'm' => { + // 检查是否是 "ms" + if chars.peek() == Some(&'s') { + chars.next(); + total_millis += num as i64; + } else { + total_millis += (num * 60.0 * 1000.0) as i64; + } + } + 's' => total_millis += (num * 1000.0) as i64, + _ => return None, + } + } + } + + // 处理末尾没有单位的数字(默认为秒) + if !current_num.is_empty() { + let num: f64 = current_num.parse().ok()?; + total_millis += (num * 1000.0) as i64; + } + + if total_millis > 0 { + Some(Duration::milliseconds(total_millis)) + } else { + None + } +} + +/// 从 429 响应中提取重试延迟 +/// +/// 尝试从以下位置提取: +/// 1. Retry-After 头 +/// 2. 响应体中的 retryDelay 字段 +/// 3. 响应体中的 quotaResetDelay 字段 +/// +/// # 参数 +/// - `headers`: HTTP 响应头 +/// - `body`: 响应体 JSON +/// +/// # 返回 +/// 解析后的 Duration,如果无法提取返回 None +pub fn extract_retry_delay( + headers: Option<&reqwest::header::HeaderMap>, + body: Option<&serde_json::Value>, +) -> Option { + // 1. 尝试从 Retry-After 头提取 + if let Some(hdrs) = headers { + if let Some(retry_after) = hdrs.get("retry-after").and_then(|v| v.to_str().ok()) { + // Retry-After 可以是秒数或 HTTP 日期 + if let Ok(secs) = retry_after.parse::() { + return Some(Duration::seconds(secs)); + } + // 尝试解析为 Duration 字符串 + if let Some(d) = parse_duration_string(retry_after) { + return d.into(); + } + } + } + + // 2. 尝试从响应体提取 + if let Some(json) = body { + // 尝试 error.details[].retryDelay + if let Some(details) = json + .get("error") + .and_then(|e| e.get("details")) + .and_then(|d| d.as_array()) + { + for detail in details { + if let Some(retry_delay) = detail.get("retryDelay").and_then(|r| r.as_str()) { + if let Some(d) = parse_duration_string(retry_delay) { + return Some(d); + } + } + if let Some(quota_reset) = detail.get("quotaResetDelay").and_then(|r| r.as_str()) { + if let Some(d) = parse_duration_string(quota_reset) { + return Some(d); + } + } + } + } + + // 尝试顶层 retryDelay + if let Some(retry_delay) = json.get("retryDelay").and_then(|r| r.as_str()) { + if let Some(d) = parse_duration_string(retry_delay) { + return Some(d); + } + } + } + + None +} + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn test_parse_duration_string() { + // 秒 + assert_eq!( + parse_duration_string("1.5s"), + Some(Duration::milliseconds(1500)) + ); + assert_eq!( + parse_duration_string("30s"), + Some(Duration::milliseconds(30000)) + ); + + // 分钟 + assert_eq!( + parse_duration_string("5m"), + Some(Duration::milliseconds(300000)) + ); + + // 小时 + assert_eq!( + parse_duration_string("2h"), + Some(Duration::milliseconds(7200000)) + ); + + // 复合格式 + assert_eq!( + parse_duration_string("1h16m0.667s"), + Some(Duration::milliseconds(4560667)) + ); + + // 毫秒 + assert_eq!( + parse_duration_string("500ms"), + Some(Duration::milliseconds(500)) + ); + + // 无效输入 + assert_eq!(parse_duration_string(""), None); + assert_eq!(parse_duration_string("invalid"), None); + } + + #[test] + fn test_exponential_backoff() { + let tracker = RateLimitTracker::new(5, 300); + + // 第一次失败: 5 秒 + assert_eq!( + tracker.calculate_exponential_backoff(1), + Duration::seconds(5) + ); + + // 第二次失败: 10 秒 + assert_eq!( + tracker.calculate_exponential_backoff(2), + Duration::seconds(10) + ); + + // 第三次失败: 20 秒 + assert_eq!( + tracker.calculate_exponential_backoff(3), + Duration::seconds(20) + ); + + // 第七次失败: 320 秒,但被限制为 300 秒 + assert_eq!( + tracker.calculate_exponential_backoff(7), + Duration::seconds(300) + ); + } + + #[test] + fn test_rate_limit_tracker() { + let tracker = RateLimitTracker::new(5, 300); + + // 初始状态不应该被限流 + assert!(!tracker.is_rate_limited("account1")); + + // 标记限流 + tracker.mark_rate_limited("account1", RateLimitReason::QuotaExhausted, None, None); + + // 应该被限流 + assert!(tracker.is_rate_limited("account1")); + + // 清除限流 + tracker.clear_rate_limit("account1"); + assert!(!tracker.is_rate_limited("account1")); + } +} diff --git a/src-tauri/src/session/session_manager.rs b/src-tauri/src/session/session_manager.rs new file mode 100644 index 000000000..a9424cde2 --- /dev/null +++ b/src-tauri/src/session/session_manager.rs @@ -0,0 +1,239 @@ +//! 会话管理器 +//! +//! 根据请求内容生成稳定的会话指纹(Session Fingerprint), +//! 用于实现会话粘性和 Prompt Caching 优化。 + +use crate::models::openai::ChatCompletionRequest; +use sha2::{Digest, Sha256}; + +/// 会话管理器 +pub struct SessionManager; + +impl SessionManager { + /// 根据 OpenAI 请求生成稳定的会话指纹 + /// + /// 策略: + /// 基于第一条用户消息内容 + 模型名称生成 SHA256 哈希 + /// + /// # 参数 + /// - `request`: OpenAI 格式的请求 + /// + /// # 返回 + /// 稳定的会话 ID,格式为 `sid-{hash前16位}` + pub fn extract_session_id(request: &ChatCompletionRequest) -> String { + // 智能内容指纹 (SHA256) + let mut hasher = Sha256::new(); + + // 混入模型名称增加区分度 + hasher.update(request.model.as_bytes()); + + let mut content_found = false; + for msg in &request.messages { + if msg.role != "user" { + continue; + } + + let text = msg.get_content_text(); + let clean_text = text.trim(); + + // 跳过过短的消息(可能是探测消息)或含有系统标签的消息 + if clean_text.len() > 10 && !clean_text.contains("") { + hasher.update(clean_text.as_bytes()); + content_found = true; + break; // 只取第一条关键消息作为锚点 + } + } + + if !content_found { + // 如果没找到有意义的内容,退化为对最后一条消息进行哈希 + if let Some(last_msg) = request.messages.last() { + hasher.update(last_msg.get_content_text().as_bytes()); + } + } + + let hash = format!("{:x}", hasher.finalize()); + let sid = format!("sid-{}", &hash[..16]); + + tracing::debug!( + "[SessionManager] Generated fingerprint: {} for model {}", + sid, + request.model + ); + sid + } + + /// 根据 JSON 请求生成稳定的会话指纹 + /// + /// 用于处理原始 JSON 格式的请求 + pub fn extract_session_id_from_json(request: &serde_json::Value, model: &str) -> String { + let mut hasher = Sha256::new(); + hasher.update(model.as_bytes()); + + let mut content_found = false; + + // 尝试从 messages 数组中提取用户消息 + if let Some(messages) = request.get("messages").and_then(|m| m.as_array()) { + for msg in messages { + if msg.get("role").and_then(|r| r.as_str()) != Some("user") { + continue; + } + + // 提取文本内容 + let text = if let Some(content) = msg.get("content") { + if let Some(s) = content.as_str() { + s.to_string() + } else if let Some(arr) = content.as_array() { + arr.iter() + .filter_map(|part| { + if part.get("type").and_then(|t| t.as_str()) == Some("text") { + part.get("text").and_then(|t| t.as_str()) + } else { + None + } + }) + .collect::>() + .join(" ") + } else { + String::new() + } + } else { + String::new() + }; + + let clean_text = text.trim(); + if clean_text.len() > 10 && !clean_text.contains("") { + hasher.update(clean_text.as_bytes()); + content_found = true; + break; + } + } + } + + // 尝试从 Gemini 格式的 contents 数组中提取 + if !content_found { + if let Some(contents) = request.get("contents").and_then(|c| c.as_array()) { + for content in contents { + if content.get("role").and_then(|r| r.as_str()) != Some("user") { + continue; + } + + if let Some(parts) = content.get("parts").and_then(|p| p.as_array()) { + let text: String = parts + .iter() + .filter_map(|part| part.get("text").and_then(|t| t.as_str())) + .collect::>() + .join(" "); + + let clean_text = text.trim(); + if clean_text.len() > 10 && !clean_text.contains("") { + hasher.update(clean_text.as_bytes()); + content_found = true; + break; + } + } + } + } + } + + if !content_found { + // 兜底:对整个请求进行摘要 + hasher.update(request.to_string().as_bytes()); + } + + let hash = format!("{:x}", hasher.finalize()); + let sid = format!("sid-{}", &hash[..16]); + + tracing::debug!( + "[SessionManager] Generated fingerprint from JSON: {} for model {}", + sid, + model + ); + sid + } +} + +#[cfg(test)] +mod tests { + use super::*; + use crate::models::openai::{ChatCompletionRequest, ChatMessage}; + + #[test] + fn test_session_id_stability() { + let request = ChatCompletionRequest { + model: "gpt-4".to_string(), + messages: vec![ChatMessage { + role: "user".to_string(), + content: Some(crate::models::openai::MessageContent::Text( + "Hello, how are you?".to_string(), + )), + tool_calls: None, + tool_call_id: None, + }], + temperature: None, + max_tokens: None, + top_p: None, + stream: false, + tools: None, + tool_choice: None, + reasoning_effort: None, + }; + + let sid1 = SessionManager::extract_session_id(&request); + let sid2 = SessionManager::extract_session_id(&request); + + assert_eq!(sid1, sid2, "Same request should generate same session ID"); + assert!( + sid1.starts_with("sid-"), + "Session ID should start with 'sid-'" + ); + } + + #[test] + fn test_different_content_different_sid() { + let request1 = ChatCompletionRequest { + model: "gpt-4".to_string(), + messages: vec![ChatMessage { + role: "user".to_string(), + content: Some(crate::models::openai::MessageContent::Text( + "Hello, how are you?".to_string(), + )), + tool_calls: None, + tool_call_id: None, + }], + temperature: None, + max_tokens: None, + top_p: None, + stream: false, + tools: None, + tool_choice: None, + reasoning_effort: None, + }; + + let request2 = ChatCompletionRequest { + model: "gpt-4".to_string(), + messages: vec![ChatMessage { + role: "user".to_string(), + content: Some(crate::models::openai::MessageContent::Text( + "What is the weather today?".to_string(), + )), + tool_calls: None, + tool_call_id: None, + }], + temperature: None, + max_tokens: None, + top_p: None, + stream: false, + tools: None, + tool_choice: None, + reasoning_effort: None, + }; + + let sid1 = SessionManager::extract_session_id(&request1); + let sid2 = SessionManager::extract_session_id(&request2); + + assert_ne!( + sid1, sid2, + "Different content should generate different session IDs" + ); + } +} diff --git a/src-tauri/src/session/signature_store.rs b/src-tauri/src/session/signature_store.rs new file mode 100644 index 000000000..91e7a7f88 --- /dev/null +++ b/src-tauri/src/session/signature_store.rs @@ -0,0 +1,117 @@ +//! thoughtSignature 全局存储 +//! +//! 用于在流式响应中捕获 thoughtSignature,并在后续请求中注入。 +//! 这对于 Gemini 3 Pro 的 Tool Use 功能至关重要。 + +use std::sync::RwLock; + +/// 最小有效签名长度 +const MIN_SIGNATURE_LENGTH: usize = 50; + +/// 全局 thoughtSignature 存储 +static THOUGHT_SIGNATURE: RwLock> = RwLock::new(None); + +/// 存储 thoughtSignature 到全局存储 +/// +/// 只有当新签名长度大于等于最小长度时才会存储。 +/// 如果已有签名,只有当新签名更长时才会替换。 +/// +/// # 参数 +/// - `sig`: 要存储的签名 +pub fn store_thought_signature(sig: &str) { + if sig.len() < MIN_SIGNATURE_LENGTH { + tracing::debug!( + "[SignatureStore] Ignoring short signature (length: {} < {})", + sig.len(), + MIN_SIGNATURE_LENGTH + ); + return; + } + + let mut store = THOUGHT_SIGNATURE.write().unwrap(); + + // 只有当新签名更长时才替换 + let should_replace = match &*store { + Some(existing) => sig.len() > existing.len(), + None => true, + }; + + if should_replace { + tracing::debug!( + "[SignatureStore] Storing thought_signature (length: {})", + sig.len() + ); + *store = Some(sig.to_string()); + } +} + +/// 获取存储的 thoughtSignature(不清除) +/// +/// # 返回 +/// 存储的签名,如果没有则返回 None +pub fn get_thought_signature() -> Option { + let store = THOUGHT_SIGNATURE.read().unwrap(); + store.clone() +} + +/// 获取并清除存储的 thoughtSignature +/// +/// # 返回 +/// 存储的签名,如果没有则返回 None +pub fn take_thought_signature() -> Option { + let mut store = THOUGHT_SIGNATURE.write().unwrap(); + store.take() +} + +/// 清除存储的 thoughtSignature +pub fn clear_thought_signature() { + let mut store = THOUGHT_SIGNATURE.write().unwrap(); + *store = None; + tracing::debug!("[SignatureStore] Cleared thought_signature"); +} + +/// 检查是否有有效的 thoughtSignature +pub fn has_valid_signature() -> bool { + let store = THOUGHT_SIGNATURE.read().unwrap(); + store + .as_ref() + .map(|s| s.len() >= MIN_SIGNATURE_LENGTH) + .unwrap_or(false) +} + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn test_signature_store() { + // 清除之前的状态 + clear_thought_signature(); + + // 初始状态应该为空 + assert!(get_thought_signature().is_none()); + + // 存储短签名应该被忽略 + store_thought_signature("short"); + assert!(get_thought_signature().is_none()); + + // 存储有效签名 + let valid_sig = "a".repeat(MIN_SIGNATURE_LENGTH); + store_thought_signature(&valid_sig); + assert_eq!(get_thought_signature(), Some(valid_sig.clone())); + + // 存储更长的签名应该替换 + let longer_sig = "b".repeat(MIN_SIGNATURE_LENGTH + 10); + store_thought_signature(&longer_sig); + assert_eq!(get_thought_signature(), Some(longer_sig.clone())); + + // 存储更短的签名不应该替换 + store_thought_signature(&valid_sig); + assert_eq!(get_thought_signature(), Some(longer_sig.clone())); + + // take 应该返回并清除 + let taken = take_thought_signature(); + assert_eq!(taken, Some(longer_sig)); + assert!(get_thought_signature().is_none()); + } +} diff --git a/src-tauri/src/session/sticky_config.rs b/src-tauri/src/session/sticky_config.rs new file mode 100644 index 000000000..8a3d88fa7 --- /dev/null +++ b/src-tauri/src/session/sticky_config.rs @@ -0,0 +1,105 @@ +//! 会话粘性配置 +//! +//! 提供调度模式配置,用于控制账号选择策略。 + +use serde::{Deserialize, Serialize}; + +/// 调度模式枚举 +#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)] +pub enum SchedulingMode { + /// 缓存优先 (Cache-first): 尽可能锁定同一账号,限流时优先等待,极大提升 Prompt Caching 命中率 + CacheFirst, + /// 平衡模式 (Balance): 锁定同一账号,限流时立即切换到备选账号,兼顾成功率和性能 + Balance, + /// 性能优先 (Performance-first): 纯轮询模式 (Round-robin),账号负载最均衡,但不利用缓存 + PerformanceFirst, +} + +impl Default for SchedulingMode { + fn default() -> Self { + Self::Balance + } +} + +impl std::fmt::Display for SchedulingMode { + fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + match self { + Self::CacheFirst => write!(f, "CacheFirst"), + Self::Balance => write!(f, "Balance"), + Self::PerformanceFirst => write!(f, "PerformanceFirst"), + } + } +} + +/// 粘性会话配置 +#[derive(Debug, Clone, Serialize, Deserialize)] +pub struct StickySessionConfig { + /// 当前调度模式 + pub mode: SchedulingMode, + /// 缓存优先模式下的最大等待时间 (秒) + pub max_wait_seconds: u64, + /// 60 秒全局锁定窗口(用于无 session_id 情况的默认保护) + pub global_lock_window_seconds: u64, +} + +impl Default for StickySessionConfig { + fn default() -> Self { + Self { + mode: SchedulingMode::Balance, + max_wait_seconds: 60, + global_lock_window_seconds: 60, + } + } +} + +impl StickySessionConfig { + /// 创建缓存优先配置 + pub fn cache_first() -> Self { + Self { + mode: SchedulingMode::CacheFirst, + max_wait_seconds: 120, + global_lock_window_seconds: 60, + } + } + + /// 创建性能优先配置 + pub fn performance_first() -> Self { + Self { + mode: SchedulingMode::PerformanceFirst, + max_wait_seconds: 0, + global_lock_window_seconds: 0, + } + } + + /// 是否启用会话粘性 + pub fn is_sticky_enabled(&self) -> bool { + self.mode != SchedulingMode::PerformanceFirst + } +} + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn test_default_config() { + let config = StickySessionConfig::default(); + assert_eq!(config.mode, SchedulingMode::Balance); + assert_eq!(config.max_wait_seconds, 60); + assert!(config.is_sticky_enabled()); + } + + #[test] + fn test_cache_first_config() { + let config = StickySessionConfig::cache_first(); + assert_eq!(config.mode, SchedulingMode::CacheFirst); + assert!(config.is_sticky_enabled()); + } + + #[test] + fn test_performance_first_config() { + let config = StickySessionConfig::performance_first(); + assert_eq!(config.mode, SchedulingMode::PerformanceFirst); + assert!(!config.is_sticky_enabled()); + } +} diff --git a/src-tauri/src/session/sticky_manager.rs b/src-tauri/src/session/sticky_manager.rs new file mode 100644 index 000000000..4aafe2799 --- /dev/null +++ b/src-tauri/src/session/sticky_manager.rs @@ -0,0 +1,346 @@ +//! 会话粘性管理器 +//! +//! 实现会话与账号的映射,支持: +//! - 会话绑定到特定账号 +//! - 60 秒全局锁定窗口 +//! - 订阅等级排序 + +use super::rate_limit::RateLimitTracker; +use super::sticky_config::{SchedulingMode, StickySessionConfig}; +use dashmap::DashMap; +use std::sync::atomic::{AtomicUsize, Ordering}; +use std::sync::Arc; +use std::time::Instant; +use tokio::sync::RwLock; + +/// 账号信息 +#[derive(Debug, Clone)] +pub struct AccountInfo { + /// 账号 ID + pub account_id: String, + /// 邮箱 + pub email: String, + /// 订阅等级 (ULTRA, PRO, FREE) + pub subscription_tier: Option, + /// 是否被禁用 + pub disabled: bool, +} + +impl AccountInfo { + /// 获取订阅等级优先级(数字越小优先级越高) + pub fn tier_priority(&self) -> u8 { + match self.subscription_tier.as_deref() { + Some("ULTRA") => 0, + Some("PRO") => 1, + Some("FREE") => 2, + _ => 3, + } + } +} + +/// 会话粘性管理器 +pub struct StickySessionManager { + /// 会话与账号映射 (session_id -> account_id) + session_accounts: DashMap, + /// 最后使用的账号 (account_id, timestamp) + last_used_account: Arc>>, + /// 当前轮询索引 + current_index: AtomicUsize, + /// 限流追踪器 + rate_limit_tracker: Arc, + /// 粘性配置 + sticky_config: Arc>, +} + +impl Default for StickySessionManager { + fn default() -> Self { + Self::new(Arc::new(RateLimitTracker::default())) + } +} + +impl StickySessionManager { + /// 创建新的会话粘性管理器 + pub fn new(rate_limit_tracker: Arc) -> Self { + Self { + session_accounts: DashMap::new(), + last_used_account: Arc::new(tokio::sync::Mutex::new(None)), + current_index: AtomicUsize::new(0), + rate_limit_tracker, + sticky_config: Arc::new(RwLock::new(StickySessionConfig::default())), + } + } + + /// 获取当前配置 + pub async fn get_config(&self) -> StickySessionConfig { + self.sticky_config.read().await.clone() + } + + /// 设置配置 + pub async fn set_config(&self, config: StickySessionConfig) { + *self.sticky_config.write().await = config; + } + + /// 绑定会话到账号 + pub fn bind_session(&self, session_id: &str, account_id: &str) { + self.session_accounts + .insert(session_id.to_string(), account_id.to_string()); + tracing::debug!( + "[StickySession] 绑定会话 {} 到账号 {}", + session_id, + account_id + ); + } + + /// 解绑会话 + pub fn unbind_session(&self, session_id: &str) { + if let Some((_, account_id)) = self.session_accounts.remove(session_id) { + tracing::debug!( + "[StickySession] 解绑会话 {} (原账号: {})", + session_id, + account_id + ); + } + } + + /// 获取会话绑定的账号 + pub fn get_bound_account(&self, session_id: &str) -> Option { + self.session_accounts.get(session_id).map(|v| v.clone()) + } + + /// 选择账号(支持粘性会话和智能调度) + /// + /// # 参数 + /// - `accounts`: 可用账号列表 + /// - `session_id`: 会话 ID(可选) + /// - `force_rotate`: 是否强制轮换 + /// - `quota_group`: 配额组(如 "claude", "gemini", "image_gen") + /// + /// # 返回 + /// 选中的账号,如果没有可用账号返回 None + pub async fn select_account( + &self, + accounts: &[AccountInfo], + session_id: Option<&str>, + force_rotate: bool, + quota_group: &str, + ) -> Option { + if accounts.is_empty() { + return None; + } + + // 按订阅等级排序(ULTRA > PRO > FREE) + let mut sorted_accounts = accounts.to_vec(); + sorted_accounts.sort_by_key(|a| a.tier_priority()); + + let config = self.sticky_config.read().await.clone(); + let total = sorted_accounts.len(); + + // 模式 A: 粘性会话处理 + if !force_rotate && session_id.is_some() && config.mode != SchedulingMode::PerformanceFirst + { + let sid = session_id.unwrap(); + + // 检查会话是否已绑定账号 + if let Some(bound_id) = self.get_bound_account(sid) { + // 找到绑定的账号 + if let Some(bound_account) = + sorted_accounts.iter().find(|a| a.account_id == bound_id) + { + // 检查是否被限流 + if !self + .rate_limit_tracker + .is_rate_limited(&bound_account.email) + { + tracing::debug!( + "[StickySession] 复用绑定账号 {} (会话: {})", + bound_account.email, + sid + ); + return Some(bound_account.clone()); + } else { + // 账号被限流,解绑并切换 + tracing::warn!( + "[StickySession] 绑定账号 {} 被限流,解绑会话 {}", + bound_account.email, + sid + ); + self.unbind_session(sid); + } + } else { + // 绑定的账号不存在,解绑 + self.unbind_session(sid); + } + } + } + + // 模式 B: 60 秒全局锁定(针对无 session_id 情况) + if !force_rotate && quota_group != "image_gen" && config.global_lock_window_seconds > 0 { + let last_used = self.last_used_account.lock().await; + if let Some((account_id, last_time)) = &*last_used { + if last_time.elapsed().as_secs() < config.global_lock_window_seconds { + // 找到最后使用的账号 + if let Some(account) = + sorted_accounts.iter().find(|a| &a.account_id == account_id) + { + if !self.rate_limit_tracker.is_rate_limited(&account.email) { + tracing::debug!("[StickySession] 60s 窗口内复用账号 {}", account.email); + return Some(account.clone()); + } + } + } + } + drop(last_used); + } + + // 模式 C: 轮询选择 + let start_idx = self.current_index.fetch_add(1, Ordering::SeqCst) % total; + for offset in 0..total { + let idx = (start_idx + offset) % total; + let candidate = &sorted_accounts[idx]; + + // 跳过被禁用的账号 + if candidate.disabled { + continue; + } + + // 跳过被限流的账号 + if self.rate_limit_tracker.is_rate_limited(&candidate.email) { + continue; + } + + // 找到可用账号 + tracing::debug!( + "[StickySession] 轮询选择账号 {} (索引: {})", + candidate.email, + idx + ); + + // 更新最后使用的账号 + { + let mut last_used = self.last_used_account.lock().await; + *last_used = Some((candidate.account_id.clone(), Instant::now())); + } + + // 如果有会话 ID 且启用粘性,绑定会话 + if let Some(sid) = session_id { + if config.mode != SchedulingMode::PerformanceFirst { + self.bind_session(sid, &candidate.account_id); + } + } + + return Some(candidate.clone()); + } + + // 没有可用账号 + tracing::warn!("[StickySession] 没有可用账号"); + None + } + + /// 标记账号请求成功(清除限流状态) + pub fn mark_success(&self, account_id: &str) { + self.rate_limit_tracker.clear_rate_limit(account_id); + } + + /// 获取限流追踪器 + pub fn rate_limit_tracker(&self) -> &Arc { + &self.rate_limit_tracker + } + + /// 清理过期的会话绑定 + pub fn cleanup_expired_sessions(&self, max_age_seconds: u64) { + // 这里可以添加会话过期清理逻辑 + // 目前简单实现,不做过期清理 + let _ = max_age_seconds; + } +} + +#[cfg(test)] +mod tests { + use super::*; + + #[tokio::test] + async fn test_session_binding() { + let manager = StickySessionManager::default(); + + // 绑定会话 + manager.bind_session("session1", "account1"); + assert_eq!( + manager.get_bound_account("session1"), + Some("account1".to_string()) + ); + + // 解绑会话 + manager.unbind_session("session1"); + assert_eq!(manager.get_bound_account("session1"), None); + } + + #[tokio::test] + async fn test_account_selection() { + let manager = StickySessionManager::default(); + + let accounts = vec![ + AccountInfo { + account_id: "acc1".to_string(), + email: "user1@example.com".to_string(), + subscription_tier: Some("FREE".to_string()), + disabled: false, + }, + AccountInfo { + account_id: "acc2".to_string(), + email: "user2@example.com".to_string(), + subscription_tier: Some("PRO".to_string()), + disabled: false, + }, + AccountInfo { + account_id: "acc3".to_string(), + email: "user3@example.com".to_string(), + subscription_tier: Some("ULTRA".to_string()), + disabled: false, + }, + ]; + + // 应该优先选择 ULTRA 账号 + let selected = manager + .select_account(&accounts, None, false, "claude") + .await; + assert!(selected.is_some()); + assert_eq!( + selected.unwrap().subscription_tier, + Some("ULTRA".to_string()) + ); + } + + #[tokio::test] + async fn test_sticky_session() { + let manager = StickySessionManager::default(); + + let accounts = vec![ + AccountInfo { + account_id: "acc1".to_string(), + email: "user1@example.com".to_string(), + subscription_tier: Some("PRO".to_string()), + disabled: false, + }, + AccountInfo { + account_id: "acc2".to_string(), + email: "user2@example.com".to_string(), + subscription_tier: Some("PRO".to_string()), + disabled: false, + }, + ]; + + // 第一次选择,应该绑定会话 + let selected1 = manager + .select_account(&accounts, Some("session1"), false, "claude") + .await; + assert!(selected1.is_some()); + let account_id = selected1.unwrap().account_id; + + // 第二次选择同一会话,应该返回相同账号 + let selected2 = manager + .select_account(&accounts, Some("session1"), false, "claude") + .await; + assert!(selected2.is_some()); + assert_eq!(selected2.unwrap().account_id, account_id); + } +}