69 lines
2.3 KiB
Python
69 lines
2.3 KiB
Python
import logging
|
|
from datetime import datetime, timezone
|
|
|
|
from sqlalchemy import func, select
|
|
from sqlalchemy.ext.asyncio import AsyncSession
|
|
|
|
from app.models.api.api_generation_task import ApiGenerationTask
|
|
from app.models.api.api_key import ApiKey
|
|
|
|
logger = logging.getLogger("videogen")
|
|
|
|
|
|
async def check_quota(key: ApiKey) -> bool:
|
|
"""检查 API Key 配额是否充足。
|
|
|
|
Returns:
|
|
True = 配额充足或无限额, False = 已超限。
|
|
"""
|
|
if key.quota_limit is None:
|
|
return True
|
|
return key.quota_used < key.quota_limit
|
|
|
|
|
|
async def get_active_video_tasks_count(api_key_id: str, db: AsyncSession) -> int:
|
|
"""统计 API Key 当前活跃的视频任务数。
|
|
|
|
活跃 = status IN ('pending', 'generating', 'processing') AND gen_type='video'
|
|
"""
|
|
result = await db.execute(
|
|
select(func.count(ApiGenerationTask.id)).where(
|
|
ApiGenerationTask.api_key_id == api_key_id,
|
|
ApiGenerationTask.gen_type == "video",
|
|
ApiGenerationTask.status.in_(["pending", "generating", "processing"]),
|
|
ApiGenerationTask.deleted_at.is_(None),
|
|
)
|
|
)
|
|
return result.scalar_one() or 0
|
|
|
|
|
|
async def can_start_video_task(key: ApiKey, db: AsyncSession) -> bool:
|
|
"""检查是否可以立即启动新的视频任务。
|
|
|
|
Returns:
|
|
True = 可以立即启动, False = 需要排队。
|
|
"""
|
|
if key.max_concurrent_video_tasks is None:
|
|
return True # 无限制
|
|
current = await get_active_video_tasks_count(key.id, db)
|
|
return current < key.max_concurrent_video_tasks
|
|
|
|
|
|
async def get_queued_video_tasks(key: ApiKey, db: AsyncSession, limit: int = 10) -> list[ApiGenerationTask]:
|
|
"""获取排队的视频任务列表(按创建时间排序)。"""
|
|
result = await db.execute(
|
|
select(ApiGenerationTask).where(
|
|
ApiGenerationTask.api_key_id == key.id,
|
|
ApiGenerationTask.gen_type == "video",
|
|
ApiGenerationTask.status == "queued",
|
|
ApiGenerationTask.deleted_at.is_(None),
|
|
).order_by(ApiGenerationTask.created_at.asc()).limit(limit)
|
|
)
|
|
return list(result.scalars().all())
|
|
|
|
|
|
async def increment_quota(db: AsyncSession, key: ApiKey, credits_cost: float) -> None:
|
|
"""原子性增加配额使用量。"""
|
|
key.quota_used = round((key.quota_used or 0.0) + credits_cost, 2)
|
|
await db.flush()
|