mirror of
https://github.com/fawney19/Aether.git
synced 2026-09-02 01:10:23 +08:00
feat: 优化批量余额查询和用户模型权限检查
- 添加批量余额查询并发限制,避免数据库连接池耗尽 - 支持余额加载 pending 状态和前端自动重试机制 - 添加用户可用模型 API,统一使用 AccessRestrictions - 修复用户表单编辑时数组引用共享导致的数据覆盖问题 - 添加数据库连接池配置说明到 .env.example
This commit is contained in:
@@ -19,11 +19,13 @@ from src.database import get_db
|
||||
from src.models.api import (
|
||||
ChangePasswordRequest,
|
||||
CreateMyApiKeyRequest,
|
||||
PublicGlobalModelListResponse,
|
||||
PublicGlobalModelResponse,
|
||||
UpdateApiKeyProvidersRequest,
|
||||
UpdatePreferencesRequest,
|
||||
UpdateProfileRequest,
|
||||
)
|
||||
from src.models.database import ApiKey, Provider, Usage, User
|
||||
from src.models.database import ApiKey, GlobalModel, Model, Provider, Usage, User
|
||||
from src.services.usage.service import UsageService
|
||||
from src.services.user.apikey import ApiKeyService
|
||||
from src.services.user.preference import PreferenceService
|
||||
@@ -261,6 +263,34 @@ async def list_available_providers(request: Request, db: Session = Depends(get_d
|
||||
return await pipeline.run(adapter=adapter, http_request=request, db=db, mode=adapter.mode)
|
||||
|
||||
|
||||
@router.get("/available-models")
|
||||
async def list_available_models(
|
||||
request: Request,
|
||||
skip: int = Query(0, ge=0, description="跳过记录数"),
|
||||
limit: int = Query(100, ge=1, le=1000, description="返回记录数限制"),
|
||||
search: Optional[str] = Query(None, description="搜索关键词"),
|
||||
db: Session = Depends(get_db),
|
||||
):
|
||||
"""
|
||||
获取用户可用的模型列表
|
||||
|
||||
根据用户权限返回可用的 GlobalModel 列表。
|
||||
- 管理员:可以看到所有活跃提供商的模型
|
||||
- 普通用户:只能看到关联提供商的模型
|
||||
|
||||
**查询参数**:
|
||||
- skip: 跳过的记录数,用于分页,默认 0
|
||||
- limit: 返回记录数限制,默认 100,范围 1-1000
|
||||
- search: 可选,搜索关键词,支持模糊匹配模型名称
|
||||
|
||||
**返回字段**:
|
||||
- models: 模型列表
|
||||
- total: 符合条件的模型总数
|
||||
"""
|
||||
adapter = ListAvailableModelsAdapter(skip=skip, limit=limit, search=search)
|
||||
return await pipeline.run(adapter=adapter, http_request=request, db=db, mode=adapter.mode)
|
||||
|
||||
|
||||
@router.get("/endpoint-status")
|
||||
async def get_endpoint_status(request: Request, db: Session = Depends(get_db)):
|
||||
"""
|
||||
@@ -981,13 +1011,109 @@ class GetMyActivityHeatmapAdapter(AuthenticatedApiAdapter):
|
||||
return result
|
||||
|
||||
|
||||
@dataclass
|
||||
class ListAvailableModelsAdapter(AuthenticatedApiAdapter):
|
||||
"""获取用户可用模型列表的适配器"""
|
||||
|
||||
skip: int
|
||||
limit: int
|
||||
search: Optional[str]
|
||||
|
||||
async def handle(self, context): # type: ignore[override]
|
||||
from sqlalchemy import or_
|
||||
|
||||
from src.api.base.models_service import AccessRestrictions
|
||||
|
||||
db = context.db
|
||||
user = context.user
|
||||
|
||||
# 使用 AccessRestrictions 类来处理限制(与 /v1/models 逻辑一致)
|
||||
restrictions = AccessRestrictions.from_api_key_and_user(api_key=None, user=user)
|
||||
|
||||
# 获取所有活跃的 Provider ID
|
||||
all_active_provider_ids = {
|
||||
p.id for p in db.query(Provider.id).filter(Provider.is_active == True).all()
|
||||
}
|
||||
|
||||
if not all_active_provider_ids:
|
||||
return {"models": [], "total": 0}
|
||||
|
||||
# 查询所有活跃的 GlobalModel 及其关联的 Model
|
||||
id_query = (
|
||||
db.query(GlobalModel.id, GlobalModel.name, Model.provider_id)
|
||||
.join(Model, Model.global_model_id == GlobalModel.id)
|
||||
.filter(
|
||||
and_(
|
||||
Model.provider_id.in_(all_active_provider_ids),
|
||||
Model.is_active == True,
|
||||
GlobalModel.is_active == True,
|
||||
)
|
||||
)
|
||||
)
|
||||
|
||||
# 搜索过滤
|
||||
if self.search:
|
||||
search_term = f"%{self.search}%"
|
||||
id_query = id_query.filter(
|
||||
or_(
|
||||
GlobalModel.name.ilike(search_term),
|
||||
GlobalModel.display_name.ilike(search_term),
|
||||
)
|
||||
)
|
||||
|
||||
# 获取所有匹配的记录
|
||||
all_matches = id_query.all()
|
||||
|
||||
# 应用访问限制过滤
|
||||
allowed_global_model_ids = set()
|
||||
for global_model_id, model_name, provider_id in all_matches:
|
||||
# 使用 AccessRestrictions.is_model_allowed 检查模型是否可访问
|
||||
# 它会同时检查 allowed_providers 和 allowed_models
|
||||
if restrictions.is_model_allowed(model_name, provider_id):
|
||||
allowed_global_model_ids.add(global_model_id)
|
||||
|
||||
# 统计总数
|
||||
total = len(allowed_global_model_ids)
|
||||
|
||||
if not allowed_global_model_ids:
|
||||
return {"models": [], "total": 0}
|
||||
|
||||
# 分页并获取完整的 GlobalModel 对象
|
||||
models = (
|
||||
db.query(GlobalModel)
|
||||
.filter(GlobalModel.id.in_(allowed_global_model_ids))
|
||||
.order_by(GlobalModel.name)
|
||||
.offset(self.skip)
|
||||
.limit(self.limit)
|
||||
.all()
|
||||
)
|
||||
|
||||
# 转换为响应格式(复用 PublicGlobalModelResponse schema)
|
||||
model_responses = [
|
||||
PublicGlobalModelResponse(
|
||||
id=gm.id,
|
||||
name=gm.name,
|
||||
display_name=gm.display_name,
|
||||
is_active=gm.is_active,
|
||||
default_price_per_request=gm.default_price_per_request,
|
||||
default_tiered_pricing=gm.default_tiered_pricing,
|
||||
supported_capabilities=gm.supported_capabilities,
|
||||
config=gm.config,
|
||||
)
|
||||
for gm in models
|
||||
]
|
||||
|
||||
logger.debug(f"用户 {user.email} 可用模型: {len(model_responses)} 个")
|
||||
return PublicGlobalModelListResponse(models=model_responses, total=total)
|
||||
|
||||
|
||||
class ListAvailableProvidersAdapter(AuthenticatedApiAdapter):
|
||||
"""获取可用提供商列表的适配器"""
|
||||
|
||||
async def handle(self, context): # type: ignore[override]
|
||||
from sqlalchemy.orm import selectinload
|
||||
|
||||
from src.models.database import Model, ProviderEndpoint
|
||||
from src.models.database import ProviderEndpoint
|
||||
|
||||
db = context.db
|
||||
|
||||
|
||||
@@ -5,12 +5,14 @@ Provider 操作服务
|
||||
"""
|
||||
|
||||
import asyncio
|
||||
import os
|
||||
from dataclasses import asdict
|
||||
from datetime import datetime, timezone
|
||||
from typing import Any, Dict, List, Optional
|
||||
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from src.config import config
|
||||
from src.core.cache_service import CacheService
|
||||
from src.core.crypto import CryptoService
|
||||
from src.core.logger import logger
|
||||
@@ -33,6 +35,36 @@ from src.services.provider_ops.types import (
|
||||
BALANCE_CACHE_TTL = 86400
|
||||
|
||||
|
||||
def _get_batch_balance_concurrency() -> int:
|
||||
"""
|
||||
动态计算批量余额查询的并发限制
|
||||
|
||||
计算逻辑:
|
||||
1. 优先使用环境变量 BATCH_BALANCE_CONCURRENCY
|
||||
2. 否则根据连接池大小自动计算(取 40% 的连接池容量)
|
||||
3. 限制在 [3, 15] 范围内
|
||||
|
||||
Returns:
|
||||
并发限制数
|
||||
"""
|
||||
# 优先使用环境变量
|
||||
env_value = os.getenv("BATCH_BALANCE_CONCURRENCY")
|
||||
if env_value:
|
||||
try:
|
||||
return max(1, int(env_value))
|
||||
except ValueError:
|
||||
pass
|
||||
|
||||
# 根据连接池大小自动计算
|
||||
# 连接池容量 = pool_size + max_overflow
|
||||
pool_capacity = config.db_pool_size + config.db_max_overflow
|
||||
|
||||
# 取 40% 的连接池容量,保留 60% 给其他请求
|
||||
# 最小 3(保证基本并发),最大 15(避免过度并发)
|
||||
calculated = int(pool_capacity * 0.4)
|
||||
return max(3, min(calculated, 15))
|
||||
|
||||
|
||||
class ProviderOpsService:
|
||||
"""
|
||||
Provider 操作服务
|
||||
@@ -366,6 +398,7 @@ class ProviderOpsService:
|
||||
self,
|
||||
provider_id: str,
|
||||
trigger_refresh: bool = True,
|
||||
allow_sync_query: bool = True,
|
||||
) -> ActionResult:
|
||||
"""
|
||||
查询余额(优先返回缓存,可触发异步刷新)
|
||||
@@ -373,6 +406,7 @@ class ProviderOpsService:
|
||||
Args:
|
||||
provider_id: Provider ID
|
||||
trigger_refresh: 是否触发后台异步刷新
|
||||
allow_sync_query: 缓存未命中时是否允许同步查询(False 时仅返回缓存或触发异步刷新)
|
||||
|
||||
Returns:
|
||||
操作结果(可能是缓存的)
|
||||
@@ -387,19 +421,43 @@ class ProviderOpsService:
|
||||
asyncio.create_task(self._refresh_balance_async(provider_id))
|
||||
return cached
|
||||
|
||||
# 没有缓存,同步查询一次(首次访问)
|
||||
logger.info(f"余额缓存未命中,同步查询: provider_id={provider_id}")
|
||||
return await self.query_balance(provider_id)
|
||||
# 没有缓存
|
||||
if allow_sync_query:
|
||||
# 同步查询一次(首次访问)
|
||||
logger.info(f"余额缓存未命中,同步查询: provider_id={provider_id}")
|
||||
return await self.query_balance(provider_id)
|
||||
else:
|
||||
# 仅触发异步刷新,立即返回
|
||||
logger.debug(f"余额缓存未命中,触发异步刷新: provider_id={provider_id}")
|
||||
asyncio.create_task(self._refresh_balance_async(provider_id))
|
||||
return ActionResult(
|
||||
status=ActionStatus.PENDING,
|
||||
action_type=ProviderActionType.QUERY_BALANCE,
|
||||
message="余额数据加载中,请稍后刷新",
|
||||
)
|
||||
|
||||
async def _refresh_balance_async(self, provider_id: str) -> None:
|
||||
"""后台异步刷新余额(使用独立的数据库 session)"""
|
||||
"""
|
||||
后台异步刷新余额(使用独立的数据库 session)
|
||||
|
||||
注意:这是一个后台任务,使用独立的短生命周期 session,
|
||||
避免长时间占用连接池资源。
|
||||
"""
|
||||
db = None
|
||||
try:
|
||||
# 后台任务需要创建独立的 session,因为原请求的 session 可能已关闭
|
||||
with create_session() as db:
|
||||
service = ProviderOpsService(db)
|
||||
await service.query_balance(provider_id)
|
||||
db = create_session()
|
||||
service = ProviderOpsService(db)
|
||||
await service.query_balance(provider_id)
|
||||
except Exception as e:
|
||||
logger.warning(f"异步刷新余额失败: provider_id={provider_id}, error={e}")
|
||||
finally:
|
||||
# 确保 session 被关闭,归还连接到连接池
|
||||
if db is not None:
|
||||
try:
|
||||
db.close()
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
async def _clear_balance_cache(self, provider_id: str) -> None:
|
||||
"""清除余额缓存(认证失败时调用)"""
|
||||
@@ -643,6 +701,8 @@ class ProviderOpsService:
|
||||
"""
|
||||
批量查询余额(优先返回缓存,后台异步刷新)
|
||||
|
||||
使用信号量限制并发数,避免数据库连接池耗尽。
|
||||
|
||||
Args:
|
||||
provider_ids: Provider ID 列表,None 表示查询所有已配置的
|
||||
|
||||
@@ -658,26 +718,36 @@ class ProviderOpsService:
|
||||
if p.config and p.config.get("provider_ops")
|
||||
]
|
||||
|
||||
# 并行查询,使用缓存优先策略
|
||||
tasks = [
|
||||
self.query_balance_with_cache(provider_id, trigger_refresh=True)
|
||||
for provider_id in provider_ids
|
||||
]
|
||||
results_list = await asyncio.gather(*tasks, return_exceptions=True)
|
||||
if not provider_ids:
|
||||
return {}
|
||||
|
||||
results = {}
|
||||
for provider_id, result in zip(provider_ids, results_list):
|
||||
if isinstance(result, Exception):
|
||||
logger.warning(f"查询余额失败: provider_id={provider_id}, error={result}")
|
||||
results[provider_id] = ActionResult(
|
||||
status=ActionStatus.UNKNOWN_ERROR,
|
||||
action_type=ProviderActionType.QUERY_BALANCE,
|
||||
message=str(result),
|
||||
)
|
||||
else:
|
||||
results[provider_id] = result
|
||||
# 使用信号量限制并发数,避免同时发起过多请求耗尽连接池
|
||||
concurrency = _get_batch_balance_concurrency()
|
||||
semaphore = asyncio.Semaphore(concurrency)
|
||||
logger.debug(f"批量余额查询: providers={len(provider_ids)}, concurrency={concurrency}")
|
||||
|
||||
return results
|
||||
async def _query_with_limit(provider_id: str) -> tuple[str, ActionResult]:
|
||||
async with semaphore:
|
||||
try:
|
||||
# 批量查询时禁用同步查询,避免阻塞请求
|
||||
# 缓存未命中时会触发异步刷新,前端可稍后重试
|
||||
result = await self.query_balance_with_cache(
|
||||
provider_id, trigger_refresh=True, allow_sync_query=False
|
||||
)
|
||||
return provider_id, result
|
||||
except Exception as e:
|
||||
logger.warning(f"查询余额失败: provider_id={provider_id}, error={e}")
|
||||
return provider_id, ActionResult(
|
||||
status=ActionStatus.UNKNOWN_ERROR,
|
||||
action_type=ProviderActionType.QUERY_BALANCE,
|
||||
message=str(e),
|
||||
)
|
||||
|
||||
# 并行查询,但受信号量限制
|
||||
tasks = [_query_with_limit(provider_id) for provider_id in provider_ids]
|
||||
results_list = await asyncio.gather(*tasks)
|
||||
|
||||
return dict(results_list)
|
||||
|
||||
# ==================== 认证验证 ====================
|
||||
|
||||
|
||||
@@ -34,6 +34,7 @@ class ActionStatus(str, Enum):
|
||||
"""操作执行状态"""
|
||||
|
||||
SUCCESS = "success" # 成功
|
||||
PENDING = "pending" # 处理中(异步任务已触发,尚未完成)
|
||||
AUTH_FAILED = "auth_failed" # 认证失败
|
||||
AUTH_EXPIRED = "auth_expired" # 认证过期
|
||||
RATE_LIMITED = "rate_limited" # 频率限制
|
||||
|
||||
@@ -408,28 +408,28 @@ class UserService:
|
||||
"""获取用户可用的模型
|
||||
|
||||
通过 GlobalModel + Model 关联查询用户可用模型
|
||||
逻辑:用户可用提供商 -> Provider 的 Model 实现 -> 关联的 GlobalModel
|
||||
逻辑:使用 AccessRestrictions 统一处理 allowed_providers 和 allowed_models 限制
|
||||
"""
|
||||
# 获取用户可用的提供商
|
||||
if user.role == UserRole.ADMIN:
|
||||
# 管理员可以使用所有活动提供商
|
||||
provider_ids = [
|
||||
p.id for p in db.query(Provider.id).filter(Provider.is_active == True).all()
|
||||
]
|
||||
else:
|
||||
# 普通用户使用关联的提供商
|
||||
provider_ids = [p.id for p in user.providers]
|
||||
from src.api.base.models_service import AccessRestrictions
|
||||
|
||||
if not provider_ids:
|
||||
# 使用 AccessRestrictions 类来处理限制(与 /v1/models 逻辑一致)
|
||||
restrictions = AccessRestrictions.from_api_key_and_user(api_key=None, user=user)
|
||||
|
||||
# 获取所有活跃的 Provider ID
|
||||
all_active_provider_ids = [
|
||||
p.id for p in db.query(Provider.id).filter(Provider.is_active == True).all()
|
||||
]
|
||||
|
||||
if not all_active_provider_ids:
|
||||
return []
|
||||
|
||||
# 查询这些提供商的所有活跃 Model(关联 GlobalModel)
|
||||
models = (
|
||||
# 查询所有活跃的 Model(关联 GlobalModel)
|
||||
all_models = (
|
||||
db.query(Model)
|
||||
.join(GlobalModel, Model.global_model_id == GlobalModel.id)
|
||||
.filter(
|
||||
and_(
|
||||
Model.provider_id.in_(provider_ids),
|
||||
Model.provider_id.in_(all_active_provider_ids),
|
||||
Model.is_active == True,
|
||||
GlobalModel.is_active == True,
|
||||
)
|
||||
@@ -437,6 +437,14 @@ class UserService:
|
||||
.all()
|
||||
)
|
||||
|
||||
logger.debug(f"用户 {user.email} 可用模型: {len(models)} 个 (提供商数: {len(provider_ids)})")
|
||||
# 应用访问限制过滤
|
||||
filtered_models = []
|
||||
for model in all_models:
|
||||
model_name = model.global_model.name if model.global_model else model.provider_model_name
|
||||
# 使用 AccessRestrictions.is_model_allowed 检查模型是否可访问
|
||||
if restrictions.is_model_allowed(model_name, model.provider_id):
|
||||
filtered_models.append(model)
|
||||
|
||||
return models
|
||||
logger.debug(f"用户 {user.email} 可用模型: {len(filtered_models)} 个")
|
||||
|
||||
return filtered_models
|
||||
|
||||
Reference in New Issue
Block a user