213 lines
6.5 KiB
Python
213 lines
6.5 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}失败积分回退",
|
|
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
|