mirror of
https://github.com/fawney19/Aether.git
synced 2026-09-02 01:10:23 +08:00
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:
30
_deprecated_py_src/services/scheduling/__init__.py
Normal file
30
_deprecated_py_src/services/scheduling/__init__.py
Normal 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",
|
||||
]
|
||||
757
_deprecated_py_src/services/scheduling/affinity_manager.py
Normal file
757
_deprecated_py_src/services/scheduling/affinity_manager.py
Normal 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 ID(api_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
|
||||
924
_deprecated_py_src/services/scheduling/aware_scheduler.py
Normal file
924
_deprecated_py_src/services/scheduling/aware_scheduler.py
Normal 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
|
||||
749
_deprecated_py_src/services/scheduling/candidate_builder.py
Normal file
749
_deprecated_py_src/services/scheduling/candidate_builder.py
Normal 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 signature(family: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
|
||||
347
_deprecated_py_src/services/scheduling/candidate_sorter.py
Normal file
347
_deprecated_py_src/services/scheduling/candidate_sorter.py
Normal 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
|
||||
144
_deprecated_py_src/services/scheduling/concurrency_checker.py
Normal file
144
_deprecated_py_src/services/scheduling/concurrency_checker.py
Normal 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()
|
||||
146
_deprecated_py_src/services/scheduling/protocols.py
Normal file
146
_deprecated_py_src/services/scheduling/protocols.py
Normal 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]: ...
|
||||
17
_deprecated_py_src/services/scheduling/quota_skipper.py
Normal file
17
_deprecated_py_src/services/scheduling/quota_skipper.py
Normal 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
|
||||
@@ -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
|
||||
82
_deprecated_py_src/services/scheduling/scheduling_config.py
Normal file
82
_deprecated_py_src/services/scheduling/scheduling_config.py
Normal 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)
|
||||
109
_deprecated_py_src/services/scheduling/schemas.py
Normal file
109
_deprecated_py_src/services/scheduling/schemas.py
Normal 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})"
|
||||
)
|
||||
53
_deprecated_py_src/services/scheduling/utils.py
Normal file
53
_deprecated_py_src/services/scheduling/utils.py
Normal 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
|
||||
Reference in New Issue
Block a user