import logging from sqlalchemy import select from sqlalchemy.ext.asyncio import AsyncSession from app.models.api.api_model_pricing import ApiModelPricing from app.models.image_engine import ImageEngine from app.models.video_engine import VideoEngine logger = logging.getLogger("videogen") class PricingNotConfiguredError(Exception): """模型+分辨率组合未配置价格。""" def __init__(self, *, model_name: str, resolution: str): self.model_name = model_name self.resolution = resolution super().__init__( f"模型或引擎 '{self.model_name}' 在分辨率 '{self.resolution}' 下未配置,无法生成" ) async def resolve_engine_display_name(db: AsyncSession, engine_id: str) -> str: """根据引擎 ID 解析展示名称(找不到时原样返回 ID)。""" if not engine_id: return engine_id or "unknown" result = await db.execute( select(VideoEngine.name).where(VideoEngine.id == engine_id, VideoEngine.deleted_at.is_(None)).limit(1) ) name = result.scalar_one_or_none() if name: return name result = await db.execute( select(ImageEngine.name).where(ImageEngine.id == engine_id, ImageEngine.deleted_at.is_(None)).limit(1) ) name = result.scalar_one_or_none() return name or engine_id async def _get_api_pricing( db: AsyncSession, *, gen_type: str, resolution: str, engine_id: str | None = None, ) -> ApiModelPricing | None: """按引擎精确规则优先获取定价;找不到时回退到同类型同分辨率。 查询优先级: 1. gen_type + engine_id + resolution 精确规则 2. gen_type + resolution 下 base_price 最高规则 """ gen_type = (gen_type or "").lower().strip() resolution = (resolution or "").strip() engine_id = (engine_id or "").strip() or None if engine_id: result = await db.execute( select(ApiModelPricing) .where(ApiModelPricing.gen_type == gen_type) .where(ApiModelPricing.model_config_id == engine_id) .where(ApiModelPricing.resolution == resolution) .order_by(ApiModelPricing.base_price.desc(), ApiModelPricing.per_second_price.desc()) .limit(1) ) pricing = result.scalar_one_or_none() if pricing: return pricing result = await db.execute( select(ApiModelPricing) .where(ApiModelPricing.gen_type == gen_type) .where(ApiModelPricing.resolution == resolution) .order_by(ApiModelPricing.base_price.desc(), ApiModelPricing.per_second_price.desc()) .limit(1) ) return result.scalar_one_or_none() async def calc_api_video_price( db: AsyncSession, duration: int, resolution: str, engine_id: str | None = None, input_video_duration: float = 0, input_image_count: int = 0, ) -> float: """计算 API 视频生成价格(元)。 未配置价格时抛出 PricingNotConfiguredError。 公式(与 credit_ratios 一致): base_cost = (base_price + per_second_price × duration) × price_ratio if 传入视频: += (input_video_base_price + input_video_per_second_price × input_video_duration) × input_video_ratio if 传入图片: += (input_image_base_price + input_image_per_image_price × input_image_count) × input_image_ratio """ if not engine_id: result = await db.execute( select(VideoEngine.id) .where(VideoEngine.is_active == True, VideoEngine.deleted_at.is_(None)) .order_by(VideoEngine.priority.desc()) .limit(1) ) engine_id = result.scalar_one_or_none() pricing = await _get_api_pricing(db, gen_type="video", resolution=resolution, engine_id=engine_id) if not pricing: raise PricingNotConfiguredError( model_name=await resolve_engine_display_name(db, engine_id), resolution=resolution, ) # 基础价格 base_cost = (pricing.base_price + pricing.per_second_price * duration) * pricing.price_ratio # 传入视频附加费(每秒 × 倍率) if input_video_duration > 0: base_cost += (pricing.input_video_base_price + pricing.input_video_per_second_price * input_video_duration) * pricing.input_video_ratio # 传入图片附加费(每张 × 倍率) if input_image_count > 0: base_cost += (pricing.input_image_base_price + pricing.input_image_per_image_price * input_image_count) * pricing.input_image_ratio return round(base_cost, 2) async def calc_api_image_price( db: AsyncSession, image_size: str, engine_id: str | None = None, input_image_count: int = 0, ) -> float: """计算 API 图片生成价格(元)。 未配置价格时抛出 PricingNotConfiguredError。 公式(与 credit_ratios 一致): base_cost = base_price × price_ratio if 传入图片: += (input_image_base_price + input_image_per_image_price × input_image_count) × input_image_ratio """ if not engine_id: result = await db.execute( select(ImageEngine.id) .where(ImageEngine.is_active == True, ImageEngine.deleted_at.is_(None)) .order_by(ImageEngine.priority.desc()) .limit(1) ) engine_id = result.scalar_one_or_none() pricing = await _get_api_pricing(db, gen_type="image", resolution=image_size, engine_id=engine_id) if not pricing: raise PricingNotConfiguredError( model_name=await resolve_engine_display_name(db, engine_id), resolution=image_size, ) # 基础价格 base_cost = pricing.base_price * pricing.price_ratio # 传入图片附加费(每张 × 倍率) if input_image_count > 0: base_cost += (pricing.input_image_base_price + pricing.input_image_per_image_price * input_image_count) * pricing.input_image_ratio return round(base_cost, 2) async def get_priced_models(db: AsyncSession) -> set[str]: """获取所有已配置价格的引擎 ID 集合。 用于过滤 /api/v3/models 接口,仅返回已定价的模型。 """ result = await db.execute( select(ApiModelPricing.model_config_id).distinct() ) return {row[0] for row in result.all()}