refactor: 移除 Python 后端源码,全面迁移至 Rust gateway 架构

- 删除全部 Python 源码 (src/) 及 Alembic 迁移脚本,归档至 _deprecated_py_src/
- 重构 Rust gateway ai_pipeline: 拆分 planner/finalize 模块,新增 contracts/adaptation 层
- 重组 handlers 模块为 admin/public/proxy/internal/shared 子模块结构
- 新增 executor 模块,引入 Rust 原生数据库迁移 (aether-data/migrations)
- 简化 CI/Docker 构建流程,移除 base image 二级构建,统一为单一 app image
- 移除 Python 相关基础设施文件 (entrypoint.sh, gunicorn_conf.py, Dockerfile.base)
This commit is contained in:
fawney19
2026-04-03 16:26:16 +08:00
parent 8f26e1a31f
commit 1d9c77522a
868 changed files with 1735 additions and 2433 deletions

View File

@@ -0,0 +1,30 @@
"""调度系统(候选/排序/亲和性/并发检查)。
原 `src.services.cache` 中与调度相关的代码已迁移到此包;
`src.services.cache` 现在只保留通用缓存 backend/sync/*_cache。
"""
from src.services.scheduling.affinity_manager import CacheAffinityManager, get_affinity_manager
from src.services.scheduling.aware_scheduler import (
CacheAwareScheduler,
ConcurrencySnapshot,
ProviderCandidate,
get_cache_aware_scheduler,
)
from src.services.scheduling.candidate_builder import CandidateBuilder
from src.services.scheduling.candidate_sorter import CandidateSorter
from src.services.scheduling.concurrency_checker import ConcurrencyChecker
from src.services.scheduling.scheduling_config import SchedulingConfig
__all__ = [
"CacheAffinityManager",
"CandidateBuilder",
"CandidateSorter",
"CacheAwareScheduler",
"ConcurrencyChecker",
"ConcurrencySnapshot",
"ProviderCandidate",
"SchedulingConfig",
"get_affinity_manager",
"get_cache_aware_scheduler",
]

View File

@@ -0,0 +1,757 @@
"""
缓存亲和性管理器 (Cache Affinity Manager) - 支持 Redis 或内存存储
职责:
1. 跟踪请求API Key的Provider+Key缓存状态
2. 管理缓存有效期
3. 提供缓存统计和分析
4. 自动失效不支持缓存的Provider
设计原理:
- 每个API Key使用某个Provider的Key后在缓存TTL期内应该继续使用同一个Key
- 这样可以最大化利用提供商的Prompt Caching机制
- 当Key故障时自动失效该Key的缓存亲和性
- 当Provider关闭缓存支持时自动失效所有相关亲和性
注意:
- affinity_key 参数通常为请求使用的 API Key IDapi_key_id
- 这样可以支持"独立余额Key"场景每个Key有自己的缓存亲和性
"""
from __future__ import annotations
import asyncio
import json
import os
import time
from contextlib import asynccontextmanager
from typing import Any, NamedTuple
from src.config.constants import CacheTTL
from src.core.logger import logger
from src.core.redis_utils import delete_redis_keys, scan_delete_pattern
class CacheAffinity(NamedTuple):
"""缓存亲和性信息"""
provider_id: str
endpoint_id: str
key_id: str
api_format: str # API格式 (claude/openai)
model_name: str # 模型名称
created_at: float # 创建时间戳
expire_at: float # 过期时间戳
request_count: int # 使用次数
class CacheAffinityManager:
"""
缓存亲和性管理器(支持 Redis 或内存存储)
存储结构:
----------------------
Key格式: cache_affinity:{affinity_key}:{api_format}:{model_name}
- affinity_key: 通常为请求使用的 API Key ID支持独立余额Key场景
- api_format: API格式 (claude/openai)
- model_name: 模型名称(区分不同模型的缓存亲和性)
Value格式: JSON/Dict
{
"provider_id": "xxx",
"endpoint_id": "yyy",
"key_id": "zzz",
"model_name": "claude-3-5-sonnet-20241022",
"created_at": 1234567890.123,
"expire_at": 1234567890.123,
"request_count": 5
}
TTL: 自动过期
设计改进:
- 每个API Key可以对多个API格式和模型分别维护缓存亲和性
- 不同模型请求使用独立的缓存亲和性,避免模型切换导致的缓存失效
- 某个端点故障切换不会影响其他端点的亲和性
- 更精确的缓存命中率统计
- 支持"独立余额Key"场景每个Key有独立的缓存亲和性
"""
# 默认缓存TTL- 使用统一常量
DEFAULT_CACHE_TTL = CacheTTL.CACHE_AFFINITY
REDIS_SCAN_BATCH_SIZE = 200
REDIS_DELETE_BATCH_SIZE = 500
def __init__(
self, redis_client: Any | None = None, default_ttl: int = DEFAULT_CACHE_TTL
) -> None:
"""
初始化缓存亲和性管理器
Args:
redis_client: Redis客户端可选
default_ttl: 默认缓存TTL
"""
self.redis = redis_client
self.default_ttl = default_ttl
self._memory_store: dict[str, dict[str, Any]] = {}
self._memory_store_max_size: int = 5000 # 内存模式最大条目数
self._memory_lock: asyncio.Lock | None = None
# L1 缓存(即使使用 Redis 也启用,减少网络往返)
# 注意L1 是本地进程内缓存多实例部署时存在短暂不一致窗口TTL 秒级)。
# 当前 TTL 默认 3 秒,对于亲和性路由来说可接受:最坏情况是短暂路由到
# 旧 provider下次请求即可自动修正。如果需要严格一致性将 TTL 设为 0
# 以禁用 L1 缓存,或通过 CacheSyncService 接收 pub/sub 主动失效。
self._l1_cache_ttl = int(os.getenv("CACHE_AFFINITY_L1_TTL", str(CacheTTL.L1_LOCAL)))
self._l1_cache: dict[str, tuple[float, dict[str, Any]]] = {}
self._l1_lock = asyncio.Lock()
self._l1_max_size = int(os.getenv("CACHE_AFFINITY_L1_MAX_SIZE", "1000")) # 最大缓存条目数
self._l1_last_cleanup = time.time()
# 请求级别锁,避免同一用户+端点同时更新造成抖动
self._request_locks: dict[str, asyncio.Lock] = {}
self._request_locks_max_size: int = 500 # 锁字典上限,防止无界增长
# 统计信息
self._stats = {
"total_affinities": 0,
"cache_hits": 0,
"cache_misses": 0,
"cache_invalidations": 0,
"provider_switches": 0,
"key_switches": 0,
}
if self.redis:
logger.debug("CacheAffinityManager: 使用Redis存储")
else:
logger.debug(
"CacheAffinityManager: Redis不可用回退到内存存储(仅适用于单实例/开发环境)"
)
def _is_memory_backend(self) -> bool:
"""是否处于内存模式"""
return self.redis is None
def _get_memory_lock(self) -> asyncio.Lock:
"""懒初始化内存锁"""
if self._memory_lock is None:
self._memory_lock = asyncio.Lock()
return self._memory_lock
def _get_cache_key(self, affinity_key: str, api_format: str, model_name: str) -> str:
"""
生成Redis Key
Args:
affinity_key: 亲和性标识符通常为API Key ID
api_format: API格式 (claude/openai)
model_name: 模型名称
Returns:
格式化的缓存键: cache_affinity:{affinity_key}:{api_format}:{model_name}
"""
return f"cache_affinity:{affinity_key}:{api_format}:{model_name}"
async def _get_l1_entry(self, cache_key: str) -> dict[str, Any] | None:
async with self._l1_lock:
record = self._l1_cache.get(cache_key)
if not record:
return None
expire_at, payload = record
if time.time() > expire_at:
self._l1_cache.pop(cache_key, None)
return None
return dict(payload)
async def _set_l1_entry(self, cache_key: str, payload: dict[str, Any] | None) -> None:
async with self._l1_lock:
if not payload:
self._l1_cache.pop(cache_key, None)
return
expire_at = time.time() + max(1, self._l1_cache_ttl)
self._l1_cache[cache_key] = (expire_at, dict(payload))
# 定期清理过期条目(每 60 秒最多一次)
current_time = time.time()
if current_time - self._l1_last_cleanup > 60:
self._cleanup_l1_cache_unlocked(current_time)
self._l1_last_cleanup = current_time
def _cleanup_l1_cache_unlocked(self, current_time: float) -> int:
"""清理过期的 L1 缓存条目(需要在持有锁的情况下调用)
Returns:
清理的条目数量
"""
expired_keys = [
key for key, (expire_at, _) in self._l1_cache.items() if current_time > expire_at
]
for key in expired_keys:
self._l1_cache.pop(key, None)
# 如果缓存仍然过大,按过期时间排序移除最旧的条目
if len(self._l1_cache) > self._l1_max_size:
sorted_items = sorted(
self._l1_cache.items(), key=lambda x: x[1][0] # 按 expire_at 排序
)
# 移除最旧的 20% 条目
remove_count = len(self._l1_cache) - int(self._l1_max_size * 0.8)
for key, _ in sorted_items[:remove_count]:
self._l1_cache.pop(key, None)
expired_keys.extend([k for k, _ in sorted_items[:remove_count]])
if expired_keys:
logger.debug(
f"L1 缓存清理: 移除 {len(expired_keys)} 个条目,当前 {len(self._l1_cache)}"
)
return len(expired_keys)
@asynccontextmanager
async def _acquire_request_lock(self, cache_key: str) -> None:
lock = self._request_locks.get(cache_key)
if lock is None:
# 超出上限时淘汰未被持有的锁,防止无界增长
if len(self._request_locks) >= self._request_locks_max_size:
free_keys = [k for k, lk in self._request_locks.items() if not lk.locked()]
for k in free_keys[: len(free_keys) // 2 or 1]: # 清理一半空闲锁
del self._request_locks[k]
lock = asyncio.Lock()
self._request_locks[cache_key] = lock
await lock.acquire()
try:
yield
finally:
lock.release()
async def _load_affinity_dict(self, cache_key: str) -> dict[str, Any] | None:
"""读取缓存亲和性字典"""
# 先尝试L1缓存
l1_value = await self._get_l1_entry(cache_key)
if l1_value is not None:
return l1_value
if not self._is_memory_backend():
data = await self.redis.get(cache_key)
if not data:
return None
value = json.loads(data)
await self._set_l1_entry(cache_key, value)
return value
lock = self._get_memory_lock()
async with lock:
record = self._memory_store.get(cache_key)
if record:
await self._set_l1_entry(cache_key, record)
return dict(record) if record else None
async def _save_affinity_dict(
self, cache_key: str, ttl: int, affinity_dict: dict[str, Any]
) -> None:
"""存储缓存亲和性字典"""
if not self._is_memory_backend():
await self.redis.setex(cache_key, ttl, json.dumps(affinity_dict))
await self._set_l1_entry(cache_key, affinity_dict)
return
lock = self._get_memory_lock()
async with lock:
self._memory_store[cache_key] = dict(affinity_dict)
# 超出上限时清理过期条目
if len(self._memory_store) > self._memory_store_max_size:
now = time.time()
expired = [k for k, v in self._memory_store.items() if now > v.get("expire_at", 0)]
for k in expired:
del self._memory_store[k]
await self._set_l1_entry(cache_key, affinity_dict)
async def _delete_affinity_key(self, cache_key: str) -> None:
"""删除缓存亲和性"""
if not self._is_memory_backend():
await self.redis.delete(cache_key)
else:
lock = self._get_memory_lock()
async with lock:
self._memory_store.pop(cache_key, None)
await self._set_l1_entry(cache_key, None)
async def _delete_redis_keys(self, keys: list[str]) -> int:
if self._is_memory_backend() or not keys:
return 0
return await delete_redis_keys(self.redis, keys)
async def _scan_delete_pattern(self, pattern: str) -> int:
if self._is_memory_backend():
return 0
return await scan_delete_pattern(
self.redis,
pattern,
scan_batch_size=self.REDIS_SCAN_BATCH_SIZE,
delete_batch_size=self.REDIS_DELETE_BATCH_SIZE,
)
async def _clear_l1_entries_by_prefix(self, prefix: str) -> int:
"""清理匹配前缀的 L1 本地缓存。"""
async with self._l1_lock:
keys_to_remove = [key for key in self._l1_cache if key.startswith(prefix)]
for key in keys_to_remove:
self._l1_cache.pop(key, None)
return len(keys_to_remove)
async def _snapshot_memory_items(self) -> dict[str, dict[str, Any]]:
"""复制内存存储内容(仅内存模式使用)"""
lock = self._get_memory_lock()
async with lock:
return {k: dict(v) for k, v in self._memory_store.items()}
async def get_affinity(
self, affinity_key: str, api_format: str, model_name: str
) -> CacheAffinity | None:
"""
获取指定亲和性标识符对特定API格式和模型的缓存亲和性
Args:
affinity_key: 亲和性标识符通常为API Key ID
api_format: API格式 (claude/openai)
model_name: 模型名称
Returns:
CacheAffinity对象如果不存在或已过期则返回None
"""
try:
cache_key = self._get_cache_key(affinity_key, api_format, model_name)
async with self._acquire_request_lock(cache_key):
affinity_dict = await self._load_affinity_dict(cache_key)
if not affinity_dict:
self._stats["cache_misses"] += 1
return None
# 检查是否过期双重检查防止TTL未及时清理
current_time = time.time()
if current_time > affinity_dict["expire_at"]:
await self._delete_affinity_key(cache_key)
self._stats["cache_misses"] += 1
return None
self._stats["cache_hits"] += 1
return CacheAffinity(
provider_id=affinity_dict["provider_id"],
endpoint_id=affinity_dict["endpoint_id"],
key_id=affinity_dict["key_id"],
api_format=affinity_dict.get("api_format", api_format),
model_name=affinity_dict.get("model_name", model_name),
created_at=affinity_dict["created_at"],
expire_at=affinity_dict["expire_at"],
request_count=affinity_dict["request_count"],
)
except Exception as e:
logger.exception(f"获取缓存亲和性失败: {e}")
self._stats["cache_misses"] += 1
return None
async def set_affinity(
self,
affinity_key: str,
provider_id: str,
endpoint_id: str,
key_id: str,
api_format: str,
model_name: str,
supports_caching: bool = True,
ttl: int | None = None,
) -> None:
"""
设置指定亲和性标识符对特定API格式和模型的缓存亲和性
Args:
affinity_key: 亲和性标识符通常为API Key ID
provider_id: Provider ID
endpoint_id: Endpoint ID
key_id: Key ID
api_format: API格式 (claude/openai)
model_name: 模型名称
supports_caching: 该Provider是否支持缓存
ttl: 缓存有效期(秒),如果不提供则使用默认值
注意每次调用都会刷新过期时间滑动窗口机制以保持对同一个Provider/Endpoint/Key的亲和性
"""
if not supports_caching:
# 不支持缓存的Provider不记录亲和性
logger.debug(f"Provider {provider_id[:8]}... 不支持缓存,跳过亲和性记录")
return
ttl = ttl or self.default_ttl
current_time = time.time()
expire_at = current_time + ttl # 每次都刷新过期时间
cache_key = self._get_cache_key(affinity_key, api_format, model_name)
try:
async with self._acquire_request_lock(cache_key):
existing_dict = await self._load_affinity_dict(cache_key)
existing_affinity: CacheAffinity | None = None
if existing_dict and current_time <= existing_dict.get("expire_at", 0):
existing_affinity = CacheAffinity(
provider_id=existing_dict["provider_id"],
endpoint_id=existing_dict["endpoint_id"],
key_id=existing_dict["key_id"],
api_format=existing_dict.get("api_format", api_format),
model_name=existing_dict.get("model_name", model_name),
created_at=existing_dict["created_at"],
expire_at=existing_dict["expire_at"],
request_count=existing_dict.get("request_count", 0),
)
if existing_affinity:
created_at = existing_affinity.created_at
request_count = existing_affinity.request_count + 1
# 检查是否切换了 Provider/Endpoint/Key
if (
existing_affinity.provider_id != provider_id
or existing_affinity.endpoint_id != endpoint_id
or existing_affinity.key_id != key_id
):
self._stats["key_switches"] += 1
logger.debug(
f"Key {affinity_key[:8]}... 在 {api_format} 格式下切换后端: "
f"[{existing_affinity.provider_id[:8]}.../{existing_affinity.endpoint_id[:8]}.../"
f"{existing_affinity.key_id[:8]}...] → "
f"[{provider_id[:8]}.../{endpoint_id[:8]}.../{key_id[:8]}...], 重置计数器"
)
created_at = current_time
request_count = 1
else:
logger.debug(
f"刷新缓存亲和性: key={affinity_key[:8]}..., api_format={api_format}, "
f"provider={provider_id[:8]}..., endpoint={endpoint_id[:8]}..., "
f"provider_key={key_id[:8]}..., ttl+={ttl}s"
)
else:
created_at = current_time
request_count = 1
self._stats["total_affinities"] += 1
affinity_dict = {
"provider_id": provider_id,
"endpoint_id": endpoint_id,
"key_id": key_id,
"api_format": api_format,
"model_name": model_name,
"created_at": created_at,
"expire_at": expire_at,
"request_count": request_count,
}
await self._save_affinity_dict(cache_key, ttl, affinity_dict)
logger.debug(
f"设置缓存亲和性: key={affinity_key[:8]}..., api_format={api_format}, "
f"model={model_name}, provider={provider_id[:8]}..., endpoint={endpoint_id[:8]}..., "
f"provider_key={key_id[:8]}..., ttl={ttl}s"
)
except Exception as e:
logger.exception(f"设置缓存亲和性失败: {e}")
async def invalidate_affinity(
self,
affinity_key: str,
api_format: str,
model_name: str,
key_id: str | None = None,
provider_id: str | None = None,
endpoint_id: str | None = None,
) -> None:
"""
失效指定亲和性标识符对特定API格式和模型的缓存亲和性
Args:
affinity_key: 亲和性标识符通常为API Key ID
api_format: API格式 (claude/openai)
model_name: 模型名称
key_id: Provider Key ID可选如果提供则只在Key匹配时失效
provider_id: Provider ID可选如果提供则只在Provider匹配时失效
endpoint_id: Endpoint ID可选如果提供则只在Endpoint匹配时失效
"""
existing_affinity = await self.get_affinity(affinity_key, api_format, model_name)
if not existing_affinity:
return
# 检查是否匹配过滤条件
should_invalidate = True
if key_id and existing_affinity.key_id != key_id:
should_invalidate = False
if provider_id and existing_affinity.provider_id != provider_id:
should_invalidate = False
if endpoint_id and existing_affinity.endpoint_id != endpoint_id:
should_invalidate = False
if not should_invalidate:
logger.debug(
f"跳过失效: affinity_key={affinity_key[:8]}..., api_format={api_format}, "
f"model={model_name}, 过滤条件不匹配 (key={key_id}, provider={provider_id}, endpoint={endpoint_id})"
)
return
try:
cache_key = self._get_cache_key(affinity_key, api_format, model_name)
async with self._acquire_request_lock(cache_key):
await self._delete_affinity_key(cache_key)
self._stats["cache_invalidations"] += 1
logger.debug(
f"失效缓存亲和性: affinity_key={affinity_key[:8]}..., api_format={api_format}, "
f"model={model_name}, provider={existing_affinity.provider_id[:8]}..., "
f"endpoint={existing_affinity.endpoint_id[:8]}..., "
f"provider_key={existing_affinity.key_id[:8]}..."
)
except Exception as e:
logger.exception(f"删除缓存亲和性失败: {e}")
async def invalidate_all_for_provider(self, provider_id: str) -> int:
"""
失效所有与指定Provider相关的缓存亲和性
用途当Provider关闭缓存支持时调用
Args:
provider_id: Provider ID
Returns:
失效的亲和性数量
"""
try:
invalidated_count = 0
if not self._is_memory_backend():
cursor: int | str = 0
while True:
cursor, scan_keys = await self.redis.scan(
cursor=cursor,
match="cache_affinity:*",
count=self.REDIS_SCAN_BATCH_SIZE,
)
if scan_keys:
# Pipeline batch GET to reduce round-trips.
values = await self.redis.mget(scan_keys)
keys_to_delete: list[str] = []
for key, raw in zip(scan_keys, values):
if not raw:
continue
try:
data = json.loads(raw)
except Exception:
continue
if data.get("provider_id") == provider_id:
keys_to_delete.append(key)
if keys_to_delete:
deleted = await self._delete_redis_keys(keys_to_delete)
invalidated_count += deleted
self._stats["cache_invalidations"] += deleted
# Clear L1 for deleted keys.
for key in keys_to_delete:
await self._set_l1_entry(key, None)
if int(cursor) == 0:
break
else:
keys = list((await self._snapshot_memory_items()).keys())
for key in keys:
affinity_dict = await self._load_affinity_dict(key)
if not affinity_dict:
continue
if affinity_dict.get("provider_id") == provider_id:
await self._delete_affinity_key(key)
invalidated_count += 1
self._stats["cache_invalidations"] += 1
if invalidated_count > 0:
logger.debug(
f"批量失效Provider缓存亲和性: provider={provider_id[:8]}..., "
f"失效数量={invalidated_count}"
)
return invalidated_count
except Exception as e:
logger.exception(f"批量失效Provider缓存亲和性失败: {e}")
return 0
async def clear_all(self) -> int:
"""
清除所有缓存亲和性(管理功能)
Returns:
清除的数量
"""
try:
if not self._is_memory_backend():
count = await self._scan_delete_pattern("cache_affinity:*")
await self._clear_l1_entries_by_prefix("cache_affinity:")
if count:
logger.debug(f"清除所有Redis缓存亲和性: {count}")
return count
lock = self._get_memory_lock()
async with lock:
count = len(self._memory_store)
self._memory_store.clear()
await self._clear_l1_entries_by_prefix("cache_affinity:")
if count:
logger.debug(f"清除所有内存缓存亲和性: {count}")
return count
except Exception as e:
logger.exception(f"清除缓存亲和性失败: {e}")
return 0
def get_stats(self) -> dict[str, Any]:
"""获取统计信息"""
cache_hit_rate = 0.0
total_requests = self._stats["cache_hits"] + self._stats["cache_misses"]
if total_requests > 0:
cache_hit_rate = self._stats["cache_hits"] / total_requests
storage_type = "redis" if not self._is_memory_backend() else "memory"
return {
"storage_type": storage_type,
"total_affinities": self._stats["total_affinities"],
"cache_hits": self._stats["cache_hits"],
"cache_misses": self._stats["cache_misses"],
"cache_hit_rate": cache_hit_rate,
"cache_invalidations": self._stats["cache_invalidations"],
"provider_switches": self._stats["provider_switches"],
"key_switches": self._stats["key_switches"],
"config": {
"default_ttl": self.default_ttl,
},
}
async def list_affinities(self) -> list[dict[str, Any]]:
"""获取所有缓存亲和性列表
返回的每条记录包含:
- affinity_key: 亲和性标识符(通常是 API Key ID
- provider_id, endpoint_id, key_id: Provider 相关信息
- api_format, model_name: API 格式和模型名称
- created_at, expire_at, request_count: 缓存元数据
"""
results: list[dict[str, Any]] = []
try:
pattern = "cache_affinity:*"
cursor = 0
if not self._is_memory_backend():
while True:
cursor, keys = await self.redis.scan(cursor=cursor, match=pattern, count=200)
if keys:
values = await self.redis.mget(*keys)
for cache_key, data in zip(keys, values):
if not data:
continue
try:
affinity = json.loads(data)
# 解析 cache_affinity:{affinity_key}:{api_format}:{model_name}
parts = cache_key.split(":")
affinity_key_value = parts[1] if len(parts) > 1 else cache_key
api_format = (
parts[2]
if len(parts) > 2
else affinity.get("api_format", "unknown")
)
model_name = (
parts[3]
if len(parts) > 3
else affinity.get("model_name", "unknown")
)
affinity["affinity_key"] = affinity_key_value
if "api_format" not in affinity:
affinity["api_format"] = api_format
if "model_name" not in affinity:
affinity["model_name"] = model_name
results.append(affinity)
except json.JSONDecodeError as e:
logger.exception(f"解析缓存亲和性记录失败: {cache_key} - {e}")
if cursor == 0:
break
else:
snapshot = await self._snapshot_memory_items()
expired_keys: list[str] = []
current_time = time.time()
for cache_key, affinity in snapshot.items():
if current_time > affinity["expire_at"]:
expired_keys.append(cache_key)
continue
# 解析 cache_affinity:{affinity_key}:{api_format}:{model_name}
parts = cache_key.split(":")
affinity_key_value = parts[1] if len(parts) > 1 else cache_key
api_format = (
parts[2] if len(parts) > 2 else affinity.get("api_format", "unknown")
)
model_name = (
parts[3] if len(parts) > 3 else affinity.get("model_name", "unknown")
)
affinity_with_key = dict(affinity)
affinity_with_key["affinity_key"] = affinity_key_value
if "api_format" not in affinity_with_key:
affinity_with_key["api_format"] = api_format
if "model_name" not in affinity_with_key:
affinity_with_key["model_name"] = model_name
results.append(affinity_with_key)
# 清理过期的键
if expired_keys:
async with self._get_memory_lock():
for key in expired_keys:
self._memory_store.pop(key, None)
except Exception as e:
logger.exception(f"获取缓存亲和性列表失败: {e}")
return results
# 全局单例
_affinity_manager: CacheAffinityManager | None = None
async def get_affinity_manager(redis_client: Any | None = None) -> CacheAffinityManager:
"""
获取全局CacheAffinityManager实例若Redis不可用则降级为内存模式
Args:
redis_client: Redis客户端可选
Returns:
CacheAffinityManager实例
"""
global _affinity_manager
if _affinity_manager is None:
_affinity_manager = CacheAffinityManager(redis_client)
elif redis_client and _affinity_manager.redis is None:
# 当最初使用内存后 Redis 可用时,升级为 Redis 存储
_affinity_manager = CacheAffinityManager(redis_client)
return _affinity_manager

View File

@@ -0,0 +1,924 @@
"""
缓存感知调度器 (Cache-Aware Scheduler)
职责:
1. 统一管理Provider/Endpoint/Key的选择逻辑
2. 集成缓存亲和性管理优先使用有缓存的Provider+Key
3. 协调并发控制和缓存优先级
4. 实现故障转移机制同Endpoint内优先跨Provider按优先级
核心设计思想:
===============
1. 用户首次请求: 按 provider_priority 选择最优 Provider+Endpoint+Key
2. 用户后续请求:
- 优先使用缓存的Endpoint+Key (利用Prompt Caching)
- 如果缓存的Key并发满尝试同Endpoint其他Key
- 如果Endpoint不可用按 provider_priority 切换到其他Provider
3. 并发控制(动态预留机制):
- 探测阶段使用低预留10%),让系统快速学习真实并发限制
- 稳定阶段根据置信度和负载动态调整预留比例10%-35%
- 置信度因素连续成功次数、429冷却时间、调整历史稳定性
- 缓存用户可使用全部槽位,新用户只能用 (1-预留比例) 的槽位
4. 故障转移:
- Key故障: 同Endpoint内切换其他Key检查模型支持
- Endpoint故障: 按 provider_priority 切换到其他Provider
- 注意不同Endpoint的协议完全不兼容不能在同Provider内切换Endpoint
- 失效缓存亲和性,避免重复选择故障资源
"""
from __future__ import annotations
import time
from typing import TYPE_CHECKING, Any
if TYPE_CHECKING:
from src.services.scheduling.protocols import (
CacheAffinityManagerProtocol,
CandidateBuilderProtocol,
CandidateSorterProtocol,
ConcurrencyCheckerProtocol,
)
from sqlalchemy.orm import Session
from src.core.exceptions import ModelNotSupportedException, ProviderNotAvailableException
from src.core.logger import logger
from src.core.model_permissions import (
check_model_allowed,
get_allowed_models_preview,
)
from src.models.database import (
ApiKey,
Provider,
ProviderAPIKey,
ProviderEndpoint,
)
from src.services.cache.model_cache import ModelCacheService
from src.services.provider.format import normalize_endpoint_signature
from src.services.rate_limit.adaptive_reservation import (
get_adaptive_reservation_manager,
)
from src.services.rate_limit.concurrency_manager import get_concurrency_manager
from src.services.scheduling.affinity_manager import (
get_affinity_manager,
)
from src.services.scheduling.candidate_builder import (
CandidateBuilder,
)
from src.services.scheduling.candidate_builder import (
_sort_endpoints_by_family_priority as _sort_endpoints_by_family_priority,
)
from src.services.scheduling.candidate_sorter import CandidateSorter
from src.services.scheduling.concurrency_checker import ConcurrencyChecker
from src.services.scheduling.restriction_checker import get_effective_restrictions
from src.services.scheduling.scheduling_config import SchedulingConfig
from src.services.scheduling.schemas import ConcurrencySnapshot as ConcurrencySnapshot # re-export
from src.services.scheduling.schemas import PoolCandidate as PoolCandidate # re-export
from src.services.scheduling.schemas import ProviderCandidate as ProviderCandidate # re-export
from src.services.scheduling.utils import affinity_hash as _affinity_hash # re-export compat
from src.services.scheduling.utils import (
release_db_connection_before_await,
)
from src.services.system.config import SystemConfigService
class CacheAwareScheduler:
"""
缓存感知调度器 - 薄协调层
编排以下子组件:
- SchedulingConfig: 调度模式和优先级模式管理
- CandidateBuilder: 候选构建(查询 Provider/Endpoint/Key
- CandidateSorter: 候选排序(优先级/负载均衡)
- ConcurrencyChecker: 并发控制RPM + 动态预留)
- CacheAffinityManager: 缓存亲和性管理
"""
# 类常量 re-export保持外部访问兼容性
PRIORITY_MODE_PROVIDER = SchedulingConfig.PRIORITY_MODE_PROVIDER
PRIORITY_MODE_GLOBAL_KEY = SchedulingConfig.PRIORITY_MODE_GLOBAL_KEY
ALLOWED_PRIORITY_MODES = SchedulingConfig.ALLOWED_PRIORITY_MODES
SCHEDULING_MODE_FIXED_ORDER = SchedulingConfig.SCHEDULING_MODE_FIXED_ORDER
SCHEDULING_MODE_CACHE_AFFINITY = SchedulingConfig.SCHEDULING_MODE_CACHE_AFFINITY
SCHEDULING_MODE_LOAD_BALANCE = SchedulingConfig.SCHEDULING_MODE_LOAD_BALANCE
ALLOWED_SCHEDULING_MODES = SchedulingConfig.ALLOWED_SCHEDULING_MODES
def __init__(
self,
redis_client: Any | None = None,
priority_mode: str | None = None,
scheduling_mode: str | None = None,
*,
candidate_builder: CandidateBuilderProtocol | None = None,
candidate_sorter: CandidateSorterProtocol | None = None,
concurrency_checker: ConcurrencyCheckerProtocol | None = None,
affinity_manager: CacheAffinityManagerProtocol | None = None,
) -> None:
"""
初始化调度器
注意: 不再持久化 db Session,避免跨请求使用已关闭的会话
每个方法调用时需要传入当前请求的 db Session
Args:
redis_client: Redis客户端可选
priority_mode: 候选排序策略provider | global_key
scheduling_mode: 调度模式fixed_order | cache_affinity
"""
self.redis = redis_client
self._config = SchedulingConfig(priority_mode, scheduling_mode)
# 异步子组件(将在第一次使用时初始化,可通过构造函数注入)
self._affinity_manager: CacheAffinityManagerProtocol | None = affinity_manager
self._concurrency_checker: ConcurrencyCheckerProtocol | None = concurrency_checker
self._metrics: dict[str, Any] = {
"total_batches": 0,
"last_batch_size": 0,
"total_candidates": 0,
"last_candidate_count": 0,
"cache_hits": 0,
"cache_misses": 0,
"concurrency_denied": 0,
"last_api_format": None,
"last_model_name": None,
"last_updated_at": None,
# 动态预留相关指标
"reservation_probe_count": 0,
"reservation_stable_count": 0,
"avg_reservation_ratio": 0.0,
"last_reservation_result": None,
}
# 初始化子模块(不传 self解除反向引用可通过构造函数注入
self._candidate_sorter: CandidateSorterProtocol = candidate_sorter or CandidateSorter(
self._config
)
self._candidate_builder: CandidateBuilderProtocol = candidate_builder or CandidateBuilder(
self._candidate_sorter
)
# ── 属性代理(保持外部访问兼容性)──────────────────────────
@property
def priority_mode(self) -> str:
return self._config.priority_mode
@priority_mode.setter
def priority_mode(self, value: str) -> None:
self._config.priority_mode = value
@property
def scheduling_mode(self) -> str:
return self._config.scheduling_mode
@scheduling_mode.setter
def scheduling_mode(self, value: str) -> None:
self._config.scheduling_mode = value
def set_priority_mode(self, mode: str | None) -> None:
"""运行时更新候选排序策略"""
self._config.set_priority_mode(mode)
def set_scheduling_mode(self, mode: str | None) -> None:
"""运行时更新调度模式"""
self._config.set_scheduling_mode(mode)
# ── 静态方法兼容壳 ───────────────────────────────────────
@staticmethod
def _release_db_connection_before_await(db: Session) -> None:
release_db_connection_before_await(db)
@staticmethod
def _affinity_hash(affinity_key: str, identifier: str) -> int:
return _affinity_hash(affinity_key, identifier)
# ── 异步初始化 ───────────────────────────────────────────
async def _ensure_initialized(self) -> None:
"""确保所有异步组件已初始化"""
if self._affinity_manager is None:
self._affinity_manager = await get_affinity_manager(self.redis)
if self._concurrency_checker is None:
concurrency_manager = await get_concurrency_manager()
reservation_manager = get_adaptive_reservation_manager()
self._concurrency_checker = ConcurrencyChecker(
concurrency_manager=concurrency_manager,
reservation_manager=reservation_manager,
)
# ── 核心编排方法 ─────────────────────────────────────────
async def select_with_cache_affinity(
self,
db: Session,
affinity_key: str,
api_format: str,
model_name: str,
excluded_endpoints: list[str] | None = None,
excluded_keys: list[str] | None = None,
provider_batch_size: int = 20,
max_candidates_per_batch: int | None = None,
) -> tuple[Provider, ProviderEndpoint, ProviderAPIKey]:
"""
缓存感知选择 - 核心方法
逻辑:一次性获取所有候选(缓存命中优先),按顺序检查
排除列表和并发限制,返回首个可用组合,并在需要时刷新缓存亲和性。
Args:
db: 数据库会话
affinity_key: 亲和性标识符通常为API Key ID
api_format: API格式
model_name: 模型名称
excluded_endpoints: 排除的Endpoint ID列表
excluded_keys: 排除的Provider Key ID列表
provider_batch_size: Provider批量大小
max_candidates_per_batch: 每批最大候选数
"""
await self._ensure_initialized()
excluded_endpoints_set = set(excluded_endpoints or [])
excluded_keys_set = set(excluded_keys or [])
normalized_format = normalize_endpoint_signature(api_format)
logger.debug(
"[CacheAwareScheduler] select_with_cache_affinity: "
"affinity_key={}..., api_format={}, model={}",
affinity_key[:8],
normalized_format,
model_name,
)
self._metrics["last_api_format"] = normalized_format
self._metrics["last_model_name"] = model_name
provider_offset = 0
global_model_id = None # 用于缓存亲和性
while True:
candidates, resolved_global_model_id, provider_batch_count = (
await self.list_all_candidates(
db=db,
api_format=normalized_format,
model_name=model_name,
affinity_key=affinity_key,
provider_offset=provider_offset,
provider_limit=provider_batch_size,
max_candidates=max_candidates_per_batch,
)
)
if resolved_global_model_id and global_model_id is None:
global_model_id = resolved_global_model_id
if provider_batch_count == 0:
if provider_offset == 0:
raise ProviderNotAvailableException("请求的模型当前不可用")
break
self._metrics["total_batches"] += 1
self._metrics["last_batch_size"] = len(candidates)
self._metrics["last_updated_at"] = int(time.time())
for candidate in candidates:
provider = candidate.provider
endpoint = candidate.endpoint
key = candidate.key
if endpoint.id in excluded_endpoints_set:
logger.debug(" └─ Endpoint {}... 在排除列表,跳过", endpoint.id[:8])
continue
if key.id in excluded_keys_set:
logger.debug(" └─ Key {}... 在排除列表,跳过", key.id[:8])
continue
is_cached_user = bool(candidate.is_cached)
can_use, snapshot = await self._concurrency_checker.check_available(
key,
is_cached_user=is_cached_user,
)
# 更新预留指标
self._update_reservation_metrics(snapshot)
if not can_use:
logger.debug(" └─ Key {}... 并发已满 ({})", key.id[:8], snapshot.describe())
self._metrics["concurrency_denied"] += 1
continue
logger.debug(
" └─ 选择 Provider={}, Endpoint={}..., " "Key={}, 缓存命中={}, 并发状态[{}]",
provider.name,
endpoint.id[:8],
key.name,
is_cached_user,
snapshot.describe(),
)
if key.cache_ttl_minutes > 0 and global_model_id:
ttl = key.cache_ttl_minutes * 60
await self.set_cache_affinity(
affinity_key=affinity_key,
provider_id=str(provider.id),
endpoint_id=str(endpoint.id),
key_id=str(key.id),
api_format=normalized_format,
global_model_id=global_model_id,
ttl=ttl,
)
if is_cached_user:
self._metrics["cache_hits"] += 1
else:
self._metrics["cache_misses"] += 1
return provider, endpoint, key
provider_offset += provider_batch_size
if provider_batch_count < provider_batch_size:
break
raise ProviderNotAvailableException("服务暂时繁忙,请稍后重试")
async def list_all_candidates(
self,
db: Session,
api_format: str,
model_name: str,
affinity_key: str | None = None,
user_api_key: ApiKey | None = None,
provider_offset: int = 0,
provider_limit: int | None = None,
max_candidates: int | None = None,
is_stream: bool = False,
capability_requirements: dict[str, bool] | None = None,
request_body: dict | None = None,
) -> tuple[list[ProviderCandidate], str, int]:
"""
预先获取所有可用的 Provider/Endpoint/Key 组合
编排流程:
1. 解析 GlobalModel
2. 检查访问限制
3. 查询 Providers委托给 CandidateBuilder
4. 构建候选列表(委托给 CandidateBuilder
5. 应用排序和缓存亲和性
Args:
db: 数据库会话
api_format: API 格式
model_name: 模型名称
affinity_key: 亲和性标识符通常为API Key ID用于缓存亲和性
user_api_key: 用户 API Key用于访问限制过滤同时考虑 User 级别限制)
provider_offset: Provider 分页偏移
provider_limit: Provider 分页限制
max_candidates: 最大候选数量
is_stream: 是否是流式请求,如果为 True 则过滤不支持流式的 Provider
capability_requirements: 能力需求(用于过滤不满足能力要求的 Key
Returns:
(候选列表, global_model_id, provider_batch_count)
- global_model_id 用于缓存亲和性
- provider_batch_count 表示本次查询到的 Provider 数量(未应用 allowed_providers 过滤前)
"""
# If the caller already touched the DB, release the connection before we do async work.
release_db_connection_before_await(db)
await self._ensure_initialized()
target_format = normalize_endpoint_signature(api_format)
logger.debug(
"[Scheduler] list_all_candidates: model={}, api_format={}",
model_name,
target_format,
)
# 0. 解析 model_name 到 GlobalModel仅接受 GlobalModel.name
normalized_name = model_name.strip() if isinstance(model_name, str) else ""
if not normalized_name:
logger.warning("GlobalModel not found: <empty model name>")
raise ModelNotSupportedException(model=model_name)
global_model = await ModelCacheService.get_global_model_by_name(db, normalized_name)
if not global_model or not global_model.is_active:
logger.warning("GlobalModel not found or inactive: {}", normalized_name)
raise ModelNotSupportedException(model=model_name)
logger.debug(
"[Scheduler] GlobalModel resolved: id={}, name={}",
global_model.id,
global_model.name,
)
# 使用 GlobalModel.id 作为缓存亲和性的模型标识,确保映射名和规范名都能命中同一个缓存
global_model_id: str = str(global_model.id)
queried_provider_count = 0
# 提取模型映射(用于 Provider Key 的 allowed_models 匹配)
model_mappings: list[str] = (global_model.config or {}).get("model_mappings", [])
if model_mappings:
logger.debug(
"[Scheduler] GlobalModel={} 配置了映射规则: {}",
global_model.name,
model_mappings,
)
# 获取合并后的访问限制ApiKey + User
restrictions = get_effective_restrictions(user_api_key)
allowed_api_formats = restrictions["allowed_api_formats"]
allowed_providers = restrictions["allowed_providers"]
allowed_models = restrictions["allowed_models"]
# 0.1 检查 API 格式是否被允许
if allowed_api_formats is not None:
allowed_norm = {normalize_endpoint_signature(f) for f in allowed_api_formats if f}
if target_format not in allowed_norm:
logger.debug(
"API Key {}... 不允许使用 API 格式 {}, 允许的格式: {}",
user_api_key.id[:8] if user_api_key else "N/A",
target_format,
allowed_api_formats,
)
return [], global_model_id, queried_provider_count
# 0.2 检查模型是否被允许
if not check_model_allowed(
model_name=model_name,
allowed_models=allowed_models,
):
logger.debug(
"用户/API Key 不允许使用模型 {}, 允许的模型: {}",
model_name,
get_allowed_models_preview(allowed_models),
)
return [], global_model_id, queried_provider_count
# 1. 查询 Providers委托给 CandidateBuilder
providers = []
if allowed_providers is not None:
provider_refs = self._candidate_builder._query_provider_refs(
db=db,
provider_offset=provider_offset,
provider_limit=provider_limit,
)
queried_provider_count = len(provider_refs)
allowed_values = {value for value in allowed_providers if value}
matched_provider_ids = [
provider_id
for provider_id, provider_name in provider_refs
if provider_id in allowed_values or provider_name in allowed_values
]
if queried_provider_count != len(matched_provider_ids):
logger.debug(
"用户/API Key 过滤 Provider 预加载范围: {} -> {}",
queried_provider_count,
len(matched_provider_ids),
)
if matched_provider_ids:
providers = self._candidate_builder._query_providers(
db=db,
provider_ids=matched_provider_ids,
)
else:
providers = self._candidate_builder._query_providers(
db=db,
provider_offset=provider_offset,
provider_limit=provider_limit,
)
queried_provider_count = len(providers)
# Provider query starts a transaction; release connection before entering async candidate build.
release_db_connection_before_await(db)
logger.debug(
"[Scheduler] Found {} active providers: {}",
len(providers),
", ".join(p.name for p in providers),
)
if not providers:
return [], global_model_id, queried_provider_count
# 2. 构建候选列表(委托给 CandidateBuilder
# 格式转换总开关(数据库配置):关闭时禁止任何跨格式候选进入队列
global_conversion_enabled = SystemConfigService.is_format_conversion_enabled(db)
candidates = await self._candidate_builder._build_candidates(
db=db,
providers=providers,
client_format=target_format,
model_name=model_name,
model_mappings=model_mappings,
affinity_key=affinity_key,
max_candidates=max_candidates,
is_stream=is_stream,
capability_requirements=capability_requirements,
global_conversion_enabled=global_conversion_enabled,
request_body=request_body,
)
# 3. 应用优先级模式排序 + 调度模式排序
candidates = await self.reorder_candidates(
candidates=candidates,
db=db,
affinity_key=affinity_key,
api_format=target_format,
global_model_id=global_model_id,
)
# 更新指标
self._metrics["total_candidates"] += len(candidates)
self._metrics["last_candidate_count"] = len(candidates)
logger.debug(
"预先获取到 {} 个可用组合 (api_format={}, model={})",
len(candidates),
target_format,
model_name,
)
return candidates, global_model_id, queried_provider_count
async def reorder_candidates(
self,
candidates: list[ProviderCandidate],
db: Session,
affinity_key: str | None = None,
api_format: str | None = None,
global_model_id: str | None = None,
) -> list[ProviderCandidate]:
"""对候选列表应用优先级模式排序和调度模式排序。
在分页汇总后调用此方法可修正跨页排序失真。
Args:
candidates: 候选列表
db: 数据库会话
affinity_key: 亲和性标识符
api_format: API 格式
global_model_id: GlobalModel ID缓存亲和模式需要
Returns:
重排序后的候选列表
"""
if not candidates:
return candidates
# 1. 优先级模式排序(委托给 CandidateSorter
candidates = self._candidate_sorter._apply_priority_mode_sort(
candidates, db, affinity_key, api_format
)
# 排序完成后释放 DB 连接,避免后续 Redis 操作期间占用连接
release_db_connection_before_await(db)
# 2. 调度模式排序
if self.scheduling_mode == self.SCHEDULING_MODE_CACHE_AFFINITY:
if affinity_key and candidates and global_model_id:
candidates = await self._apply_cache_affinity(
candidates=candidates,
db=db,
affinity_key=affinity_key,
api_format=api_format or "",
global_model_id=global_model_id,
)
elif self.scheduling_mode == self.SCHEDULING_MODE_LOAD_BALANCE:
candidates = self._candidate_sorter._apply_load_balance(candidates, api_format)
for candidate in candidates:
candidate.is_cached = False
else:
for candidate in candidates:
candidate.is_cached = False
return candidates
async def _apply_cache_affinity(
self,
candidates: list[ProviderCandidate],
db: Session,
affinity_key: str,
api_format: str,
global_model_id: str,
) -> list[ProviderCandidate]:
"""
应用缓存亲和性排序
缓存命中的候选会被提升到列表前面
Args:
candidates: 候选列表
affinity_key: 亲和性标识符通常为API Key ID
api_format: API 格式
global_model_id: GlobalModel ID规范化的模型标识
Returns:
重排序后的候选列表
"""
try:
# 查询该亲和性标识符在当前 API 格式和模型下的缓存亲和性
api_format_str = str(api_format)
affinity = await self._affinity_manager.get_affinity(
affinity_key, api_format_str, global_model_id
)
if not affinity:
# 没有缓存亲和性,所有候选都标记为非缓存
for candidate in candidates:
candidate.is_cached = False
return candidates
# 判断候选是否应该被降级(用于分组)
global_keep_priority = SystemConfigService.is_keep_priority_on_conversion(db)
def should_demote(c: ProviderCandidate) -> bool:
"""判断候选是否应该被降级"""
if global_keep_priority:
return False # 全局开启时,所有候选都不降级
if not c.needs_conversion:
return False # exact 候选不降级
if getattr(c.provider, "keep_priority_on_conversion", False):
return False # 提供商配置了保持优先级
return True # 需要降级
# 按是否匹配缓存亲和性分类候选,同时记录是否降级
matched_candidate: ProviderCandidate | None = None
matched = False
for candidate in candidates:
provider = candidate.provider
endpoint = candidate.endpoint
key = candidate.key
is_pool_candidate = isinstance(candidate, PoolCandidate)
pool_matched = (
is_pool_candidate
and provider.id == affinity.provider_id
and endpoint.id == affinity.endpoint_id
)
key_matched = (
(not is_pool_candidate)
and provider.id == affinity.provider_id
and endpoint.id == affinity.endpoint_id
and key.id == affinity.key_id
)
if pool_matched or key_matched:
candidate.is_cached = True
matched_candidate = candidate
matched = True
logger.debug(
"检测到缓存亲和性: affinity_key={}..., "
"api_format={}, global_model_id={}..., "
"provider={}, endpoint={}..., "
"provider_key={}, 使用次数={}",
affinity_key[:8],
api_format_str,
global_model_id[:8],
provider.name,
endpoint.id[:8],
key.name,
affinity.request_count,
)
else:
candidate.is_cached = False
if not matched:
logger.debug("API格式 {} 的缓存亲和性存在但组合不可用", api_format_str)
return candidates
# 缓存亲和性命中且该候选可用(未被跳过)时,无条件优先使用
# 理由1) 它之前成功过2) 它有 prompt cache 优势
# 只有当缓存亲和性的候选被跳过(健康度太低/熔断)时,才按 exact 优先排序
assert matched_candidate is not None # guaranteed by matched=True
if not matched_candidate.is_skipped:
# 缓存命中且健康,无条件提升到最前面
other_candidates = [c for c in candidates if c is not matched_candidate]
result = [matched_candidate] + other_candidates
logger.debug(
"缓存亲和性命中且健康,无条件优先使用 (needs_conversion={})",
matched_candidate.needs_conversion,
)
return result
# 缓存命中但被跳过(不健康),按 exact 优先排序
# 缓存候选在其所属类别内提升到最前面
logger.debug(
"缓存亲和性命中但不健康 (skip_reason={}),按 exact 优先排序",
matched_candidate.skip_reason,
)
matched_should_demote = should_demote(matched_candidate)
# 分组:非降级类 和 降级类
keep_priority_candidates: list[ProviderCandidate] = []
demote_candidates: list[ProviderCandidate] = []
for c in candidates:
if c is matched_candidate:
continue # 先跳过缓存命中的候选
if should_demote(c):
demote_candidates.append(c)
else:
keep_priority_candidates.append(c)
# 将缓存命中的候选插入到其所属类别的最前面
if matched_should_demote:
# 缓存命中的是降级类,插入到降级类最前面
demote_candidates.insert(0, matched_candidate)
else:
# 缓存命中的是非降级类,插入到非降级类最前面
keep_priority_candidates.insert(0, matched_candidate)
result = keep_priority_candidates + demote_candidates
logger.debug("缓存组合已提升至其类别内优先级 (demote={})", matched_should_demote)
return result
except Exception as e:
logger.warning("检查缓存亲和性失败: {},继续使用默认排序", e)
return candidates
# ── 委托方法(外部 API 兼容)──────────────────────────────
async def invalidate_cache(
self,
affinity_key: str,
api_format: str,
global_model_id: str,
endpoint_id: str | None = None,
key_id: str | None = None,
provider_id: str | None = None,
) -> Any:
"""
失效指定亲和性标识符对特定API格式和模型的缓存亲和性
Args:
affinity_key: 亲和性标识符通常为API Key ID
api_format: API格式 (claude/openai)
global_model_id: GlobalModel ID规范化的模型标识
endpoint_id: 端点ID可选如果提供则只在Endpoint匹配时失效
key_id: Provider Key ID可选
provider_id: Provider ID可选
"""
await self._ensure_initialized()
await self._affinity_manager.invalidate_affinity(
affinity_key=affinity_key,
api_format=api_format,
model_name=global_model_id,
endpoint_id=endpoint_id,
key_id=key_id,
provider_id=provider_id,
)
async def set_cache_affinity(
self,
affinity_key: str,
provider_id: str,
endpoint_id: str,
key_id: str,
api_format: str,
global_model_id: str,
ttl: int | None = None,
) -> Any:
"""
记录缓存亲和性(供编排器调用)
Args:
affinity_key: 亲和性标识符通常为API Key ID
provider_id: Provider ID
endpoint_id: Endpoint ID
key_id: Provider Key ID
api_format: API格式
global_model_id: GlobalModel ID规范化的模型标识
ttl: 缓存TTL
注意:每次调用都会刷新过期时间,实现滑动窗口机制
"""
await self._ensure_initialized()
await self._affinity_manager.set_affinity(
affinity_key=affinity_key,
provider_id=provider_id,
endpoint_id=endpoint_id,
key_id=key_id,
api_format=api_format,
model_name=global_model_id,
supports_caching=True,
ttl=ttl,
)
# ── 指标 ────────────────────────────────────────────────
def _update_reservation_metrics(self, snapshot: ConcurrencySnapshot) -> None:
"""根据并发检查结果更新预留相关指标"""
if snapshot.reservation_phase == "probe":
self._metrics["reservation_probe_count"] += 1
elif snapshot.reservation_phase != "unknown":
self._metrics["reservation_stable_count"] += 1
# 计算移动平均预留比例
total_reservations = (
self._metrics["reservation_probe_count"] + self._metrics["reservation_stable_count"]
)
if total_reservations > 0:
alpha = 0.1
self._metrics["avg_reservation_ratio"] = (
alpha * snapshot.reservation_ratio
+ (1 - alpha) * self._metrics["avg_reservation_ratio"]
)
self._metrics["last_reservation_result"] = {
"ratio": snapshot.reservation_ratio,
"phase": snapshot.reservation_phase,
"confidence": snapshot.reservation_confidence,
"load_factor": snapshot.load_factor,
}
async def get_stats(self) -> dict:
"""获取调度器统计信息"""
await self._ensure_initialized()
affinity_stats = self._affinity_manager.get_stats()
metrics = dict(self._metrics)
cache_total = metrics["cache_hits"] + metrics["cache_misses"]
metrics["cache_hit_rate"] = metrics["cache_hits"] / cache_total if cache_total else 0.0
metrics["avg_candidates_per_batch"] = (
metrics["total_candidates"] / metrics["total_batches"]
if metrics["total_batches"]
else 0.0
)
# 动态预留统计
reservation_stats = self._concurrency_checker.get_reservation_stats()
total_reservation_checks = (
metrics["reservation_probe_count"] + metrics["reservation_stable_count"]
)
if total_reservation_checks > 0:
probe_ratio = metrics["reservation_probe_count"] / total_reservation_checks
else:
probe_ratio = 0.0
return {
"scheduler": "cache_aware",
"dynamic_reservation": {
"enabled": True,
"config": reservation_stats["config"],
"current_avg_ratio": round(metrics["avg_reservation_ratio"], 3),
"probe_phase_ratio": round(probe_ratio, 3),
"total_checks": total_reservation_checks,
"last_result": metrics["last_reservation_result"],
},
"affinity_stats": affinity_stats,
"scheduler_metrics": metrics,
}
# 全局单例
_scheduler: CacheAwareScheduler | None = None
async def get_cache_aware_scheduler(
redis_client: Any | None = None,
priority_mode: str | None = None,
scheduling_mode: str | None = None,
) -> CacheAwareScheduler:
"""
获取全局CacheAwareScheduler实例
注意: 不再接受 db 参数,避免持久化请求级别的 Session
每次调用 scheduler 方法时需要传入当前请求的 db Session
Args:
redis_client: Redis客户端可选
priority_mode: 外部覆盖的优先级模式provider | global_key
scheduling_mode: 外部覆盖的调度模式fixed_order | cache_affinity
Returns:
CacheAwareScheduler实例
"""
global _scheduler
if _scheduler is None:
_scheduler = CacheAwareScheduler(
redis_client, priority_mode=priority_mode, scheduling_mode=scheduling_mode
)
else:
if priority_mode:
_scheduler.set_priority_mode(priority_mode)
if scheduling_mode:
_scheduler.set_scheduling_mode(scheduling_mode)
return _scheduler

View File

@@ -0,0 +1,749 @@
"""
候选构建器 (CandidateBuilder)
从 CacheAwareScheduler 拆分出的候选构建逻辑,负责:
- 查询活跃 Provider
- 检查模型支持
- 检查 Key 可用性
- 构建候选列表
"""
from __future__ import annotations
import re
from collections.abc import Sequence
from typing import TYPE_CHECKING
from sqlalchemy import or_
from sqlalchemy.orm import Session, selectinload
from src.core.api_format.conversion.compatibility import is_format_compatible
from src.core.api_format.enums import EndpointKind
from src.core.api_format.signature import make_signature_key, parse_signature_key
from src.core.key_capabilities import (
CapabilityMatchMode,
check_capability_match,
compute_capability_score,
get_capability,
)
from src.core.logger import logger
from src.core.model_permissions import check_model_allowed_with_mappings
from src.models.database import (
Model,
Provider,
ProviderAPIKey,
ProviderEndpoint,
)
from src.services.health.monitor import get_health_monitor
from src.services.provider.format import normalize_endpoint_signature
from src.services.provider.pool.account_state import (
resolve_pool_account_state as _resolve_pool_account_state,
)
from src.services.scheduling.quota_skipper import is_key_quota_exhausted
from src.services.scheduling.utils import release_db_connection_before_await
if TYPE_CHECKING:
from src.models.database import GlobalModel
from src.services.provider.pool.config import PoolConfig
from src.services.scheduling.protocols import CandidateSorterProtocol
from src.services.scheduling.schemas import ProviderCandidate
from src.services.cache.model_cache import ModelCacheService
def _get_pool_config(provider: Provider) -> "PoolConfig | None":
"""Return parsed PoolConfig if the provider has pool enabled, else None."""
from src.services.provider.pool.config import parse_pool_config
return parse_pool_config(getattr(provider, "config", None))
def _sort_endpoints_by_family_priority(
eps: Sequence[ProviderEndpoint],
) -> list[ProviderEndpoint]:
"""按 ApiFamily 优先级对端点排序(同分组内使用)。"""
from src.core.api_format.enums import ApiFamily
def sort_key(ep: ProviderEndpoint) -> int:
family_str = str(getattr(ep, "api_family", "") or "").strip().lower()
try:
return ApiFamily(family_str).priority
except ValueError:
return 99
return sorted(eps, key=sort_key)
class CandidateBuilder:
"""候选构建器,负责查询 Provider、检查模型支持和 Key 可用性、构建候选列表。"""
def __init__(self, candidate_sorter: CandidateSorterProtocol) -> None:
self._sorter = candidate_sorter
def _query_provider_refs(
self,
db: Session,
provider_offset: int = 0,
provider_limit: int | None = None,
) -> list[tuple[str, str]]:
"""仅查询当前分页内 Provider 的轻量引用信息。"""
provider_query = (
db.query(Provider.id, Provider.name)
.filter(Provider.is_active.is_(True))
.order_by(Provider.provider_priority.asc())
)
if provider_offset:
provider_query = provider_query.offset(provider_offset)
if provider_limit:
provider_query = provider_query.limit(provider_limit)
return [
(str(provider_id), str(provider_name))
for provider_id, provider_name in provider_query.all()
]
def _query_providers(
self,
db: Session,
provider_offset: int = 0,
provider_limit: int | None = None,
allowed_providers: list[str] | None = None,
provider_ids: list[str] | None = None,
) -> list[Provider]:
"""
查询活跃的 Providers带预加载
Args:
db: 数据库会话
provider_offset: 分页偏移
provider_limit: 分页限制
Returns:
Provider 列表
"""
provider_query = (
db.query(Provider)
.options(
# 预加载 Provider 级别的 api_keys
# defer 排除调度热路径不需要的大 JSON 字段,减少号池场景内存占用
# - 凭证类: api_key/auth_config 在执行阶段由 get_provider_auth() 按需加载
# - adjustment_history/utilization_samples: 仅 AdaptiveReservationManager
# 在并发检查时读取单个 key号池/非号池均可 lazy load
# - upstream_metadata: 号池模式在 _build_candidates 中预计算账号封禁
# 状态并挂到 key._pool_account_state排序阶段不再需要原始 JSON
# 非号池模式 key 少lazy load 可忽略
selectinload(Provider.api_keys)
.defer(ProviderAPIKey.api_key)
.defer(ProviderAPIKey.auth_config)
.defer(ProviderAPIKey.note)
.defer(ProviderAPIKey.last_error_msg)
.defer(ProviderAPIKey.auto_fetch_models)
.defer(ProviderAPIKey.locked_models)
.defer(ProviderAPIKey.model_include_patterns)
.defer(ProviderAPIKey.model_exclude_patterns)
.defer(ProviderAPIKey.last_models_fetch_at)
.defer(ProviderAPIKey.last_models_fetch_error)
.defer(ProviderAPIKey.max_probe_interval_minutes)
.defer(ProviderAPIKey.expires_at)
.defer(ProviderAPIKey.adjustment_history)
.defer(ProviderAPIKey.utilization_samples)
.defer(ProviderAPIKey.upstream_metadata),
# 预加载 endpoints用于按 api_format 选择请求配置)
selectinload(Provider.endpoints),
# 同时加载 models 和 global_model 关系
selectinload(Provider.models).selectinload(Model.global_model),
)
.filter(Provider.is_active.is_(True))
.order_by(Provider.provider_priority.asc())
)
if allowed_providers:
allowed_values = [value for value in allowed_providers if value]
if allowed_values:
provider_query = provider_query.filter(
or_(Provider.id.in_(allowed_values), Provider.name.in_(allowed_values))
)
if provider_ids is not None:
if not provider_ids:
return []
provider_query = provider_query.filter(Provider.id.in_(provider_ids))
if provider_ids is None and provider_offset:
provider_query = provider_query.offset(provider_offset)
if provider_ids is None and provider_limit:
provider_query = provider_query.limit(provider_limit)
providers = provider_query.all()
if provider_ids is None:
return providers
order_map = {provider_id: index for index, provider_id in enumerate(provider_ids)}
providers.sort(key=lambda provider: order_map.get(str(provider.id), len(order_map)))
return providers
async def _check_model_support(
self,
db: Session,
provider: Provider,
model_name: str,
api_format: str | None = None,
is_stream: bool = False,
capability_requirements: dict[str, bool] | None = None,
) -> tuple[bool, str | None, list[str] | None, set[str] | None]:
"""
检查 Provider 是否支持指定模型(可选检查流式支持和能力需求)
模型能力检查在这里进行(而不是在 Key 级别),因为:
- 模型支持的能力是全局的,与具体的 Key 无关
- 如果模型不支持某能力,整个 Provider 的所有 Key 都应该被跳过
仅支持直接匹配 GlobalModel.name外部请求不接受映射名
Args:
db: 数据库会话
provider: Provider 对象
model_name: 模型名称(必须是 GlobalModel.name
is_stream: 是否是流式请求,如果为 True 则同时检查流式支持
capability_requirements: 能力需求(可选),用于检查模型是否支持所需能力
Returns:
(is_supported, skip_reason, supported_capabilities, provider_model_names)
- is_supported: 是否支持
- skip_reason: 跳过原因
- supported_capabilities: 模型支持的能力列表
- provider_model_names: Provider 侧可用的模型名称集合(主名称 + 映射名称,按 api_format 过滤)
"""
# Avoid holding a DB connection while awaiting cache/Redis inside ModelCacheService.
release_db_connection_before_await(db)
# 仅接受 GlobalModel.name不允许映射名
normalized_name = model_name.strip() if isinstance(model_name, str) else ""
if not normalized_name:
return False, "模型不存在或名称无效", None, None
global_model = await ModelCacheService.get_global_model_by_name(db, normalized_name)
if not global_model or not global_model.is_active:
return False, "模型不存在或已停用", None, None
# 找到 GlobalModel 后,检查当前 Provider 是否支持
is_supported, skip_reason, caps, provider_model_names = (
await self._check_model_support_for_global_model(
db,
provider,
global_model,
model_name,
api_format,
is_stream,
capability_requirements,
)
)
return is_supported, skip_reason, caps, provider_model_names
async def _check_model_support_for_global_model(
self,
db: Session,
provider: Provider,
global_model: GlobalModel,
model_name: str,
api_format: str | None = None,
is_stream: bool = False,
capability_requirements: dict[str, bool] | None = None,
) -> tuple[bool, str | None, list[str] | None, set[str] | None]:
"""
检查 Provider 是否支持指定的 GlobalModel
Args:
db: 数据库会话
provider: Provider 对象
global_model: GlobalModel 对象
model_name: 用户请求的模型名称(用于错误消息)
is_stream: 是否是流式请求
capability_requirements: 能力需求
Returns:
(is_supported, skip_reason, supported_capabilities, provider_model_names)
"""
# 确保 global_model 附加到当前 Session
# 注意:从缓存重建的对象是 transient 状态,不能使用 load=False
# 使用 load=True默认允许 SQLAlchemy 正确处理 transient 对象
from sqlalchemy import inspect
insp = inspect(global_model)
if insp.transient or insp.detached:
# transient/detached 对象:使用默认 merge会查询 DB 检查是否存在)
global_model = db.merge(global_model)
else:
# persistent 对象:已经附加到 session无需 merge
pass
# 获取模型支持的能力列表
model_supported_capabilities: list[str] = list(global_model.supported_capabilities or [])
# 查询该 Provider 是否有实现这个 GlobalModel
for model in provider.models:
if model.global_model_id == global_model.id and model.is_active:
# 检查流式支持
if is_stream:
supports_streaming = model.get_effective_supports_streaming()
if not supports_streaming:
return False, f"模型 {model_name} 在此 Provider 不支持流式", None, None
# 检查模型是否支持所需的能力(在 Provider 级别检查,而不是 Key 级别)
# 只有当 model_supported_capabilities 非空时才进行检查
# 空列表意味着模型没有配置能力限制,默认支持所有能力
# COMPATIBLE 能力跳过模型级硬过滤(交由排序阶段处理)
if capability_requirements and model_supported_capabilities:
for cap_name, is_required in capability_requirements.items():
if is_required and cap_name not in model_supported_capabilities:
cap_def = get_capability(cap_name)
if cap_def and cap_def.match_mode == CapabilityMatchMode.COMPATIBLE:
continue
return (
False,
f"模型 {model_name} 不支持能力: {cap_name}",
list(model_supported_capabilities),
None,
)
provider_model_names: set[str] = {model.provider_model_name}
raw_mappings = model.provider_model_mappings
if isinstance(raw_mappings, list):
for raw in raw_mappings:
if not isinstance(raw, dict):
continue
name = raw.get("name")
if not isinstance(name, str) or not name.strip():
continue
mapping_api_formats = raw.get("api_formats")
if api_format and mapping_api_formats:
# 新模式endpoint signaturefamily:kind按小写 canonical 比较
if isinstance(mapping_api_formats, list):
target = str(api_format).strip().lower()
allowed = {
str(fmt).strip().lower() for fmt in mapping_api_formats if fmt
}
if target not in allowed:
continue
provider_model_names.add(name.strip())
return True, None, list(model_supported_capabilities), provider_model_names
return False, "Provider 未实现此模型", None, None
def _check_key_availability(
self,
key: ProviderAPIKey,
api_format: str | None,
model_name: str,
capability_requirements: dict[str, bool] | None = None,
model_mappings: list[str] | None = None,
candidate_models: set[str] | None = None,
*,
provider_type: str | None = None,
) -> tuple[bool, str | None, str | None]:
"""
检查 API Key 的可用性
注意:模型能力检查已移到 _check_model_support 中进行Provider 级别),
这里只检查 Key 级别的能力匹配。
Args:
key: API Key 对象
model_name: 模型名称GlobalModel.name
capability_requirements: 能力需求(可选)
model_mappings: GlobalModel 的映射列表(用于通配符匹配)
candidate_models: Provider 侧可用的模型名称集合(用于限制映射匹配范围)
Returns:
(is_available, skip_reason, mapping_matched_model)
- is_available: Key 是否可用
- skip_reason: 不可用时的原因
- mapping_matched_model: 通过映射匹配到的模型名(用于实际请求)
"""
# 检查熔断器状态(使用详细状态方法获取更丰富的跳过原因,按 API 格式)
is_available, circuit_reason = get_health_monitor().get_circuit_breaker_status(
key, api_format=api_format
)
if not is_available:
return False, circuit_reason or "熔断器已打开", None
# 模型权限检查:使用 allowed_models 白名单
# None = 允许所有模型,[] = 拒绝所有模型,["a","b"] = 只允许指定模型
# 支持通配符映射匹配(通过 model_mappings
try:
is_allowed, mapping_matched_model = check_model_allowed_with_mappings(
model_name=model_name,
allowed_models=key.allowed_models,
model_mappings=model_mappings,
candidate_models=candidate_models,
)
if mapping_matched_model:
logger.debug(
"[Scheduler] Key {}... 模型名匹配: model={} -> {}, allowed_models={}",
key.id[:8],
model_name,
mapping_matched_model,
key.allowed_models,
)
except TimeoutError:
# 正则匹配超时(可能是 ReDoS 攻击或复杂模式)
logger.warning("映射匹配超时: key_id={}, model={}", key.id, model_name)
return False, "映射匹配超时,请简化配置", None
except re.error as e:
# 正则语法错误(配置问题)
logger.warning("映射规则无效: key_id={}, model={}, error={}", key.id, model_name, e)
return False, f"映射规则无效: {str(e)}", None
except Exception as e:
# 其他未知异常
logger.error(
"映射匹配异常: key_id={}, model={}, error={}", key.id, model_name, e, exc_info=True
)
# 异常时保守处理:不允许使用该 Key
return False, "映射匹配失败", None
if not is_allowed:
return (
False,
f"Key 不支持 {model_name}",
None,
)
# Key 级别的能力匹配检查
# 注意:模型级别的能力检查已在 _check_model_support 中完成
# 始终执行检查,即使 capability_requirements 为空
# 因为 check_capability_match 会检查 Key 的 EXCLUSIVE 能力是否被浪费
key_caps: dict[str, bool] = dict(key.capabilities or {})
is_match, skip_reason = check_capability_match(key_caps, capability_requirements)
if not is_match:
return False, skip_reason, None
effective_model_name = mapping_matched_model or model_name
quota_exhausted, quota_reason = is_key_quota_exhausted(
provider_type,
key,
model_name=effective_model_name,
)
if quota_exhausted:
return False, quota_reason, mapping_matched_model
return True, None, mapping_matched_model
async def _build_candidates(
self,
db: Session,
providers: list[Provider],
client_format: str,
model_name: str,
affinity_key: str | None,
model_mappings: list[str] | None = None,
max_candidates: int | None = None,
is_stream: bool = False,
capability_requirements: dict[str, bool] | None = None,
global_conversion_enabled: bool = True,
request_body: dict | None = None,
) -> "list[ProviderCandidate]":
"""
构建候选列表
Key 直属 Provider通过 api_formats 筛选符合端点格式的 Key。
Args:
db: 数据库会话
providers: Provider 列表
client_format: 客户端请求的 API 格式
model_name: 模型名称GlobalModel.name
affinity_key: 亲和性标识符通常为API Key ID
model_mappings: GlobalModel 的映射列表(用于 Key.allowed_models 通配符匹配)
max_candidates: 最大候选数
is_stream: 是否是流式请求,如果为 True 则过滤不支持流式的 Provider
capability_requirements: 能力需求(可选)
global_conversion_enabled: 格式转换全局开关(数据库配置),关闭时回退到 Provider/Endpoint 精细化配置
Returns:
候选列表
"""
from src.services.scheduling.schemas import PoolCandidate, ProviderCandidate
candidates: list[ProviderCandidate] = []
client_format_str = normalize_endpoint_signature(client_format)
client_sig = parse_signature_key(client_format_str)
client_family, client_kind = client_sig.api_family, client_sig.endpoint_kind
# 提取 GlobalModel 配置的 output_limit用于跨格式转换时的 max_tokens 默认值)
output_limit: int | None = None
normalized_name = model_name.strip() if isinstance(model_name, str) else ""
if normalized_name:
gm = await ModelCacheService.get_global_model_by_name(db, normalized_name)
if gm and isinstance(gm.config, dict):
raw = gm.config.get("output_limit")
if isinstance(raw, int) and raw > 0:
output_limit = raw
# chat/cli 互相可回退用于同协议族下的端点变体compact 可回退到 cli。
# video/image 等不跨类回退。
if client_kind in {EndpointKind.CHAT, EndpointKind.CLI}:
allowed_kinds = {EndpointKind.CHAT, EndpointKind.CLI}
elif client_kind == EndpointKind.COMPACT:
allowed_kinds = {EndpointKind.COMPACT, EndpointKind.CLI}
else:
allowed_kinds = {client_kind}
for provider in providers:
# 按端点格式分别判断兼容性与模型/Key 可用性:
# - 同格式端点优先needs_conversion=False
# - 跨格式端点次之needs_conversion=True
model_support_cache: dict[
str, tuple[bool, str | None, list[str] | None, set[str] | None]
] = {}
exact_candidates: list[ProviderCandidate] = []
convertible_candidates: list[ProviderCandidate] = []
pool_cfg = _get_pool_config(provider)
# 使用新架构字段 (api_family, endpoint_kind) 进行预过滤与排序:
# - family/kind 匹配的 endpoint 排在前面(但不做硬过滤,避免破坏格式转换路径)
# - chat/cli 请求允许互相回退(优先同 kind
# - video 等请求只允许同 kind
endpoints = list(provider.endpoints or [])
allowed_kind_values = {k.value for k in allowed_kinds}
preferred: list[ProviderEndpoint] = []
preferred_other_family: list[ProviderEndpoint] = []
fallback: list[ProviderEndpoint] = []
fallback_other_family: list[ProviderEndpoint] = []
for ep in endpoints:
if not getattr(ep, "is_active", False):
continue
raw_family = getattr(ep, "api_family", None)
raw_kind = getattr(ep, "endpoint_kind", None)
if not isinstance(raw_family, str) or not raw_family.strip():
continue
if not isinstance(raw_kind, str) or not raw_kind.strip():
continue
ep_family = raw_family.strip().lower()
ep_kind = raw_kind.strip().lower()
if allowed_kind_values and ep_kind not in allowed_kind_values:
continue
same_family = ep_family == client_family.value
same_kind = ep_kind == client_kind.value
if same_kind and same_family:
preferred.append(ep)
elif same_kind:
preferred_other_family.append(ep)
elif same_family:
fallback.append(ep)
else:
fallback_other_family.append(ep)
endpoints = (
_sort_endpoints_by_family_priority(preferred)
+ _sort_endpoints_by_family_priority(preferred_other_family)
+ _sort_endpoints_by_family_priority(fallback)
+ _sort_endpoints_by_family_priority(fallback_other_family)
)
for endpoint in endpoints:
if not endpoint.is_active:
continue
endpoint_format_str = make_signature_key(
str(getattr(endpoint, "api_family", "")).strip().lower(),
str(getattr(endpoint, "endpoint_kind", "")).strip().lower(),
)
# 格式转换开关(从高到低):
# 1) 全局开关 enable_format_conversion=ON -> 允许跨格式(跳过端点检查)
# 2) 全局开关 OFF -> Provider.enable_format_conversion=ON -> 允许跨格式(跳过端点检查)
# 3) 否则 -> 需 Endpoint.format_acceptance_config 显式允许
provider_conversion_enabled = bool(
getattr(provider, "enable_format_conversion", False)
)
skip_endpoint_check = global_conversion_enabled or provider_conversion_enabled
is_compatible, needs_conversion, _compat_reason = is_format_compatible(
client_format_str,
endpoint_format_str,
getattr(endpoint, "format_acceptance_config", None),
is_stream,
global_conversion_enabled,
skip_endpoint_check=skip_endpoint_check,
)
if not is_compatible:
continue
# 检查模型支持(按端点格式过滤 provider_model_mappings
if endpoint_format_str not in model_support_cache:
model_support_cache[endpoint_format_str] = await self._check_model_support(
db,
provider,
model_name,
api_format=endpoint_format_str,
is_stream=is_stream,
capability_requirements=capability_requirements,
)
supports_model, skip_reason, _model_caps, provider_model_names = (
model_support_cache[endpoint_format_str]
)
if not supports_model:
continue
# Key 直属 Provider通过 api_formats 按端点格式筛选
# api_formats=None 视为"全支持"(兼容历史数据)
active_keys = [
key
for key in provider.api_keys
if key.is_active
and (key.api_formats is None or endpoint_format_str in key.api_formats)
]
if not active_keys:
continue
use_random = all((key.cache_ttl_minutes or 0) == 0 for key in active_keys)
if pool_cfg is not None:
use_random = False
elif use_random and len(active_keys) > 1:
logger.debug(
" Provider {} 启用 Key 轮换模式 (endpoint_format={}, {} keys)",
provider.name,
endpoint_format_str,
len(active_keys),
)
keys_to_check = self._sorter.shuffle_keys_by_internal_priority(
active_keys, affinity_key, use_random
)
if pool_cfg is not None:
# 号池优化:跳过逐 key 的 _check_key_availability 检查,
# 直接收集全部 active key将检查推迟到 PoolManager 排序后分页执行。
pool_keys = list(keys_to_check)
if not pool_keys:
continue
# 在释放 DB 连接前预计算账号封禁状态并挂到 Key 对象上。
# upstream_metadata 是 deferred 字段,逐条 lazy load 会产生 N+1 查询;
# 这里集中触发后PoolManager 排序时直接读取 _pool_account_state 即可。
provider_type_str = (
str(getattr(provider, "provider_type", "") or "").strip().lower() or None
)
for pk in pool_keys:
setattr(
pk,
"_pool_account_state",
_resolve_pool_account_state(
provider_type=provider_type_str,
upstream_metadata=getattr(pk, "upstream_metadata", None),
oauth_invalid_reason=getattr(pk, "oauth_invalid_reason", None),
),
)
# 释放 DB 连接,因为后续的 PoolManager 排序涉及大量 Redis 操作,
# 避免在 Redis I/O 期间长时间占用 DB 连接池。
release_db_connection_before_await(db)
provider_priority_raw = getattr(provider, "provider_priority", None)
try:
provider_priority = (
int(provider_priority_raw)
if provider_priority_raw is not None
else 999999
)
except Exception:
provider_priority = 999999
try:
pool_priority = (
int(pool_cfg.global_priority)
if pool_cfg.global_priority is not None
else provider_priority
)
except Exception:
pool_priority = provider_priority
pool_candidate = PoolCandidate(
provider=provider,
endpoint=endpoint,
key=pool_keys[0],
pool_keys=pool_keys,
pool_config=pool_cfg,
pool_priority=pool_priority,
needs_conversion=needs_conversion,
provider_api_format=str(endpoint_format_str or ""),
output_limit=output_limit,
capability_miss_count=0,
)
# 打包延迟检查参数,供 PoolManager 排序后分页调用
pool_candidate._deferred_check_params = {
"endpoint_format": endpoint_format_str,
"model_name": model_name,
"capability_requirements": capability_requirements,
"model_mappings": model_mappings,
"candidate_models": provider_model_names,
"provider_type": getattr(provider, "provider_type", None),
}
if needs_conversion:
convertible_candidates.append(pool_candidate)
else:
exact_candidates.append(pool_candidate)
break
for key in keys_to_check:
# Key 级别检查(健康度/熔断按 provider_format bucket
# 传入 provider_model_names 作为 candidate_models
# 用于检查 Key 的 allowed_models 是否支持 Provider 定义的模型名称
is_available, key_skip_reason, mapping_matched_model = (
self._check_key_availability(
key,
endpoint_format_str,
model_name,
capability_requirements,
model_mappings=model_mappings,
candidate_models=provider_model_names,
provider_type=getattr(provider, "provider_type", None),
)
)
candidate = ProviderCandidate(
provider=provider,
endpoint=endpoint,
key=key,
is_skipped=not is_available,
skip_reason=key_skip_reason,
mapping_matched_model=mapping_matched_model,
needs_conversion=needs_conversion,
provider_api_format=str(endpoint_format_str or ""),
output_limit=output_limit,
# is_skipped 候选不参与排序miss_count 无意义,置 0 避免干扰
capability_miss_count=(
compute_capability_score(
key.capabilities or {},
capability_requirements,
)
if is_available
else 0
),
)
if needs_conversion:
convertible_candidates.append(candidate)
else:
exact_candidates.append(candidate)
candidates.extend(exact_candidates)
candidates.extend(convertible_candidates)
# max_candidates 截断应在所有候选收集完成后统一处理,确保优先级排序正确
if max_candidates and len(candidates) > max_candidates:
candidates = candidates[:max_candidates]
return candidates

View File

@@ -0,0 +1,347 @@
"""
候选排序器 (CandidateSorter)
从 CacheAwareScheduler 拆分出的候选排序逻辑,负责:
- 优先级模式排序provider / global_key
- 负载均衡模式排序
- Key 内部按优先级分组打乱
"""
from __future__ import annotations
import random
from collections import defaultdict
from typing import TYPE_CHECKING
from src.services.scheduling.scheduling_config import SchedulingConfig
from src.services.scheduling.schemas import PoolCandidate
from src.services.scheduling.utils import affinity_hash
from src.services.system.config import SystemConfigService
if TYPE_CHECKING:
from collections.abc import Callable
from sqlalchemy.orm import Session
from src.models.database import ProviderAPIKey
from src.services.scheduling.schemas import ProviderCandidate
class CandidateSorter:
"""候选排序器,负责优先级模式排序、负载均衡排序和 Key 内部打乱。"""
def __init__(self, config: SchedulingConfig) -> None:
self._config = config
@staticmethod
def _split_by_capability_match(
candidates: list[ProviderCandidate],
) -> tuple[list[ProviderCandidate], list[ProviderCandidate]]:
"""按 capability_miss_count 分组:完全匹配(0)在前,部分匹配(>0)在后"""
full_match = [c for c in candidates if c.capability_miss_count == 0]
partial_match = [c for c in candidates if c.capability_miss_count > 0]
return full_match, partial_match
def _with_capability_split(
self,
candidates: list[ProviderCandidate],
sort_fn: Callable[..., list[ProviderCandidate]],
*args: object,
**kwargs: object,
) -> list[ProviderCandidate]:
"""通用包装:先按 capability_miss_count 分组,再分别排序后合并"""
if not candidates:
return candidates
full_match, partial_match = self._split_by_capability_match(candidates)
return sort_fn(full_match, *args, **kwargs) + sort_fn(partial_match, *args, **kwargs)
def _apply_priority_mode_sort(
self,
candidates: list[ProviderCandidate],
db: Session,
affinity_key: str | None = None,
api_format: str | None = None,
) -> list[ProviderCandidate]:
"""
根据优先级模式对候选列表排序(数字越小越优先)
排序规则:
0. 按 capability_miss_count 分组:完全匹配(0)在前,部分匹配(>0)在后
1. 如果全局配置 keep_priority_on_conversion=True所有候选保持原优先级
2. 否则,按 needs_conversion 和 provider.keep_priority_on_conversion 分组:
- 保持优先级的候选exact 或 provider.keep_priority_on_conversion=True按原优先级排序
- 需要降级的候选convertible 且 provider.keep_priority_on_conversion=False整体排在后面
3. 在同一组内,按优先级模式排序:
- provider: 按 Provider.provider_priority -> Key.internal_priority 排序
- global_key: 按 Key.global_priority_by_format 排序
"""
if not candidates:
return candidates
return self._with_capability_split(
candidates, self._apply_priority_mode_sort_inner, db, affinity_key, api_format
)
def _apply_priority_mode_sort_inner(
self,
candidates: list[ProviderCandidate],
db: Session,
affinity_key: str | None = None,
api_format: str | None = None,
) -> list[ProviderCandidate]:
"""优先级模式排序的内部实现(不含 capability_miss_count 分组)"""
if not candidates:
return candidates
# 全局配置:如果开启,所有候选保持原优先级
global_keep_priority = SystemConfigService.is_keep_priority_on_conversion(db)
if global_keep_priority:
# 全局开启:不分组,直接按优先级模式排序
if self._config.priority_mode == SchedulingConfig.PRIORITY_MODE_GLOBAL_KEY:
return self._sort_by_global_priority_with_hash(candidates, affinity_key, api_format)
# 提供商优先模式:保持构建时的顺序(已按 provider_priority 排序)
return candidates
# 全局未开启:按是否需要降级分组
# - 不需要降级exact 候选 或 provider.keep_priority_on_conversion=True 的 convertible 候选
# - 需要降级convertible 且 provider.keep_priority_on_conversion=False
keep_priority_candidates: list[ProviderCandidate] = []
demote_candidates: list[ProviderCandidate] = []
for c in candidates:
if not c.needs_conversion:
# exact 候选:不需要降级
keep_priority_candidates.append(c)
elif getattr(c.provider, "keep_priority_on_conversion", False):
# convertible 但提供商配置了保持优先级
keep_priority_candidates.append(c)
else:
# convertible 且未配置保持优先级:降级
demote_candidates.append(c)
if self._config.priority_mode == SchedulingConfig.PRIORITY_MODE_GLOBAL_KEY:
# 全局 Key 优先模式:分别对两组排序后合并
sorted_keep = self._sort_by_global_priority_with_hash(
keep_priority_candidates, affinity_key, api_format
)
sorted_demote = self._sort_by_global_priority_with_hash(
demote_candidates, affinity_key, api_format
)
return sorted_keep + sorted_demote
# 提供商优先模式:保持优先级的在前,降级的在后(各组内部顺序已由构建时保证)
return keep_priority_candidates + demote_candidates
def _sort_by_global_priority_with_hash(
self,
candidates: list[ProviderCandidate],
affinity_key: str | None = None,
api_format: str | None = None,
) -> list[ProviderCandidate]:
"""
按 global_priority_by_format 分组排序,同优先级内通过哈希分散实现负载均衡
排序逻辑:
1. 按 global_priority_by_format[api_format] 分组数字小的优先NULL 排后面)
2. 同优先级组内,使用 affinity_key 哈希分散
3. 确保同一用户请求稳定选择同一个 Key缓存亲和性
"""
def get_priority(candidate: ProviderCandidate) -> int:
"""获取候选的优先级"""
if isinstance(candidate, PoolCandidate):
return int(getattr(candidate, "pool_priority", 999999) or 999999)
if not candidate.key:
return 999999
priority_by_format = candidate.key.global_priority_by_format or {}
if api_format and api_format in priority_by_format:
return priority_by_format[api_format]
return 999999 # NULL 排在后面
# 按优先级分组
priority_groups: dict[int, list[ProviderCandidate]] = defaultdict(list)
for candidate in candidates:
priority = get_priority(candidate)
priority_groups[priority].append(candidate)
result = []
for priority in sorted(priority_groups.keys()): # 数字小的优先级高
group = priority_groups[priority]
if len(group) > 1 and affinity_key:
# 同优先级内哈希分散负载均衡
scored_candidates = []
for candidate in group:
if isinstance(candidate, PoolCandidate):
hash_id = str(getattr(candidate.provider, "id", "") or "")
else:
hash_id = candidate.key.id if candidate.key else ""
hash_value = affinity_hash(affinity_key, hash_id)
scored_candidates.append((hash_value, candidate))
# 按哈希值排序
sorted_group = [c for _, c in sorted(scored_candidates, key=lambda x: x[0])]
result.extend(sorted_group)
else:
# 单个候选或没有 affinity_key按次要排序条件排序
def secondary_sort(c: ProviderCandidate) -> tuple[int, int, str]:
pp = c.provider.provider_priority
if isinstance(c, PoolCandidate):
ip = int(getattr(c, "pool_priority", 999999) or 999999)
key_id = str(getattr(c.provider, "id", "") or "")
else:
ip = c.key.internal_priority if c.key else None
key_id = c.key.id if c.key else ""
return (
pp if pp is not None else 999999,
ip if ip is not None else 999999,
key_id,
)
result.extend(sorted(group, key=secondary_sort))
return result
def _apply_load_balance(
self, candidates: list[ProviderCandidate], api_format: str | None = None
) -> list[ProviderCandidate]:
"""
负载均衡模式:同优先级内随机轮换
排序逻辑:
0. 按 capability_miss_count 分组:完全匹配(0)在前,部分匹配(>0)在后
1. 按优先级分组provider_priority, internal_priority 或 global_priority_by_format
2. 同优先级组内随机打乱
3. 不考虑缓存亲和性
"""
if not candidates:
return candidates
return self._with_capability_split(candidates, self._apply_load_balance_inner, api_format)
def _apply_load_balance_inner(
self, candidates: list[ProviderCandidate], api_format: str | None = None
) -> list[ProviderCandidate]:
"""负载均衡排序的内部实现(不含 capability_miss_count 分组)"""
if not candidates:
return candidates
priority_groups: dict[tuple, list[ProviderCandidate]] = defaultdict(list)
# 根据优先级模式选择分组方式
if self._config.priority_mode == SchedulingConfig.PRIORITY_MODE_GLOBAL_KEY:
# 全局 Key 优先模式:按格式特定优先级分组
for candidate in candidates:
if isinstance(candidate, PoolCandidate):
priority = int(getattr(candidate, "pool_priority", 999999) or 999999)
# -1 使号池候选独立成组,不与普通 key 候选 (0) 混组打乱
priority_groups[(priority, -1)].append(candidate)
continue
else:
priority = 999999
if candidate.key:
priority_by_format = candidate.key.global_priority_by_format or {}
if api_format and api_format in priority_by_format:
priority = priority_by_format[api_format]
priority_groups[(priority, 0)].append(candidate)
else:
# 提供商优先模式:按 (provider_priority, internal_priority) 分组
for candidate in candidates:
pp = candidate.provider.provider_priority
if isinstance(candidate, PoolCandidate):
# 号池候选独立成组,不与普通 key 候选混组打乱。
ip = -1
else:
ip = candidate.key.internal_priority if candidate.key else None
key = (
pp if pp is not None else 999999,
ip if ip is not None else 999999,
)
priority_groups[key].append(candidate)
result: list[ProviderCandidate] = []
for priority in sorted(priority_groups.keys()):
group = priority_groups[priority]
if len(group) > 1:
# 同优先级内随机打乱
shuffled = list(group)
random.shuffle(shuffled)
result.extend(shuffled)
else:
result.extend(group)
return result
def shuffle_keys_by_internal_priority(
self,
keys: list[ProviderAPIKey],
affinity_key: str | None = None,
use_random: bool = False,
) -> list[ProviderAPIKey]:
return self._shuffle_keys_by_internal_priority(keys, affinity_key, use_random)
def _shuffle_keys_by_internal_priority(
self,
keys: list[ProviderAPIKey],
affinity_key: str | None = None,
use_random: bool = False,
) -> list[ProviderAPIKey]:
"""
对 API Key 按 internal_priority 分组,同优先级内部基于 affinity_key 进行确定性打乱
目的:
- 数字越小越优先使用
- 同优先级 Key 之间实现负载均衡
- 使用 affinity_key 哈希确保同一请求 Key 的请求稳定(避免破坏缓存亲和性)
- 当 use_random=True 时,使用随机排序实现轮换(用于 TTL=0 的场景)
Args:
keys: API Key 列表
affinity_key: 亲和性标识符(通常为 API Key ID用于确定性打乱
use_random: 是否使用随机排序TTL=0 时为 True
Returns:
排序后的 Key 列表
"""
if not keys:
return []
# 按 internal_priority 分组
priority_groups: dict[int, list[ProviderAPIKey]] = defaultdict(list)
for key in keys:
priority = key.internal_priority if key.internal_priority is not None else 999999
priority_groups[priority].append(key)
# 对每个优先级组内的 Key 进行打乱
result = []
for priority in sorted(priority_groups.keys()): # 数字小的优先级高,排前面
group_keys = priority_groups[priority]
if len(group_keys) > 1:
should_randomize = (
use_random
or self._config.scheduling_mode == SchedulingConfig.SCHEDULING_MODE_LOAD_BALANCE
or not affinity_key
)
if should_randomize:
# 随机排序TTL=0 / 负载均衡模式 / 无 affinity_key
shuffled = list(group_keys)
random.shuffle(shuffled)
result.extend(shuffled)
else:
# 缓存亲和模式使用哈希确定性排序should_randomize=False 蕴含 affinity_key 非空)
key_scores = []
for key in group_keys:
hash_value = affinity_hash(affinity_key, key.id) # type: ignore[arg-type]
key_scores.append((hash_value, key))
sorted_group = [key for _, key in sorted(key_scores, key=lambda x: x[0])]
result.extend(sorted_group)
else:
# 单个 Key 直接添加
result.extend(group_keys)
return result

View File

@@ -0,0 +1,144 @@
"""
并发控制检查器 (ConcurrencyChecker)
从 CacheAwareScheduler 提取的 RPM 限流和动态预留逻辑。
"""
from __future__ import annotations
import math
from typing import Any
from src.core.logger import logger
from src.models.database import ProviderAPIKey
from src.services.rate_limit.adaptive_reservation import AdaptiveReservationManager
from src.services.rate_limit.adaptive_rpm import get_adaptive_rpm_manager
from src.services.scheduling.schemas import ConcurrencySnapshot
class ConcurrencyChecker:
"""并发控制检查器,封装 RPM 限流和动态预留逻辑。"""
def __init__(
self,
concurrency_manager: Any,
reservation_manager: AdaptiveReservationManager,
) -> None:
self._concurrency_manager = concurrency_manager
self._reservation_manager = reservation_manager
@staticmethod
def get_effective_rpm_limit(key: ProviderAPIKey) -> int | None:
"""获取有效的 RPM 限制(委托给 AdaptiveRPMManager 统一逻辑)"""
return get_adaptive_rpm_manager().get_effective_limit(key)
async def check_available(
self,
key: ProviderAPIKey,
is_cached_user: bool = False,
) -> tuple[bool, ConcurrencySnapshot]:
"""
检查 RPM 限制是否可用(使用动态预留机制)
核心逻辑 - 动态缓存预留机制:
- 总槽位: 有效 RPM 限制(固定值或学习到的值)
- 预留比例: 由 AdaptiveReservationManager 根据置信度和负载动态计算
- 缓存用户可用: 全部槽位
- 新用户可用: 总槽位 x (1 - 动态预留比例)
Args:
key: ProviderAPIKey对象
is_cached_user: 是否是缓存用户
Returns:
(是否可用, 并发快照)
"""
# 获取有效的并发限制
effective_key_limit = self.get_effective_rpm_limit(key)
logger.debug(
" -> 并发检查: _concurrency_manager={}, "
"is_cached_user={}, effective_limit={}",
self._concurrency_manager is not None,
is_cached_user,
effective_key_limit,
)
if not self._concurrency_manager:
# 并发管理器不可用直接返回True
logger.debug(" -> 无并发管理器,直接通过")
snapshot = ConcurrencySnapshot(
key_current=0,
key_limit=effective_key_limit,
is_cached_user=is_cached_user,
)
return True, snapshot
# 获取当前 RPM 计数
key_count = await self._concurrency_manager.get_key_rpm_count(
key_id=str(key.id),
)
can_use = True
# 计算动态预留比例
reservation_result = self._reservation_manager.calculate_reservation(
key=key,
current_usage=key_count,
effective_limit=effective_key_limit,
)
available_for_new = None
reservation_ratio = reservation_result.ratio
# 检查Key级别限制使用动态预留比例
if effective_key_limit is not None:
if is_cached_user:
# 缓存用户: 可以使用全部槽位
if key_count >= effective_key_limit:
can_use = False
else:
# 新用户: 只能使用 (1 - 动态预留比例) 的槽位
# 使用 max 确保至少有 1 个槽位可用
# 与 ConcurrencyManager 的 Lua 脚本保持一致:使用 floor 计算新用户可用槽位
available_for_new = max(
1, math.floor(effective_key_limit * (1 - reservation_ratio))
)
if key_count >= available_for_new:
logger.debug(
"Key {}... 新用户配额已满 " "({}/{}, 总{}, 预留{:.0%}[{}])",
key.id[:8],
key_count,
available_for_new,
effective_key_limit,
reservation_ratio,
reservation_result.phase,
)
can_use = False
key_limit_for_snapshot: int | None
if is_cached_user:
key_limit_for_snapshot = effective_key_limit
elif effective_key_limit is not None:
key_limit_for_snapshot = (
available_for_new if available_for_new is not None else effective_key_limit
)
else:
key_limit_for_snapshot = None
snapshot = ConcurrencySnapshot(
key_current=key_count,
key_limit=key_limit_for_snapshot,
is_cached_user=is_cached_user,
reservation_ratio=reservation_ratio,
reservation_phase=reservation_result.phase,
reservation_confidence=reservation_result.confidence,
load_factor=reservation_result.load_factor,
)
return can_use, snapshot
def get_reservation_stats(self) -> dict[str, Any]:
"""获取动态预留管理器的统计信息"""
return self._reservation_manager.get_stats()

View File

@@ -0,0 +1,146 @@
"""调度/候选子组件的协议接口。
目的:
- 用 `Protocol` 固化 CacheAwareScheduler 的子组件契约
- 便于单测注入 stub/mocks减少对具体实现类的耦合
说明:这里的协议面向“调度器内部协作”,因此保留了部分 `_` 前缀方法。
后续如果要对外暴露更稳定的 API可再抽出无下划线的 facade。
"""
from __future__ import annotations
from typing import TYPE_CHECKING, Any, Protocol
if TYPE_CHECKING:
from sqlalchemy.orm import Session
from src.models.database import GlobalModel, Provider, ProviderAPIKey
from src.services.scheduling.affinity_manager import CacheAffinity
from src.services.scheduling.schemas import ConcurrencySnapshot, ProviderCandidate
class CandidateSorterProtocol(Protocol):
def _apply_priority_mode_sort(
self,
candidates: list[ProviderCandidate],
db: Session,
affinity_key: str | None = None,
api_format: str | None = None,
) -> list[ProviderCandidate]: ...
def _apply_load_balance(
self, candidates: list[ProviderCandidate], api_format: str | None = None
) -> list[ProviderCandidate]: ...
def shuffle_keys_by_internal_priority(
self,
keys: list[ProviderAPIKey],
affinity_key: str | None = None,
use_random: bool = False,
) -> list[ProviderAPIKey]: ...
class CandidateBuilderProtocol(Protocol):
def _query_provider_refs(
self,
db: Session,
provider_offset: int = 0,
provider_limit: int | None = None,
) -> list[tuple[str, str]]: ...
def _query_providers(
self,
db: Session,
provider_offset: int = 0,
provider_limit: int | None = None,
allowed_providers: list[str] | None = None,
provider_ids: list[str] | None = None,
) -> list[Provider]: ...
async def _build_candidates(
self,
db: Session,
providers: list[Provider],
client_format: str,
model_name: str,
affinity_key: str | None,
model_mappings: list[str] | None = None,
max_candidates: int | None = None,
is_stream: bool = False,
capability_requirements: dict[str, bool] | None = None,
global_conversion_enabled: bool = True,
) -> list[ProviderCandidate]: ...
async def _check_model_support(
self,
db: Session,
provider: Provider,
model_name: str,
api_format: str | None = None,
is_stream: bool = False,
capability_requirements: dict[str, bool] | None = None,
) -> tuple[bool, str | None, list[str] | None, set[str] | None]: ...
async def _check_model_support_for_global_model(
self,
db: Session,
provider: Provider,
global_model: GlobalModel,
model_name: str,
api_format: str | None = None,
is_stream: bool = False,
capability_requirements: dict[str, bool] | None = None,
) -> tuple[bool, str | None, list[str] | None, set[str] | None]: ...
def _check_key_availability(
self,
key: ProviderAPIKey,
api_format: str | None,
model_name: str,
capability_requirements: dict[str, bool] | None = None,
model_mappings: list[str] | None = None,
candidate_models: set[str] | None = None,
*,
provider_type: str | None = None,
) -> tuple[bool, str | None, str | None]: ...
class ConcurrencyCheckerProtocol(Protocol):
async def check_available(
self,
key: ProviderAPIKey,
is_cached_user: bool = False,
) -> tuple[bool, ConcurrencySnapshot]: ...
def get_reservation_stats(self) -> dict[str, Any]: ...
class CacheAffinityManagerProtocol(Protocol):
async def get_affinity(
self, affinity_key: str, api_format: str, model_name: str
) -> CacheAffinity | None: ...
async def set_affinity(
self,
affinity_key: str,
provider_id: str,
endpoint_id: str,
key_id: str,
api_format: str,
model_name: str,
supports_caching: bool = True,
ttl: int | None = None,
) -> None: ...
async def invalidate_affinity(
self,
affinity_key: str,
api_format: str,
model_name: str,
key_id: str | None = None,
provider_id: str | None = None,
endpoint_id: str | None = None,
) -> None: ...
def get_stats(self) -> dict[str, Any]: ...

View File

@@ -0,0 +1,17 @@
from __future__ import annotations
from src.models.database import ProviderAPIKey
from src.services.provider_keys.quota_reader import get_quota_reader
def is_key_quota_exhausted(
provider_type: str | None,
key: ProviderAPIKey,
*,
model_name: str,
) -> tuple[bool, str | None]:
"""Check ProviderAPIKey.upstream_metadata quota and decide whether to skip."""
reader = get_quota_reader(provider_type, getattr(key, "upstream_metadata", None))
result = reader.is_exhausted(model_name)
return result.exhausted, result.reason

View File

@@ -0,0 +1,82 @@
"""
访问限制检查
从 CacheAwareScheduler 提取的 ApiKey + User 访问限制合并逻辑。
"""
from __future__ import annotations
from typing import Any
from src.core.logger import logger
from src.core.model_permissions import merge_allowed_models
from src.models.database import ApiKey
def get_effective_restrictions(user_api_key: ApiKey | None) -> dict[str, Any]:
"""
获取有效的访问限制(合并 ApiKey 和 User 的限制)
逻辑:
- 如果 ApiKey 和 User 都有限制,取交集
- 如果只有一方有限制,使用该方的限制
- 如果都没有限制,返回 None表示不限制
Args:
user_api_key: 用户 API Key 对象(可能包含 user relationship
Returns:
包含 allowed_providers, allowed_models, allowed_api_formats 的字典
"""
result: dict[str, Any] = {
"allowed_providers": None,
"allowed_models": None,
"allowed_api_formats": None,
}
if not user_api_key:
return result
# 获取 User 的限制
# 注意:这里可能触发 lazy loading需要确保 session 仍然有效
try:
user = user_api_key.user if hasattr(user_api_key, "user") else None
except Exception as e:
logger.warning("无法加载 ApiKey 关联的 User: {},仅使用 ApiKey 级别的限制", e)
user = None
# 调试日志
logger.debug(
"[_get_effective_restrictions] ApiKey={}..., User={}..., "
"ApiKey.allowed_models={}, User.allowed_models={}",
user_api_key.id[:8],
user.id[:8] if user else "None",
user_api_key.allowed_models,
user.allowed_models if user else "N/A",
)
# 合并 allowed_providers
result["allowed_providers"] = merge_restriction_sets(
user_api_key.allowed_providers, user.allowed_providers if user else None
)
# 合并 allowed_models取交集
result["allowed_models"] = merge_allowed_models(
user_api_key.allowed_models, user.allowed_models if user else None
)
# 合并 allowed_api_formats
result["allowed_api_formats"] = merge_restriction_sets(
user_api_key.allowed_api_formats, user.allowed_api_formats if user else None
)
return result
def merge_restriction_sets(key_restriction: Any, user_restriction: Any) -> set[Any] | None:
"""合并两个限制列表,取交集;任一方为空则使用另一方;均空返回 None"""
key_set = set(key_restriction) if key_restriction else None
user_set = set(user_restriction) if user_restriction else None
if key_set and user_set:
return key_set & user_set
return key_set or user_set

View File

@@ -0,0 +1,82 @@
"""
调度配置 (SchedulingConfig)
从 CacheAwareScheduler 提取的常量定义和模式管理逻辑。
"""
from __future__ import annotations
from src.core.logger import logger
class SchedulingConfig:
"""调度配置:管理优先级模式和调度模式的常量、归一化和运行时更新。"""
# 优先级模式常量
PRIORITY_MODE_PROVIDER = "provider" # 提供商优先模式
PRIORITY_MODE_GLOBAL_KEY = "global_key" # 全局 Key 优先模式
ALLOWED_PRIORITY_MODES = {
PRIORITY_MODE_PROVIDER,
PRIORITY_MODE_GLOBAL_KEY,
}
# 调度模式常量
SCHEDULING_MODE_FIXED_ORDER = "fixed_order" # 固定顺序模式:严格按优先级,忽略缓存
SCHEDULING_MODE_CACHE_AFFINITY = "cache_affinity" # 缓存亲和模式:优先缓存,同优先级哈希分散
SCHEDULING_MODE_LOAD_BALANCE = "load_balance" # 负载均衡模式:忽略缓存,同优先级随机轮换
ALLOWED_SCHEDULING_MODES = {
SCHEDULING_MODE_FIXED_ORDER,
SCHEDULING_MODE_CACHE_AFFINITY,
SCHEDULING_MODE_LOAD_BALANCE,
}
def __init__(
self,
priority_mode: str | None = None,
scheduling_mode: str | None = None,
) -> None:
self.priority_mode = self._normalize_priority_mode(
priority_mode or self.PRIORITY_MODE_PROVIDER
)
self.scheduling_mode = self._normalize_scheduling_mode(
scheduling_mode or self.SCHEDULING_MODE_CACHE_AFFINITY
)
logger.debug(
"[SchedulingConfig] 初始化优先级模式: {}, 调度模式: {}",
self.priority_mode,
self.scheduling_mode,
)
def _normalize_priority_mode(self, mode: str | None) -> str:
normalized = (mode or "").strip().lower()
if normalized not in self.ALLOWED_PRIORITY_MODES:
if normalized:
logger.warning("[SchedulingConfig] 无效的优先级模式 '{}',回退为 provider", mode)
return self.PRIORITY_MODE_PROVIDER
return normalized
def _normalize_scheduling_mode(self, mode: str | None) -> str:
normalized = (mode or "").strip().lower()
if normalized not in self.ALLOWED_SCHEDULING_MODES:
if normalized:
logger.warning(
"[SchedulingConfig] 无效的调度模式 '{}',回退为 cache_affinity", mode
)
return self.SCHEDULING_MODE_CACHE_AFFINITY
return normalized
def set_priority_mode(self, mode: str | None) -> None:
"""运行时更新候选排序策略"""
normalized = self._normalize_priority_mode(mode)
if normalized == self.priority_mode:
return
self.priority_mode = normalized
logger.debug("[SchedulingConfig] 切换优先级模式为: {}", self.priority_mode)
def set_scheduling_mode(self, mode: str | None) -> None:
"""运行时更新调度模式"""
normalized = self._normalize_scheduling_mode(mode)
if normalized == self.scheduling_mode:
return
self.scheduling_mode = normalized
logger.debug("[SchedulingConfig] 切换调度模式为: {}", self.scheduling_mode)

View File

@@ -0,0 +1,109 @@
"""
调度器核心数据类型
从 CacheAwareScheduler 提取的共享数据结构,被 24+ 个模块使用。
"""
from __future__ import annotations
from dataclasses import dataclass, field
from typing import TYPE_CHECKING, Any
from src.models.database import (
Provider,
ProviderAPIKey,
ProviderEndpoint,
)
if TYPE_CHECKING:
from src.services.provider.pool.config import PoolConfig
@dataclass
class ProviderCandidate:
"""候选 provider 组合及是否命中缓存"""
provider: Provider
endpoint: ProviderEndpoint
key: ProviderAPIKey
is_cached: bool = False
is_skipped: bool = False # 是否被跳过
skip_reason: str | None = None # 跳过原因
mapping_matched_model: str | None = None # 通过映射匹配到的模型名(用于实际请求)
needs_conversion: bool = False # 是否需要格式转换
provider_api_format: str = "" # Provider 端点实际格式(用于健康度/熔断 bucket
output_limit: int | None = None # GlobalModel 配置的模型输出上限
capability_miss_count: int = 0 # COMPATIBLE 能力不匹配数0=完全匹配,用于排序)
def _stable_order_key(self) -> tuple[int, int, str, str, str]:
"""
为排序/优先队列提供稳定的比较键。
说明:
- 运行时偶发会出现对 ProviderCandidate 做 tuple 排序/heap 排序的场景;
当主键相同需要比较候选本身时,若候选不可比较会触发:
TypeError: '<' not supported between instances of 'ProviderCandidate' and 'ProviderCandidate'
- 这里提供一个与调度逻辑无关、但足够稳定且可比的兜底顺序。
"""
provider_priority_raw = getattr(self.provider, "provider_priority", None)
internal_priority_raw = getattr(self.key, "internal_priority", None)
try:
provider_priority = (
int(provider_priority_raw) if provider_priority_raw is not None else 999999
)
except Exception:
provider_priority = 999999
try:
internal_priority = (
int(internal_priority_raw) if internal_priority_raw is not None else 999999
)
except Exception:
internal_priority = 999999
provider_id = str(getattr(self.provider, "id", "") or "")
endpoint_id = str(getattr(self.endpoint, "id", "") or "")
key_id = str(getattr(self.key, "id", "") or "")
return (provider_priority, internal_priority, provider_id, endpoint_id, key_id)
def __lt__(self, other: object) -> bool:
if not isinstance(other, ProviderCandidate):
return NotImplemented
return self._stable_order_key() < other._stable_order_key()
@dataclass
class PoolCandidate(ProviderCandidate):
"""号池候选。
排序阶段作为单个候选参与;执行阶段再在 pool_keys 内部选择/切换 key。
"""
pool_keys: list[ProviderAPIKey] = field(default_factory=list)
pool_config: PoolConfig | None = None
pool_priority: int = 999999
_pool_key_index: int = 0
# 延迟可用性检查参数(号池优化:先排序再分页检查)
_deferred_check_params: dict[str, Any] | None = field(default=None, repr=False)
@dataclass
class ConcurrencySnapshot:
key_current: int
key_limit: int | None
is_cached_user: bool = False
# 动态预留信息
reservation_ratio: float = 0.0
reservation_phase: str = "unknown"
reservation_confidence: float = 0.0
load_factor: float = 0.0
def describe(self) -> str:
key_limit_text = str(self.key_limit) if self.key_limit is not None else "inf"
reservation_text = f"{self.reservation_ratio:.0%}" if self.reservation_ratio > 0 else "N/A"
return (
f"key={self.key_current}/{key_limit_text}, "
f"cached={self.is_cached_user}, "
f"reserve={reservation_text}({self.reservation_phase})"
)

View File

@@ -0,0 +1,53 @@
"""
调度器工具函数
从 CacheAwareScheduler 提取的静态工具方法。
"""
from __future__ import annotations
import hashlib
from sqlalchemy.orm import Session
def affinity_hash(affinity_key: str, identifier: str) -> int:
"""基于 affinity_key 和标识符的确定性哈希(用于同优先级内分散负载均衡)"""
return int(hashlib.sha256(f"{affinity_key}:{identifier}".encode()).hexdigest()[:16], 16)
def release_db_connection_before_await(db: Session) -> None:
"""
Best-effort: end a read-only transaction before awaiting async I/O.
This scheduler does a lot of async work (cache/Redis) mixed with sync SQLAlchemy reads.
If a SELECT has already started a transaction, the pooled connection can remain checked
out while we await, causing pool pressure under concurrency.
Safety:
- Only commits when the Session has no ORM pending changes.
- Temporarily disables expire_on_commit to keep already-loaded ORM objects usable.
"""
try:
if db is None:
return
has_pending_changes = bool(db.new) or bool(db.dirty) or bool(db.deleted)
if has_pending_changes:
return
if not db.in_transaction():
return
original_expire_on_commit = getattr(db, "expire_on_commit", True)
db.expire_on_commit = False
try:
db.commit()
except Exception:
try:
db.rollback()
except Exception:
pass
finally:
db.expire_on_commit = original_expire_on_commit
except Exception:
# Never let this optimization break scheduling
return