From 1079a5a7d163b46ff0b757aa369f86ef7d929399 Mon Sep 17 00:00:00 2001 From: GinHa <15201596918@163.com> Date: Wed, 3 Jun 2026 16:15:45 +0800 Subject: [PATCH] =?UTF-8?q?=E7=94=9F=E6=88=90=E4=BB=BB=E5=8A=A1=E5=A4=B1?= =?UTF-8?q?=E8=B4=A5=E7=A7=AF=E5=88=86=E5=9B=9E=E9=80=80?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- .../516b84e3bf17_add_credit_record_biz_key.py | 105 +++++++ video-gen-api/app/api/v1/admin.py | 97 +++++-- video-gen-api/app/api/v1/generation.py | 140 ++++++--- video-gen-api/app/api/v1/generation_ai.py | 99 +++++-- .../app/models/chat_generation_task.py | 8 +- video-gen-api/app/models/credit_record.py | 17 +- video-gen-api/app/services/credits.py | 109 ++++++- .../app/services/generation_ai_service.py | 9 +- .../services/generation_billing_service.py | 268 +++++++++++------- .../app/services/generation_refund_service.py | 205 ++++++++++++++ video-gen-api/app/services/video_queue.py | 44 ++- .../app/tasks/generation_create_tasks.py | 23 +- .../app/tasks/generation_download_tasks.py | 15 +- .../app/tasks/generation_poll_tasks.py | 51 ++-- 14 files changed, 931 insertions(+), 259 deletions(-) create mode 100644 video-gen-api/alembic/versions/516b84e3bf17_add_credit_record_biz_key.py create mode 100644 video-gen-api/app/services/generation_refund_service.py diff --git a/video-gen-api/alembic/versions/516b84e3bf17_add_credit_record_biz_key.py b/video-gen-api/alembic/versions/516b84e3bf17_add_credit_record_biz_key.py new file mode 100644 index 00000000..24cf009d --- /dev/null +++ b/video-gen-api/alembic/versions/516b84e3bf17_add_credit_record_biz_key.py @@ -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") \ No newline at end of file diff --git a/video-gen-api/app/api/v1/admin.py b/video-gen-api/app/api/v1/admin.py index 26ca72ae..43bb0e5f 100644 --- a/video-gen-api/app/api/v1/admin.py +++ b/video-gen-api/app/api/v1/admin.py @@ -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} diff --git a/video-gen-api/app/api/v1/generation.py b/video-gen-api/app/api/v1/generation.py index b157217a..0a05c617 100644 --- a/video-gen-api/app/api/v1/generation.py +++ b/video-gen-api/app/api/v1/generation.py @@ -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) diff --git a/video-gen-api/app/api/v1/generation_ai.py b/video-gen-api/app/api/v1/generation_ai.py index 1aaebcca..309798eb 100644 --- a/video-gen-api/app/api/v1/generation_ai.py +++ b/video-gen-api/app/api/v1/generation_ai.py @@ -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="任务已重新投递", - ) \ No newline at end of file + message="任务已重新扣费并重新投递", + ) diff --git a/video-gen-api/app/models/chat_generation_task.py b/video-gen-api/app/models/chat_generation_task.py index da17fd3f..3283f589 100644 --- a/video-gen-api/app/models/chat_generation_task.py +++ b/video-gen-api/app/models/chat_generation_task.py @@ -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( diff --git a/video-gen-api/app/models/credit_record.py b/video-gen-api/app/models/credit_record.py index c3ba2343..908e3c84 100644 --- a/video-gen-api/app/models/credit_record.py +++ b/video-gen-api/app/models/credit_record.py @@ -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) diff --git a/video-gen-api/app/services/credits.py b/video-gen-api/app/services/credits.py index eab2e46f..ab4de901 100644 --- a/video-gen-api/app/services/credits.py +++ b/video-gen-api/app/services/credits.py @@ -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) diff --git a/video-gen-api/app/services/generation_ai_service.py b/video-gen-api/app/services/generation_ai_service.py index 3ceaed5a..a9c5d711 100644 --- a/video-gen-api/app/services/generation_ai_service.py +++ b/video-gen-api/app/services/generation_ai_service.py @@ -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( diff --git a/video-gen-api/app/services/generation_billing_service.py b/video-gen-api/app/services/generation_billing_service.py index 30c33938..ae542eeb 100644 --- a/video-gen-api/app/services/generation_billing_service.py +++ b/video-gen-api/app/services/generation_billing_service.py @@ -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[^:]+):(?P[^:]+):attempt:(?P\d+):(?P[^:]+):(?Pcharge|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, ) diff --git a/video-gen-api/app/services/generation_refund_service.py b/video-gen-api/app/services/generation_refund_service.py new file mode 100644 index 00000000..0bec6f5b --- /dev/null +++ b/video-gen-api/app/services/generation_refund_service.py @@ -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 diff --git a/video-gen-api/app/services/video_queue.py b/video-gen-api/app/services/video_queue.py index 2578564c..f94388db 100644 --- a/video-gen-api/app/services/video_queue.py +++ b/video-gen-api/app/services/video_queue.py @@ -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}") diff --git a/video-gen-api/app/tasks/generation_create_tasks.py b/video-gen-api/app/tasks/generation_create_tasks.py index dd5262b9..46e30c6d 100644 --- a/video-gen-api/app/tasks/generation_create_tasks.py +++ b/video-gen-api/app/tasks/generation_create_tasks.py @@ -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) diff --git a/video-gen-api/app/tasks/generation_download_tasks.py b/video-gen-api/app/tasks/generation_download_tasks.py index 993e22b2..9c974f25 100644 --- a/video-gen-api/app/tasks/generation_download_tasks.py +++ b/video-gen-api/app/tasks/generation_download_tasks.py @@ -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( diff --git a/video-gen-api/app/tasks/generation_poll_tasks.py b/video-gen-api/app/tasks/generation_poll_tasks.py index 5bf04f36..8999d625 100644 --- a/video-gen-api/app/tasks/generation_poll_tasks.py +++ b/video-gen-api/app/tasks/generation_poll_tasks.py @@ -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: