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()