mirror of
https://github.com/aiclientproxy/proxycast.git
synced 2026-09-24 23:10:56 +08:00
feat: add session management module
Add comprehensive session management capabilities: - Stable SessionId generation based on request content hash - Global thoughtSignature caching - Sticky session management (session-to-account mapping) - Scheduling mode configuration - Enhanced rate limiting with duration parsing and exponential backoff Co-Authored-By: Claude Opus 4.5 <noreply@anthropic.com>
This commit is contained in:
@@ -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<SafetySetting> {
|
||||
/// 模型名称映射
|
||||
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<String> = 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<serde_json::Value> = 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()),
|
||||
};
|
||||
|
||||
|
||||
@@ -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;
|
||||
|
||||
@@ -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",
|
||||
|
||||
@@ -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")
|
||||
|
||||
@@ -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};
|
||||
@@ -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<Utc>,
|
||||
/// 限流结束时间(预计)
|
||||
pub reset_at: DateTime<Utc>,
|
||||
/// 连续失败次数
|
||||
pub consecutive_failures: u32,
|
||||
/// 模型名称(如果是模型级别限流)
|
||||
pub model: Option<String>,
|
||||
}
|
||||
|
||||
/// 限流追踪器
|
||||
#[derive(Debug)]
|
||||
pub struct RateLimitTracker {
|
||||
/// 账号级别限流记录
|
||||
account_limits: DashMap<String, RateLimitRecord>,
|
||||
/// 模型级别限流记录 (account_id:model -> record)
|
||||
model_limits: DashMap<String, RateLimitRecord>,
|
||||
/// 连续失败计数 (account_id -> count)
|
||||
failure_counts: DashMap<String, AtomicU32>,
|
||||
/// 基础退避时间(秒)
|
||||
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<Duration>,
|
||||
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<String> {
|
||||
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<Duration> {
|
||||
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<Duration> {
|
||||
// 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::<i64>() {
|
||||
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"));
|
||||
}
|
||||
}
|
||||
@@ -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("<system-reminder>") {
|
||||
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::<Vec<_>>()
|
||||
.join(" ")
|
||||
} else {
|
||||
String::new()
|
||||
}
|
||||
} else {
|
||||
String::new()
|
||||
};
|
||||
|
||||
let clean_text = text.trim();
|
||||
if clean_text.len() > 10 && !clean_text.contains("<system-reminder>") {
|
||||
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::<Vec<_>>()
|
||||
.join(" ");
|
||||
|
||||
let clean_text = text.trim();
|
||||
if clean_text.len() > 10 && !clean_text.contains("<system-reminder>") {
|
||||
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"
|
||||
);
|
||||
}
|
||||
}
|
||||
@@ -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<Option<String>> = 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<String> {
|
||||
let store = THOUGHT_SIGNATURE.read().unwrap();
|
||||
store.clone()
|
||||
}
|
||||
|
||||
/// 获取并清除存储的 thoughtSignature
|
||||
///
|
||||
/// # 返回
|
||||
/// 存储的签名,如果没有则返回 None
|
||||
pub fn take_thought_signature() -> Option<String> {
|
||||
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());
|
||||
}
|
||||
}
|
||||
@@ -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());
|
||||
}
|
||||
}
|
||||
@@ -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<String>,
|
||||
/// 是否被禁用
|
||||
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<String, String>,
|
||||
/// 最后使用的账号 (account_id, timestamp)
|
||||
last_used_account: Arc<tokio::sync::Mutex<Option<(String, Instant)>>>,
|
||||
/// 当前轮询索引
|
||||
current_index: AtomicUsize,
|
||||
/// 限流追踪器
|
||||
rate_limit_tracker: Arc<RateLimitTracker>,
|
||||
/// 粘性配置
|
||||
sticky_config: Arc<RwLock<StickySessionConfig>>,
|
||||
}
|
||||
|
||||
impl Default for StickySessionManager {
|
||||
fn default() -> Self {
|
||||
Self::new(Arc::new(RateLimitTracker::default()))
|
||||
}
|
||||
}
|
||||
|
||||
impl StickySessionManager {
|
||||
/// 创建新的会话粘性管理器
|
||||
pub fn new(rate_limit_tracker: Arc<RateLimitTracker>) -> 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<String> {
|
||||
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<AccountInfo> {
|
||||
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<RateLimitTracker> {
|
||||
&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);
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user