mirror of
https://github.com/fawney19/Aether.git
synced 2026-09-01 17:00:21 +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:
15
_deprecated_py_src/services/cache/__init__.py
vendored
Normal file
15
_deprecated_py_src/services/cache/__init__.py
vendored
Normal 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",
|
||||
]
|
||||
338
_deprecated_py_src/services/cache/backend.py
vendored
Normal file
338
_deprecated_py_src/services/cache/backend.py
vendored
Normal 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)
|
||||
155
_deprecated_py_src/services/cache/invalidation.py
vendored
Normal file
155
_deprecated_py_src/services/cache/invalidation.py
vendored
Normal 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
|
||||
671
_deprecated_py_src/services/cache/model_cache.py
vendored
Normal file
671
_deprecated_py_src/services/cache/model_cache.py
vendored
Normal 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
|
||||
31
_deprecated_py_src/services/cache/model_list_cache.py
vendored
Normal file
31
_deprecated_py_src/services/cache/model_list_cache.py
vendored
Normal 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)
|
||||
209
_deprecated_py_src/services/cache/provider_cache.py
vendored
Normal file
209
_deprecated_py_src/services/cache/provider_cache.py
vendored
Normal 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]}...")
|
||||
222
_deprecated_py_src/services/cache/sync.py
vendored
Normal file
222
_deprecated_py_src/services/cache/sync.py
vendored
Normal 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
|
||||
178
_deprecated_py_src/services/cache/user_cache.py
vendored
Normal file
178
_deprecated_py_src/services/cache/user_cache.py
vendored
Normal 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
|
||||
Reference in New Issue
Block a user