from __future__ import annotations from datetime import datetime, timezone from typing import Iterable from sqlalchemy import select from sqlalchemy.ext.asyncio import AsyncSession from app.models.chat_generation_task import ChatGenerationTask from app.models.credit_record import CreditRecord from app.models.generation_record import GenerationRecord from app.services.credits import refund_credits from app.services.generation_billing_service import ( CHARGE_MEDIA, OWNER_CHAT_GENERATION_TASK, OWNER_GENERATION_RECORD, build_credit_biz_key, parse_credit_biz_key, ) TERMINAL_FAILED_STAGES = {"failed", "timeout", "download_failed"} def _round2(value: float | int | None) -> float: return round(float(value or 0), 2) async def _has_refund_for_biz_key(db: AsyncSession, *, user_id: str, charge_biz_key: str) -> bool: result = await db.execute( select(CreditRecord.id) .where( CreditRecord.user_id == user_id, CreditRecord.type == "refund", CreditRecord.refund_for_biz_key == charge_biz_key, ) .limit(1) ) return result.scalar_one_or_none() is not None async def _find_unrefunded_media_charges( db: AsyncSession, *, user_id: str, owner_type: str, owner_id: str, ) -> list[CreditRecord]: """查找当前任务下所有未退款的媒体生成扣费流水。""" pattern = f"{owner_type}:{owner_id}:attempt:%:{CHARGE_MEDIA}:charge" result = await db.execute( select(CreditRecord) .where( CreditRecord.user_id == user_id, CreditRecord.related_id == owner_id, CreditRecord.type == "consume", CreditRecord.biz_key.like(pattern), ) .order_by(CreditRecord.created_at.asc()) ) charges = list(result.scalars().all()) unrefunded: list[CreditRecord] = [] for charge in charges: if not charge.biz_key: continue if not await _has_refund_for_biz_key(db, user_id=user_id, charge_biz_key=charge.biz_key): unrefunded.append(charge) return unrefunded async def refund_unrefunded_media_charges( db: AsyncSession, *, user_id: str, owner_type: str, owner_id: str, description_prefix: str, ) -> float: """回退当前任务所有未退款媒体扣费流水。 容灾考虑: - 不依赖 retry_count 推断当前轮次。 - 如果扣费成功后 worker 崩溃,最终失败时会扫出未退款的 media charge 并补偿。 - refund_for_biz_key 保证同一轮扣费不会重复退款。 """ total_refunded = 0.0 charges = await _find_unrefunded_media_charges( db, user_id=user_id, owner_type=owner_type, owner_id=owner_id, ) for charge in charges: parsed = parse_credit_biz_key(charge.biz_key) if not parsed: continue attempt_no = int(parsed["attempt_no"]) refund_biz_key = build_credit_biz_key( owner_type=owner_type, owner_id=owner_id, attempt_no=attempt_no, charge_kind=CHARGE_MEDIA, action="refund", ) amount = abs(_round2(charge.amount)) if amount <= 0: continue await refund_credits( db, user_id=user_id, amount=amount, description=f"{description_prefix}失败积分回退", related_id=owner_id, biz_key=refund_biz_key, refund_for_biz_key=charge.biz_key, ) total_refunded = round(total_refunded + amount, 2) return total_refunded async def mark_generation_record_failed_and_refund_once( db: AsyncSession, *, record_id: str | None = None, record: GenerationRecord | None = None, error_message: str | None = None, ) -> GenerationRecord | None: """把 GenerationRecord 标记为最终失败并幂等退回媒体生成积分。 调用方负责 commit;本函数只 flush,确保状态和退款在同一事务内提交。 """ if record is None: if not record_id: return None result = await db.execute( select(GenerationRecord) .where(GenerationRecord.id == record_id, GenerationRecord.deleted_at.is_(None)) .with_for_update() .limit(1) ) record = result.scalar_one_or_none() if not record: return None if record.status == "completed": return record record.status = "failed" if error_message: record.error_message = error_message refunded_amount = await refund_unrefunded_media_charges( db, user_id=record.user_id, owner_type=OWNER_GENERATION_RECORD, owner_id=record.id, description_prefix="生成记录", ) if refunded_amount > 0: # GenerationRecord.credits_cost 只代表视频/图片生成媒体积分。 # 提示词优化积分在 text_credits_cost 中记录,优化成功后不参与生成失败退款。 record.credits_cost = max(0.0, round(float(record.credits_cost or 0) - refunded_amount, 2)) await db.flush() return record async def mark_chat_generation_task_failed_and_refund_once( db: AsyncSession, *, task_id: str | None = None, task: ChatGenerationTask | None = None, error_message: str | None = None, pipeline_stage: str = "failed", ) -> ChatGenerationTask | None: """把 ChatGenerationTask 标记为最终失败并幂等退回媒体生成积分。""" if task is None: if not task_id: return None result = await db.execute( select(ChatGenerationTask) .where( ChatGenerationTask.id == task_id, ChatGenerationTask.generation_mode == "chatapi_async", ChatGenerationTask.deleted_at.is_(None), ) .with_for_update() .limit(1) ) task = result.scalar_one_or_none() if not task: return None if task.status == "completed": return task task.status = "failed" task.pipeline_stage = pipeline_stage if pipeline_stage in TERMINAL_FAILED_STAGES else "failed" if error_message: task.error_message = error_message refunded_amount = await refund_unrefunded_media_charges( db, user_id=task.user_id, owner_type=OWNER_CHAT_GENERATION_TASK, owner_id=task.id, description_prefix="任务生成", ) if refunded_amount > 0: task.credits_cost = max(0.0, round(float(task.credits_cost or 0) - refunded_amount, 2)) await db.flush() return task