1
This commit is contained in:
@@ -0,0 +1,212 @@
|
||||
from __future__ import annotations
|
||||
|
||||
|
||||
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.credit_record_meta_service import build_refund_meta_from_charge
|
||||
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,
|
||||
record_meta=build_refund_meta_from_charge(charge, attempt_no=attempt_no),
|
||||
)
|
||||
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.in_(["chatapi_async", "chatapi_main", "chatapi_child", "hot_opening_replicate", "shot_replicate"]),
|
||||
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
|
||||
Reference in New Issue
Block a user