mirror of
https://github.com/fawney19/Aether.git
synced 2026-09-02 01:10:23 +08:00
refactor: 将上游模型缓存从前端迁移到后端 Redis
- 后端:定时任务刷新时将上游模型写入 Redis 缓存 - 后端:provider_query 优先从缓存读取,支持 force_refresh 参数 - 后端:auto_fetch_models 开启时改为同步获取,确保前端能立即看到数据 - 前端:移除本地缓存,只保留并发请求去重逻辑 - 前端:KeyAllowedModelsEditDialog 添加刷新上游模型按钮
This commit is contained in:
@@ -256,16 +256,14 @@ class AdminUpdateEndpointKeyAdapter(AdminApiAdapter):
|
||||
|
||||
# 处理 auto_fetch_models 的开启和关闭
|
||||
if not auto_fetch_enabled_before and auto_fetch_enabled_after:
|
||||
# 刚刚开启了 auto_fetch_models,立即触发一次模型获取
|
||||
logger.info("[AUTO_FETCH] Key %s 开启自动获取模型,立即触发模型获取", self.key_id)
|
||||
# 刚刚开启了 auto_fetch_models,同步执行模型获取
|
||||
logger.info("[AUTO_FETCH] Key %s 开启自动获取模型,同步执行模型获取", self.key_id)
|
||||
try:
|
||||
from src.services.model.fetch_scheduler import get_model_fetch_scheduler
|
||||
|
||||
scheduler = get_model_fetch_scheduler()
|
||||
# 在后台异步执行,不阻塞当前请求
|
||||
import asyncio
|
||||
|
||||
asyncio.create_task(scheduler._fetch_models_for_key_by_id(self.key_id))
|
||||
# 同步等待模型获取完成,确保前端刷新时能看到最新数据
|
||||
await scheduler._fetch_models_for_key_by_id(self.key_id)
|
||||
except Exception as e:
|
||||
logger.error(f"触发模型获取失败: {e}")
|
||||
# 不抛出异常,避免影响 Key 更新操作
|
||||
@@ -630,17 +628,15 @@ class AdminCreateProviderKeyAdapter(AdminApiAdapter):
|
||||
f"Formats={self.key_data.api_formats}, Key=***{self.key_data.api_key[-4:]}, ID={new_key.id}"
|
||||
)
|
||||
|
||||
# 如果开启了 auto_fetch_models,立即触发一次模型获取
|
||||
# 如果开启了 auto_fetch_models,同步执行模型获取
|
||||
if self.key_data.auto_fetch_models:
|
||||
logger.info("[AUTO_FETCH] 新 Key %s 开启自动获取模型,立即触发模型获取", new_key.id)
|
||||
logger.info("[AUTO_FETCH] 新 Key %s 开启自动获取模型,同步执行模型获取", new_key.id)
|
||||
try:
|
||||
from src.services.model.fetch_scheduler import get_model_fetch_scheduler
|
||||
|
||||
scheduler = get_model_fetch_scheduler()
|
||||
# 在后台异步执行,不阻塞当前请求
|
||||
import asyncio
|
||||
|
||||
asyncio.create_task(scheduler._fetch_models_for_key_by_id(new_key.id))
|
||||
# 同步等待模型获取完成,确保前端刷新时能看到最新数据
|
||||
await scheduler._fetch_models_for_key_by_id(new_key.id)
|
||||
except Exception as e:
|
||||
logger.error(f"触发模型获取失败: {e}")
|
||||
# 不抛出异常,避免影响 Key 创建操作
|
||||
|
||||
@@ -33,6 +33,7 @@ class ModelsQueryRequest(BaseModel):
|
||||
|
||||
provider_id: str
|
||||
api_key_id: Optional[str] = None
|
||||
force_refresh: bool = False # 强制刷新,跳过缓存
|
||||
|
||||
|
||||
class TestModelRequest(BaseModel):
|
||||
@@ -73,6 +74,8 @@ async def query_available_models(
|
||||
"""
|
||||
查询提供商可用模型
|
||||
|
||||
优先从缓存获取(缓存由定时任务刷新),缓存未命中时实时调用上游 API。
|
||||
|
||||
遍历所有活跃端点,根据端点的 API 格式选择正确的 Adapter 进行请求:
|
||||
- OPENAI/OPENAI_CLI: 使用 OpenAIChatAdapter.fetch_models
|
||||
- CLAUDE/CLAUDE_CLI: 使用 ClaudeChatAdapter.fetch_models
|
||||
@@ -84,7 +87,12 @@ async def query_available_models(
|
||||
Returns:
|
||||
所有端点的模型列表(合并)
|
||||
"""
|
||||
# 获取提供商及其端点和 API Keys
|
||||
from src.services.model.fetch_scheduler import (
|
||||
get_upstream_models_from_cache,
|
||||
set_upstream_models_to_cache,
|
||||
)
|
||||
|
||||
# 获取提供商基本信息
|
||||
provider = (
|
||||
db.query(Provider)
|
||||
.options(
|
||||
@@ -98,6 +106,26 @@ async def query_available_models(
|
||||
if not provider:
|
||||
raise HTTPException(status_code=404, detail="Provider not found")
|
||||
|
||||
# 如果指定了 api_key_id 且不是强制刷新,优先从缓存获取
|
||||
# 注:不指定 api_key_id 时(Provider 级别查询)不使用缓存,因为:
|
||||
# 1. Provider 级别查询会遍历多个 Key,结果不稳定
|
||||
# 2. 缓存按 Key 粒度存储,与定时任务的刷新逻辑一致
|
||||
if request.api_key_id and not request.force_refresh:
|
||||
cached_models = await get_upstream_models_from_cache(
|
||||
request.provider_id, request.api_key_id
|
||||
)
|
||||
if cached_models is not None:
|
||||
return {
|
||||
"success": True,
|
||||
"data": {"models": cached_models, "error": None, "from_cache": True},
|
||||
"provider": {
|
||||
"id": provider.id,
|
||||
"name": provider.name,
|
||||
},
|
||||
}
|
||||
|
||||
# 缓存未命中或强制刷新,实时获取
|
||||
|
||||
# 收集所有活跃端点的配置
|
||||
endpoint_configs: list[dict] = []
|
||||
|
||||
@@ -235,9 +263,15 @@ async def query_available_models(
|
||||
if not unique_models and not error:
|
||||
error = "No models returned from any endpoint"
|
||||
|
||||
# 如果指定了 api_key_id 且获取成功,写入缓存
|
||||
if request.api_key_id and unique_models:
|
||||
await set_upstream_models_to_cache(
|
||||
request.provider_id, request.api_key_id, unique_models
|
||||
)
|
||||
|
||||
return {
|
||||
"success": len(unique_models) > 0,
|
||||
"data": {"models": unique_models, "error": error},
|
||||
"data": {"models": unique_models, "error": error, "from_cache": False},
|
||||
"provider": {
|
||||
"id": provider.id,
|
||||
"name": provider.name,
|
||||
|
||||
@@ -18,6 +18,7 @@ from typing import Optional
|
||||
import httpx
|
||||
from sqlalchemy.orm import Session, joinedload
|
||||
|
||||
from src.core.cache_service import CacheService
|
||||
from src.core.crypto import crypto_service
|
||||
from src.core.headers import get_extra_headers_from_endpoint
|
||||
from src.core.logger import logger
|
||||
@@ -35,6 +36,35 @@ MAX_CONCURRENT_REQUESTS = 5
|
||||
# 单个 Key 处理的超时时间(秒)
|
||||
KEY_FETCH_TIMEOUT_SECONDS = 120
|
||||
|
||||
# 上游模型缓存 TTL(与定时任务间隔保持一致)
|
||||
UPSTREAM_MODELS_CACHE_TTL_SECONDS = MODEL_FETCH_INTERVAL_MINUTES * 60
|
||||
|
||||
|
||||
def _get_upstream_models_cache_key(provider_id: str, api_key_id: str) -> str:
|
||||
"""生成上游模型缓存的 key"""
|
||||
return f"upstream_models:{provider_id}:{api_key_id}"
|
||||
|
||||
|
||||
async def get_upstream_models_from_cache(
|
||||
provider_id: str, api_key_id: str
|
||||
) -> Optional[list[dict]]:
|
||||
"""从缓存获取上游模型列表"""
|
||||
cache_key = _get_upstream_models_cache_key(provider_id, api_key_id)
|
||||
cached = await CacheService.get(cache_key)
|
||||
if cached is not None:
|
||||
logger.debug(f"上游模型缓存命中: {cache_key}")
|
||||
return cached # type: ignore[no-any-return]
|
||||
return None
|
||||
|
||||
|
||||
async def set_upstream_models_to_cache(
|
||||
provider_id: str, api_key_id: str, models: list[dict]
|
||||
) -> None:
|
||||
"""将上游模型列表写入缓存"""
|
||||
cache_key = _get_upstream_models_cache_key(provider_id, api_key_id)
|
||||
await CacheService.set(cache_key, models, UPSTREAM_MODELS_CACHE_TTL_SECONDS)
|
||||
logger.debug(f"上游模型已缓存: {cache_key}, 数量={len(models)}")
|
||||
|
||||
|
||||
def _get_adapter_for_format(api_format: str) -> Optional[type]:
|
||||
"""根据 API 格式获取对应的 Adapter 类"""
|
||||
@@ -307,6 +337,22 @@ class ModelFetchScheduler:
|
||||
f"Provider {provider.name} Key {key.id} 获取到 {len(fetched_model_ids)} 个唯一模型"
|
||||
)
|
||||
|
||||
# 写入上游模型缓存(按 model id + api_format 去重后的完整模型信息)
|
||||
seen_keys: set[str] = set()
|
||||
unique_models: list[dict] = []
|
||||
for model in all_models:
|
||||
model_id = model.get("id")
|
||||
api_format = model.get("api_format", "")
|
||||
unique_key = f"{model_id}:{api_format}"
|
||||
if model_id and unique_key not in seen_keys:
|
||||
seen_keys.add(unique_key)
|
||||
unique_models.append(model)
|
||||
await set_upstream_models_to_cache(
|
||||
provider_id, # type: ignore[arg-type]
|
||||
key.id, # type: ignore[arg-type]
|
||||
unique_models,
|
||||
)
|
||||
|
||||
# 更新 allowed_models(保留 locked_models)
|
||||
has_changed = self._update_key_allowed_models(key, fetched_model_ids)
|
||||
|
||||
|
||||
Reference in New Issue
Block a user