Files
video-gen/video-gen-api/app/services/generation_refund_service.py
T
2026-06-03 16:15:45 +08:00

206 lines
6.1 KiB
Python

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}失败积分回退 attempt:{attempt_no}",
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
await refund_unrefunded_media_charges(
db,
user_id=record.user_id,
owner_type=OWNER_GENERATION_RECORD,
owner_id=record.id,
description_prefix="生成记录",
)
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
await refund_unrefunded_media_charges(
db,
user_id=task.user_id,
owner_type=OWNER_CHAT_GENERATION_TASK,
owner_id=task.id,
description_prefix="Chat生成任务",
)
await db.flush()
return task