Files
Aether/_deprecated_py_src/services/scheduling/aware_scheduler.py
fawney19 1d9c77522a 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)
2026-04-03 16:26:16 +08:00

925 lines
36 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
"""
缓存感知调度器 (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