mirror of
https://github.com/fawney19/Aether.git
synced 2026-09-02 01:10:23 +08:00
feat(models): Key 白名单变更时自动关联 GlobalModel
- Key 创建/更新/自动获取模型后,检查 allowed_models 变化 - 根据 GlobalModel 的映射规则自动创建 Provider 关联 - 添加缓存失效机制确保映射规则能匹配新白名单
This commit is contained in:
@@ -19,12 +19,12 @@ from src.core.key_capabilities import get_capability
|
|||||||
from src.core.logger import logger
|
from src.core.logger import logger
|
||||||
from src.database import get_db
|
from src.database import get_db
|
||||||
from src.models.database import Provider, ProviderAPIKey, ProviderEndpoint
|
from src.models.database import Provider, ProviderAPIKey, ProviderEndpoint
|
||||||
from src.services.cache.provider_cache import ProviderCacheService
|
|
||||||
from src.models.endpoint_models import (
|
from src.models.endpoint_models import (
|
||||||
EndpointAPIKeyCreate,
|
EndpointAPIKeyCreate,
|
||||||
EndpointAPIKeyResponse,
|
EndpointAPIKeyResponse,
|
||||||
EndpointAPIKeyUpdate,
|
EndpointAPIKeyUpdate,
|
||||||
)
|
)
|
||||||
|
from src.services.cache.provider_cache import ProviderCacheService
|
||||||
|
|
||||||
router = APIRouter(tags=["Provider Keys"])
|
router = APIRouter(tags=["Provider Keys"])
|
||||||
pipeline = ApiRequestPipeline()
|
pipeline = ApiRequestPipeline()
|
||||||
@@ -214,7 +214,14 @@ class AdminUpdateEndpointKeyAdapter(AdminApiAdapter):
|
|||||||
|
|
||||||
# 检查是否开启了 auto_fetch_models(用于后续立即获取模型)
|
# 检查是否开启了 auto_fetch_models(用于后续立即获取模型)
|
||||||
auto_fetch_enabled_before = key.auto_fetch_models
|
auto_fetch_enabled_before = key.auto_fetch_models
|
||||||
auto_fetch_enabled_after = self.key_data.auto_fetch_models if "auto_fetch_models" in self.key_data.model_fields_set else auto_fetch_enabled_before
|
auto_fetch_enabled_after = (
|
||||||
|
self.key_data.auto_fetch_models
|
||||||
|
if "auto_fetch_models" in self.key_data.model_fields_set
|
||||||
|
else auto_fetch_enabled_before
|
||||||
|
)
|
||||||
|
|
||||||
|
# 记录 allowed_models 变化前的值
|
||||||
|
allowed_models_before = set(key.allowed_models or [])
|
||||||
|
|
||||||
update_data = self.key_data.model_dump(exclude_unset=True)
|
update_data = self.key_data.model_dump(exclude_unset=True)
|
||||||
if "api_key" in update_data:
|
if "api_key" in update_data:
|
||||||
@@ -252,9 +259,11 @@ class AdminUpdateEndpointKeyAdapter(AdminApiAdapter):
|
|||||||
logger.info("[AUTO_FETCH] Key %s 开启自动获取模型,立即触发模型获取", self.key_id)
|
logger.info("[AUTO_FETCH] Key %s 开启自动获取模型,立即触发模型获取", self.key_id)
|
||||||
try:
|
try:
|
||||||
from src.services.model.fetch_scheduler import get_model_fetch_scheduler
|
from src.services.model.fetch_scheduler import get_model_fetch_scheduler
|
||||||
|
|
||||||
scheduler = get_model_fetch_scheduler()
|
scheduler = get_model_fetch_scheduler()
|
||||||
# 在后台异步执行,不阻塞当前请求
|
# 在后台异步执行,不阻塞当前请求
|
||||||
import asyncio
|
import asyncio
|
||||||
|
|
||||||
asyncio.create_task(scheduler._fetch_models_for_key_by_id(self.key_id))
|
asyncio.create_task(scheduler._fetch_models_for_key_by_id(self.key_id))
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
logger.error(f"触发模型获取失败: {e}")
|
logger.error(f"触发模型获取失败: {e}")
|
||||||
@@ -282,6 +291,17 @@ class AdminUpdateEndpointKeyAdapter(AdminApiAdapter):
|
|||||||
# 包括 is_active、allowed_models、capabilities 等影响权限和行为的字段
|
# 包括 is_active、allowed_models、capabilities 等影响权限和行为的字段
|
||||||
await ProviderCacheService.invalidate_provider_api_key_cache(self.key_id)
|
await ProviderCacheService.invalidate_provider_api_key_cache(self.key_id)
|
||||||
|
|
||||||
|
# 检查 allowed_models 是否有变化,触发缓存失效和自动关联
|
||||||
|
allowed_models_after = set(key.allowed_models or [])
|
||||||
|
if allowed_models_before != allowed_models_after and key.provider_id:
|
||||||
|
from src.services.model.global_model import on_key_allowed_models_changed
|
||||||
|
|
||||||
|
on_key_allowed_models_changed(
|
||||||
|
db=db,
|
||||||
|
provider_id=key.provider_id,
|
||||||
|
allowed_models=list(key.allowed_models or []),
|
||||||
|
)
|
||||||
|
|
||||||
logger.info("[OK] 更新 Key: ID=%s, Updates=%s", self.key_id, list(update_data.keys()))
|
logger.info("[OK] 更新 Key: ID=%s, Updates=%s", self.key_id, list(update_data.keys()))
|
||||||
|
|
||||||
return _build_key_response(key)
|
return _build_key_response(key)
|
||||||
@@ -613,12 +633,24 @@ class AdminCreateProviderKeyAdapter(AdminApiAdapter):
|
|||||||
logger.info("[AUTO_FETCH] 新 Key %s 开启自动获取模型,立即触发模型获取", new_key.id)
|
logger.info("[AUTO_FETCH] 新 Key %s 开启自动获取模型,立即触发模型获取", new_key.id)
|
||||||
try:
|
try:
|
||||||
from src.services.model.fetch_scheduler import get_model_fetch_scheduler
|
from src.services.model.fetch_scheduler import get_model_fetch_scheduler
|
||||||
|
|
||||||
scheduler = get_model_fetch_scheduler()
|
scheduler = get_model_fetch_scheduler()
|
||||||
# 在后台异步执行,不阻塞当前请求
|
# 在后台异步执行,不阻塞当前请求
|
||||||
import asyncio
|
import asyncio
|
||||||
|
|
||||||
asyncio.create_task(scheduler._fetch_models_for_key_by_id(new_key.id))
|
asyncio.create_task(scheduler._fetch_models_for_key_by_id(new_key.id))
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
logger.error(f"触发模型获取失败: {e}")
|
logger.error(f"触发模型获取失败: {e}")
|
||||||
# 不抛出异常,避免影响 Key 创建操作
|
# 不抛出异常,避免影响 Key 创建操作
|
||||||
|
|
||||||
|
# 如果创建时指定了 allowed_models,触发自动关联检查
|
||||||
|
if new_key.allowed_models:
|
||||||
|
from src.services.model.global_model import on_key_allowed_models_changed
|
||||||
|
|
||||||
|
on_key_allowed_models_changed(
|
||||||
|
db=db,
|
||||||
|
provider_id=self.provider_id,
|
||||||
|
allowed_models=list(new_key.allowed_models),
|
||||||
|
)
|
||||||
|
|
||||||
return _build_key_response(new_key, api_key_plain=self.key_data.api_key)
|
return _build_key_response(new_key, api_key_plain=self.key_data.api_key)
|
||||||
|
|||||||
17
src/services/cache/invalidation.py
vendored
17
src/services/cache/invalidation.py
vendored
@@ -53,6 +53,23 @@ class CacheInvalidationService:
|
|||||||
|
|
||||||
def on_model_changed(self, provider_id: str, global_model_id: str):
|
def on_model_changed(self, provider_id: str, global_model_id: str):
|
||||||
"""Model 变更时的缓存失效"""
|
"""Model 变更时的缓存失效"""
|
||||||
|
self._refresh_provider_cache(provider_id)
|
||||||
|
|
||||||
|
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)
|
||||||
|
|
||||||
|
def _refresh_provider_cache(self, provider_id: str) -> None:
|
||||||
|
"""刷新指定 Provider 的 ModelMapper 缓存"""
|
||||||
for mapper in self._model_mappers:
|
for mapper in self._model_mappers:
|
||||||
mapper.refresh_cache(provider_id)
|
mapper.refresh_cache(provider_id)
|
||||||
|
|
||||||
|
|||||||
@@ -307,11 +307,27 @@ class ModelFetchScheduler:
|
|||||||
)
|
)
|
||||||
|
|
||||||
# 更新 allowed_models(保留 locked_models)
|
# 更新 allowed_models(保留 locked_models)
|
||||||
self._update_key_allowed_models(key, fetched_model_ids)
|
has_changed = self._update_key_allowed_models(key, fetched_model_ids)
|
||||||
|
|
||||||
|
# 如果白名单有变化,触发缓存失效和自动关联检查
|
||||||
|
if has_changed and provider_id:
|
||||||
|
from src.services.model.global_model import on_key_allowed_models_changed
|
||||||
|
|
||||||
|
on_key_allowed_models_changed(
|
||||||
|
db=db,
|
||||||
|
provider_id=provider_id,
|
||||||
|
allowed_models=list(key.allowed_models or []),
|
||||||
|
)
|
||||||
|
|
||||||
return "success"
|
return "success"
|
||||||
|
|
||||||
def _update_key_allowed_models(self, key: ProviderAPIKey, fetched_model_ids: set[str]) -> None:
|
def _update_key_allowed_models(self, key: ProviderAPIKey, fetched_model_ids: set[str]) -> bool:
|
||||||
"""更新 Key 的 allowed_models,保留 locked_models"""
|
"""
|
||||||
|
更新 Key 的 allowed_models,保留 locked_models
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
bool: 是否有变化
|
||||||
|
"""
|
||||||
# 获取当前锁定的模型
|
# 获取当前锁定的模型
|
||||||
locked_models = set(key.locked_models or [])
|
locked_models = set(key.locked_models or [])
|
||||||
|
|
||||||
@@ -333,8 +349,10 @@ class ModelFetchScheduler:
|
|||||||
logger.info(f"Key {key.id} 移除模型: {sorted(removed)}")
|
logger.info(f"Key {key.id} 移除模型: {sorted(removed)}")
|
||||||
|
|
||||||
key.allowed_models = new_allowed_models
|
key.allowed_models = new_allowed_models
|
||||||
|
return True
|
||||||
else:
|
else:
|
||||||
logger.debug(f"Key {key.id} 模型列表无变化")
|
logger.debug(f"Key {key.id} 模型列表无变化")
|
||||||
|
return False
|
||||||
|
|
||||||
async def _fetch_models_from_endpoints(
|
async def _fetch_models_from_endpoints(
|
||||||
self, endpoint_configs: list[dict]
|
self, endpoint_configs: list[dict]
|
||||||
|
|||||||
@@ -4,7 +4,7 @@ GlobalModel 服务层
|
|||||||
提供 GlobalModel 的 CRUD 操作、查询和统计功能
|
提供 GlobalModel 的 CRUD 操作、查询和统计功能
|
||||||
"""
|
"""
|
||||||
|
|
||||||
from typing import Dict, List, Optional
|
from typing import Dict, List, Optional, Set
|
||||||
|
|
||||||
from sqlalchemy import and_, func
|
from sqlalchemy import and_, func
|
||||||
from sqlalchemy.orm import Session, joinedload
|
from sqlalchemy.orm import Session, joinedload
|
||||||
@@ -15,6 +15,37 @@ from src.models.database import GlobalModel, Model
|
|||||||
from src.models.pydantic_models import GlobalModelUpdate
|
from src.models.pydantic_models import GlobalModelUpdate
|
||||||
|
|
||||||
|
|
||||||
|
def on_key_allowed_models_changed(
|
||||||
|
db: Session,
|
||||||
|
provider_id: str,
|
||||||
|
allowed_models: List[str],
|
||||||
|
) -> None:
|
||||||
|
"""
|
||||||
|
Key 的 allowed_models 变更后的统一处理
|
||||||
|
|
||||||
|
包括:
|
||||||
|
1. 触发缓存失效
|
||||||
|
2. 检查并自动关联匹配的 GlobalModel
|
||||||
|
|
||||||
|
Args:
|
||||||
|
db: 数据库 Session
|
||||||
|
provider_id: Provider ID
|
||||||
|
allowed_models: 更新后的 allowed_models 列表
|
||||||
|
"""
|
||||||
|
from src.services.cache.invalidation import get_cache_invalidation_service
|
||||||
|
|
||||||
|
# 1. 触发缓存失效
|
||||||
|
cache_service = get_cache_invalidation_service()
|
||||||
|
cache_service.on_key_allowed_models_changed(provider_id)
|
||||||
|
|
||||||
|
# 2. 检查并自动关联 GlobalModel
|
||||||
|
if allowed_models:
|
||||||
|
GlobalModelService.auto_associate_provider_by_key_whitelist(
|
||||||
|
db=db,
|
||||||
|
provider_id=provider_id,
|
||||||
|
allowed_models=allowed_models,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
class GlobalModelService:
|
class GlobalModelService:
|
||||||
"""GlobalModel 服务"""
|
"""GlobalModel 服务"""
|
||||||
@@ -163,7 +194,9 @@ class GlobalModelService:
|
|||||||
|
|
||||||
# 级联删除所有关联的 Provider 模型实现
|
# 级联删除所有关联的 Provider 模型实现
|
||||||
if associated_models:
|
if associated_models:
|
||||||
logger.info(f"删除 GlobalModel {global_model.name} 的 {len(associated_models)} 个关联 Provider 模型")
|
logger.info(
|
||||||
|
f"删除 GlobalModel {global_model.name} 的 {len(associated_models)} 个关联 Provider 模型"
|
||||||
|
)
|
||||||
for model in associated_models:
|
for model in associated_models:
|
||||||
db.delete(model)
|
db.delete(model)
|
||||||
|
|
||||||
@@ -286,3 +319,137 @@ class GlobalModelService:
|
|||||||
|
|
||||||
db.commit()
|
db.commit()
|
||||||
return results
|
return results
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def auto_associate_provider_by_key_whitelist(
|
||||||
|
db: Session,
|
||||||
|
provider_id: str,
|
||||||
|
allowed_models: List[str],
|
||||||
|
) -> Dict:
|
||||||
|
"""
|
||||||
|
根据 Key 白名单自动关联 Provider 到匹配的 GlobalModel
|
||||||
|
|
||||||
|
当 Key 的 allowed_models 更新后调用此方法,检查所有 GlobalModel 的映射规则,
|
||||||
|
如果有映射规则匹配到 Key 白名单中的模型,且 Provider 尚未关联到该 GlobalModel,
|
||||||
|
则自动创建关联。
|
||||||
|
|
||||||
|
Args:
|
||||||
|
db: 数据库 Session
|
||||||
|
provider_id: Provider ID
|
||||||
|
allowed_models: Key 的白名单模型列表
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
Dict: 包含 success 和 errors 列表
|
||||||
|
"""
|
||||||
|
from src.core.model_permissions import match_model_with_pattern
|
||||||
|
from src.models.database import Provider
|
||||||
|
|
||||||
|
results: Dict[str, List[Dict]] = {
|
||||||
|
"success": [],
|
||||||
|
"errors": [],
|
||||||
|
}
|
||||||
|
|
||||||
|
if not allowed_models:
|
||||||
|
return results
|
||||||
|
|
||||||
|
# 获取 Provider
|
||||||
|
provider = db.query(Provider).filter(Provider.id == provider_id).first()
|
||||||
|
if not provider:
|
||||||
|
logger.warning(f"Provider {provider_id} not found for auto-association")
|
||||||
|
return results
|
||||||
|
|
||||||
|
# 获取该 Provider 已关联的 GlobalModel ID 集合
|
||||||
|
existing_associations = (
|
||||||
|
db.query(Model.global_model_id, Model.provider_model_name)
|
||||||
|
.filter(Model.provider_id == provider_id)
|
||||||
|
.all()
|
||||||
|
)
|
||||||
|
linked_global_model_ids: Set[str] = {row[0] for row in existing_associations if row[0]}
|
||||||
|
# 同时获取已存在的 provider_model_name 集合,避免唯一约束冲突
|
||||||
|
existing_provider_model_names: Set[str] = {
|
||||||
|
row[1] for row in existing_associations if row[1]
|
||||||
|
}
|
||||||
|
|
||||||
|
# 获取所有活跃的 GlobalModel(带映射规则)
|
||||||
|
global_models = db.query(GlobalModel).filter(GlobalModel.is_active == True).all()
|
||||||
|
|
||||||
|
allowed_models_set = set(allowed_models)
|
||||||
|
|
||||||
|
for global_model in global_models:
|
||||||
|
# 跳过已关联的
|
||||||
|
if global_model.id in linked_global_model_ids:
|
||||||
|
continue
|
||||||
|
|
||||||
|
# 跳过 provider_model_name 已存在的(避免唯一约束冲突)
|
||||||
|
if global_model.name in existing_provider_model_names:
|
||||||
|
logger.debug(
|
||||||
|
f"Skipping auto-association for GlobalModel {global_model.name}: "
|
||||||
|
f"provider_model_name already exists for Provider {provider.name}"
|
||||||
|
)
|
||||||
|
continue
|
||||||
|
|
||||||
|
# 提取映射规则
|
||||||
|
model_mappings: List[str] = []
|
||||||
|
if global_model.config and isinstance(global_model.config, dict):
|
||||||
|
mappings = global_model.config.get("model_mappings")
|
||||||
|
if isinstance(mappings, list):
|
||||||
|
model_mappings = [m for m in mappings if isinstance(m, str)]
|
||||||
|
|
||||||
|
if not model_mappings:
|
||||||
|
continue
|
||||||
|
|
||||||
|
# 检查是否有映射规则匹配到 Key 白名单
|
||||||
|
matched = False
|
||||||
|
for mapping_pattern in model_mappings:
|
||||||
|
for allowed_model in allowed_models_set:
|
||||||
|
if match_model_with_pattern(mapping_pattern, allowed_model):
|
||||||
|
matched = True
|
||||||
|
break
|
||||||
|
if matched:
|
||||||
|
break
|
||||||
|
|
||||||
|
if not matched:
|
||||||
|
continue
|
||||||
|
|
||||||
|
# 自动创建关联
|
||||||
|
try:
|
||||||
|
new_model = Model(
|
||||||
|
provider_id=provider_id,
|
||||||
|
global_model_id=global_model.id,
|
||||||
|
provider_model_name=global_model.name,
|
||||||
|
is_active=True,
|
||||||
|
)
|
||||||
|
db.add(new_model)
|
||||||
|
db.flush()
|
||||||
|
|
||||||
|
# 添加到已存在集合,避免后续循环重复创建
|
||||||
|
existing_provider_model_names.add(global_model.name)
|
||||||
|
|
||||||
|
results["success"].append(
|
||||||
|
{
|
||||||
|
"global_model_id": global_model.id,
|
||||||
|
"global_model_name": global_model.name,
|
||||||
|
"model_id": new_model.id,
|
||||||
|
}
|
||||||
|
)
|
||||||
|
logger.info(
|
||||||
|
f"Auto-associated Provider {provider.name} to GlobalModel {global_model.name} "
|
||||||
|
f"via mapping rule match"
|
||||||
|
)
|
||||||
|
except Exception as e:
|
||||||
|
db.rollback()
|
||||||
|
logger.error(
|
||||||
|
f"Failed to auto-associate Provider {provider.name} to GlobalModel {global_model.name}: {e}"
|
||||||
|
)
|
||||||
|
results["errors"].append(
|
||||||
|
{
|
||||||
|
"global_model_id": global_model.id,
|
||||||
|
"global_model_name": global_model.name,
|
||||||
|
"error": str(e),
|
||||||
|
}
|
||||||
|
)
|
||||||
|
|
||||||
|
if results["success"]:
|
||||||
|
db.commit()
|
||||||
|
|
||||||
|
return results
|
||||||
|
|||||||
Reference in New Issue
Block a user