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

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

View File

@@ -0,0 +1,15 @@
"""通用缓存模块。
包含缓存后端、缓存失效与缓存同步等能力backend/sync/*_cache
调度/候选/缓存亲和性相关逻辑已迁移到 `src.services.scheduling`。
"""
from src.services.cache.backend import BaseCacheBackend, LocalCache, RedisCache, get_cache_backend
__all__ = [
"BaseCacheBackend",
"LocalCache",
"RedisCache",
"get_cache_backend",
]

View File

@@ -0,0 +1,338 @@
"""
缓存后端抽象层
提供统一的缓存接口,支持多种后端实现:
1. LocalCache: 内存缓存(单实例,线程安全)
2. RedisCache: Redis 缓存(分布式)
使用场景:
- ModelCacheService: 模型解析缓存
- 其他需要缓存的服务
"""
from __future__ import annotations
import asyncio
import json
import time
from abc import ABC, abstractmethod
from collections import OrderedDict
from typing import Any
import redis.asyncio as aioredis
from src.clients.redis_client import get_redis_client_sync
from src.core.logger import logger
class BaseCacheBackend(ABC):
"""缓存后端抽象基类"""
@abstractmethod
async def get(self, key: str) -> Any | None:
"""获取缓存值"""
pass
@abstractmethod
async def set(self, key: str, value: Any, ttl: int = 300) -> None:
"""设置缓存值"""
pass
@abstractmethod
async def delete(self, key: str) -> None:
"""删除缓存值"""
pass
@abstractmethod
async def clear(self, pattern: str | None = None) -> None:
"""清空缓存(支持模式匹配)"""
pass
@abstractmethod
async def exists(self, key: str) -> bool:
"""检查键是否存在"""
pass
class LocalCache(BaseCacheBackend):
"""本地内存缓存后端LRU + TTL线程安全"""
def __init__(self, max_size: int = 1000, default_ttl: int = 300):
"""
初始化本地缓存
Args:
max_size: 最大缓存条目数
default_ttl: 默认过期时间(秒)
"""
self._cache: OrderedDict = OrderedDict()
self._expiry: dict[str, float] = {}
self._max_size = max_size
self._default_ttl = default_ttl
self._lock = asyncio.Lock()
async def get(self, key: str) -> Any | None:
"""获取缓存值(线程安全)"""
async with self._lock:
if key not in self._cache:
return None
# 检查过期
if key in self._expiry and time.time() > self._expiry[key]:
# 过期,删除
del self._cache[key]
del self._expiry[key]
return None
# 更新访问顺序LRU
self._cache.move_to_end(key)
return self._cache[key]
async def set(self, key: str, value: Any, ttl: int | None = None) -> None:
"""设置缓存值(线程安全)"""
async with self._lock:
if ttl is None:
ttl = self._default_ttl
# 如果键已存在,更新访问顺序
if key in self._cache:
self._cache.move_to_end(key)
elif len(self._cache) >= self._max_size:
# 插入新键前淘汰最旧项,确保容量不超过 max_size
oldest_key = next(iter(self._cache))
del self._cache[oldest_key]
if oldest_key in self._expiry:
del self._expiry[oldest_key]
self._cache[key] = value
self._expiry[key] = time.time() + ttl
async def delete(self, key: str) -> None:
"""删除缓存值(线程安全)"""
async with self._lock:
if key in self._cache:
del self._cache[key]
if key in self._expiry:
del self._expiry[key]
async def clear(self, pattern: str | None = None) -> None:
"""清空缓存(线程安全)"""
async with self._lock:
if pattern is None:
# 清空所有
self._cache.clear()
self._expiry.clear()
else:
# 模式匹配删除(简单实现:支持前缀匹配)
prefix = pattern.rstrip("*")
keys_to_delete = [k for k in self._cache.keys() if k.startswith(prefix)]
for key in keys_to_delete:
del self._cache[key]
if key in self._expiry:
del self._expiry[key]
async def exists(self, key: str) -> bool:
"""检查键是否存在(线程安全)"""
async with self._lock:
if key not in self._cache:
return False
# 检查过期
if key in self._expiry and time.time() > self._expiry[key]:
del self._cache[key]
del self._expiry[key]
return False
return True
def get_stats(self) -> dict[str, Any]:
"""获取缓存统计信息"""
return {
"backend": "local",
"size": len(self._cache),
"max_size": self._max_size,
"default_ttl": self._default_ttl,
}
class RedisCache(BaseCacheBackend):
"""Redis 缓存后端(分布式)"""
def __init__(
self, redis_client: aioredis.Redis, key_prefix: str = "cache", default_ttl: int = 300
):
"""
初始化 Redis 缓存
Args:
redis_client: Redis 客户端实例
key_prefix: 缓存键前缀
default_ttl: 默认过期时间(秒)
"""
self._redis = redis_client
self._key_prefix = key_prefix
self._default_ttl = default_ttl
def _make_key(self, key: str) -> str:
"""构造完整的 Redis 键"""
return f"{self._key_prefix}:{key}"
async def get(self, key: str) -> Any | None:
"""获取缓存值"""
try:
redis_key = self._make_key(key)
value = await self._redis.get(redis_key)
if value is None:
return None
# 尝试 JSON 反序列化
try:
return json.loads(value)
except (json.JSONDecodeError, TypeError):
# 如果不是 JSON直接返回字符串
return value
except Exception as e:
logger.error(f"[RedisCache] 获取缓存失败: {key}, 错误: {e}")
return None
async def set(self, key: str, value: Any, ttl: int | None = None) -> None:
"""设置缓存值"""
if ttl is None:
ttl = self._default_ttl
try:
redis_key = self._make_key(key)
# 序列化值
if isinstance(value, (dict, list, tuple)):
serialized = json.dumps(value)
elif isinstance(value, (int, float, bool)):
serialized = json.dumps(value)
else:
serialized = str(value)
await self._redis.setex(redis_key, ttl, serialized)
except Exception as e:
logger.error(f"[RedisCache] 设置缓存失败: {key}, 错误: {e}")
async def delete(self, key: str) -> None:
"""删除缓存值"""
try:
redis_key = self._make_key(key)
await self._redis.delete(redis_key)
except Exception as e:
logger.error(f"[RedisCache] 删除缓存失败: {key}, 错误: {e}")
async def clear(self, pattern: str | None = None) -> None:
"""清空缓存"""
try:
if pattern is None:
# 清空所有带前缀的键
pattern = "*"
redis_pattern = self._make_key(pattern)
cursor = 0
deleted_count = 0
while True:
cursor, keys = await self._redis.scan(cursor, match=redis_pattern, count=100)
if keys:
await self._redis.delete(*keys)
deleted_count += len(keys)
if cursor == 0:
break
logger.info(f"[RedisCache] 清空缓存: {redis_pattern}, 删除 {deleted_count} 个键")
except Exception as e:
logger.error(f"[RedisCache] 清空缓存失败: {pattern}, 错误: {e}")
async def exists(self, key: str) -> bool:
"""检查键是否存在"""
try:
redis_key = self._make_key(key)
return await self._redis.exists(redis_key) > 0
except Exception as e:
logger.error(f"[RedisCache] 检查键存在失败: {key}, 错误: {e}")
return False
async def publish_invalidation(self, channel: str, key: str) -> None:
"""发布缓存失效消息(用于分布式同步)"""
try:
message = json.dumps({"key": key, "timestamp": time.time()})
await self._redis.publish(channel, message)
logger.debug(f"[RedisCache] 发布缓存失效: {channel} -> {key}")
except Exception as e:
logger.error(f"[RedisCache] 发布缓存失效失败: {channel}, {key}, 错误: {e}")
def get_stats(self) -> dict[str, Any]:
"""获取缓存统计信息"""
return {
"backend": "redis",
"key_prefix": self._key_prefix,
"default_ttl": self._default_ttl,
}
# 缓存后端工厂
_cache_backends: dict[str, BaseCacheBackend] = {}
_cache_backend_lock = asyncio.Lock()
async def get_cache_backend(
name: str, backend_type: str = "auto", max_size: int = 1000, ttl: int = 300
) -> BaseCacheBackend:
"""
获取缓存后端实例
Args:
name: 缓存名称(用于区分不同的缓存实例)
backend_type: 后端类型 (auto/local/redis)
max_size: LocalCache 的最大容量
ttl: 默认过期时间(秒)
Returns:
BaseCacheBackend 实例
"""
cache_key = f"{name}:{backend_type}"
# 无锁快路径
if cache_key in _cache_backends:
return _cache_backends[cache_key]
async with _cache_backend_lock:
# Double-check: 锁内再检查一次,避免重复创建
if cache_key in _cache_backends:
return _cache_backends[cache_key]
backend = _create_cache_backend(name, backend_type, max_size, ttl)
_cache_backends[cache_key] = backend
return backend
def _create_cache_backend(
name: str, backend_type: str, max_size: int, ttl: int
) -> BaseCacheBackend:
"""根据类型创建缓存后端实例"""
if backend_type == "redis":
redis_client = get_redis_client_sync()
if redis_client is None:
logger.warning(f"[CacheBackend] Redis 未初始化,{name} 降级为本地缓存")
return LocalCache(max_size=max_size, default_ttl=ttl)
else:
logger.info(f"[CacheBackend] {name} 使用 Redis 缓存")
return RedisCache(redis_client=redis_client, key_prefix=name, default_ttl=ttl)
elif backend_type == "local":
logger.info(f"[CacheBackend] {name} 使用本地缓存")
return LocalCache(max_size=max_size, default_ttl=ttl)
else: # auto
redis_client = get_redis_client_sync()
if redis_client is not None:
logger.debug(f"[CacheBackend] {name} 自动选择 Redis 缓存")
return RedisCache(redis_client=redis_client, key_prefix=name, default_ttl=ttl)
else:
logger.debug(f"[CacheBackend] {name} 自动选择本地缓存Redis 不可用)")
return LocalCache(max_size=max_size, default_ttl=ttl)

View File

@@ -0,0 +1,155 @@
"""
缓存失效服务
统一管理各种缓存的失效逻辑
"""
from __future__ import annotations
import hashlib
import json
from typing import Any
from src.core.logger import logger
_PROVIDER_MAPPING_PREVIEW_CACHE_KEY_PREFIX = "admin:providers:mapping-preview:global"
_PROVIDER_MAPPING_PREVIEW_CACHE_PATTERN = f"{_PROVIDER_MAPPING_PREVIEW_CACHE_KEY_PREFIX}:v:*"
class CacheInvalidationService:
"""缓存失效服务"""
def __init__(self) -> None:
pass
def register_model_mapper(self, model_mapper: Any) -> None:
"""注册 ModelMapper 实例(已弃用,保留兼容性,不执行任何操作)"""
pass
async def on_global_model_changed(
self, model_name: str, global_model_id: str | None = None
) -> None:
"""
GlobalModel 变更时的缓存失效
Args:
model_name: 变更的 GlobalModel.name
global_model_id: GlobalModel ID可选
"""
logger.info(f"[CacheInvalidation] GlobalModel 变更: {model_name}")
# 1. 清空正则缓存
from src.core.model_permissions import clear_regex_cache
clear_regex_cache()
# 2. 清空 ModelMapper 共享缓存
from src.services.model.mapper import ModelMapperMiddleware
ModelMapperMiddleware.clear_cache()
# 3. 清空 ModelCacheService 缓存
from src.services.cache.model_cache import ModelCacheService
try:
await ModelCacheService.invalidate_global_model_cache(
global_model_id=global_model_id or "", name=model_name
)
except Exception as e:
logger.error(f"[CacheInvalidation] 失效 ModelCacheService 缓存失败: {e}")
# 4. 清除 /v1/models 列表缓存
from src.services.cache.model_list_cache import invalidate_models_list_cache
try:
await invalidate_models_list_cache()
except Exception as e:
logger.error(f"[CacheInvalidation] 失效 models list 缓存失败: {e}")
# 5. 清除 Provider 映射预览缓存GlobalModel 映射规则变化会影响所有 Provider
try:
await self._invalidate_all_provider_mapping_preview_cache()
except Exception as e:
logger.error(f"[CacheInvalidation] 失效 Provider 映射预览缓存失败: {e}")
def on_model_changed(self, provider_id: str, global_model_id: str) -> Any:
"""Model 变更时的缓存失效"""
self._refresh_provider_cache(provider_id)
async def on_key_allowed_models_changed(self, provider_id: str) -> None:
"""
Key 的 allowed_models 变更时的缓存失效
当 Key 的模型白名单变化时(如自动获取更新),需要刷新相关缓存,
以便正则映射规则能够重新匹配到新的白名单模型。
Args:
provider_id: 变更的 Key 所属的 Provider ID
"""
logger.info(f"[CacheInvalidation] Key allowed_models 变更: provider_id={provider_id}")
self._refresh_provider_cache(provider_id)
# 清除该 Provider 的映射预览缓存(详情页模型映射依赖)
try:
await self._invalidate_provider_mapping_preview_cache(provider_id)
except Exception as e:
logger.error(f"[CacheInvalidation] 失效 Provider 映射预览缓存失败: {e}")
# 清除 /v1/models 列表缓存allowed_models 变更会影响模型可用性)
from src.services.cache.model_list_cache import invalidate_models_list_cache
try:
await invalidate_models_list_cache()
except Exception as e:
logger.error(f"[CacheInvalidation] 失效 models list 缓存失败: {e}")
def _refresh_provider_cache(self, provider_id: str) -> None:
"""刷新指定 Provider 的 ModelMapper 缓存"""
from src.services.model.mapper import ModelMapperMiddleware
ModelMapperMiddleware.refresh_cache(provider_id)
@staticmethod
def _build_provider_mapping_preview_cache_key(provider_id: str) -> str:
"""构建指定 Provider 的 mapping-preview 缓存键。"""
raw = json.dumps(
{"provider_id": provider_id},
sort_keys=True,
ensure_ascii=False,
default=str,
)
vary_hash = hashlib.sha1(raw.encode("utf-8")).hexdigest()[:16]
return f"{_PROVIDER_MAPPING_PREVIEW_CACHE_KEY_PREFIX}:v:{vary_hash}"
async def _invalidate_provider_mapping_preview_cache(self, provider_id: str) -> None:
"""失效指定 Provider 的 mapping-preview 缓存。"""
from src.core.cache_service import CacheService
cache_key = self._build_provider_mapping_preview_cache_key(provider_id)
await CacheService.delete(cache_key)
async def _invalidate_all_provider_mapping_preview_cache(self) -> None:
"""失效所有 Provider 的 mapping-preview 缓存。"""
from src.core.cache_service import CacheService
await CacheService.delete_pattern(_PROVIDER_MAPPING_PREVIEW_CACHE_PATTERN)
def clear_all_caches(self) -> None:
"""清空所有缓存"""
from src.services.model.mapper import ModelMapperMiddleware
ModelMapperMiddleware.clear_cache()
# 全局单例
_cache_invalidation_service: CacheInvalidationService | None = None
def get_cache_invalidation_service() -> CacheInvalidationService:
"""获取全局缓存失效服务实例"""
global _cache_invalidation_service
if _cache_invalidation_service is None:
_cache_invalidation_service = CacheInvalidationService()
return _cache_invalidation_service

View File

@@ -0,0 +1,671 @@
"""
Model 映射缓存服务 - 减少模型查询
架构说明
========
本服务采用混合 async/sync 模式:
- 缓存操作CacheService真正的 async使用 aioredis
- 数据库查询db.query同步的 SQLAlchemy Session
设计决策
--------
1. 保持 async 方法签名:因为缓存命中时完全异步,性能最优
2. 缓存未命中时的同步查询FastAPI 会在线程池中执行,不会阻塞事件循环
3. 调用方必须在 async 上下文中使用 await
使用示例
--------
global_model = await ModelCacheService.resolve_global_model_by_name_or_mapping(db, "gpt-4")
"""
import time
from sqlalchemy.orm import Session
from src.config.constants import CacheTTL
from src.core.cache_service import CacheService
from src.core.logger import logger
from src.core.metrics import (
model_mapping_conflict_total,
model_mapping_resolution_duration_seconds,
model_mapping_resolution_total,
)
from src.models.database import GlobalModel, Model
class ModelCacheService:
"""Model 映射缓存服务
提供 GlobalModel 和 Model 的缓存查询功能,减少数据库访问。
所有公开方法均为 async需要在 async 上下文中调用。
"""
# 缓存 TTL- 使用统一常量
CACHE_TTL = CacheTTL.MODEL
PROVIDER_MAPPING_INDEX_CACHE_KEY = "global_model:resolve_index:provider_model_mappings"
MODEL_MAPPING_RULES_CACHE_KEY = "global_model:resolve_index:model_mappings"
@staticmethod
async def get_model_by_id(db: Session, model_id: str) -> Model | None:
"""
获取 Model带缓存
Args:
db: 数据库会话
model_id: Model ID
Returns:
Model 对象或 None
"""
cache_key = f"model:id:{model_id}"
# 1. 尝试从缓存获取
cached_data = await CacheService.get(cache_key)
if cached_data:
logger.debug(f"Model 缓存命中: {model_id}")
return ModelCacheService._dict_to_model(cached_data)
# 2. 缓存未命中,查询数据库
model = db.query(Model).filter(Model.id == model_id).first()
# 3. 写入缓存
if model:
model_dict = ModelCacheService._model_to_dict(model)
await CacheService.set(cache_key, model_dict, ttl_seconds=ModelCacheService.CACHE_TTL)
logger.debug(f"Model 已缓存: {model_id}")
return model
@staticmethod
async def get_global_model_by_id(db: Session, global_model_id: str) -> GlobalModel | None:
"""
获取 GlobalModel带缓存
Args:
db: 数据库会话
global_model_id: GlobalModel ID
Returns:
GlobalModel 对象或 None
"""
cache_key = f"global_model:id:{global_model_id}"
# 1. 尝试从缓存获取
cached_data = await CacheService.get(cache_key)
if cached_data:
logger.debug(f"GlobalModel 缓存命中: {global_model_id}")
return ModelCacheService._dict_to_global_model(cached_data)
# 2. 缓存未命中,查询数据库
global_model = db.query(GlobalModel).filter(GlobalModel.id == global_model_id).first()
# 3. 写入缓存
if global_model:
global_model_dict = ModelCacheService._global_model_to_dict(global_model)
await CacheService.set(
cache_key, global_model_dict, ttl_seconds=ModelCacheService.CACHE_TTL
)
logger.debug(f"GlobalModel 已缓存: {global_model_id}")
return global_model
@staticmethod
async def get_model_by_provider_and_global_model(
db: Session, provider_id: str, global_model_id: str
) -> Model | None:
"""
通过 Provider ID 和 GlobalModel ID 获取 Model带缓存
Args:
db: 数据库会话
provider_id: Provider ID
global_model_id: GlobalModel ID
Returns:
Model 对象或 None
"""
cache_key = f"model:provider_global:{provider_id}:{global_model_id}"
hit_count_key = f"model:provider_global:hits:{provider_id}:{global_model_id}"
# 1. 尝试从缓存获取
cached_data = await CacheService.get(cache_key)
if cached_data:
logger.debug(
f"Model 缓存命中(provider+global): {provider_id[:8]}...+{global_model_id[:8]}..."
)
# 递增命中计数,同时刷新 TTL
await CacheService.incr(hit_count_key, ttl_seconds=ModelCacheService.CACHE_TTL)
return ModelCacheService._dict_to_model(cached_data)
# 2. 缓存未命中,查询数据库
model = (
db.query(Model)
.filter(
Model.provider_id == provider_id,
Model.global_model_id == global_model_id,
Model.is_active == True,
)
.first()
)
# 3. 写入缓存
if model:
model_dict = ModelCacheService._model_to_dict(model)
await CacheService.set(cache_key, model_dict, ttl_seconds=ModelCacheService.CACHE_TTL)
# 重置命中计数新缓存从1开始
await CacheService.set(hit_count_key, 1, ttl_seconds=ModelCacheService.CACHE_TTL)
logger.debug(
f"Model 已缓存(provider+global): {provider_id[:8]}...+{global_model_id[:8]}..."
)
return model
@staticmethod
async def get_global_model_by_name(db: Session, name: str) -> GlobalModel | None:
"""
通过名称获取 GlobalModel带缓存
Args:
db: 数据库会话
name: GlobalModel 名称
Returns:
GlobalModel 对象或 None
"""
cache_key = f"global_model:name:{name}"
# 1. 尝试从缓存获取
cached_data = await CacheService.get(cache_key)
if cached_data:
logger.debug(f"GlobalModel 缓存命中(名称): {name}")
return ModelCacheService._dict_to_global_model(cached_data)
# 2. 缓存未命中,查询数据库
global_model = db.query(GlobalModel).filter(GlobalModel.name == name).first()
# 3. 写入缓存
if global_model:
global_model_dict = ModelCacheService._global_model_to_dict(global_model)
await CacheService.set(
cache_key, global_model_dict, ttl_seconds=ModelCacheService.CACHE_TTL
)
logger.debug(f"GlobalModel 已缓存(名称): {name}")
return global_model
@staticmethod
async def invalidate_model_cache(
model_id: str,
provider_id: str | None = None,
global_model_id: str | None = None,
provider_model_name: str | None = None,
provider_model_mappings: list | None = None,
) -> None:
"""清除 Model 缓存
Args:
model_id: Model ID
provider_id: Provider ID用于清除 provider_global 缓存)
global_model_id: GlobalModel ID用于清除 provider_global 缓存)
provider_model_name: provider_model_name用于清除 resolve 缓存)
provider_model_mappings: 映射名称列表(用于清除 resolve 缓存)
"""
# 清除 model:id 缓存
await CacheService.delete(f"model:id:{model_id}")
# 清除 provider_global 缓存及其命中计数(如果提供了必要参数)
if provider_id and global_model_id:
await CacheService.delete(f"model:provider_global:{provider_id}:{global_model_id}")
await CacheService.delete(f"model:provider_global:hits:{provider_id}:{global_model_id}")
logger.debug(
f"Model 缓存已清除: {model_id}, provider_global:{provider_id[:8]}...:{global_model_id[:8]}..."
)
else:
logger.debug(f"Model 缓存已清除: {model_id}")
# 清除 resolve 缓存provider_model_name 和 mappings 可能都被用作解析 key
resolve_keys_to_clear = []
if provider_model_name:
resolve_keys_to_clear.append(provider_model_name)
if provider_model_mappings:
for mapping_entry in provider_model_mappings:
if isinstance(mapping_entry, dict):
mapping_name = mapping_entry.get("name", "").strip()
if mapping_name:
resolve_keys_to_clear.append(mapping_name)
for key in resolve_keys_to_clear:
await CacheService.delete(f"global_model:resolve:{key}")
if resolve_keys_to_clear:
logger.debug(f"Model resolve 缓存已清除: {resolve_keys_to_clear}")
# provider_model_mappings 更新后,需要重建映射索引缓存。
await CacheService.delete(ModelCacheService.PROVIDER_MAPPING_INDEX_CACHE_KEY)
@staticmethod
async def invalidate_global_model_cache(global_model_id: str, name: str | None = None) -> None:
"""清除 GlobalModel 缓存"""
await CacheService.delete(f"global_model:id:{global_model_id}")
if name:
await CacheService.delete(f"global_model:name:{name}")
# 同时清除 resolve 缓存,因为 GlobalModel.name 也是一个 resolve key
await CacheService.delete(f"global_model:resolve:{name}")
# 全量清除 resolve 缓存,确保映射规则变更后不命中旧缓存
try:
await CacheService.delete_pattern("global_model:resolve:*")
await CacheService.delete(ModelCacheService.PROVIDER_MAPPING_INDEX_CACHE_KEY)
await CacheService.delete(ModelCacheService.MODEL_MAPPING_RULES_CACHE_KEY)
except Exception as e:
logger.error(f"GlobalModel resolve 缓存清除失败,可能导致映射不一致: {e}")
logger.debug(f"GlobalModel 缓存已清除: {global_model_id}")
@staticmethod
async def invalidate_all_resolve_cache() -> None:
"""
清除所有 GlobalModel 解析缓存
在 Provider 启用/禁用时调用,因为 Provider 状态变更会影响模型解析结果。
"""
try:
deleted = await CacheService.delete_pattern("global_model:resolve:*")
await CacheService.delete(ModelCacheService.PROVIDER_MAPPING_INDEX_CACHE_KEY)
await CacheService.delete(ModelCacheService.MODEL_MAPPING_RULES_CACHE_KEY)
logger.debug(f"已清除 {deleted} 个 GlobalModel resolve 缓存")
except Exception as e:
logger.error(f"GlobalModel resolve 缓存清除失败: {e}")
@staticmethod
async def _get_provider_mapping_index(
db: Session,
) -> dict[str, list[dict[str, object]]]:
cached_data = await CacheService.get(ModelCacheService.PROVIDER_MAPPING_INDEX_CACHE_KEY)
if isinstance(cached_data, dict):
return {
str(name): value
for name, value in cached_data.items()
if isinstance(name, str) and isinstance(value, list)
}
from src.models.database import Provider
rows = (
db.query(Model, GlobalModel)
.join(Provider, Model.provider_id == Provider.id)
.join(GlobalModel, Model.global_model_id == GlobalModel.id)
.filter(
Provider.is_active == True,
Model.is_active == True,
GlobalModel.is_active == True,
Model.provider_model_mappings.isnot(None),
)
.all()
)
index: dict[str, list[dict[str, object]]] = {}
seen_pairs: set[tuple[str, str]] = set()
for model, global_model in rows:
raw_mappings = getattr(model, "provider_model_mappings", None)
if not isinstance(raw_mappings, list):
continue
global_model_dict = ModelCacheService._global_model_to_dict(global_model)
for raw in raw_mappings:
if not isinstance(raw, dict):
continue
name = raw.get("name")
if not isinstance(name, str):
continue
normalized_name = name.strip()
if not normalized_name:
continue
pair_key = (normalized_name, str(global_model.id))
if pair_key in seen_pairs:
continue
seen_pairs.add(pair_key)
index.setdefault(normalized_name, []).append(global_model_dict)
await CacheService.set(
ModelCacheService.PROVIDER_MAPPING_INDEX_CACHE_KEY,
index,
ttl_seconds=ModelCacheService.CACHE_TTL,
)
return index
@staticmethod
async def _get_model_mapping_rules(
db: Session,
) -> list[dict[str, object]]:
cached_data = await CacheService.get(ModelCacheService.MODEL_MAPPING_RULES_CACHE_KEY)
if isinstance(cached_data, list):
return [entry for entry in cached_data if isinstance(entry, dict)]
rows = (
db.query(GlobalModel)
.filter(GlobalModel.is_active == True, GlobalModel.config.isnot(None))
.all()
)
rules: list[dict[str, object]] = []
for global_model in rows:
config = getattr(global_model, "config", None) or {}
mappings = config.get("model_mappings")
if not isinstance(mappings, list) or not mappings:
continue
patterns = [
pattern for pattern in mappings if isinstance(pattern, str) and pattern.strip()
]
if not patterns:
continue
rules.append(
{
"global_model": ModelCacheService._global_model_to_dict(global_model),
"patterns": patterns,
}
)
await CacheService.set(
ModelCacheService.MODEL_MAPPING_RULES_CACHE_KEY,
rules,
ttl_seconds=ModelCacheService.CACHE_TTL,
)
return rules
@staticmethod
async def resolve_global_model_by_name_or_mapping(
db: Session, model_name: str
) -> GlobalModel | None:
"""
通过名称解析 GlobalModel带缓存
查找顺序:
1. 检查缓存
2. 直接匹配 GlobalModel.name
3. 通过 provider_model_name 匹配(查询 Model 表)
4. 通过 provider_model_mappings 匹配(查询 Model 表)
5. 通过 GlobalModel.config.model_mappings 匹配(支持正则)
注意provider_model_mappings 是 Provider 级别的映射配置,可能存在跨 Provider 冲突;
如匹配到多个 GlobalModel将记录告警并选择第一个匹配结果。
Args:
db: 数据库会话
model_name: 模型名称(可以是 GlobalModel.name 或 provider_model_name
Returns:
GlobalModel 对象或 None
"""
start_time = time.time()
resolution_method = "not_found"
cache_hit = False
normalized_name = model_name.strip()
if not normalized_name:
return None
cache_key = f"global_model:resolve:{normalized_name}"
try:
# 1. 尝试从缓存获取
cached_data = await CacheService.get(cache_key)
if cached_data:
if cached_data == "NOT_FOUND":
# 缓存的负结果
cache_hit = True
resolution_method = "not_found"
logger.debug(f"GlobalModel 缓存命中(映射解析-未找到): {normalized_name}")
return None
if isinstance(cached_data, dict) and "supported_capabilities" not in cached_data:
# 兼容旧缓存:字段不全时视为未命中,走 DB 刷新
logger.debug(f"GlobalModel 缓存命中但 schema 过旧,刷新: {normalized_name}")
else:
cache_hit = True
resolution_method = "direct_match" # 缓存命中时无法区分原始解析方式
logger.debug(f"GlobalModel 缓存命中(映射解析): {normalized_name}")
return ModelCacheService._dict_to_global_model(cached_data)
# 2. 直接通过 GlobalModel.name 匹配(优先级最高)
# 说明:如果存在同名 GlobalModel应优先解析为 GlobalModel 本身,
# 避免被某个 Provider 的 provider_model_name 误导导致解析到错误的 GlobalModel。
global_model = (
db.query(GlobalModel)
.filter(GlobalModel.name == normalized_name, GlobalModel.is_active == True)
.first()
)
if global_model:
resolution_method = "direct_match"
global_model_dict = ModelCacheService._global_model_to_dict(global_model)
await CacheService.set(
cache_key, global_model_dict, ttl_seconds=ModelCacheService.CACHE_TTL
)
logger.debug(f"GlobalModel 已缓存(映射解析-直接匹配): {normalized_name}")
return global_model
# 3. 通过 provider_model_name 匹配
from src.models.database import Provider
models_with_global = (
db.query(Model, GlobalModel)
.join(Provider, Model.provider_id == Provider.id)
.join(GlobalModel, Model.global_model_id == GlobalModel.id)
.filter(
Provider.is_active == True,
Model.is_active == True,
GlobalModel.is_active == True,
Model.provider_model_name == normalized_name,
)
.all()
)
# 收集匹配的 GlobalModel只通过 provider_model_name 匹配)
matched_global_models: list[GlobalModel] = []
seen_global_model_ids: set[str] = set()
for model, gm in models_with_global:
if gm.id not in seen_global_model_ids:
seen_global_model_ids.add(gm.id)
matched_global_models.append(gm)
logger.debug(
f"模型名称 '{normalized_name}' 通过 provider_model_name 匹配到 "
f"GlobalModel: {gm.name} (Model: {model.id[:8]}...)"
)
# 如果通过 provider_model_name 找到了,返回
if matched_global_models:
resolution_method = "provider_model_name"
if len(matched_global_models) > 1:
# 检测到冲突(多个不同的 GlobalModel 有相同的 provider_model_name
model_names = [gm.name for gm in matched_global_models if gm.name]
logger.warning(
f"模型映射冲突: 名称 '{normalized_name}' 匹配到多个不同的 GlobalModel: "
f"{', '.join(model_names)},使用第一个匹配结果"
)
# 记录冲突指标
model_mapping_conflict_total.inc()
# 返回第一个匹配的 GlobalModel
result_global_model = matched_global_models[0]
global_model_dict = ModelCacheService._global_model_to_dict(result_global_model)
await CacheService.set(
cache_key, global_model_dict, ttl_seconds=ModelCacheService.CACHE_TTL
)
logger.debug(
f"GlobalModel 已缓存(映射解析-{resolution_method}): {normalized_name} -> {result_global_model.name}"
)
return result_global_model
# 4. 通过 provider_model_mappings 匹配
provider_mapping_index = await ModelCacheService._get_provider_mapping_index(db)
mapping_matched_global_models = [
ModelCacheService._dict_to_global_model(global_model_dict)
for global_model_dict in provider_mapping_index.get(normalized_name, [])
if isinstance(global_model_dict, dict)
]
for gm in mapping_matched_global_models:
logger.debug(
f"模型名称 '{normalized_name}' 通过 provider_model_mappings 匹配到 "
f"GlobalModel: {gm.name}"
)
if mapping_matched_global_models:
resolution_method = "provider_model_mappings"
if len(mapping_matched_global_models) > 1:
model_names = [gm.name for gm in mapping_matched_global_models if gm.name]
logger.warning(
f"模型映射冲突: 名称 '{normalized_name}' 匹配到多个不同的 GlobalModel: "
f"{', '.join(model_names)},使用第一个匹配结果"
)
model_mapping_conflict_total.inc()
# 按名称排序确保确定性
result_global_model = sorted(
mapping_matched_global_models, key=lambda gm: gm.name or ""
)[0]
global_model_dict = ModelCacheService._global_model_to_dict(result_global_model)
await CacheService.set(
cache_key, global_model_dict, ttl_seconds=ModelCacheService.CACHE_TTL
)
logger.debug(
f"GlobalModel 已缓存(映射解析-{resolution_method}): "
f"{normalized_name} -> {result_global_model.name}"
)
return result_global_model
# 5. 通过 GlobalModel.config.model_mappings 匹配(支持正则)
from src.core.model_permissions import match_model_with_pattern
mapping_matches: list[GlobalModel] = []
for entry in await ModelCacheService._get_model_mapping_rules(db):
global_model_dict = entry.get("global_model")
patterns = entry.get("patterns")
if not isinstance(global_model_dict, dict) or not isinstance(patterns, list):
continue
for pattern in patterns:
if isinstance(pattern, str) and match_model_with_pattern(
pattern, normalized_name
):
mapping_matches.append(
ModelCacheService._dict_to_global_model(global_model_dict)
)
break
if mapping_matches:
resolution_method = "model_mappings"
if len(mapping_matches) > 1:
model_names = [gm.name for gm in mapping_matches if gm.name]
logger.warning(
f"模型映射冲突: 名称 '{normalized_name}' 匹配到多个不同的 GlobalModel: "
f"{', '.join(model_names)},使用第一个匹配结果"
)
model_mapping_conflict_total.inc()
# 按名称排序确保确定性
result_global_model = sorted(mapping_matches, key=lambda gm: gm.name or "")[0]
global_model_dict = ModelCacheService._global_model_to_dict(result_global_model)
await CacheService.set(
cache_key, global_model_dict, ttl_seconds=ModelCacheService.CACHE_TTL
)
logger.debug(
f"GlobalModel 已缓存(映射解析-{resolution_method}): "
f"{normalized_name} -> {result_global_model.name}"
)
return result_global_model
# 6. 完全未找到
resolution_method = "not_found"
# 未找到匹配,缓存负结果
await CacheService.set(cache_key, "NOT_FOUND", ttl_seconds=ModelCacheService.CACHE_TTL)
logger.debug(f"GlobalModel 未找到(映射解析): {normalized_name}")
return None
finally:
# 记录监控指标
duration = time.time() - start_time
model_mapping_resolution_total.labels(
method=resolution_method, cache_hit=str(cache_hit).lower()
).inc()
model_mapping_resolution_duration_seconds.labels(method=resolution_method).observe(
duration
)
@staticmethod
def _model_to_dict(model: Model) -> dict:
"""将 Model 对象转换为字典"""
return {
"id": model.id,
"provider_id": model.provider_id,
"global_model_id": model.global_model_id,
"provider_model_name": model.provider_model_name,
"provider_model_mappings": getattr(model, "provider_model_mappings", None),
"is_active": model.is_active,
"is_available": model.is_available if hasattr(model, "is_available") else True,
"price_per_request": (
float(model.price_per_request) if model.price_per_request is not None else None
),
"tiered_pricing": model.tiered_pricing,
"supports_vision": model.supports_vision,
"supports_function_calling": model.supports_function_calling,
"supports_streaming": model.supports_streaming,
"supports_extended_thinking": model.supports_extended_thinking,
"supports_image_generation": getattr(model, "supports_image_generation", None),
"config": model.config,
}
@staticmethod
def _dict_to_model(model_dict: dict) -> Model:
"""从字典重建 Model 对象"""
model = Model(
id=model_dict["id"],
provider_id=model_dict["provider_id"],
global_model_id=model_dict["global_model_id"],
provider_model_name=model_dict["provider_model_name"],
provider_model_mappings=model_dict.get("provider_model_mappings"),
is_active=model_dict["is_active"],
is_available=model_dict.get("is_available", True),
price_per_request=model_dict.get("price_per_request"),
tiered_pricing=model_dict.get("tiered_pricing"),
supports_vision=model_dict.get("supports_vision"),
supports_function_calling=model_dict.get("supports_function_calling"),
supports_streaming=model_dict.get("supports_streaming"),
supports_extended_thinking=model_dict.get("supports_extended_thinking"),
supports_image_generation=model_dict.get("supports_image_generation"),
config=model_dict.get("config"),
)
return model
@staticmethod
def _global_model_to_dict(global_model: GlobalModel) -> dict:
"""将 GlobalModel 对象转换为字典"""
return {
"id": global_model.id,
"name": global_model.name,
"display_name": global_model.display_name,
"supported_capabilities": global_model.supported_capabilities,
"config": global_model.config,
"default_tiered_pricing": global_model.default_tiered_pricing,
"default_price_per_request": (
float(global_model.default_price_per_request)
if global_model.default_price_per_request is not None
else None
),
"is_active": global_model.is_active,
}
@staticmethod
def _dict_to_global_model(global_model_dict: dict) -> GlobalModel:
"""从字典重建 GlobalModel 对象"""
global_model = GlobalModel(
id=global_model_dict["id"],
name=global_model_dict["name"],
display_name=global_model_dict.get("display_name"),
supported_capabilities=global_model_dict.get("supported_capabilities") or [],
config=global_model_dict.get("config"),
default_tiered_pricing=global_model_dict.get("default_tiered_pricing"),
default_price_per_request=global_model_dict.get("default_price_per_request"),
is_active=global_model_dict.get("is_active", True),
)
return global_model

View File

@@ -0,0 +1,31 @@
"""
/v1/models 列表缓存管理。
从 api/base/models_service.py 迁移到 services 层,
消除 services→api 的反向依赖。
"""
from __future__ import annotations
from src.core.cache_service import CacheService
from src.core.logger import logger
# 缓存 key 前缀models_service.py 也使用此常量)
MODELS_LIST_CACHE_PREFIX = "models:list"
async def invalidate_models_list_cache() -> None:
"""
清除所有 /v1/models 列表缓存
在模型创建、更新、删除时调用,确保模型列表实时更新
"""
try:
# 使用通配符删除所有 models:list:* 缓存(包括多格式组合的 key
deleted = await CacheService.delete_pattern(f"{MODELS_LIST_CACHE_PREFIX}:*")
if deleted > 0:
logger.info("[ModelsService] 已清除 {}{} 缓存", deleted, MODELS_LIST_CACHE_PREFIX)
else:
logger.debug("[ModelsService] 无 {} 缓存需要清除", MODELS_LIST_CACHE_PREFIX)
except Exception as e:
logger.warning("[ModelsService] 清除缓存失败: {}", e)

View File

@@ -0,0 +1,209 @@
"""
Provider 缓存服务 - 减少 Provider 和 ProviderAPIKey 查询
用于缓存 Provider 的 billing_type 和 ProviderAPIKey 的 rate_multiplier
这些数据在 UsageService.record_usage() 中被频繁查询但变化不频繁。
"""
from __future__ import annotations
from sqlalchemy.orm import Session
from src.config.constants import CacheTTL
from src.core.cache_service import CacheService
from src.core.enums import ProviderBillingType
from src.core.logger import logger
from src.models.database import Provider, ProviderAPIKey
class ProviderCacheService:
"""Provider 缓存服务
提供 Provider 和 ProviderAPIKey 的缓存查询功能,减少数据库访问。
主要用于 UsageService 中获取费率倍数和计费类型。
"""
CACHE_TTL = CacheTTL.PROVIDER # 5 分钟
@staticmethod
def compute_rate_multiplier(
rate_multipliers: dict | None,
api_format: str | None = None,
) -> float:
"""
计算 rate_multiplier 的纯函数(无数据库/缓存依赖)
返回指定 API 格式的倍率,如果没有则返回 1.0。
Args:
rate_multipliers: 按 API 格式的倍率配置字典
api_format: API 格式(可选),如 "CLAUDE""OPENAI"
Returns:
计算后的 rate_multiplier
"""
if api_format and rate_multipliers:
format_key = str(api_format).strip().lower()
if format_key in rate_multipliers:
return float(rate_multipliers[format_key])
return 1.0
@staticmethod
async def get_provider_api_key_rate_multiplier(
db: Session, provider_api_key_id: str, api_format: str | None = None
) -> float | None:
"""
获取 ProviderAPIKey 的 rate_multiplier带缓存
优先返回指定 API 格式的倍率,如果没有则返回默认倍率。
Args:
db: 数据库会话
provider_api_key_id: ProviderAPIKey ID
api_format: API 格式(可选),如 "CLAUDE""OPENAI"
Returns:
rate_multiplier 或 None如果找不到
"""
# 缓存键包含 api_format
format_suffix = str(api_format).strip().lower() if api_format else "default"
cache_key = f"provider_api_key:rate_multiplier:{provider_api_key_id}:{format_suffix}"
# 1. 尝试从缓存获取
cached_data = await CacheService.get(cache_key)
if cached_data is not None:
logger.debug(
f"ProviderAPIKey rate_multiplier 缓存命中: {provider_api_key_id[:8]}... format={format_suffix}"
)
# 缓存的 "NOT_FOUND" 表示数据库中不存在
if cached_data == "NOT_FOUND":
return None
return float(cached_data)
# 2. 缓存未命中,查询数据库
provider_key = (
db.query(ProviderAPIKey.rate_multipliers)
.filter(ProviderAPIKey.id == provider_api_key_id)
.first()
)
# 3. 计算倍率并写入缓存
if provider_key:
rate_multiplier = ProviderCacheService.compute_rate_multiplier(
provider_key.rate_multipliers, api_format
)
await CacheService.set(
cache_key, rate_multiplier, ttl_seconds=ProviderCacheService.CACHE_TTL
)
logger.debug(
f"ProviderAPIKey rate_multiplier 已缓存: {provider_api_key_id[:8]}... format={format_suffix} value={rate_multiplier}"
)
return rate_multiplier
else:
# 缓存负结果
await CacheService.set(
cache_key, "NOT_FOUND", ttl_seconds=ProviderCacheService.CACHE_TTL
)
return None
@staticmethod
async def get_provider_billing_type(
db: Session, provider_id: str
) -> ProviderBillingType | None:
"""
获取 Provider 的 billing_type带缓存
Args:
db: 数据库会话
provider_id: Provider ID
Returns:
billing_type 或 None如果找不到
"""
cache_key = f"provider:billing_type:{provider_id}"
# 1. 尝试从缓存获取
cached_data = await CacheService.get(cache_key)
if cached_data is not None:
logger.debug(f"Provider billing_type 缓存命中: {provider_id[:8]}...")
if cached_data == "NOT_FOUND":
return None
try:
return ProviderBillingType(cached_data)
except ValueError:
# 缓存值无效,删除并重新查询
await CacheService.delete(cache_key)
# 2. 缓存未命中,查询数据库
provider = db.query(Provider.billing_type).filter(Provider.id == provider_id).first()
# 3. 写入缓存
if provider:
billing_type = provider.billing_type
await CacheService.set(
cache_key, billing_type.value, ttl_seconds=ProviderCacheService.CACHE_TTL
)
logger.debug(f"Provider billing_type 已缓存: {provider_id[:8]}...")
return billing_type
else:
# 缓存负结果
await CacheService.set(
cache_key, "NOT_FOUND", ttl_seconds=ProviderCacheService.CACHE_TTL
)
return None
@staticmethod
async def get_rate_multiplier_and_free_tier(
db: Session,
provider_api_key_id: str | None,
provider_id: str | None,
api_format: str | None = None,
) -> tuple[float, bool]:
"""
获取费率倍数和是否免费套餐(带缓存)
这是 UsageService._get_rate_multiplier_and_free_tier 的缓存版本。
Args:
db: 数据库会话
provider_api_key_id: ProviderAPIKey ID可选
provider_id: Provider ID可选
api_format: API 格式(可选),用于获取按格式配置的倍率
Returns:
(rate_multiplier, is_free_tier) 元组
"""
actual_rate_multiplier = 1.0
is_free_tier = False
# 获取费率倍数(支持按 API 格式查询)
if provider_api_key_id:
rate_multiplier = await ProviderCacheService.get_provider_api_key_rate_multiplier(
db, provider_api_key_id, api_format
)
if rate_multiplier is not None:
actual_rate_multiplier = rate_multiplier
# 获取计费类型
if provider_id:
billing_type = await ProviderCacheService.get_provider_billing_type(db, provider_id)
if billing_type == ProviderBillingType.FREE_TIER:
is_free_tier = True
return actual_rate_multiplier, is_free_tier
@staticmethod
async def invalidate_provider_api_key_cache(provider_api_key_id: str) -> None:
"""清除 ProviderAPIKey 缓存(包括所有 API 格式的缓存)"""
# 使用模式匹配删除所有格式的缓存
await CacheService.delete_pattern(
f"provider_api_key:rate_multiplier:{provider_api_key_id}:*"
)
logger.debug(f"ProviderAPIKey 缓存已清除: {provider_api_key_id[:8]}...")
@staticmethod
async def invalidate_provider_cache(provider_id: str) -> None:
"""清除 Provider 缓存"""
await CacheService.delete(f"provider:billing_type:{provider_id}")
logger.debug(f"Provider 缓存已清除: {provider_id[:8]}...")

View File

@@ -0,0 +1,222 @@
"""
缓存同步服务Redis Pub/Sub
提供分布式缓存失效同步功能,用于多实例部署场景。
当一个实例修改数据并失效本地缓存时,通过 Redis pub/sub 通知其他实例同步失效。
使用场景:
1. 多实例部署时,确保所有实例的缓存一致性
2. GlobalModel/Model 变更时,同步失效所有实例的缓存
"""
from __future__ import annotations
import asyncio
import json
from collections.abc import Callable
from typing import Any
import redis.asyncio as aioredis
from src.clients.redis_client import get_redis_client_sync
from src.core.logger import logger
class CacheSyncService:
"""
缓存同步服务
通过 Redis pub/sub 实现分布式缓存失效同步
"""
# Redis 频道名称
CHANNEL_GLOBAL_MODEL = "cache:invalidate:global_model"
CHANNEL_MODEL = "cache:invalidate:model"
CHANNEL_CLEAR_ALL = "cache:invalidate:clear_all"
def __init__(self, redis_client: aioredis.Redis):
"""
初始化缓存同步服务
Args:
redis_client: Redis 客户端实例
"""
self._redis = redis_client
self._pubsub: aioredis.client.PubSub | None = None
self._listener_task: asyncio.Task | None = None
self._handlers: dict[str, Callable] = {}
self._running = False
async def start(self) -> Any:
"""启动缓存同步服务(订阅 Redis 频道)"""
if self._running:
logger.warning("[CacheSync] 服务已在运行")
return
try:
self._pubsub = self._redis.pubsub()
# 订阅所有缓存失效频道
await self._pubsub.subscribe(
self.CHANNEL_GLOBAL_MODEL,
self.CHANNEL_MODEL,
self.CHANNEL_CLEAR_ALL,
)
# 启动监听任务
self._listener_task = asyncio.create_task(self._listen())
self._running = True
logger.info(
"[CacheSync] 缓存同步服务已启动,订阅频道: "
f"{self.CHANNEL_GLOBAL_MODEL}, "
f"{self.CHANNEL_MODEL}, {self.CHANNEL_CLEAR_ALL}"
)
except Exception as e:
logger.error(f"[CacheSync] 启动失败: {e}")
raise
async def stop(self) -> Any:
"""停止缓存同步服务"""
if not self._running:
return
self._running = False
# 取消监听任务
if self._listener_task:
self._listener_task.cancel()
try:
await self._listener_task
except asyncio.CancelledError:
pass
# 取消订阅
if self._pubsub:
await self._pubsub.unsubscribe()
await self._pubsub.close()
logger.info("[CacheSync] 缓存同步服务已停止")
def register_handler(self, channel: str, handler: Callable) -> None:
"""
注册缓存失效处理器
Args:
channel: Redis 频道名称
handler: 处理函数(接收消息数据作为参数)
"""
self._handlers[channel] = handler
logger.debug(f"[CacheSync] 注册处理器: {channel}")
async def _listen(self) -> None:
"""监听 Redis pub/sub 消息(含断线重连)"""
logger.info("[CacheSync] 开始监听缓存失效消息")
consecutive_failures = 0
max_consecutive_failures = 10
reconnect_interval = 5.0
while self._running:
try:
async for message in self._pubsub.listen():
consecutive_failures = 0 # 收到消息即重置
if message["type"] == "message":
channel = message["channel"]
data = message["data"]
try:
payload = json.loads(data)
logger.debug(f"[CacheSync] 收到消息: {channel} -> {payload}")
if channel in self._handlers:
handler = self._handlers[channel]
await handler(payload)
else:
logger.warning(f"[CacheSync] 未找到处理器: {channel}")
except json.JSONDecodeError as e:
logger.error(f"[CacheSync] 消息解析失败: {data}, 错误: {e}")
except Exception as e:
logger.error(f"[CacheSync] 处理消息失败: {channel}, 错误: {e}")
except asyncio.CancelledError:
logger.info("[CacheSync] 监听任务已取消")
return
except Exception as e:
consecutive_failures += 1
logger.error(
f"[CacheSync] 监听失败 ({consecutive_failures}/{max_consecutive_failures}): {e}"
)
if consecutive_failures >= max_consecutive_failures:
logger.error("[CacheSync] 连续失败次数过多,停止重连")
return
await asyncio.sleep(reconnect_interval)
async def publish_global_model_changed(self, model_name: str) -> Any:
"""发布 GlobalModel 变更通知"""
await self._publish(self.CHANNEL_GLOBAL_MODEL, {"model_name": model_name})
async def publish_model_changed(self, provider_id: str, global_model_id: str) -> Any:
"""发布 Model 变更通知"""
await self._publish(
self.CHANNEL_MODEL, {"provider_id": provider_id, "global_model_id": global_model_id}
)
async def publish_clear_all(self) -> Any:
"""发布清空所有缓存通知"""
await self._publish(self.CHANNEL_CLEAR_ALL, {})
async def _publish(self, channel: str, data: dict) -> None:
"""发布消息到 Redis 频道(含简单重试)"""
message = json.dumps(data)
last_error: Exception | None = None
for attempt in range(2):
try:
await self._redis.publish(channel, message)
logger.debug(f"[CacheSync] 发布消息: {channel} -> {data}")
return
except Exception as e:
last_error = e
if attempt == 0:
await asyncio.sleep(0.5)
logger.error(f"[CacheSync] 发布消息失败(已重试): {channel}, 错误: {last_error}")
# 全局单例
_cache_sync_service: CacheSyncService | None = None
async def get_cache_sync_service(
redis_client: aioredis.Redis | None = None,
) -> CacheSyncService | None:
"""
获取缓存同步服务实例
Args:
redis_client: Redis 客户端实例(首次调用时需要提供)
Returns:
CacheSyncService 实例,如果 Redis 不可用返回 None
"""
global _cache_sync_service
if _cache_sync_service is None:
if redis_client is None:
# 尝试获取全局 Redis 客户端
redis_client = get_redis_client_sync()
if redis_client is None:
logger.warning("[CacheSync] Redis 不可用,分布式缓存同步已禁用")
return None
_cache_sync_service = CacheSyncService(redis_client)
logger.info("[CacheSync] 缓存同步服务已初始化")
return _cache_sync_service
async def close_cache_sync_service() -> None:
"""关闭缓存同步服务"""
global _cache_sync_service
if _cache_sync_service:
await _cache_sync_service.stop()
_cache_sync_service = None

View File

@@ -0,0 +1,178 @@
"""
用户缓存服务 - 减少数据库查询
架构说明
========
本服务采用混合 async/sync 模式:
- 缓存操作CacheService真正的 async使用 aioredis
- 数据库查询db.query同步的 SQLAlchemy Session
设计决策
--------
1. 保持 async 方法签名:因为缓存命中时完全异步,性能最优
2. 缓存未命中时的同步查询FastAPI 会在线程池中执行,不会阻塞事件循环
3. 调用方必须在 async 上下文中使用 await
使用示例
--------
user = await UserCacheService.get_user_by_id(db, user_id)
await UserCacheService.invalidate_user_cache(user_id, email)
"""
from __future__ import annotations
from typing import Any
from sqlalchemy.orm import Session
from src.config.constants import CacheTTL
from src.core.cache_service import CacheKeys, CacheService
from src.core.logger import logger
from src.models.database import User
class UserCacheService:
"""用户缓存服务
提供 User 的缓存查询功能,减少数据库访问。
所有公开方法均为 async需要在 async 上下文中调用。
"""
# 缓存 TTL- 使用统一常量
CACHE_TTL = CacheTTL.USER
@staticmethod
async def get_user_by_id(db: Session, user_id: str) -> User | None:
"""
获取用户(带缓存)
Args:
db: 数据库会话
user_id: 用户ID
Returns:
User 对象或 None
"""
cache_key = CacheKeys.user_by_id(user_id)
# 1. 尝试从缓存获取
cached_data = await CacheService.get(cache_key)
if cached_data:
logger.debug(f"用户缓存命中: {user_id}")
# 从缓存数据重建 User 对象
return UserCacheService._dict_to_user(db, cached_data)
# 2. 缓存未命中,查询数据库
user = db.query(User).filter(User.id == user_id).first()
# 3. 写入缓存
if user:
user_dict = UserCacheService._user_to_dict(user)
await CacheService.set(cache_key, user_dict, ttl_seconds=UserCacheService.CACHE_TTL)
logger.debug(f"用户已缓存: {user_id}")
return user
@staticmethod
async def get_user_by_email(db: Session, email: str) -> User | None:
"""
通过邮箱获取用户(带缓存)
Args:
db: 数据库会话
email: 用户邮箱
Returns:
User 对象或 None
"""
cache_key = CacheKeys.user_by_email(email)
# 1. 尝试从缓存获取
cached_data = await CacheService.get(cache_key)
if cached_data:
logger.debug(f"用户缓存命中(邮箱): {email}")
return UserCacheService._dict_to_user(db, cached_data)
# 2. 缓存未命中,查询数据库
user = db.query(User).filter(User.email == email).first()
# 3. 写入缓存
if user:
user_dict = UserCacheService._user_to_dict(user)
await CacheService.set(cache_key, user_dict, ttl_seconds=UserCacheService.CACHE_TTL)
logger.debug(f"用户已缓存(邮箱): {email}")
return user
@staticmethod
async def invalidate_user_cache(user_id: str, email: str | None = None) -> Any:
"""
清除用户缓存
Args:
user_id: 用户ID
email: 用户邮箱(可选)
"""
# 删除 ID 缓存
await CacheService.delete(CacheKeys.user_by_id(user_id))
# 删除邮箱缓存
if email:
await CacheService.delete(CacheKeys.user_by_email(email))
logger.debug(f"用户缓存已清除: {user_id}")
@staticmethod
def _user_to_dict(user: User) -> dict:
"""将 User 对象转换为字典(用于缓存)"""
return {
"id": user.id,
"email": user.email,
"email_verified": user.email_verified,
"username": user.username,
"role": user.role.value if user.role else None,
"is_active": user.is_active,
"auth_source": user.auth_source.value if user.auth_source else None,
"created_at": user.created_at.isoformat() if user.created_at else None,
"last_login_at": user.last_login_at.isoformat() if user.last_login_at else None,
"model_capability_settings": user.model_capability_settings,
}
@staticmethod
def _dict_to_user(db: Session, user_dict: dict) -> User:
"""
从字典重建 User 对象
注意:这是一个"分离"的对象,不在 Session 中
如果需要修改,需要使用 db.merge() 或重新查询
"""
from datetime import datetime
from src.core.enums import AuthSource
from src.models.database import UserRole
user = User(
id=user_dict["id"],
email=user_dict.get("email"),
email_verified=user_dict.get("email_verified", False),
username=user_dict["username"],
is_active=user_dict["is_active"],
)
# 设置可选字段
if user_dict.get("role"):
user.role = UserRole(user_dict["role"])
if user_dict.get("auth_source"):
user.auth_source = AuthSource(user_dict["auth_source"])
if user_dict.get("created_at"):
user.created_at = datetime.fromisoformat(user_dict["created_at"])
if user_dict.get("last_login_at"):
user.last_login_at = datetime.fromisoformat(user_dict["last_login_at"])
if user_dict.get("model_capability_settings") is not None:
user.model_capability_settings = user_dict["model_capability_settings"]
return user