测试极速成片

This commit is contained in:
Azmat@qq.com
2026-08-26 15:18:18 +08:00
parent 245525ec53
commit 0ee498d807
30 changed files with 1963 additions and 367 deletions
+50 -2
View File
@@ -1,5 +1,6 @@
import uuid
from django.core.exceptions import ObjectDoesNotExist
from rest_framework import serializers
from apps.assets.serializers import AssetFileSerializer
@@ -341,6 +342,7 @@ class QuickCreateJobSerializer(serializers.ModelSerializer):
phase_index = serializers.SerializerMethodField()
settings = serializers.SerializerMethodField()
result = serializers.SerializerMethodField()
error_message = serializers.SerializerMethodField()
class Meta:
model = QuickCreateJob
@@ -365,6 +367,15 @@ class QuickCreateJobSerializer(serializers.ModelSerializer):
]
read_only_fields = fields
def get_error_message(self, obj) -> str:
from apps.projects.services.quick_create import _public_error
stored = obj.error_message or ""
if obj.status != QuickCreateJob.Status.FAILED:
return stored
hidden = str((obj.metadata or {}).get("internal_error") or "")
return _public_error(" ".join(part for part in (stored, hidden) if part)) or stored
def get_product_images(self, obj) -> list[dict]:
product = getattr(obj.project, "product", None)
if product is None:
@@ -403,7 +414,8 @@ class QuickCreateJobSerializer(serializers.ModelSerializer):
def get_result(self, obj) -> dict | None:
project = obj.project
settings = self.get_settings(obj)
video_url = _final_video_url(project)
final_video_url = _final_video_url(project)
video_url = final_video_url
segments = list(project.video_segments.all())
if not video_url:
for segment in sorted(segments, key=lambda item: item.sort_order):
@@ -446,6 +458,7 @@ class QuickCreateJobSerializer(serializers.ModelSerializer):
)
return {
"video_url": video_url,
"final_video_url": final_video_url,
"poster_url": poster_url,
"duration_seconds": duration or 15,
"aspect_ratio": settings["aspect_ratio"],
@@ -479,6 +492,23 @@ class ScriptVersionSerializer(serializers.ModelSerializer):
read_only_fields = fields
def _quick_create_job(obj: Project):
try:
return obj.quick_create_job
except ObjectDoesNotExist:
return None
def _quick_create_status(obj: Project) -> str:
job = _quick_create_job(obj)
return job.status if job else ""
def _quick_create_job_id(obj: Project) -> str:
job = _quick_create_job(obj)
return str(job.id) if job else ""
class ProjectListSerializer(serializers.ModelSerializer):
"""列表/仪表盘/侧栏用的轻量项目序列化:不嵌套 阶段/片段/故事板/时间线(那些只详情页要)。
脚本数/镜数走 annotate 计数(见 ProjectViewSet.get_queryset),避免逐项目拉全套关联(原列表 2-3s)。"""
@@ -490,13 +520,15 @@ class ProjectListSerializer(serializers.ModelSerializer):
# 合成成片地址:项目列表的播放按钮据此直接播成片(没合成过为空 → 退回进流水线)
final_video_url = serializers.SerializerMethodField()
quick_create = serializers.SerializerMethodField()
quick_create_status = serializers.SerializerMethodField()
quick_create_job_id = serializers.SerializerMethodField()
class Meta:
model = Project
fields = [
"id", "name", "product", "product_title", "cover_preview_url",
"status", "current_stage", "script_version_count", "video_segment_count",
"final_video_url", "quick_create",
"final_video_url", "quick_create", "quick_create_status", "quick_create_job_id",
"is_deleted", "purged_at", "created_at", "updated_at",
]
@@ -509,6 +541,12 @@ class ProjectListSerializer(serializers.ModelSerializer):
def get_quick_create(self, obj) -> bool:
return bool((obj.metadata or {}).get("quick_create"))
def get_quick_create_status(self, obj) -> str:
return _quick_create_status(obj)
def get_quick_create_job_id(self, obj) -> str:
return _quick_create_job_id(obj)
class ProjectSerializer(serializers.ModelSerializer):
stages = ProjectStageSerializer(many=True, read_only=True)
@@ -520,6 +558,8 @@ class ProjectSerializer(serializers.ModelSerializer):
timeline = TimelineSerializer(read_only=True)
# 合成成片地址(最新一次成功拼接):视频阶段的「播放成片 / 下载成片」直接用它
final_video_url = serializers.SerializerMethodField()
quick_create_status = serializers.SerializerMethodField()
quick_create_job_id = serializers.SerializerMethodField()
class Meta:
model = Project
@@ -542,6 +582,8 @@ class ProjectSerializer(serializers.ModelSerializer):
"video_segments",
"timeline",
"final_video_url",
"quick_create_status",
"quick_create_job_id",
"created_at",
"updated_at",
]
@@ -550,6 +592,12 @@ class ProjectSerializer(serializers.ModelSerializer):
def get_final_video_url(self, obj) -> str:
return _final_video_url(obj)
def get_quick_create_status(self, obj) -> str:
return _quick_create_status(obj)
def get_quick_create_job_id(self, obj) -> str:
return _quick_create_job_id(obj)
class ScriptTemplateSerializer(serializers.ModelSerializer):
"""套路模板 · 列表与详情共用。写入只开放 name(其余字段由存模板端点从脚本抽)。"""
File diff suppressed because it is too large Load Diff
+25 -24
View File
@@ -47,9 +47,15 @@ def run_export_job_task(self, export_job_id: str) -> str:
QUICK_CREATE_QUEUE = "airshelf.quick"
@app.task(bind=True, max_retries=0, soft_time_limit=240, time_limit=270, queue=QUICK_CREATE_QUEUE)
# 豆包长思考脚本允许完整跑 30 分钟;硬上限额外留 60 秒让 soft timeout 的
# 收尾/落库完成,避免 15 分钟时仍在正常输出却被 worker 强制中断。
@app.task(bind=True, max_retries=0, soft_time_limit=1800, time_limit=1860, queue=QUICK_CREATE_QUEUE)
def run_quick_script_task(self, quick_job_id: str) -> str:
"""脚本生成单独跑,避免把整条极速成片编排堵在一次 SSE 消费里。"""
"""脚本生成单独跑,避免把整条极速成片编排堵在一次 SSE 消费里。
软超时必须长于豆包思考流(允许最长 30 分钟)。短于 HTTP 流超时会 SIGUSR1 掐连接,
任务监视器就记成 stream aborted (client disconnected)。
"""
from celery.exceptions import SoftTimeLimitExceeded
from apps.projects.models import QuickCreateJob
@@ -57,41 +63,36 @@ def run_quick_script_task(self, quick_job_id: str) -> str:
try:
consume_quick_script(quick_job_id)
except SoftTimeLimitExceeded:
except SoftTimeLimitExceeded as exc:
job = QuickCreateJob.objects.select_related("project").filter(id=quick_job_id).first()
if job is not None and job.status not in {
QuickCreateJob.Status.SUCCEEDED,
QuickCreateJob.Status.FAILED,
QuickCreateJob.Status.CANCELLED,
}:
fail_quick_create(job, "脚本生成超时,请稍后重试或进入专业模式查看")
raise
fail_quick_create(
job,
"脚本生成时间较长,系统会自动重试",
internal_error=f"SoftTimeLimitExceeded: stream aborted (client disconnected); {exc}",
)
return quick_job_id
return quick_job_id
@app.task(bind=True, max_retries=0, queue=QUICK_CREATE_QUEUE)
@app.task(bind=True, max_retries=0)
def advance_quick_create_task(self, quick_job_id: str) -> str:
"""一次只推进一个可重入状态,等待型阶段通过重新入队轮询,不占 worker 睡眠。"""
from apps.projects.services.quick_create import advance_quick_create
"""一次只推进一个可重入状态,等待型阶段通过重新入队轮询,不占 worker 睡眠。
编排走默认 celery 队列:worker 即使没听 airshelf.quick,也不会停在「等待开始」。
"""
from apps.projects.models import QuickCreateJob
from apps.projects.services.quick_create import _claim_next_advance, _enqueue_advance, advance_quick_create
next_delay = advance_quick_create(quick_job_id)
if next_delay is not None:
from apps.projects.services.quick_create import _claim_next_advance
if not _claim_next_advance(quick_job_id, int(next_delay)):
return quick_job_id
try:
advance_quick_create_task.apply_async(
args=[quick_job_id],
countdown=max(1, int(next_delay)),
queue=QUICK_CREATE_QUEUE,
)
except Exception as exc: # noqa: BLE001 — 重排失败不能留下永久“生成中”
from apps.projects.models import QuickCreateJob
from apps.projects.services.quick_create import fail_quick_create
job = QuickCreateJob.objects.select_related("project").filter(id=quick_job_id).first()
if job is not None:
fail_quick_create(job, "生成队列暂时中断,请稍后重试", internal_error=str(exc))
raise
job = QuickCreateJob.objects.filter(id=quick_job_id).first()
if job is not None:
_enqueue_advance(job, countdown=max(1, int(next_delay)))
return quick_job_id
+425 -21
View File
@@ -12,12 +12,14 @@ from apps.accounts.models import Team, TeamMember, User
from apps.ai.models import AITask, ModelConfig, ModelProvider
from apps.assets.models import Asset, AssetFile
from apps.products.models import Product, ProductImage
from apps.projects.models import Project, ProjectStage, QuickCreateJob, ScriptSegment, ScriptVersion, VideoSegment, VideoSegmentVersion
from apps.projects.models import BaseAssetGroup, Project, ProjectStage, QuickCreateJob, ScriptSegment, ScriptVersion, VideoSegment, VideoSegmentVersion
from apps.projects.serializers import ProjectListSerializer, QuickCreateJobSerializer
from apps.projects.services.pipeline import initialize_project_pipeline
from apps.projects.services.quick_create import (
REVIEW_FAIL_MESSAGE,
_consume_script_agent,
_reviews_ready,
_safe_error,
_start_videos,
advance_quick_create,
cancel_quick_create,
@@ -94,7 +96,7 @@ class QuickCreateApiTests(TestCase):
self.assertEqual(str(job.id), response.data["id"])
enqueue.assert_called_once()
self.assertEqual(enqueue.call_args.kwargs["args"], [str(job.id)])
self.assertEqual(enqueue.call_args.kwargs["queue"], "airshelf.quick")
self.assertIn(enqueue.call_args.kwargs["queue"], {"celery", "airshelf.quick"})
require_worker_task.assert_called_once_with("apps.projects.tasks.advance_quick_create_task")
self.assertEqual(get_model.call_count, 2)
get_quick_model.assert_called_once()
@@ -290,7 +292,7 @@ class QuickCreateApiTests(TestCase):
self.assertEqual(job.status, QuickCreateJob.Status.RUNNING)
enqueue.assert_called_once()
def test_history_lists_team_jobs_including_failed(self):
def test_history_lists_completed_jobs_only(self):
product = Product.objects.create(team=self.team, created_by=self.user, title="历史商品")
project = Project.objects.create(team=self.team, created_by=self.user, product=product, name="历史商品 · 极速成片")
QuickCreateJob.objects.create(
@@ -334,12 +336,45 @@ class QuickCreateApiTests(TestCase):
response = self.client.get("/api/projects/quick-create-history/")
self.assertEqual(response.status_code, 200)
self.assertEqual(response.data["count"], 2)
self.assertEqual(response.data["count"], 1)
titles = [item["title"] for item in response.data["results"]]
self.assertIn("历史商品 · 极速成片", titles)
self.assertIn("失败商品 · 极速成片", titles)
self.assertNotIn("失败商品 · 极速成片", titles)
self.assertNotIn("进行中商品 · 极速成片", titles)
def test_history_hides_deleted_projects(self):
product = Product.objects.create(team=self.team, created_by=self.user, title="还在的商品")
alive = Project.objects.create(team=self.team, created_by=self.user, product=product, name="还在的商品 · 极速成片")
QuickCreateJob.objects.create(
team=self.team,
created_by=self.user,
project=alive,
status=QuickCreateJob.Status.SUCCEEDED,
phase=QuickCreateJob.Phase.COMPLETE,
)
deleted_product = Product.objects.create(team=self.team, created_by=self.user, title="已删商品")
deleted = Project.objects.create(
team=self.team,
created_by=self.user,
product=deleted_product,
name="已删商品 · 极速成片",
is_deleted=True,
)
QuickCreateJob.objects.create(
team=self.team,
created_by=self.user,
project=deleted,
status=QuickCreateJob.Status.CANCELLED,
phase=QuickCreateJob.Phase.SCRIPT,
)
response = self.client.get("/api/projects/quick-create-history/")
self.assertEqual(response.status_code, 200)
titles = [item["title"] for item in response.data["results"]]
self.assertEqual(response.data["count"], 1)
self.assertIn("还在的商品 · 极速成片", titles)
self.assertNotIn("已删商品 · 极速成片", titles)
def test_list_serializer_flags_quick_create_projects(self):
product = Product.objects.create(team=self.team, created_by=self.user, title="列表商品")
quick = Project.objects.create(
@@ -357,6 +392,44 @@ class QuickCreateApiTests(TestCase):
)
self.assertTrue(ProjectListSerializer(quick).data["quick_create"])
self.assertFalse(ProjectListSerializer(normal).data["quick_create"])
self.assertEqual(ProjectListSerializer(quick).data["quick_create_status"], "")
self.assertEqual(ProjectListSerializer(quick).data["quick_create_job_id"], "")
QuickCreateJob.objects.create(team=self.team, created_by=self.user, project=quick, status=QuickCreateJob.Status.RUNNING)
self.assertEqual(ProjectListSerializer(quick).data["quick_create_status"], "running")
self.assertTrue(ProjectListSerializer(quick).data["quick_create_job_id"])
def test_running_quick_create_blocks_professional_edits(self):
product = Product.objects.create(team=self.team, created_by=self.user, title="锁单商品")
project = Project.objects.create(
team=self.team,
created_by=self.user,
product=product,
name="锁单商品 · 极速成片",
metadata={"quick_create": True},
)
job = QuickCreateJob.objects.create(
team=self.team,
created_by=self.user,
project=project,
status=QuickCreateJob.Status.RUNNING,
)
blocked = self.client.patch(f"/api/projects/{project.id}/", {"name": "不该改"}, format="json")
self.assertEqual(blocked.status_code, 409)
self.assertIn("极速成片", str(blocked.data))
allowed = self.client.get(f"/api/projects/{project.id}/")
self.assertEqual(allowed.status_code, 200)
self.assertEqual(allowed.data["quick_create_status"], "running")
self.assertEqual(allowed.data["quick_create_job_id"], str(job.id))
job.status = QuickCreateJob.Status.FAILED
job.save(update_fields=["status", "updated_at"])
resumed = self.client.patch(f"/api/projects/{project.id}/", {"name": "专业模式可改"}, format="json")
self.assertEqual(resumed.status_code, 200)
project.refresh_from_db()
self.assertEqual(project.name, "专业模式可改")
class QuickCreateCoordinatorTests(TestCase):
@@ -410,7 +483,7 @@ class QuickCreateCoordinatorTests(TestCase):
self.assertTrue(self.job.metadata.get("script_started"))
start_script.assert_called_once()
self.assertEqual(start_script.call_args.kwargs["args"], [str(self.job.id)])
self.assertEqual(start_script.call_args.kwargs["queue"], "airshelf.quick")
self.assertIn(start_script.call_args.kwargs["queue"], {"celery", "airshelf.quick"})
def test_adopted_script_moves_to_assets_without_rerunning(self):
script = ScriptVersion.objects.create(project=self.project, title="脚本", content="口播", is_adopted=True)
@@ -431,8 +504,8 @@ class QuickCreateCoordinatorTests(TestCase):
self.assertIsNone(delay)
self.assertEqual(self.job.status, QuickCreateJob.Status.CANCELLED)
@patch("apps.projects.services.quick_create._run_quick_script_in_thread")
def test_recover_reruns_script_locally_when_queue_drops_it(self, run_local):
@patch("apps.projects.services.quick_create._enqueue_script")
def test_recover_requeues_script_when_queue_drops_it(self, enqueue_script):
self.job.status = QuickCreateJob.Status.RUNNING
self.job.phase = QuickCreateJob.Phase.SCRIPT
self.job.metadata = {
@@ -443,9 +516,8 @@ class QuickCreateCoordinatorTests(TestCase):
recover_quick_create(self.job)
self.job.refresh_from_db()
self.assertTrue(self.job.metadata.get("script_local"))
self.assertEqual(self.job.status, QuickCreateJob.Status.RUNNING)
run_local.assert_called_once_with(str(self.job.id))
enqueue_script.assert_called_once_with(str(self.job.id))
def test_recover_starts_stale_queued_job_without_waiting_for_quick_queue(self):
QuickCreateJob.objects.filter(id=self.job.id).update(
@@ -455,13 +527,13 @@ class QuickCreateCoordinatorTests(TestCase):
updated_at=timezone.now() - timedelta(seconds=30),
)
self.job.refresh_from_db()
with patch("apps.projects.services.quick_create._advance_without_quick_queue") as advance:
with patch("apps.projects.services.quick_create._enqueue_advance") as enqueue:
recover_quick_create(self.job)
advance.assert_called_once_with(str(self.job.id))
enqueue.assert_called_once()
@patch("apps.common.celery_health.worker_consumes_queue", return_value=False)
@patch("apps.projects.services.quick_create._run_quick_script_in_thread")
def test_advance_script_runs_locally_when_quick_queue_has_no_consumer(self, run_local, _listens):
@patch("apps.projects.tasks.run_quick_script_task.apply_async")
def test_advance_script_falls_back_to_celery_when_quick_queue_has_no_consumer(self, start_script, _listens):
self.job.status = QuickCreateJob.Status.RUNNING
self.job.phase = QuickCreateJob.Phase.SCRIPT
self.job.save(update_fields=["status", "phase", "updated_at"])
@@ -470,11 +542,11 @@ class QuickCreateCoordinatorTests(TestCase):
self.job.refresh_from_db()
self.assertEqual(delay, 5)
self.assertTrue(self.job.metadata.get("script_started"))
self.assertTrue(self.job.metadata.get("script_local"))
run_local.assert_called_once_with(str(self.job.id))
start_script.assert_called_once()
self.assertEqual(start_script.call_args.kwargs["queue"], "celery")
@patch("apps.projects.services.quick_create._run_quick_script_in_thread")
def test_recover_does_not_steal_script_when_model_call_is_inflight(self, run_local):
@patch("apps.projects.services.quick_create._enqueue_script")
def test_recover_does_not_steal_script_when_model_call_is_inflight(self, enqueue_script):
self._script_task(AITask.Status.SUBMITTED, key="script-inflight-1")
self.job.status = QuickCreateJob.Status.RUNNING
self.job.phase = QuickCreateJob.Phase.SCRIPT
@@ -486,9 +558,8 @@ class QuickCreateCoordinatorTests(TestCase):
recover_quick_create(self.job)
self.job.refresh_from_db()
self.assertFalse(self.job.metadata.get("script_local"))
self.assertEqual(self.job.status, QuickCreateJob.Status.RUNNING)
run_local.assert_not_called()
enqueue_script.assert_not_called()
def test_advance_script_ignores_old_failed_task_while_latest_is_running(self):
self._script_task(AITask.Status.FAILED, key="script-old-failed", error_message="stream aborted (client disconnected)")
@@ -506,6 +577,25 @@ class QuickCreateCoordinatorTests(TestCase):
self.assertEqual(delay, 5)
self.assertEqual(self.job.status, QuickCreateJob.Status.RUNNING)
@patch("apps.projects.services.quick_create._enqueue_script")
def test_advance_script_retries_retryable_failure_without_failing_job(self, enqueue_script):
self._script_task(AITask.Status.FAILED, key="script-aborted", error_message="stream aborted (client disconnected)")
self.job.status = QuickCreateJob.Status.RUNNING
self.job.phase = QuickCreateJob.Phase.SCRIPT
self.job.metadata = {
"script_started": True,
"script_started_at": timezone.now().isoformat(),
}
self.job.save(update_fields=["status", "phase", "metadata", "updated_at"])
delay = advance_quick_create(str(self.job.id))
self.job.refresh_from_db()
self.assertEqual(delay, 5)
self.assertEqual(self.job.status, QuickCreateJob.Status.RUNNING)
self.assertTrue(self.job.metadata.get("script_started"))
self.assertEqual(self.job.metadata.get("transient_retries"), 1)
enqueue_script.assert_called_once_with(str(self.job.id))
@patch("apps.projects.tasks.advance_quick_create_task.apply_async")
def test_recover_restarts_script_after_stream_abort(self, enqueue):
self.job.status = QuickCreateJob.Status.FAILED
@@ -525,6 +615,42 @@ class QuickCreateCoordinatorTests(TestCase):
self.assertEqual(self.job.metadata.get("transient_retries"), 1)
enqueue.assert_called_once()
@patch("apps.projects.tasks.advance_quick_create_task.apply_async")
def test_recover_retries_soft_time_limit_as_transient(self, enqueue):
self.job.status = QuickCreateJob.Status.FAILED
self.job.phase = QuickCreateJob.Phase.SCRIPT
self.job.error_message = "脚本生成时间较长,系统会自动重试"
self.job.metadata = {
"script_started": True,
"script_started_at": timezone.now().isoformat(),
"internal_error": "SoftTimeLimitExceeded: stream aborted (client disconnected)",
}
self.job.save(update_fields=["status", "phase", "error_message", "metadata", "updated_at"])
recover_quick_create(self.job)
self.job.refresh_from_db()
self.assertEqual(self.job.status, QuickCreateJob.Status.RUNNING)
self.assertFalse(self.job.metadata.get("script_started"))
enqueue.assert_called_once()
@patch("apps.projects.tasks.advance_quick_create_task.apply_async")
def test_recover_retries_failed_script_task_without_internal_error(self, enqueue):
self._script_task(
AITask.Status.FAILED,
key="script-soft-limit",
error_message="SoftTimeLimitExceeded()",
)
self.job.status = QuickCreateJob.Status.FAILED
self.job.phase = QuickCreateJob.Phase.SCRIPT
self.job.error_message = "脚本生成超时,请稍后重试或进入专业模式查看"
self.job.metadata = {"script_started": True}
self.job.save(update_fields=["status", "phase", "error_message", "metadata", "updated_at"])
recover_quick_create(self.job)
self.job.refresh_from_db()
self.assertEqual(self.job.status, QuickCreateJob.Status.RUNNING)
enqueue.assert_called_once()
@patch("apps.projects.services.quick_create.get_quick_script_model")
@patch("apps.projects.services.quick_create.stream_script_agent")
def test_consume_script_does_not_hit_db_on_every_sse_frame(self, stream_fn, get_model):
@@ -581,12 +707,29 @@ class QuickCreateCoordinatorTests(TestCase):
self.assertEqual(self.project.status, Project.Status.VIDEOING)
self.assertEqual(self.project.failure_reason, "")
@patch("apps.projects.tasks.advance_quick_create_task.apply_async")
def test_recover_resumes_production_after_orchestrator_timeout(self, enqueue):
self.project.status = Project.Status.VIDEOING
self.project.current_stage = ProjectStage.Stage.VIDEO
self.project.save(update_fields=["status", "current_stage", "updated_at"])
self.job.status = QuickCreateJob.Status.FAILED
self.job.phase = QuickCreateJob.Phase.PRODUCTION
self.job.error_message = "脚本已保留,后续步骤遇到网络波动。点重试会从上次进度继续"
self.job.metadata = {"storyboard_started": True, "internal_error": "Timeout reading from socket"}
self.job.save(update_fields=["status", "phase", "error_message", "metadata", "updated_at"])
recover_quick_create(self.job)
self.job.refresh_from_db()
self.assertEqual(self.job.status, QuickCreateJob.Status.RUNNING)
self.assertEqual(self.job.metadata.get("transient_retries"), 1)
enqueue.assert_called_once()
@patch("apps.projects.tasks.advance_quick_create_task.apply_async")
def test_resume_failed_production_job_keeps_progress(self, enqueue):
self.job.status = QuickCreateJob.Status.FAILED
self.job.phase = QuickCreateJob.Phase.PRODUCTION
self.job.error_message = "极速成片暂未完成,请稍后重试或进入专业模式查看"
self.job.metadata = {"storyboard_started": True, "transient_retries": 8}
self.job.metadata = {"storyboard_started": True, "transient_retries": 8, "video_fail_retries": 2}
self.job.save(update_fields=["status", "phase", "error_message", "metadata", "updated_at"])
resume_quick_create(self.job)
@@ -594,6 +737,7 @@ class QuickCreateCoordinatorTests(TestCase):
self.assertEqual(self.job.status, QuickCreateJob.Status.RUNNING)
self.assertEqual(self.job.error_message, "")
self.assertIsNone(self.job.metadata.get("transient_retries"))
self.assertIsNone(self.job.metadata.get("video_fail_retries"))
enqueue.assert_called_once()
@patch("apps.projects.services.quick_create.assets_client.is_enabled", return_value=True)
@@ -610,6 +754,40 @@ class QuickCreateCoordinatorTests(TestCase):
self.job.refresh_from_db()
self.assertTrue(self.job.metadata.get("reviews_skipped"))
def test_safe_error_maps_image_moderation_to_review_copy(self):
self.assertEqual(
_safe_error(ValueError("400 moderation_blocked safety_violations=[sexual]")),
REVIEW_FAIL_MESSAGE,
)
self.assertEqual(
_safe_error(RuntimeError("InputImageSensitiveContentDetected")),
REVIEW_FAIL_MESSAGE,
)
@patch("apps.projects.services.quick_create.assets_client.is_enabled", return_value=True)
@patch("apps.projects.services.quick_create.poll_team_reviews", return_value={})
@patch(
"apps.projects.services.quick_create.collect_video_review_blockers",
return_value=[{"review_status": "failed", "name": "女主立绘"}],
)
def test_failed_asset_review_stops_job_with_clear_message(self, _blockers, _poll, _enabled):
self.job.status = QuickCreateJob.Status.RUNNING
self.job.phase = QuickCreateJob.Phase.PRODUCTION
self.job.save(update_fields=["status", "phase", "updated_at"])
self.assertIsNone(_reviews_ready(self.job))
self.job.refresh_from_db()
self.assertEqual(self.job.status, QuickCreateJob.Status.FAILED)
self.assertEqual(self.job.error_message, REVIEW_FAIL_MESSAGE)
def test_failed_job_serializer_rewrites_hidden_moderation_error(self):
self.job.status = QuickCreateJob.Status.FAILED
self.job.phase = QuickCreateJob.Phase.PRODUCTION
self.job.error_message = "极速成片暂未完成,请稍后重试或进入专业模式查看"
self.job.metadata = {"internal_error": "400 moderation_blocked safety_violations=[sexual]"}
self.job.save(update_fields=["status", "phase", "error_message", "metadata", "updated_at"])
data = QuickCreateJobSerializer(self.job).data
self.assertEqual(data["error_message"], REVIEW_FAIL_MESSAGE)
def test_recover_marks_success_when_video_already_finished(self):
self.project.video_segments.exclude(sort_order=0).delete()
segment = self.project.video_segments.get(sort_order=0)
@@ -635,6 +813,24 @@ class QuickCreateCoordinatorTests(TestCase):
self.assertEqual(self.job.status, QuickCreateJob.Status.SUCCEEDED)
self.assertEqual(self.job.phase, QuickCreateJob.Phase.COMPLETE)
def test_recover_marks_success_when_professional_mode_already_completed(self):
"""专业模式完成后,极速任务的旧失败状态必须自动被回收。"""
self.job.status = QuickCreateJob.Status.FAILED
self.job.phase = QuickCreateJob.Phase.PRODUCTION
self.job.error_message = "极速成片暂未完成,请稍后重试或进入专业模式查看"
self.job.save(update_fields=["status", "phase", "error_message", "updated_at"])
self.project.status = Project.Status.COMPLETED
self.project.current_stage = ProjectStage.Stage.VIDEO
self.project.failure_reason = ""
self.project.save(update_fields=["status", "current_stage", "failure_reason", "updated_at"])
recover_quick_create(self.job)
self.job.refresh_from_db()
self.assertEqual(self.job.status, QuickCreateJob.Status.SUCCEEDED)
self.assertEqual(self.job.phase, QuickCreateJob.Phase.COMPLETE)
self.assertEqual(self.job.progress, 100)
@patch("apps.projects.services.quick_create.generate_base_asset", side_effect=TimeoutError("Timeout reading from socket"))
def test_asset_start_timeout_retries_without_locking_the_job(self, _generate):
script = ScriptVersion.objects.create(project=self.project, title="脚本", content="口播", is_adopted=True)
@@ -680,6 +876,213 @@ class QuickCreateCoordinatorTests(TestCase):
self.assertEqual(self.job.status, QuickCreateJob.Status.FAILED)
self.assertNotEqual(self.project.status, Project.Status.FAILED)
def _ready_asset_tasks(self):
script = ScriptVersion.objects.create(project=self.project, title="脚本", content="口播", is_adopted=True)
ScriptSegment.objects.create(script_version=script, sort_order=0, narration="开场")
ids = []
for kind, key in (
(AITask.Type.PRODUCT_IMAGE, "asset-product"),
(AITask.Type.PERSON_IMAGE, "asset-person"),
(AITask.Type.SCENE_IMAGE, "asset-scene"),
):
task = self._script_task(AITask.Status.SUCCEEDED, key=key)
task.task_type = kind
task.save(update_fields=["task_type"])
ids.append(str(task.id))
return ids
def _portrait_group(self, *, task=None, name="推荐模特"):
portrait = Asset.objects.create(
team=self.team,
created_by=self.user,
name=name,
asset_type=Asset.Type.IMAGE,
source=Asset.Source.AI_GENERATED,
category=Asset.Category.PERSON,
)
return BaseAssetGroup.objects.create(
project=self.project,
kind=BaseAssetGroup.Kind.PERSON,
task=task,
adopted_asset=portrait,
metadata={"label": name},
), portrait
@patch("apps.projects.services.quick_create.generate_person_triview")
def test_assets_skip_inflight_person_triview(self, start_triview):
base_ids = self._ready_asset_tasks()
_group, portrait = self._portrait_group()
inflight = self._script_task(AITask.Status.SUBMITTED, key="auto-triview")
inflight.request_payload = {"triview_of": str(portrait.id)}
inflight.save(update_fields=["request_payload"])
self.job.status = QuickCreateJob.Status.RUNNING
self.job.phase = QuickCreateJob.Phase.ASSETS
self.job.metadata = {"base_asset_task_ids": base_ids, "assets_started": True}
self.job.save(update_fields=["status", "phase", "metadata", "updated_at"])
delay = advance_quick_create(str(self.job.id))
self.job.refresh_from_db()
start_triview.assert_not_called()
self.assertEqual(delay, 1)
self.assertIsNone(self.job.metadata.get("triview_task_ids"))
self.assertTrue(self.job.metadata.get("triview_skipped"))
self.assertEqual(self.job.status, QuickCreateJob.Status.RUNNING)
self.assertEqual(self.job.phase, QuickCreateJob.Phase.PRODUCTION)
@patch("apps.projects.services.quick_create.generate_person_triview")
def test_assets_never_starts_triview(self, start_triview):
new_id = uuid.uuid4()
start_triview.return_value = SimpleNamespace(id=new_id)
base_ids = self._ready_asset_tasks()
group, portrait = self._portrait_group()
self.job.status = QuickCreateJob.Status.RUNNING
self.job.phase = QuickCreateJob.Phase.ASSETS
self.job.metadata = {"base_asset_task_ids": base_ids, "assets_started": True}
self.job.save(update_fields=["status", "phase", "metadata", "updated_at"])
delay = advance_quick_create(str(self.job.id))
self.job.refresh_from_db()
start_triview.assert_not_called()
self.assertEqual(delay, 1)
self.assertIsNone(self.job.metadata.get("triview_task_ids"))
self.assertTrue(self.job.metadata.get("triview_skipped"))
self.assertEqual(self.job.phase, QuickCreateJob.Phase.PRODUCTION)
self.assertIsNone(group.task_id)
def test_failed_triview_does_not_block_storyboard(self):
base_ids = self._ready_asset_tasks()
_group, portrait = self._portrait_group()
failed = self._script_task(AITask.Status.FAILED, key="triview-failed", error_message="image_edit timeout")
failed.request_payload = {"triview_of": str(portrait.id)}
failed.save(update_fields=["request_payload"])
self.job.status = QuickCreateJob.Status.RUNNING
self.job.phase = QuickCreateJob.Phase.ASSETS
self.job.metadata = {
"base_asset_task_ids": base_ids,
"assets_started": True,
"triview_task_ids": [str(failed.id)],
"triview_fail_retries": 2,
}
self.job.save(update_fields=["status", "phase", "metadata", "updated_at"])
delay = advance_quick_create(str(self.job.id))
self.job.refresh_from_db()
self.assertEqual(delay, 1)
self.assertEqual(self.job.phase, QuickCreateJob.Phase.PRODUCTION)
self.assertTrue(self.job.metadata.get("triview_skipped"))
self.assertEqual(self.job.status, QuickCreateJob.Status.RUNNING)
@patch("apps.projects.services.quick_create.generate_base_asset")
def test_assets_start_all_characters_and_scenes(self, generate):
generate.side_effect = lambda **kwargs: SimpleNamespace(id=uuid.uuid4())
ScriptVersion.objects.create(project=self.project, title="脚本", content="口播", is_adopted=True)
self.project.metadata = {
"script_entities": [
{"id": "c1", "type": "character", "name": "女主", "visual_prompt": "都市女性"},
{"id": "c2", "type": "character", "name": "闺蜜", "visual_prompt": "活泼女生"},
{"id": "s1", "type": "scene", "name": "客厅", "visual_prompt": "暖光客厅"},
{"id": "s2", "type": "scene", "name": "咖啡馆", "visual_prompt": "街边咖啡馆"},
]
}
self.project.save(update_fields=["metadata", "updated_at"])
self.job.status = QuickCreateJob.Status.RUNNING
self.job.phase = QuickCreateJob.Phase.ASSETS
self.job.save(update_fields=["status", "phase", "updated_at"])
delay = advance_quick_create(str(self.job.id))
self.job.refresh_from_db()
self.assertEqual(delay, 10)
self.assertEqual(generate.call_count, 5)
labels = [(call.kwargs["kind"], call.kwargs["label"]) for call in generate.call_args_list]
self.assertEqual(
labels,
[
(BaseAssetGroup.Kind.PRODUCT, "测试精华"),
(BaseAssetGroup.Kind.PERSON, "女主"),
(BaseAssetGroup.Kind.PERSON, "闺蜜"),
(BaseAssetGroup.Kind.SCENE, "客厅"),
(BaseAssetGroup.Kind.SCENE, "咖啡馆"),
],
)
for call in generate.call_args_list:
if call.kwargs["kind"] == BaseAssetGroup.Kind.PERSON:
self.assertFalse(call.kwargs["auto_triview"])
@patch("apps.projects.services.quick_create.generate_person_triview")
def test_assets_skip_triview_for_every_character(self, start_triview):
start_triview.side_effect = lambda **kwargs: SimpleNamespace(id=uuid.uuid4())
base_ids = self._ready_asset_tasks()
_g1, portrait_a = self._portrait_group(name="女主")
_g2, portrait_b = self._portrait_group(name="闺蜜")
self.job.status = QuickCreateJob.Status.RUNNING
self.job.phase = QuickCreateJob.Phase.ASSETS
self.job.metadata = {"base_asset_task_ids": base_ids, "assets_started": True}
self.job.save(update_fields=["status", "phase", "metadata", "updated_at"])
delay = advance_quick_create(str(self.job.id))
self.job.refresh_from_db()
self.assertEqual(delay, 1)
start_triview.assert_not_called()
self.assertIsNone(self.job.metadata.get("triview_task_ids"))
self.assertTrue(self.job.metadata.get("triview_skipped"))
self.assertEqual(self.job.phase, QuickCreateJob.Phase.PRODUCTION)
self.assertEqual(self.job.status, QuickCreateJob.Status.RUNNING)
@patch("apps.projects.services.quick_create.generate_person_triview", side_effect=ValueError("no active image model configured"))
def test_triview_submit_error_does_not_fail_job(self, _start_triview):
base_ids = self._ready_asset_tasks()
self._portrait_group()
self.job.status = QuickCreateJob.Status.RUNNING
self.job.phase = QuickCreateJob.Phase.ASSETS
self.job.metadata = {"base_asset_task_ids": base_ids, "assets_started": True}
self.job.save(update_fields=["status", "phase", "metadata", "updated_at"])
delay = advance_quick_create(str(self.job.id))
self.job.refresh_from_db()
self.assertEqual(delay, 1)
self.assertEqual(self.job.status, QuickCreateJob.Status.RUNNING)
self.assertEqual(self.job.phase, QuickCreateJob.Phase.PRODUCTION)
self.assertTrue(self.job.metadata.get("triview_skipped"))
def test_failed_triview_task_in_base_ids_does_not_fail_assets(self):
base_ids = self._ready_asset_tasks()
_group, portrait = self._portrait_group()
failed = self._script_task(AITask.Status.FAILED, key="triview-in-base", error_message="image_edit timeout")
failed.request_payload = {"kind": "person", "label": "推荐模特", "triview_of": str(portrait.id)}
failed.save(update_fields=["request_payload"])
self.job.status = QuickCreateJob.Status.RUNNING
self.job.phase = QuickCreateJob.Phase.ASSETS
self.job.metadata = {"base_asset_task_ids": [*base_ids, str(failed.id)], "assets_started": True}
self.job.save(update_fields=["status", "phase", "metadata", "updated_at"])
delay = advance_quick_create(str(self.job.id))
self.job.refresh_from_db()
self.assertEqual(delay, 1)
self.assertEqual(self.job.status, QuickCreateJob.Status.RUNNING)
self.assertEqual(self.job.phase, QuickCreateJob.Phase.PRODUCTION)
@patch("apps.projects.tasks.advance_quick_create_task.apply_async")
def test_resume_failed_assets_continues_into_storyboard(self, enqueue):
base_ids = self._ready_asset_tasks()
self.job.status = QuickCreateJob.Status.FAILED
self.job.phase = QuickCreateJob.Phase.ASSETS
self.job.error_message = "极速成片暂未完成,请稍后重试或进入专业模式查看"
self.job.metadata = {
"base_asset_task_ids": base_ids,
"assets_started": True,
"triview_task_ids": [],
"triview_ready": True,
"triview_skipped": True,
}
self.job.save(update_fields=["status", "phase", "error_message", "metadata", "updated_at"])
resume_quick_create(self.job)
self.job.refresh_from_db()
self.assertEqual(self.job.status, QuickCreateJob.Status.RUNNING)
self.assertEqual(self.job.phase, QuickCreateJob.Phase.PRODUCTION)
self.assertEqual(self.job.error_message, "")
enqueue.assert_called_once()
@patch("apps.projects.services.quick_create._reviews_ready", return_value=True)
def test_production_counts_ready_videos_without_adding_version_ids(self, _reviews):
self.project.video_segments.exclude(sort_order=0).delete()
@@ -739,6 +1142,7 @@ class QuickCreateCoordinatorTests(TestCase):
data = QuickCreateJobSerializer(self.job).data
self.assertEqual(data["phase_index"], 3)
self.assertEqual(data["result"]["video_url"], "https://cdn.example/quick.mp4")
self.assertEqual(data["result"]["final_video_url"], "")
self.assertEqual(data["result"]["duration_seconds"], 15)
def test_phase_index_matches_four_step_ui(self):
+13 -4
View File
@@ -1,4 +1,4 @@
from django.test import TestCase
from django.test import TestCase, override_settings
from unittest.mock import patch
from rest_framework.test import APIClient
@@ -351,9 +351,11 @@ class ProjectApiTests(TestCase):
group = BaseAssetGroup.objects.get(project=project, kind=BaseAssetGroup.Kind.PERSON)
self.assertEqual(group.metadata.get("label"), "女主")
@override_settings(CACHES={"default": {"BACKEND": "django.core.cache.backends.locmem.LocMemCache"}})
@patch("apps.ai.tasks.generate_base_asset_task.delay")
@patch("apps.ai.services._store_generated_media")
@patch("apps.ai.services.get_image_provider")
def test_generate_base_asset_ignores_auto_triview_request(self, get_provider, store_media):
def test_generate_person_base_asset_always_enables_auto_triview(self, get_provider, store_media, _enqueue_base_asset):
ModelConfig.objects.create(
provider=self.provider, name="img-model-auto-tri", display_name="Img Auto Tri",
capability=ModelConfig.Capability.IMAGE, endpoint="images/generations", unit_price="1.0000",
@@ -370,13 +372,13 @@ class ProjectApiTests(TestCase):
response = self.client.post(
f"/api/projects/{project.id}/generate-base-asset/",
{"kind": "person", "prompt": "portrait", "label": "hero", "auto_triview": True},
{"kind": "person", "prompt": "portrait", "label": "hero"},
format="json",
)
self.assertEqual(response.status_code, 202)
task = AITask.objects.get(id=response.data["task"]["id"])
self.assertFalse(task.request_payload.get("auto_triview"))
self.assertTrue(task.request_payload.get("auto_triview"))
@patch("apps.ai.services._store_generated_media")
@patch("apps.ai.services.get_image_provider")
@@ -1348,6 +1350,13 @@ class VideoSegmentTrueUpTests(TestCase):
self.assertEqual(task.credit_reservation.amount, video_reserve_amount(quote.points))
self.assertEqual(task.request_payload["estimated_tokens"], tokens)
def test_submit_reads_wizard_output_spec(self):
self.project.metadata = {"wizard": {"aspect_ratio": "16:9", "resolution": "480p"}}
self.project.save(update_fields=["metadata"])
task = self._submit()
self.assertEqual(task.request_payload["ratio"], "16:9")
self.assertEqual(task.request_payload["resolution"], "480p")
@patch("apps.ai.services._store_generated_media")
def test_poll_settles_by_actual_usage_tokens(self, store):
from decimal import Decimal
+41 -23
View File
@@ -8,7 +8,7 @@ from django.http import HttpResponse, JsonResponse, StreamingHttpResponse
from django.utils import timezone
from rest_framework import status
from rest_framework.decorators import action
from rest_framework.exceptions import ValidationError
from rest_framework.exceptions import APIException, ValidationError
from rest_framework.parsers import FormParser, MultiPartParser
from rest_framework.renderers import BaseRenderer
from rest_framework.response import Response
@@ -95,6 +95,12 @@ from .tasks import poll_video_segment_task
logger = logging.getLogger(__name__)
class QuickCreateInProgress(APIException):
status_code = status.HTTP_409_CONFLICT
default_detail = "该项目正在极速成片中,请到极速成片页查看进度"
default_code = "quick_create_running"
class ServerSentEventRenderer(BaseRenderer):
"""让 DRF 内容协商接受 Accept: text/event-stream(否则流式端点直接 406)。
实际响应由视图返回 StreamingHttpResponse 直接下发,这个 renderer 只用于通过协商。"""
@@ -235,7 +241,7 @@ def settle_video_completion(project: Project) -> bool:
class ProjectViewSet(TeamScopedViewSetMixin, ModelViewSet):
queryset = Project.objects.select_related("product", "timeline").prefetch_related(
queryset = Project.objects.select_related("product", "timeline", "quick_create_job").prefetch_related(
"stages",
"video_segments",
"video_segments__adopted_version__asset__files",
@@ -339,7 +345,7 @@ class ProjectViewSet(TeamScopedViewSetMixin, ModelViewSet):
# ——原列表把每个项目的 阶段/片段/故事板/时间线/资产文件全拉出,20 个项目实测 ~2s。
if self.action == "list":
qs = (
Project.objects.select_related("product", "product__cover_asset", "timeline")
Project.objects.select_related("product", "product__cover_asset", "timeline", "quick_create_job")
.prefetch_related(
"product__cover_asset__files",
# 成片地址(final_video_url)只需要「成功的导出任务」,预取到位后列表不再逐项目查库
@@ -363,6 +369,21 @@ class ProjectViewSet(TeamScopedViewSetMixin, ModelViewSet):
).order_by("-updated_at")
return super().get_queryset().filter(is_deleted=False, purged_at__isnull=True)
def initial(self, request, *args, **kwargs):
super().initial(request, *args, **kwargs)
if request.method in ("GET", "HEAD", "OPTIONS"):
return
if self.action in {"create", "destroy"}:
return
pk = kwargs.get("pk")
if not pk:
return
if QuickCreateJob.objects.filter(
project_id=pk,
status__in=[QuickCreateJob.Status.QUEUED, QuickCreateJob.Status.RUNNING],
).exists():
raise QuickCreateInProgress()
def perform_destroy(self, instance):
instance.is_deleted = True
instance.save(update_fields=["is_deleted", "updated_at"])
@@ -730,30 +751,19 @@ class ProjectViewSet(TeamScopedViewSetMixin, ModelViewSet):
job = self._quick_job_queryset().get(id=job.id)
return Response(QuickCreateJobSerializer(job).data)
@action(detail=False, methods=["post"], url_path=r"quick-create-retry/(?P<job_id>[^/.]+)")
def quick_create_retry(self, request, job_id=None):
job = self._quick_job_queryset().filter(id=job_id).first()
if job is None:
return Response({"detail": "极速成片任务不存在"}, status=status.HTTP_404_NOT_FOUND)
if job.status == QuickCreateJob.Status.SUCCEEDED:
return Response({"detail": "任务已经完成"}, status=status.HTTP_400_BAD_REQUEST)
if job.status == QuickCreateJob.Status.CANCELLED:
return Response({"detail": "已取消的任务请重新开始"}, status=status.HTTP_400_BAD_REQUEST)
from .services.quick_create import resume_quick_create
resume_quick_create(job)
job = self._quick_job_queryset().get(id=job.id)
return Response(QuickCreateJobSerializer(job).data)
@action(detail=False, methods=["get"], url_path="quick-create-history")
def quick_create_history(self, request):
from .services.quick_create import restore_false_failed_quick_creates
restore_false_failed_quick_creates(self.get_team())
# 进行中的任务看上方状态卡;列表要能找回失败后去专业模式继续的项目。
# 进行中 / 未完成的任务回填上方表单;过往列表只放已完成成片。
jobs = (
self._quick_job_queryset()
.exclude(status__in=[QuickCreateJob.Status.QUEUED, QuickCreateJob.Status.RUNNING])
.filter(
project__is_deleted=False,
project__purged_at__isnull=True,
status=QuickCreateJob.Status.SUCCEEDED,
)
.order_by("-created_at")
)
return Response({
@@ -975,8 +985,9 @@ class ProjectViewSet(TeamScopedViewSetMixin, ModelViewSet):
prompt=request.data.get("prompt", ""),
label=request.data.get("label", ""),
reference_asset_id=request.data.get("reference_asset_id") or None,
# 角色立绘不再自动接力三视图;三视图只由角色详情里的显式按钮生成。
auto_triview=False,
# 用户点击角色 AI 生成 = 生成立绘并在完成后自动接力三视图。
# 三视图任务由 worker 创建,页面刷新或离开也不会漏掉。
auto_triview=kind == BaseAssetGroup.Kind.PERSON,
)
except ValueError as exc: # 无可用模型 / 余额不足等,立即反馈
internal_kind = "user_credit_insufficient" if str(exc).strip().lower() == "insufficient credit" else ""
@@ -1453,7 +1464,14 @@ class ProjectViewSet(TeamScopedViewSetMixin, ModelViewSet):
# 火山/中转报错(如人脸需走素材库的 InputImageSensitiveContentDetected)→ 返回真实报错的 JSON,
# 任务仍保留原始错误供排障,但普通用户仅收到安全错误对象,不让 500 HTML 导致前端白屏。
try:
submit_video_segment(video_segment=segment, user=request.user, prompt=request.data.get("prompt", ""))
submit_video_segment(
video_segment=segment,
user=request.user,
prompt=request.data.get("prompt", ""),
model_config_id=request.data.get("model_config_id") or None,
aspect_ratio=request.data.get("aspect_ratio") or None,
resolution=request.data.get("resolution") or None,
)
except Exception as exc: # noqa: BLE001
public_error = classify_generation_error(exc, operation="video_generate")
return Response(