测试极速成片
This commit is contained in:
@@ -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
@@ -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
|
||||
|
||||
@@ -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):
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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(
|
||||
|
||||
Reference in New Issue
Block a user