feat: 优化批量余额查询和用户模型权限检查

- 添加批量余额查询并发限制,避免数据库连接池耗尽
- 支持余额加载 pending 状态和前端自动重试机制
- 添加用户可用模型 API,统一使用 AccessRestrictions
- 修复用户表单编辑时数组引用共享导致的数据覆盖问题
- 添加数据库连接池配置说明到 .env.example
This commit is contained in:
fawney19
2026-01-28 01:01:05 +08:00
parent 8e0695e9d9
commit 1e0255d0ed
12 changed files with 411 additions and 64 deletions

View File

@@ -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)
# ==================== 认证验证 ====================

View File

@@ -34,6 +34,7 @@ class ActionStatus(str, Enum):
"""操作执行状态"""
SUCCESS = "success" # 成功
PENDING = "pending" # 处理中(异步任务已触发,尚未完成)
AUTH_FAILED = "auth_failed" # 认证失败
AUTH_EXPIRED = "auth_expired" # 认证过期
RATE_LIMITED = "rate_limited" # 频率限制

View File

@@ -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