生成任务失败积分回退

This commit is contained in:
2026-06-03 16:15:45 +08:00
parent 7c56c02135
commit 1079a5a7d1
14 changed files with 931 additions and 259 deletions
@@ -0,0 +1,105 @@
"""add credit_record biz_key
Revision ID: 516b84e3bf17
Revises: e7bf423b248c
Create Date: 2026-06-03 16:12:46.784240
"""
from typing import Sequence, Union
from alembic import op
import sqlalchemy as sa
# revision identifiers, used by Alembic.
revision: str = "516b84e3bf17"
down_revision: Union[str, None] = "e7bf423b248c"
branch_labels: Union[str, Sequence[str], None] = None
depends_on: Union[str, Sequence[str], None] = None
def upgrade() -> None:
# 1. CreditRecord 新增正式幂等字段
op.add_column(
"credit_records",
sa.Column("biz_key", sa.String(length=160), nullable=True),
)
op.add_column(
"credit_records",
sa.Column("refund_for_biz_key", sa.String(length=160), nullable=True),
)
# 2. 普通查询索引
op.create_index(
"ix_credit_records_biz_key",
"credit_records",
["biz_key"],
unique=False,
)
op.create_index(
"ix_credit_records_refund_for_biz_key",
"credit_records",
["refund_for_biz_key"],
unique=False,
)
op.create_index(
"ix_credit_records_related_type",
"credit_records",
["related_id", "type"],
unique=False,
)
op.create_index(
"ix_credit_records_user_refund_for_biz_key",
"credit_records",
["user_id", "refund_for_biz_key"],
unique=False,
)
# 3. PostgreSQL 部分唯一索引:只限制 biz_key 不为空的记录
op.create_index(
"uq_credit_records_user_biz_key",
"credit_records",
["user_id", "biz_key"],
unique=True,
postgresql_where=sa.text("biz_key IS NOT NULL"),
)
# 4. ChatGenerationTask 幂等 key 唯一索引
# 只限制 idempotency_key 不为空且不为空字符串的数据
op.create_index(
"uq_chat_generation_tasks_user_mode_idempotency",
"chat_generation_tasks",
["user_id", "generation_mode", "idempotency_key"],
unique=True,
postgresql_where=sa.text(
"idempotency_key IS NOT NULL AND idempotency_key <> ''"
),
)
def downgrade() -> None:
op.drop_index(
"uq_chat_generation_tasks_user_mode_idempotency",
table_name="chat_generation_tasks",
)
op.drop_index(
"uq_credit_records_user_biz_key",
table_name="credit_records",
)
op.drop_index(
"ix_credit_records_user_refund_for_biz_key",
table_name="credit_records",
)
op.drop_index(
"ix_credit_records_related_type",
table_name="credit_records",
)
op.drop_index(
"ix_credit_records_refund_for_biz_key",
table_name="credit_records",
)
op.drop_index(
"ix_credit_records_biz_key",
table_name="credit_records",
)
op.drop_column("credit_records", "refund_for_biz_key")
op.drop_column("credit_records", "biz_key")
+73 -24
View File
@@ -43,6 +43,13 @@ from app.services.notification import create_notification
from app.services.auth import hash_password, verify_password
from app.services.operation_log import log_operation
from app.services.resource_signed_url_service import build_resource_signed_url
from app.services.generation_billing_service import (
OWNER_GENERATION_RECORD,
charge_generation_media_by_params,
get_next_credit_attempt_no,
)
from app.services.generation_refund_service import mark_generation_record_failed_and_refund_once
from app.utils.id_gen import generate_id
from app.schemas.generation import GenerationType, ASPECT_RATIOS, RESOLUTIONS
@@ -1054,6 +1061,7 @@ async def admin_update_generation_status(
GenerationRecord.id == record_id,
GenerationRecord.deleted_at.is_(None),
)
.with_for_update()
.limit(1)
)
record = result.scalar_one_or_none()
@@ -1064,11 +1072,21 @@ async def admin_update_generation_status(
if new_status not in ("prompt_optimized", "generating", "completed", "failed"):
raise HTTPException(status_code=400, detail="无效状态")
record.status = new_status
if new_status == "failed":
await mark_generation_record_failed_and_refund_once(
db,
record=record,
error_message=body.get("error_message") or record.error_message or "管理员设置为失败",
)
else:
record.status = new_status
if body.get("video_url"):
record.video_url = body["video_url"]
if body.get("video_cover_url"):
record.video_cover_url = body["video_cover_url"]
if body.get("image_url"):
record.image_url = body["image_url"]
if new_status == "completed":
record.generated_at = datetime.now()
await db.flush()
@@ -1082,9 +1100,8 @@ async def admin_generate_video(
admin: User = Depends(get_admin_user),
db: AsyncSession = Depends(get_db),
):
"""Admin trigger video generation for a record with specified params."""
"""Admin trigger video/image generation for a record with specified params."""
from app.models.project import Project
from app.services.credits import calc_video_credits, deduct_credits, calc_image_credits
from app.services.video_queue import task_queue
result = await db.execute(
@@ -1095,6 +1112,7 @@ async def admin_generate_video(
GenerationRecord.deleted_at.is_(None),
Project.deleted_at.is_(None),
)
.with_for_update()
)
row = result.first()
if not row:
@@ -1106,6 +1124,12 @@ async def admin_generate_video(
if record.status not in ("prompt_optimized", "failed"):
raise HTTPException(status_code=400, detail=f"当前状态不允许生成{type_str}")
attempt_no = await get_next_credit_attempt_no(
db,
owner_type=OWNER_GENERATION_RECORD,
owner_id=record.id,
)
if record.gen_type == GenerationType.video:
# Video Generation
aspect_ratio = body.get("aspect_ratio", "16:9")
@@ -1116,55 +1140,80 @@ async def admin_generate_video(
raise HTTPException(status_code=400, detail="不支持的分辨率")
duration = record.duration or 5
video_credits = await calc_video_credits(db, duration, resolution)
await deduct_credits(
db, record.user_id, video_credits,
f"视频生成(管理后台) - {project_name}",
related_id=record_id,
media_billing = await charge_generation_media_by_params(
db,
user_id=record.user_id,
record_id=record.id,
gen_type="video",
duration=duration,
resolution=resolution,
project_name=project_name,
description_prefix="视频生成(管理后台)",
owner_type=OWNER_GENERATION_RECORD,
attempt_no=attempt_no,
)
record.aspect_ratio = aspect_ratio
record.resolution = resolution
record.credits_cost = (record.credits_cost or 0) + video_credits
record.credits_cost = round(float(record.credits_cost or 0) + media_billing.total_charged, 2)
record.status = "generating"
record.error_message = None
record.video_url = None
record.video_cover_url = None
record.image_url = None
record.seedance_task_id = None
await db.flush()
try:
from app.services.video_gen import get_active_engine, submit_video_task
engine = await get_active_engine(db)
task_id = await submit_video_task(db, engine, record)
record.seedance_task_id = task_id
await db.flush()
await task_queue.enqueue(record_id)
except Exception as e:
record.status = "failed"
record.error_message = str(e)
await mark_generation_record_failed_and_refund_once(
db,
record=record,
error_message=str(e),
)
await db.flush()
elif record.gen_type == GenerationType.image:
# Image generation
post_image_size = body.get("image_size", "")
image_credits = await calc_image_credits(db, post_image_size or record.image_size or "2K")
await deduct_credits(
db, record.user_id, image_credits,
f"图片生成 - {project_name}",
related_id=record_id,
image_size = post_image_size or record.image_size or "2K"
media_billing = await charge_generation_media_by_params(
db,
user_id=record.user_id,
record_id=record.id,
gen_type="image",
image_size=image_size,
project_name=project_name,
description_prefix="图片生成(管理后台)",
owner_type=OWNER_GENERATION_RECORD,
attempt_no=attempt_no,
)
record.image_size = image_size
record.credits_cost = round(float(record.credits_cost or 0) + media_billing.total_charged, 2)
record.status = "generating"
record.error_message = None
record.image_url = None
record.video_url = None
record.video_cover_url = None
record.seedance_task_id = None
await db.flush()
try:
record.image_size = post_image_size or record.image_size or "2K"
record.credits_cost = round(image_credits, 2)
record.status = "generating"
record.error_message = None
await db.flush()
await task_queue.enqueue(record_id)
except Exception as e:
record.status = "failed"
record.error_message = str(e)
await mark_generation_record_failed_and_refund_once(
db,
record=record,
error_message=str(e),
)
await db.flush()
return {"message": "ok", "record_id": record_id}
+98 -42
View File
@@ -26,7 +26,7 @@ from app.schemas.generation import (
RESOLUTIONS,
IMAGE_SIZES,
)
from app.services.credits import deduct_credits, calc_text_credits, calc_video_credits, calc_image_credits
from app.services.credits import deduct_credits, calc_text_credits
from app.services.llm import optimize_prompt
from app.services.video_url import generate_temp_url, validate_and_get_record_id, get_video_stream_url
from app.services.resource_accounting_service import (
@@ -34,6 +34,13 @@ from app.services.resource_accounting_service import (
safe_file_size,
)
from app.services.resource_signed_url_service import build_resource_signed_url
from app.services.generation_billing_service import (
OWNER_GENERATION_RECORD,
charge_generation_media_by_params,
charge_generation_media_for_record,
get_next_credit_attempt_no,
)
from app.services.generation_refund_service import mark_generation_record_failed_and_refund_once
from app.services.video_cover_service import async_create_video_cover_for_local_video
from app.utils.id_gen import generate_id
from app.utils.exceptions import InsufficientCreditsError, RecordNotFoundError, InvalidStatusError
@@ -319,6 +326,7 @@ async def generate(
GenerationRecord.deleted_at.is_(None),
Project.deleted_at.is_(None),
)
.with_for_update()
)
row = result.first()
if not row:
@@ -328,6 +336,12 @@ async def generate(
if record.status not in ("prompt_optimized", "failed"):
raise InvalidStatusError("当前状态不允许生成")
attempt_no = await get_next_credit_attempt_no(
db,
owner_type=OWNER_GENERATION_RECORD,
owner_id=record.id,
)
if record.gen_type == GenerationType.video:
# Video generation
if req.aspect_ratio not in ASPECT_RATIOS:
@@ -336,20 +350,28 @@ async def generate(
raise HTTPException(status_code=400, detail="不支持的分辨率")
duration = record.duration or 5
video_credits = await calc_video_credits(db, duration, req.resolution)
await deduct_credits(
db, current_user.id, video_credits,
f"视频生成 - {project_name}",
related_id=record_id,
media_billing = await charge_generation_media_by_params(
db,
user_id=current_user.id,
record_id=record.id,
gen_type="video",
duration=duration,
resolution=req.resolution,
project_name=project_name,
description_prefix="视频生成",
owner_type=OWNER_GENERATION_RECORD,
attempt_no=attempt_no,
)
record.aspect_ratio = req.aspect_ratio
record.resolution = req.resolution
record.credits_cost = round(video_credits, 2)
record.credits_cost = round(float(record.credits_cost or 0) + media_billing.total_charged, 2)
record.status = "generating"
record.error_message = None
record.video_url = None
record.video_cover_url = None
record.image_url = None
record.seedance_task_id = None
await db.flush()
try:
@@ -363,34 +385,51 @@ async def generate(
await db.flush()
await task_queue.enqueue(record_id)
except Exception as e:
record.status = "failed"
record.error_message = extract_error_message(e, "视频")
await mark_generation_record_failed_and_refund_once(
db,
record=record,
error_message=extract_error_message(e, "视频"),
)
await db.flush()
elif record.gen_type == GenerationType.image:
# Image generation
image_credits = await calc_image_credits(db, req.image_size or record.image_size or "2K")
await deduct_credits(
db, current_user.id, image_credits,
f"图片生成 - {project_name}",
related_id=record_id,
image_size = req.image_size or record.image_size or "2K"
media_billing = await charge_generation_media_by_params(
db,
user_id=current_user.id,
record_id=record.id,
gen_type="image",
image_size=image_size,
project_name=project_name,
description_prefix="图片生成",
owner_type=OWNER_GENERATION_RECORD,
attempt_no=attempt_no,
)
record.image_size = req.image_size or record.image_size or "2K"
record.credits_cost = round(image_credits, 2)
record.image_size = image_size
record.credits_cost = round(float(record.credits_cost or 0) + media_billing.total_charged, 2)
record.status = "generating"
record.error_message = None
await db.commit()
record.image_url = None
record.video_url = None
record.video_cover_url = None
record.seedance_task_id = None
await db.flush()
from app.services.video_queue import task_queue
await task_queue.enqueue(record_id)
try:
from app.services.video_queue import task_queue
await task_queue.enqueue(record_id)
except Exception as e:
await mark_generation_record_failed_and_refund_once(
db,
record=record,
error_message=f"图片任务队列投递失败: {e}",
)
await db.flush()
return _record_to_out(record, project_name)
@router.post("/{record_id}/retry")
async def retry_generation(
record_id: str,
@@ -406,6 +445,7 @@ async def retry_generation(
GenerationRecord.deleted_at.is_(None),
Project.deleted_at.is_(None),
)
.with_for_update()
)
row = result.first()
if not row:
@@ -415,34 +455,45 @@ async def retry_generation(
if record.status != "failed":
raise InvalidStatusError("只有失败的记录可以重试")
# Re-deduct video credits for retry
if record.duration and record.resolution:
video_credits = await calc_video_credits(db, record.duration, record.resolution)
await deduct_credits(
db, current_user.id, video_credits,
f"视频重试 - {project_name}",
related_id=record_id,
)
record.credits_cost = round((record.credits_cost or 0) + video_credits, 2)
attempt_no = await get_next_credit_attempt_no(
db,
owner_type=OWNER_GENERATION_RECORD,
owner_id=record.id,
)
media_billing = await charge_generation_media_for_record(
db,
record=record,
project_name=project_name,
description_prefix="视频重试",
attempt_no=attempt_no,
)
record.status = "generating"
record.error_message = None
record.video_url = None
record.video_cover_url = None
record.image_url = None
record.seedance_task_id = None
record.generated_at = None
record.credits_cost = round(float(record.credits_cost or 0) + media_billing.total_charged, 2)
await db.flush()
try:
from app.services.video_gen import get_active_engine, submit_video_task, extract_error_message
from app.services.video_queue import task_queue
engine = await get_active_engine(db)
task_id = await submit_video_task(db, engine, record)
record.seedance_task_id = task_id
await db.flush()
if record.gen_type == GenerationType.video:
from app.services.video_gen import get_active_engine, submit_video_task, extract_error_message
engine = await get_active_engine(db)
task_id = await submit_video_task(db, engine, record)
record.seedance_task_id = task_id
await db.flush()
await task_queue.enqueue(record_id)
except Exception as e:
record.status = "failed"
record.error_message = extract_error_message(e)
from app.services.error_codes import extract_error_message
await mark_generation_record_failed_and_refund_once(
db,
record=record,
error_message=extract_error_message(e, "重试"),
)
await db.flush()
return _record_to_out(record, project_name)
@@ -551,6 +602,7 @@ async def seedance_callback(request: Request, db: AsyncSession = Depends(get_db)
GenerationRecord.seedance_task_id == task_id,
GenerationRecord.deleted_at.is_(None),
)
.with_for_update()
.limit(1)
)
record = result.scalar_one_or_none()
@@ -613,8 +665,12 @@ async def seedance_callback(request: Request, db: AsyncSession = Depends(get_db)
)
await push_notification_to_user(record.user_id, notif)
elif task_status == "failed":
record.status = "failed"
record.error_message = data.get("error", "视频生成失败")
error_message = data.get("error", "视频生成失败")
await mark_generation_record_failed_and_refund_once(
db,
record=record,
error_message=error_message,
)
# Log callback response
from app.services.video_gen import _log_video_response
_log_video_response(record.id, data, error=record.error_message)
+67 -30
View File
@@ -26,7 +26,13 @@ from app.services.generation_ai_service import (
record_to_out,
soft_delete_chat_generation_task,
)
from app.services.generation_billing_service import (
OWNER_CHAT_GENERATION_TASK,
charge_generation_media_by_params,
get_next_credit_attempt_no,
)
from app.services.generation_log_service import log_task_event
from app.services.generation_refund_service import mark_chat_generation_task_failed_and_refund_once
from app.tasks.celery_app import celery_app
router = APIRouter(
@@ -149,7 +155,18 @@ async def create_task(
from app.tasks.generation_create_tasks import chatapi_create_generation_task
chatapi_create_generation_task.delay(task.id)
try:
chatapi_create_generation_task.delay(task.id)
except Exception as exc:
await mark_chat_generation_task_failed_and_refund_once(
db,
task_id=task.id,
error_message=f"任务队列投递失败: {exc}",
pipeline_stage="failed",
)
await db.commit()
raise HTTPException(status_code=503, detail="任务队列投递失败,请稍后重试")
return record_to_out(task)
@@ -538,6 +555,7 @@ async def retry_task(
ChatGenerationTask.generation_mode == "chatapi_async",
ChatGenerationTask.deleted_at.is_(None),
)
.with_for_update()
.limit(1)
)
task = result.scalar_one_or_none()
@@ -546,41 +564,60 @@ async def retry_task(
if task.status != "failed":
raise HTTPException(status_code=400, detail="只有失败任务可以重试")
attempt_no = await get_next_credit_attempt_no(
db,
owner_type=OWNER_CHAT_GENERATION_TASK,
owner_id=task.id,
)
media_billing = await charge_generation_media_by_params(
db,
user_id=task.user_id,
record_id=task.id,
gen_type=task.gen_type,
image_size=task.image_size,
duration=task.duration,
resolution=task.resolution,
engine_id=task.engine_id,
project_name="AI生成任务",
description_prefix="Chat任务重试",
owner_type=OWNER_CHAT_GENERATION_TASK,
attempt_no=attempt_no,
)
task.status = "generating"
task.pipeline_stage = "queued"
task.error_message = None
task.retry_count = (task.retry_count or 0) + 1
if task.retry_count > 3:
raise HTTPException(status_code=400, detail="任务已超过最大重试次数")
# Resume from the earliest missing stage.
if not task.optimized_prompt:
task.pipeline_stage = "queued"
from app.tasks.generation_create_tasks import chatapi_create_generation_task
chatapi_create_generation_task.delay(task.id)
elif not task.seedance_task_id and not task.remote_result_url:
task.pipeline_stage = "creating_provider_task"
from app.tasks.generation_create_tasks import chatapi_create_generation_task
chatapi_create_generation_task.delay(task.id)
elif task.seedance_task_id and not task.remote_result_url:
task.pipeline_stage = "waiting_remote"
from app.tasks.generation_poll_tasks import poll_generation_task
poll_generation_task.delay(task.id)
elif task.remote_result_url and not (task.image_url or task.video_url):
task.pipeline_stage = "result_ready"
from app.tasks.generation_download_tasks import download_generation_result_task
download_generation_result_task.delay(task.id)
else:
task.status = "completed"
task.pipeline_stage = "done"
task.poll_count = 0
task.last_poll_at = None
task.provider_task_id = None
task.seedance_task_id = None
task.remote_result_url = None
task.provider_response_json = None
task.image_url = None
task.video_url = None
task.video_cover_url = None
task.generated_at = None
task.credits_cost = round(float(task.credits_cost or 0) + media_billing.total_charged, 2)
await db.commit()
from app.tasks.generation_create_tasks import chatapi_create_generation_task
try:
chatapi_create_generation_task.delay(task.id)
except Exception as exc:
await mark_chat_generation_task_failed_and_refund_once(
db,
task_id=task.id,
error_message=f"任务队列投递失败: {exc}",
pipeline_stage="failed",
)
await db.commit()
raise HTTPException(status_code=503, detail="任务队列投递失败,请稍后重试")
return GenerationAIRetryOut(
id=task.id,
status=task.status,
pipeline_stage=task.pipeline_stage,
message="任务已重新投递",
message="任务已重新扣费并重新投递",
)
@@ -1,6 +1,6 @@
from datetime import datetime
from sqlalchemy import DateTime, Float, ForeignKey, Integer, String, Text
from sqlalchemy import DateTime, Float, ForeignKey, Index, Integer, String, Text
from sqlalchemy.orm import Mapped, mapped_column
from app.models.base import Base, TimestampMixin, SoftDeleteMixin
@@ -15,6 +15,12 @@ class ChatGenerationTask(Base, TimestampMixin, SoftDeleteMixin):
"""
__tablename__ = "chat_generation_tasks"
__table_args__ = (
# 防止前端按钮连点/网络重试时同一个 idempotency_key 并发创建多条任务。
# nullable unique 兼容不传 idempotency_key 的普通请求。
Index("uq_chat_generation_tasks_user_mode_idempotency", "user_id", "generation_mode", "idempotency_key", unique=True),
)
id: Mapped[str] = mapped_column(String(32), primary_key=True)
user_id: Mapped[str] = mapped_column(
+16 -1
View File
@@ -1,4 +1,4 @@
from sqlalchemy import Float, ForeignKey, Integer, String
from sqlalchemy import Float, ForeignKey, Index, String
from sqlalchemy.orm import Mapped, mapped_column
from app.models.base import Base, TimestampMixin
@@ -6,6 +6,13 @@ from app.models.base import Base, TimestampMixin
class CreditRecord(Base, TimestampMixin):
__tablename__ = "credit_records"
__table_args__ = (
# 正式计费幂等键:同一用户同一个业务流水只能写入一次。
# PostgreSQL/MySQL/SQLite 对 nullable unique 的处理都允许多条 NULL,兼容历史数据。
Index("uq_credit_records_user_biz_key", "user_id", "biz_key", unique=True),
Index("ix_credit_records_user_refund_for_biz_key", "user_id", "refund_for_biz_key"),
Index("ix_credit_records_related_type", "related_id", "type"),
)
id: Mapped[str] = mapped_column(String(32), primary_key=True)
user_id: Mapped[str] = mapped_column(
@@ -16,3 +23,11 @@ class CreditRecord(Base, TimestampMixin):
balance_after: Mapped[float] = mapped_column(Float)
description: Mapped[str] = mapped_column(String(256))
related_id: Mapped[str | None] = mapped_column(String(64), nullable=True)
# 当前积分流水自己的业务幂等键。
# 例如:generation_record:{record_id}:attempt:1:media:charge
biz_key: Mapped[str | None] = mapped_column(String(160), nullable=True, index=True)
# 如果当前流水是退款,记录它退的是哪一次扣费。
# 例如:generation_record:{record_id}:attempt:1:media:charge
refund_for_biz_key: Mapped[str | None] = mapped_column(String(160), nullable=True, index=True)
+95 -14
View File
@@ -141,30 +141,71 @@ async def calc_image_credits(
return round(base_cost * multiplier, 2)
async def _get_existing_credit_record_by_biz_key(
db: AsyncSession,
*,
user_id: str,
biz_key: str | None,
) -> CreditRecord | None:
"""按正式业务幂等键查找已有积分流水。"""
if not biz_key:
return None
result = await db.execute(
select(CreditRecord)
.where(CreditRecord.user_id == user_id, CreditRecord.biz_key == biz_key)
.limit(1)
)
return result.scalar_one_or_none()
async def deduct_credits(
db: AsyncSession,
user_id: str,
amount: float,
description: str,
related_id: str | None = None,
*,
biz_key: str | None = None,
refund_for_biz_key: str | None = None,
) -> User:
"""Atomically deduct credits from user. Raises InsufficientCreditsError."""
result = await db.execute(
select(User).where(User.id == user_id).limit(1)
)
"""扣减用户积分,并写入消费流水。
并发安全点:
- 先用 SELECT ... FOR UPDATE 锁住 users 行,避免余额覆盖。
- biz_key 不为空时,作为正式业务幂等键;重复调用直接返回当前用户,不重复扣。
"""
amount = round(float(amount or 0), 2)
if amount <= 0:
result = await db.execute(select(User).where(User.id == user_id).with_for_update().limit(1))
user = result.scalar_one_or_none()
if not user:
raise ValueError("User not found")
return user
result = await db.execute(select(User).where(User.id == user_id).with_for_update().limit(1))
user = result.scalar_one_or_none()
if not user or user.credits < amount:
if not user:
raise ValueError("User not found")
if biz_key:
existing = await _get_existing_credit_record_by_biz_key(db, user_id=user_id, biz_key=biz_key)
if existing:
return user
if float(user.credits or 0) < amount:
raise InsufficientCreditsError()
user.credits = round(user.credits - amount, 2)
user.credits = round(float(user.credits or 0) - amount, 2)
record = CreditRecord(
id=generate_id(),
user_id=user_id,
type="consume",
amount=-round(amount, 2),
amount=-amount,
balance_after=user.credits,
description=description,
related_id=related_id,
biz_key=biz_key,
refund_for_biz_key=refund_for_biz_key,
)
db.add(record)
await db.flush()
@@ -177,30 +218,70 @@ async def add_credits(
amount: float,
description: str,
related_id: str | None = None,
*,
record_type: str = "recharge",
biz_key: str | None = None,
refund_for_biz_key: str | None = None,
) -> User:
"""Add credits to user."""
result = await db.execute(
select(User).where(User.id == user_id).limit(1)
)
"""增加用户积分,并写入流水。
record_type 默认保持原来的 recharge;生成失败回退时传 refund。
biz_key 不为空时幂等,重复调用不会重复加积分。
"""
amount = round(float(amount or 0), 2)
result = await db.execute(select(User).where(User.id == user_id).with_for_update().limit(1))
user = result.scalar_one_or_none()
if not user:
raise ValueError("User not found")
user.credits = round(user.credits + amount, 2)
if biz_key:
existing = await _get_existing_credit_record_by_biz_key(db, user_id=user_id, biz_key=biz_key)
if existing:
return user
if amount <= 0:
return user
user.credits = round(float(user.credits or 0) + amount, 2)
record = CreditRecord(
id=generate_id(),
user_id=user_id,
type="recharge",
amount=round(amount, 2),
type=record_type,
amount=amount,
balance_after=user.credits,
description=description,
related_id=related_id,
biz_key=biz_key,
refund_for_biz_key=refund_for_biz_key,
)
db.add(record)
await db.flush()
return user
async def refund_credits(
db: AsyncSession,
user_id: str,
amount: float,
description: str,
related_id: str | None = None,
*,
biz_key: str | None = None,
refund_for_biz_key: str | None = None,
) -> User:
"""生成失败积分回退。"""
return await add_credits(
db,
user_id=user_id,
amount=amount,
description=description,
related_id=related_id,
record_type="refund",
biz_key=biz_key,
refund_for_biz_key=refund_for_biz_key,
)
async def get_records(db: AsyncSession, user_id: str) -> list[CreditRecord]:
result = await db.execute(
select(CreditRecord)
@@ -24,7 +24,10 @@ from app.schemas.generation_ai import (
GenerationAITaskOut,
GenerationAIVideoEngineOptionOut,
)
from app.services.generation_billing_service import charge_generation_media_by_params
from app.services.generation_billing_service import (
OWNER_CHAT_GENERATION_TASK,
charge_generation_media_by_params,
)
from app.services.resource_accounting_service import soft_delete_chat_task_resources
from app.services.resource_signed_url_service import build_resource_signed_url
from app.utils.id_gen import generate_id
@@ -238,6 +241,8 @@ async def create_async_generation_task(db: AsyncSession, current_user: User, req
engine_id=engine.id,
project_name="AI生成任务",
description_prefix="Chat任务",
owner_type=OWNER_CHAT_GENERATION_TASK,
attempt_no=1,
)
snapshot = _build_image_snapshot(engine, size, proportion, px)
task = ChatGenerationTask(
@@ -284,6 +289,8 @@ async def create_async_generation_task(db: AsyncSession, current_user: User, req
engine_id=engine.id,
project_name="AI生成任务",
description_prefix="Chat任务",
owner_type=OWNER_CHAT_GENERATION_TASK,
attempt_no=1,
)
snapshot = _build_video_snapshot(engine, ratio, resolution, duration)
task = ChatGenerationTask(
@@ -1,25 +1,29 @@
from __future__ import annotations
from dataclasses import dataclass, asdict
import re
from dataclasses import asdict, dataclass
from typing import Any, Mapping
from sqlalchemy import select
from sqlalchemy import func, select
from sqlalchemy.ext.asyncio import AsyncSession
from app.models.credit_record import CreditRecord
from app.models.generation_record import GenerationRecord
from app.models.system_config import SystemConfig
from app.models.user import User
from app.services.credits import calc_image_credits, calc_text_credits, calc_video_credits
from app.utils.exceptions import InsufficientCreditsError
from app.utils.id_gen import generate_id
from app.services.credits import calc_image_credits, calc_text_credits, calc_video_credits, deduct_credits
CHARGE_TEXT_PROMPT = "CHATAPI_TEXT_PROMPT"
CHARGE_FILE_PARSE = "CHATAPI_FILE_PARSE"
CHARGE_VISION_INPUT = "CHATAPI_VISION_INPUT"
CHARGE_MEDIA_IMAGE = "CHATAPI_MEDIA_IMAGE"
CHARGE_MEDIA_VIDEO = "CHATAPI_MEDIA_VIDEO"
CHARGE_TEXT_PROMPT = "text_prompt"
CHARGE_FILE_PARSE = "file_parse"
CHARGE_VISION_INPUT = "vision_input"
CHARGE_MEDIA = "media"
OWNER_GENERATION_RECORD = "generation_record"
OWNER_CHAT_GENERATION_TASK = "chat_generation_task"
_BIZ_KEY_PATTERN = re.compile(
r"^(?P<owner_type>[^:]+):(?P<owner_id>[^:]+):attempt:(?P<attempt_no>\d+):(?P<charge_kind>[^:]+):(?P<action>charge|refund)$"
)
@dataclass
@@ -28,6 +32,8 @@ class BillingItem:
amount: float
charged: bool
skipped_reason: str | None = None
biz_key: str | None = None
attempt_no: int | None = None
@dataclass
@@ -62,6 +68,40 @@ def _safe_int(value: Any, default: int = 0) -> int:
return default
def build_credit_biz_key(
*,
owner_type: str,
owner_id: str,
attempt_no: int,
charge_kind: str,
action: str,
) -> str:
"""生成正式积分流水幂等键。
示例:generation_record:xxx:attempt:2:media:charge
"""
owner_type = owner_type.strip()
owner_id = owner_id.strip()
charge_kind = charge_kind.strip()
action = action.strip()
if action not in ("charge", "refund"):
raise ValueError("action 仅支持 charge/refund")
if attempt_no <= 0:
raise ValueError("attempt_no 必须大于 0")
return f"{owner_type}:{owner_id}:attempt:{attempt_no}:{charge_kind}:{action}"
def parse_credit_biz_key(biz_key: str | None) -> dict[str, Any] | None:
if not biz_key:
return None
match = _BIZ_KEY_PATTERN.match(biz_key)
if not match:
return None
data = match.groupdict()
data["attempt_no"] = int(data["attempt_no"])
return data
async def _get_config_float_or_none(db: AsyncSession, key: str) -> float | None:
result = await db.execute(select(SystemConfig).where(SystemConfig.key == key).limit(1))
config = result.scalar_one_or_none()
@@ -74,11 +114,6 @@ async def _get_config_float_or_none(db: AsyncSession, key: str) -> float | None:
async def _calc_optional_token_credits(db: AsyncSession, tokens: int, config_key: str) -> float:
"""Calculate optional token billing. Missing config means do not charge.
This prevents double-charging existing projects where uploaded file/OCR/vision
content is already included in the LLM provider's input_tokens.
"""
tokens = _safe_int(tokens)
if tokens <= 0:
return 0.0
@@ -88,41 +123,39 @@ async def _calc_optional_token_credits(db: AsyncSession, tokens: int, config_key
return round(tokens * rate / 1000, 2)
def _legacy_description_keywords(charge_key: str) -> list[str]:
# Compatibility with old patch/original project records that were inserted
# before this safe billing service added [CHARGE_KEY] prefixes.
if charge_key == CHARGE_TEXT_PROMPT:
return ["ChatAPI提示词整理", "提示词优化"]
if charge_key == CHARGE_MEDIA_IMAGE:
return ["ChatAPI异步图片生成", "图片生成"]
if charge_key == CHARGE_MEDIA_VIDEO:
return ["ChatAPI异步视频生成", "视频生成"]
if charge_key == CHARGE_FILE_PARSE:
return ["文件解析Token"]
if charge_key == CHARGE_VISION_INPUT:
return ["图片理解Token"]
return []
async def _find_existing_charge(db: AsyncSession, user_id: str, related_id: str, charge_key: str) -> CreditRecord | None:
base = (
async def _find_existing_by_biz_key(db: AsyncSession, *, user_id: str, biz_key: str) -> CreditRecord | None:
result = await db.execute(
select(CreditRecord)
.where(CreditRecord.user_id == user_id)
.where(CreditRecord.related_id == related_id)
.where(CreditRecord.type == "consume")
.where(CreditRecord.user_id == user_id, CreditRecord.biz_key == biz_key)
.limit(1)
)
return result.scalar_one_or_none()
result = await db.execute(base.where(CreditRecord.description.like(f"[{charge_key}]%")).limit(1))
existing = result.scalar_one_or_none()
if existing:
return existing
for keyword in _legacy_description_keywords(charge_key):
result = await db.execute(base.where(CreditRecord.description.like(f"%{keyword}%")).limit(1))
existing = result.scalar_one_or_none()
if existing:
return existing
return None
async def get_next_credit_attempt_no(
db: AsyncSession,
*,
owner_type: str,
owner_id: str,
charge_kind: str = CHARGE_MEDIA,
) -> int:
"""根据已有正式 biz_key 计算下一轮扣费 attempt_no。
不依赖任务表 retry_count,防止 worker 崩溃/重复消息造成状态字段不可信。
"""
prefix = f"{owner_type}:{owner_id}:attempt:%:{charge_kind}:charge"
result = await db.execute(
select(CreditRecord.biz_key)
.where(CreditRecord.related_id == owner_id)
.where(CreditRecord.type == "consume")
.where(CreditRecord.biz_key.like(prefix))
)
max_attempt = 0
for (biz_key,) in result.all():
parsed = parse_credit_biz_key(biz_key)
if parsed and parsed.get("owner_type") == owner_type and parsed.get("owner_id") == owner_id:
max_attempt = max(max_attempt, int(parsed.get("attempt_no") or 0))
return max_attempt + 1
async def deduct_credits_locked_once(
@@ -133,49 +166,38 @@ async def deduct_credits_locked_once(
description: str,
related_id: str,
charge_key: str,
biz_key: str | None = None,
attempt_no: int | None = None,
) -> BillingItem:
"""Deduct credits with row lock and idempotency.
"""按 biz_key 做幂等扣费。
- User row is locked by SELECT ... FOR UPDATE, so concurrent deductions for
the same user are serialized in PostgreSQL/MySQL.
- CreditRecord description prefix + related_id is used as an idempotency key
without changing existing table structures.
charge_key 只保留为业务分类;正式幂等以 biz_key 为准。
"""
amount = _round2(amount)
if amount <= 0:
return BillingItem(charge_key=charge_key, amount=0.0, charged=False, skipped_reason="amount_lte_zero")
return BillingItem(charge_key=charge_key, amount=0.0, charged=False, skipped_reason="amount_lte_zero", biz_key=biz_key, attempt_no=attempt_no)
result = await db.execute(select(User).where(User.id == user_id).with_for_update().limit(1))
user = result.scalar_one_or_none()
if not user:
raise ValueError("User not found")
if biz_key:
existing_charge = await _find_existing_by_biz_key(db, user_id=user_id, biz_key=biz_key)
if existing_charge:
return BillingItem(
charge_key=charge_key,
amount=abs(_round2(existing_charge.amount)),
charged=False,
skipped_reason="already_charged",
biz_key=biz_key,
attempt_no=attempt_no,
)
existing_charge = await _find_existing_charge(db, user_id, related_id, charge_key)
if existing_charge:
return BillingItem(
charge_key=charge_key,
amount=abs(_round2(existing_charge.amount)),
charged=False,
skipped_reason="already_charged",
)
if float(user.credits or 0) < amount:
raise InsufficientCreditsError()
user.credits = round(float(user.credits or 0) - amount, 2)
db.add(
CreditRecord(
id=generate_id(),
user_id=user_id,
type="consume",
amount=-amount,
balance_after=user.credits,
description=description,
related_id=related_id,
)
await deduct_credits(
db,
user_id=user_id,
amount=amount,
description=description,
related_id=related_id,
biz_key=biz_key,
)
await db.flush()
return BillingItem(charge_key=charge_key, amount=amount, charged=True)
return BillingItem(charge_key=charge_key, amount=amount, charged=True, biz_key=biz_key, attempt_no=attempt_no)
async def charge_chatapi_prompt_usage(
@@ -185,16 +207,14 @@ async def charge_chatapi_prompt_usage(
usage: Mapping[str, Any],
project_name: str | None = None,
) -> BillingSummary:
"""Charge ChatAPI prompt optimization and optional uploaded-file/vision tokens.
file_parse_credits / vision_input_credits are optional and disabled unless
SystemConfig contains these keys:
- file_parse_credits_per_1000_tokens
- vision_input_credits_per_1000_tokens
"""
"""提示词整理扣费仍按记录维度一次性幂等,不参与生成失败媒体退款。"""
project_name = project_name or "AI生成任务"
items: list[BillingItem] = []
attempt_no = 1
owner_type = OWNER_GENERATION_RECORD
owner_id = record.id
input_tokens = _safe_int(usage.get("input_tokens"))
output_tokens = _safe_int(usage.get("output_tokens"))
text_credits = await calc_text_credits(db, input_tokens, output_tokens)
@@ -206,15 +226,19 @@ async def charge_chatapi_prompt_usage(
description=f"ChatAPI提示词整理",
related_id=record.id,
charge_key=CHARGE_TEXT_PROMPT,
biz_key=build_credit_biz_key(
owner_type=owner_type,
owner_id=owner_id,
attempt_no=attempt_no,
charge_kind=CHARGE_TEXT_PROMPT,
action="charge",
),
attempt_no=attempt_no,
)
)
file_tokens = usage.get("file_parse_tokens") or usage.get("file_tokens") or usage.get("document_tokens") or 0
file_parse_credits = await _calc_optional_token_credits(
db,
_safe_int(file_tokens),
"file_parse_credits_per_1000_tokens",
)
file_parse_credits = await _calc_optional_token_credits(db, _safe_int(file_tokens), "file_parse_credits_per_1000_tokens")
items.append(
await deduct_credits_locked_once(
db,
@@ -223,15 +247,19 @@ async def charge_chatapi_prompt_usage(
description=f"文件解析Token",
related_id=record.id,
charge_key=CHARGE_FILE_PARSE,
biz_key=build_credit_biz_key(
owner_type=owner_type,
owner_id=owner_id,
attempt_no=attempt_no,
charge_kind=CHARGE_FILE_PARSE,
action="charge",
),
attempt_no=attempt_no,
)
)
vision_tokens = usage.get("vision_input_tokens") or usage.get("image_input_tokens") or usage.get("image_tokens") or 0
vision_input_credits = await _calc_optional_token_credits(
db,
_safe_int(vision_tokens),
"vision_input_credits_per_1000_tokens",
)
vision_input_credits = await _calc_optional_token_credits(db, _safe_int(vision_tokens), "vision_input_credits_per_1000_tokens")
items.append(
await deduct_credits_locked_once(
db,
@@ -240,12 +268,18 @@ async def charge_chatapi_prompt_usage(
description=f"图片理解Token",
related_id=record.id,
charge_key=CHARGE_VISION_INPUT,
biz_key=build_credit_biz_key(
owner_type=owner_type,
owner_id=owner_id,
attempt_no=attempt_no,
charge_kind=CHARGE_VISION_INPUT,
action="charge",
),
attempt_no=attempt_no,
)
)
if hasattr(record, "text_credits_cost"):
# Store expected text-side cost, even when this task is a retry and the
# actual CreditRecord was already written by an earlier attempt.
record.text_credits_cost = round(text_credits + file_parse_credits + vision_input_credits, 2)
if hasattr(record, "text_tokens_used"):
record.text_tokens_used = _safe_int(usage.get("total_tokens"), input_tokens + output_tokens)
@@ -265,10 +299,29 @@ async def charge_generation_media_by_params(
engine_id: str | None = None,
project_name: str | None = None,
description_prefix: str = "ChatAPI异步",
owner_type: str = OWNER_CHAT_GENERATION_TASK,
attempt_no: int | None = None,
) -> BillingSummary:
"""Charge image/video generation fee safely before creating provider task."""
"""图片/视频媒体生成扣费。
正式幂等由 owner_type + record_id + attempt_no + media + charge 组成。
每次用户主动重试必须传入新的 attempt_no。
"""
project_name = project_name or "AI生成任务"
gen_type = (gen_type or "").lower().strip()
attempt_no = attempt_no or await get_next_credit_attempt_no(
db,
owner_type=owner_type,
owner_id=record_id,
charge_kind=CHARGE_MEDIA,
)
biz_key = build_credit_biz_key(
owner_type=owner_type,
owner_id=record_id,
attempt_no=attempt_no,
charge_kind=CHARGE_MEDIA,
action="charge",
)
items: list[BillingItem] = []
if gen_type == "image":
@@ -281,7 +334,9 @@ async def charge_generation_media_by_params(
amount=amount,
description=f"{description_prefix}图片生成",
related_id=record_id,
charge_key=CHARGE_MEDIA_IMAGE,
charge_key=CHARGE_MEDIA,
biz_key=biz_key,
attempt_no=attempt_no,
)
)
elif gen_type == "video":
@@ -293,7 +348,9 @@ async def charge_generation_media_by_params(
amount=amount,
description=f"{description_prefix}视频生成",
related_id=record_id,
charge_key=CHARGE_MEDIA_VIDEO,
charge_key=CHARGE_MEDIA,
biz_key=biz_key,
attempt_no=attempt_no,
)
)
else:
@@ -308,6 +365,7 @@ async def charge_generation_media_for_record(
record: GenerationRecord,
project_name: str | None = None,
description_prefix: str = "ChatAPI异步",
attempt_no: int | None = None,
) -> BillingSummary:
return await charge_generation_media_by_params(
db,
@@ -319,4 +377,6 @@ async def charge_generation_media_for_record(
resolution=record.resolution,
project_name=project_name,
description_prefix=description_prefix,
owner_type=OWNER_GENERATION_RECORD,
attempt_no=attempt_no,
)
@@ -0,0 +1,205 @@
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
+32 -12
View File
@@ -16,6 +16,7 @@ from app.services.resource_accounting_service import (
)
from app.services.video_cover_service import create_video_cover_for_local_video
from app.config import settings
from app.services.generation_refund_service import mark_generation_record_failed_and_refund_once
logger = logging.getLogger("videogen")
@@ -76,6 +77,7 @@ class TaskQueue:
GenerationRecord.id == record_id,
GenerationRecord.deleted_at.is_(None),
)
.with_for_update()
.limit(1)
)
record = result.scalar_one_or_none()
@@ -92,8 +94,11 @@ class TaskQueue:
record_id = record.id
if not record.seedance_task_id:
record.status = "failed"
record.error_message = "缺少外部任务ID"
await mark_generation_record_failed_and_refund_once(
db,
record=record,
error_message="缺少外部任务ID",
)
await db.commit()
return
@@ -105,8 +110,11 @@ class TaskQueue:
count = self._active.get(record_id, 0) + 1
self._active[record_id] = count
if count >= MAX_POLLS:
record.status = "failed"
record.error_message = f"轮询超时: {e}"
await mark_generation_record_failed_and_refund_once(
db,
record=record,
error_message=f"轮询超时: {e}",
)
await db.commit()
del self._active[record_id]
else:
@@ -166,8 +174,11 @@ class TaskQueue:
logger.info(f"Video task completed: {record_id}")
elif status == "failed":
record.status = "failed"
record.error_message = poll_result.get("error", "视频生成失败")
await mark_generation_record_failed_and_refund_once(
db,
record=record,
error_message=poll_result.get("error", "视频生成失败"),
)
self._active.pop(record_id, None)
await db.commit()
logger.info(f"Video task failed: {record_id}")
@@ -176,8 +187,11 @@ class TaskQueue:
count = self._active.get(record_id, 0) + 1
self._active[record_id] = count
if count >= MAX_POLLS:
record.status = "failed"
record.error_message = "视频生成超时"
await mark_generation_record_failed_and_refund_once(
db,
record=record,
error_message="视频生成超时",
)
self._active.pop(record_id, None)
await db.commit()
logger.info(f"Video task timed out: {record_id}")
@@ -230,15 +244,21 @@ class TaskQueue:
await db.commit()
logger.info(f"Image task completed: {record_id}")
else:
record.status = "failed"
record.error_message = poll_result.get("error", "图片生成失败")
await mark_generation_record_failed_and_refund_once(
db,
record=record,
error_message=poll_result.get("error", "图片生成失败"),
)
await db.commit()
logger.info(f"Image task failed: {record_id}")
_log_image_response(record_id, poll_result)
except Exception as e:
record.status = "failed"
record.error_message = str(e)
await mark_generation_record_failed_and_refund_once(
db,
record=record,
error_message=str(e),
)
_log_image_response(record_id, {}, str(e))
await db.commit()
logger.error(f"Image task failed: {record_id}, error: {e}")
@@ -9,6 +9,7 @@ from app.models.base import async_session
from app.models.chat_generation_task import ChatGenerationTask
from app.services.error_codes import extract_error_message
from app.services.generation_log_service import log_task_event
from app.services.generation_refund_service import mark_chat_generation_task_failed_and_refund_once
from app.services.generation_provider_service import create_provider_task
from app.tasks.celery_app import celery_app
@@ -109,7 +110,7 @@ async def _run(task_id: str):
result = await db.execute(select(ChatGenerationTask).where(
ChatGenerationTask.id == task_id,
ChatGenerationTask.deleted_at.is_(None),
).limit(1))
).with_for_update().limit(1))
task = result.scalar_one_or_none()
if not task or task.generation_mode != "chatapi_async":
@@ -119,9 +120,12 @@ async def _run(task_id: str):
return
if task.deadline_at and datetime.now(timezone.utc) > task.deadline_at:
task.status = "failed"
task.pipeline_stage = "timeout"
task.error_message = "任务超时"
await mark_chat_generation_task_failed_and_refund_once(
db,
task=task,
error_message="任务超时",
pipeline_stage="timeout",
)
await db.commit()
await log_task_event(task, event_type="TASK_TIMEOUT", to_status="failed", to_stage="timeout")
return
@@ -236,12 +240,17 @@ async def _run(task_id: str):
result = await db.execute(select(ChatGenerationTask).where(
ChatGenerationTask.id == task_id,
ChatGenerationTask.deleted_at.is_(None),
).limit(1))
).with_for_update().limit(1))
task = result.scalar_one_or_none()
if task:
task.status = "failed"
task.error_message = extract_error_message(exc, "生成任务") if callable(extract_error_message) else str(exc)
error_message = extract_error_message(exc, "生成任务") if callable(extract_error_message) else str(exc)
await mark_chat_generation_task_failed_and_refund_once(
db,
task=task,
error_message=error_message,
pipeline_stage="failed",
)
await db.commit()
await log_task_event(task, event_type="TASK_FAILED", message=task.error_message)
@@ -8,6 +8,7 @@ from app.models.chat_generation_task import ChatGenerationTask
from app.services.error_codes import extract_error_message
from app.services.generation_download_service import download_generation_result
from app.services.generation_log_service import log_task_event
from app.services.generation_refund_service import mark_chat_generation_task_failed_and_refund_once
from app.services.resource_accounting_service import record_chat_task_generated_resource
from app.tasks.celery_app import celery_app
@@ -61,7 +62,7 @@ async def _reload_task(db, task_id: str) -> ChatGenerationTask | None:
select(ChatGenerationTask).where(
ChatGenerationTask.id == task_id,
ChatGenerationTask.deleted_at.is_(None),
).limit(1)
).with_for_update().limit(1)
)
return result.scalar_one_or_none()
@@ -72,7 +73,7 @@ async def _run(task_id: str):
select(ChatGenerationTask).where(
ChatGenerationTask.id == task_id,
ChatGenerationTask.deleted_at.is_(None),
).limit(1)
).with_for_update().limit(1)
)
task = result.scalar_one_or_none()
if not task or task.generation_mode != "chatapi_async":
@@ -165,9 +166,13 @@ async def _run(task_id: str):
task.retry_count = (task.retry_count or 0) + 1
if task.retry_count > 3:
task.status = "failed"
task.pipeline_stage = "download_failed"
task.error_message = extract_error_message(exc, "下载") if callable(extract_error_message) else str(exc)
error_message = extract_error_message(exc, "下载") if callable(extract_error_message) else str(exc)
await mark_chat_generation_task_failed_and_refund_once(
db,
task=task,
error_message=error_message,
pipeline_stage="download_failed",
)
await db.commit()
await log_task_event(
@@ -9,6 +9,7 @@ from app.models.base import async_session
from app.models.chat_generation_task import ChatGenerationTask
from app.services.error_codes import extract_error_message
from app.services.generation_log_service import log_task_event, log_provider_call
from app.services.generation_refund_service import mark_chat_generation_task_failed_and_refund_once
from app.services.generation_provider_service import poll_provider_task
from app.tasks.celery_app import celery_app
@@ -41,7 +42,7 @@ async def _reload_task(db, task_id: str) -> ChatGenerationTask | None:
select(ChatGenerationTask).where(
ChatGenerationTask.id == task_id,
ChatGenerationTask.deleted_at.is_(None),
).limit(1)
).with_for_update().limit(1)
)
return result.scalar_one_or_none()
@@ -51,7 +52,7 @@ async def _run(task_id: str):
result = await db.execute(select(ChatGenerationTask).where(
ChatGenerationTask.id == task_id,
ChatGenerationTask.deleted_at.is_(None),
).limit(1))
).with_for_update().limit(1))
task = result.scalar_one_or_none()
if not task or task.generation_mode != "chatapi_async":
return
@@ -61,17 +62,23 @@ async def _run(task_id: str):
return
if task.deadline_at and datetime.now(timezone.utc) > task.deadline_at:
task.status = "failed"
task.pipeline_stage = "timeout"
task.error_message = "任务轮询超时"
await mark_chat_generation_task_failed_and_refund_once(
db,
task=task,
error_message="任务轮询超时",
pipeline_stage="timeout",
)
await db.commit()
await log_task_event(task, event_type="TASK_TIMEOUT", to_status="failed", to_stage="timeout")
return
if not (task.seedance_task_id or task.provider_task_id):
task.status = "failed"
task.pipeline_stage = "failed"
task.error_message = "缺少外部任务ID"
await mark_chat_generation_task_failed_and_refund_once(
db,
task=task,
error_message="缺少外部任务ID",
pipeline_stage="failed",
)
await db.commit()
await log_task_event(task, event_type="POLL_FAILED", message=task.error_message)
return
@@ -117,9 +124,12 @@ async def _run(task_id: str):
task.provider_response_json = response_data
if not task.remote_result_url:
task.status = "failed"
task.pipeline_stage = "failed"
task.error_message = "供应商任务成功但未返回结果URL"
await mark_chat_generation_task_failed_and_refund_once(
db,
task=task,
error_message="供应商任务成功但未返回结果URL",
pipeline_stage="failed",
)
await db.commit()
await log_task_event(task, event_type="POLL_FAILED", message=task.error_message)
return
@@ -135,10 +145,13 @@ async def _run(task_id: str):
return
if _is_failed(status):
task.status = "failed"
task.pipeline_stage = "failed"
task.error_message = poll_result.get("error") or f"供应商任务失败: {status}"
task.provider_response_json = response_data
await mark_chat_generation_task_failed_and_refund_once(
db,
task=task,
error_message=poll_result.get("error") or f"供应商任务失败: {status}",
pipeline_stage="failed",
)
await db.commit()
await log_task_event(task, event_type="POLL_FAILED", message=task.error_message, detail=poll_result)
return
@@ -173,9 +186,13 @@ async def _run(task_id: str):
task.retry_count = (task.retry_count or 0) + 1
if task.retry_count > settings.CHATAPI_ASYNC_MAX_RETRIES:
task.status = "failed"
task.pipeline_stage = "failed"
task.error_message = extract_error_message(exc, "轮询") if callable(extract_error_message) else str(exc)
error_message = extract_error_message(exc, "轮询") if callable(extract_error_message) else str(exc)
await mark_chat_generation_task_failed_and_refund_once(
db,
task=task,
error_message=error_message,
pipeline_stage="failed",
)
await db.commit()
await log_task_event(task, event_type="POLL_FAILED", message=task.error_message)
else: