优化极速成品
This commit is contained in:
@@ -8,6 +8,7 @@ from __future__ import annotations
|
||||
import json
|
||||
import logging
|
||||
import threading
|
||||
import time
|
||||
from datetime import datetime, timedelta
|
||||
|
||||
from django.db import connections, transaction
|
||||
@@ -72,6 +73,22 @@ def _is_finished(job: QuickCreateJob) -> bool:
|
||||
return job.status in _FINISHED_STATUSES
|
||||
|
||||
|
||||
def _latest_script_task(project: Project) -> AITask | None:
|
||||
return (
|
||||
AITask.objects.filter(project=project, task_type=AITask.Type.SCRIPT_GENERATION)
|
||||
.order_by("-created_at")
|
||||
.first()
|
||||
)
|
||||
|
||||
|
||||
def _script_generation_inflight(project: Project) -> bool:
|
||||
return AITask.objects.filter(
|
||||
project=project,
|
||||
task_type=AITask.Type.SCRIPT_GENERATION,
|
||||
status__in=_ACTIVE_TASK_STATUSES,
|
||||
).exists()
|
||||
|
||||
|
||||
def get_quick_script_model() -> ModelConfig | None:
|
||||
"""极速成片固定复用专业创作的豆包 Seed 2.1 Pro 脚本模型。
|
||||
|
||||
@@ -145,6 +162,10 @@ def _is_retryable_exc(exc: Exception) -> bool:
|
||||
"broken pipe",
|
||||
"socket",
|
||||
"temporarily unavailable",
|
||||
"stream aborted",
|
||||
"client disconnected",
|
||||
"lost connection",
|
||||
"operationalerror",
|
||||
)
|
||||
)
|
||||
|
||||
@@ -295,10 +316,23 @@ def _consume_script_agent(job: QuickCreateJob) -> None:
|
||||
entry_source="ai",
|
||||
persona="reviewer",
|
||||
)
|
||||
# 豆包思考模型会连续吐几百帧 reasoning。原先每帧 refresh_from_db,远程 MySQL
|
||||
# 一抖动就会把生成器关掉,模型调用被误记成 stream aborted (client disconnected)。
|
||||
frames = 0
|
||||
last_cancel_check = time.monotonic()
|
||||
for frame in stream:
|
||||
job.refresh_from_db(fields=["status"])
|
||||
if _is_finished(job):
|
||||
return
|
||||
frames += 1
|
||||
now = time.monotonic()
|
||||
if frames == 1 or frames % 40 == 0 or now - last_cancel_check >= 2:
|
||||
last_cancel_check = now
|
||||
try:
|
||||
job.refresh_from_db(fields=["status"])
|
||||
except Exception: # noqa: BLE001 — 查库失败不能中断正在跑的模型流
|
||||
logger.exception("quick create script cancel check failed for job %s", job.id)
|
||||
else:
|
||||
if job.status == QuickCreateJob.Status.CANCELLED:
|
||||
stream.close()
|
||||
return
|
||||
if not frame.startswith("data:"):
|
||||
continue
|
||||
try:
|
||||
@@ -400,18 +434,20 @@ def _advance_script(job: QuickCreateJob) -> int | None:
|
||||
run_quick_script_task.apply_async(args=[str(job.id)], queue="airshelf.quick")
|
||||
return SCRIPT_POLL_SECONDS
|
||||
|
||||
failed = (
|
||||
AITask.objects.filter(
|
||||
project=job.project,
|
||||
task_type=AITask.Type.SCRIPT_GENERATION,
|
||||
status=AITask.Status.FAILED,
|
||||
)
|
||||
.order_by("-created_at")
|
||||
.first()
|
||||
)
|
||||
if failed is not None:
|
||||
fail_quick_create(job, _task_public_error(failed), internal_error=failed.error_message)
|
||||
latest = _latest_script_task(job.project)
|
||||
if latest is not None and latest.status in _ACTIVE_TASK_STATUSES:
|
||||
started_at = _parse_iso(metadata.get("script_started_at"))
|
||||
elapsed = int((timezone.now() - started_at).total_seconds()) if started_at else 0
|
||||
_save_job(job, progress=min(44, 28 + elapsed // 8), message="正在生成分镜脚本…")
|
||||
return SCRIPT_POLL_SECONDS
|
||||
if latest is not None and latest.status == AITask.Status.FAILED:
|
||||
fail_quick_create(job, _task_public_error(latest), internal_error=latest.error_message)
|
||||
return None
|
||||
if latest is not None and latest.status == AITask.Status.SUCCEEDED:
|
||||
started_at = _parse_iso(metadata.get("script_started_at"))
|
||||
elapsed = int((timezone.now() - started_at).total_seconds()) if started_at else 0
|
||||
_save_job(job, progress=min(44, 28 + elapsed // 8), message="脚本已生成,正在写入分镜…")
|
||||
return SCRIPT_POLL_SECONDS
|
||||
|
||||
started_at = _parse_iso(metadata.get("script_started_at"))
|
||||
if started_at and timezone.now() - started_at > SCRIPT_TIMEOUT:
|
||||
@@ -968,6 +1004,27 @@ def recover_quick_create(job: QuickCreateJob) -> None:
|
||||
):
|
||||
_save_job(job, status=QuickCreateJob.Status.RUNNING, error_message="", message="网络波动,正在继续生成…")
|
||||
_enqueue_advance(job)
|
||||
return
|
||||
if (
|
||||
job.status == QuickCreateJob.Status.FAILED
|
||||
and job.phase == QuickCreateJob.Phase.SCRIPT
|
||||
and _adopted_script(job.project) is None
|
||||
and _transient_internal(job)
|
||||
and retries < TRANSIENT_RETRY_LIMIT
|
||||
):
|
||||
metadata = dict(job.metadata or {})
|
||||
metadata["transient_retries"] = retries + 1
|
||||
metadata.pop("script_started", None)
|
||||
metadata.pop("script_started_at", None)
|
||||
metadata.pop("script_local", None)
|
||||
_save_job(
|
||||
job,
|
||||
status=QuickCreateJob.Status.RUNNING,
|
||||
error_message="",
|
||||
message="脚本生成中断,正在重新生成…",
|
||||
metadata=metadata,
|
||||
)
|
||||
_enqueue_advance(job)
|
||||
return
|
||||
if _is_finished(job):
|
||||
return
|
||||
@@ -983,6 +1040,7 @@ def recover_quick_create(job: QuickCreateJob) -> None:
|
||||
and timezone.now() - started_at > SCRIPT_STOLEN_AFTER
|
||||
and adopted is None
|
||||
and not local_running
|
||||
and not _script_generation_inflight(job.project)
|
||||
)
|
||||
if stolen:
|
||||
metadata["script_local"] = True
|
||||
@@ -990,7 +1048,7 @@ def recover_quick_create(job: QuickCreateJob) -> None:
|
||||
_save_job(job, metadata=metadata, message="正在继续生成分镜脚本…")
|
||||
_run_quick_script_in_thread(job_id)
|
||||
return
|
||||
if timed_out and adopted is None and not local_running:
|
||||
if timed_out and adopted is None and not local_running and not _script_generation_inflight(job.project):
|
||||
fail_quick_create(job, "脚本生成超时,请稍后重试或进入专业模式查看")
|
||||
return
|
||||
if timezone.now() - job.updated_at > STALE_AFTER and _claim_next_advance(str(job.id), SCRIPT_POLL_SECONDS):
|
||||
|
||||
@@ -9,13 +9,14 @@ from django.utils import timezone
|
||||
from rest_framework.test import APIClient
|
||||
|
||||
from apps.accounts.models import Team, TeamMember, User
|
||||
from apps.ai.models import ModelConfig, ModelProvider
|
||||
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.serializers import ProjectListSerializer, QuickCreateJobSerializer
|
||||
from apps.projects.services.pipeline import initialize_project_pipeline
|
||||
from apps.projects.services.quick_create import (
|
||||
_consume_script_agent,
|
||||
_reviews_ready,
|
||||
_start_videos,
|
||||
advance_quick_create,
|
||||
@@ -374,6 +375,25 @@ class QuickCreateCoordinatorTests(TestCase):
|
||||
)
|
||||
initialize_project_pipeline(self.project)
|
||||
self.job = QuickCreateJob.objects.create(team=self.team, created_by=self.user, project=self.project)
|
||||
self.provider = ModelProvider.objects.create(name="quick-script-abort", display_name="Quick Script")
|
||||
self.model = ModelConfig.objects.create(
|
||||
provider=self.provider,
|
||||
name="doubao-seed-2-1-pro-260628",
|
||||
display_name="豆包 2.1 Pro",
|
||||
capability=ModelConfig.Capability.TEXT,
|
||||
)
|
||||
|
||||
def _script_task(self, status, *, key, error_message=""):
|
||||
return AITask.objects.create(
|
||||
team=self.team,
|
||||
created_by=self.user,
|
||||
project=self.project,
|
||||
model_config=self.model,
|
||||
task_type=AITask.Type.SCRIPT_GENERATION,
|
||||
status=status,
|
||||
idempotency_key=key,
|
||||
error_message=error_message,
|
||||
)
|
||||
|
||||
@patch("apps.projects.tasks.run_quick_script_task.apply_async")
|
||||
def test_first_advance_recognizes_product_then_starts_script(self, start_script):
|
||||
@@ -427,6 +447,69 @@ class QuickCreateCoordinatorTests(TestCase):
|
||||
self.assertEqual(self.job.status, QuickCreateJob.Status.RUNNING)
|
||||
run_local.assert_called_once_with(str(self.job.id))
|
||||
|
||||
@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):
|
||||
self._script_task(AITask.Status.SUBMITTED, key="script-inflight-1")
|
||||
self.job.status = QuickCreateJob.Status.RUNNING
|
||||
self.job.phase = QuickCreateJob.Phase.SCRIPT
|
||||
self.job.metadata = {
|
||||
"script_started": True,
|
||||
"script_started_at": (timezone.now() - timedelta(minutes=5)).isoformat(),
|
||||
}
|
||||
self.job.save(update_fields=["status", "phase", "metadata", "updated_at"])
|
||||
|
||||
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()
|
||||
|
||||
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)")
|
||||
self._script_task(AITask.Status.SUBMITTED, key="script-latest-running")
|
||||
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)
|
||||
|
||||
@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
|
||||
self.job.phase = QuickCreateJob.Phase.SCRIPT
|
||||
self.job.error_message = "本次未生成可用结果,请重试。"
|
||||
self.job.metadata = {
|
||||
"script_started": True,
|
||||
"script_started_at": timezone.now().isoformat(),
|
||||
"internal_error": "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"))
|
||||
self.assertEqual(self.job.metadata.get("transient_retries"), 1)
|
||||
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):
|
||||
get_model.return_value = self.model
|
||||
stream_fn.return_value = (f'data: {{"type": "reasoning", "text": "{index}"}}\n\n' for index in range(80))
|
||||
with patch.object(QuickCreateJob, "refresh_from_db") as refresh:
|
||||
refresh.side_effect = lambda **kwargs: None
|
||||
with self.assertRaises(ValueError):
|
||||
_consume_script_agent(self.job)
|
||||
self.assertLessEqual(refresh.call_count, 3)
|
||||
|
||||
def test_timeout_before_video_request_does_not_fail_project(self):
|
||||
self.project.status = Project.Status.VIDEOING
|
||||
self.project.current_stage = ProjectStage.Stage.VIDEO
|
||||
|
||||
Reference in New Issue
Block a user