生成任务失败积分回退
This commit is contained in:
@@ -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")
|
||||
@@ -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}
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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(
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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
|
||||
@@ -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:
|
||||
|
||||
Reference in New Issue
Block a user