diff --git a/video-gen-api/app/api/v1/generation_ai.py b/video-gen-api/app/api/v1/generation_ai.py index f26b1742..c2461a94 100644 --- a/video-gen-api/app/api/v1/generation_ai.py +++ b/video-gen-api/app/api/v1/generation_ai.py @@ -52,6 +52,7 @@ from app.services.generation.billing_service import ( ) from app.services.generation.history_delete_service import batch_delete_generation_history_items from app.services.generation.log_service import log_task_event +from app.services.generation.media_reference_service import calculate_media_reference_usage from app.services.resource_capacity_service import assert_user_resource_capacity_available from app.services.operation_log_service import log_operation_event from app.tasks.celery_app import celery_app @@ -804,6 +805,14 @@ async def retry_task( quantity = int(target.generation_count or 1) if ( target.generation_mode == GenerationMode.CHATAPI_MAIN.value and target.gen_type == "image" ) else 1 + refs = target.media_references or "[]" + if isinstance(refs, str): + import json + try: + refs = json.loads(refs) + except Exception: + refs = [] + reference_usage = calculate_media_reference_usage(refs, include=True) media_billing = await charge_generation_media_by_params( db, user_id=target.user_id, @@ -813,6 +822,8 @@ async def retry_task( duration=target.duration, resolution=target.resolution, engine_id=target.engine_id, + input_video_duration=reference_usage.input_video_duration or None, + input_image_count=reference_usage.image_count or None, project_name="AI生成任务", description_prefix="Chat任务重试", owner_type=OWNER_CHAT_GENERATION_TASK, diff --git a/video-gen-api/app/services/generation/ai/task_create_service.py b/video-gen-api/app/services/generation/ai/task_create_service.py index aa36632d..58d16c0b 100644 --- a/video-gen-api/app/services/generation/ai/task_create_service.py +++ b/video-gen-api/app/services/generation/ai/task_create_service.py @@ -281,6 +281,7 @@ async def create_generation_task_group( gen_type=GenerationType.IMAGE.value, image_size=size, engine_id=engine.id, + input_image_count=reference_image_count if reference_image_count > 0 else None, project_name="AI生成任务", description_prefix="AI创作-", owner_type=OWNER_CHAT_GENERATION_TASK, @@ -340,6 +341,9 @@ async def create_generation_task_group( raise HTTPException(status_code=400, detail=f"视频时长不能超过 {engine.max_duration} 秒") input_video_duration = _validate_video_references(refs, max_audio_count=engine.max_audio_count) + reference_image_count = sum( + 1 for ref in refs if (ref.get("type") or "").lower() == GenerationType.IMAGE.value + ) provider_generation_resolution, upscale_enabled_snapshot, upscale_snapshot_json = await build_video_upscale_snapshot( db, target_resolution=resolution, @@ -363,6 +367,7 @@ async def create_generation_task_group( resolution=resolution, engine_id=engine.id, input_video_duration=input_video_duration if input_video_duration > 0 else None, + input_image_count=reference_image_count if reference_image_count > 0 else None, project_name="AI生成任务", description_prefix="AI创作-", owner_type=OWNER_CHAT_GENERATION_TASK, @@ -437,6 +442,7 @@ async def create_generation_task_group( resolution=resolution, engine_id=engine.id, input_video_duration=input_video_duration if input_video_duration > 0 else None, + input_image_count=reference_image_count if reference_image_count > 0 else None, project_name="AI生成任务", description_prefix=f"AI创作-第{generation_index}份-", owner_type=OWNER_CHAT_GENERATION_TASK,