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:
coso
2026-01-12 22:23:56 +08:00
co-authored by Claude Opus 4.5
parent 86e2a0db33
commit 66fa3f28f5
10 changed files with 1357 additions and 8 deletions
@@ -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()),
};
+1
View File
@@ -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;
+13 -1
View File
@@ -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")
+25
View File
@@ -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};
+450
View File
@@ -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"));
}
}
+239
View File
@@ -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"
);
}
}
+117
View File
@@ -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());
}
}
+105
View File
@@ -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());
}
}
+346
View File
@@ -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);
}
}