优化极速成品
This commit is contained in:
@@ -8,6 +8,7 @@ from __future__ import annotations
|
|||||||
import json
|
import json
|
||||||
import logging
|
import logging
|
||||||
import threading
|
import threading
|
||||||
|
import time
|
||||||
from datetime import datetime, timedelta
|
from datetime import datetime, timedelta
|
||||||
|
|
||||||
from django.db import connections, transaction
|
from django.db import connections, transaction
|
||||||
@@ -72,6 +73,22 @@ def _is_finished(job: QuickCreateJob) -> bool:
|
|||||||
return job.status in _FINISHED_STATUSES
|
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:
|
def get_quick_script_model() -> ModelConfig | None:
|
||||||
"""极速成片固定复用专业创作的豆包 Seed 2.1 Pro 脚本模型。
|
"""极速成片固定复用专业创作的豆包 Seed 2.1 Pro 脚本模型。
|
||||||
|
|
||||||
@@ -145,6 +162,10 @@ def _is_retryable_exc(exc: Exception) -> bool:
|
|||||||
"broken pipe",
|
"broken pipe",
|
||||||
"socket",
|
"socket",
|
||||||
"temporarily unavailable",
|
"temporarily unavailable",
|
||||||
|
"stream aborted",
|
||||||
|
"client disconnected",
|
||||||
|
"lost connection",
|
||||||
|
"operationalerror",
|
||||||
)
|
)
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -295,10 +316,23 @@ def _consume_script_agent(job: QuickCreateJob) -> None:
|
|||||||
entry_source="ai",
|
entry_source="ai",
|
||||||
persona="reviewer",
|
persona="reviewer",
|
||||||
)
|
)
|
||||||
|
# 豆包思考模型会连续吐几百帧 reasoning。原先每帧 refresh_from_db,远程 MySQL
|
||||||
|
# 一抖动就会把生成器关掉,模型调用被误记成 stream aborted (client disconnected)。
|
||||||
|
frames = 0
|
||||||
|
last_cancel_check = time.monotonic()
|
||||||
for frame in stream:
|
for frame in stream:
|
||||||
job.refresh_from_db(fields=["status"])
|
frames += 1
|
||||||
if _is_finished(job):
|
now = time.monotonic()
|
||||||
return
|
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:"):
|
if not frame.startswith("data:"):
|
||||||
continue
|
continue
|
||||||
try:
|
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")
|
run_quick_script_task.apply_async(args=[str(job.id)], queue="airshelf.quick")
|
||||||
return SCRIPT_POLL_SECONDS
|
return SCRIPT_POLL_SECONDS
|
||||||
|
|
||||||
failed = (
|
latest = _latest_script_task(job.project)
|
||||||
AITask.objects.filter(
|
if latest is not None and latest.status in _ACTIVE_TASK_STATUSES:
|
||||||
project=job.project,
|
started_at = _parse_iso(metadata.get("script_started_at"))
|
||||||
task_type=AITask.Type.SCRIPT_GENERATION,
|
elapsed = int((timezone.now() - started_at).total_seconds()) if started_at else 0
|
||||||
status=AITask.Status.FAILED,
|
_save_job(job, progress=min(44, 28 + elapsed // 8), message="正在生成分镜脚本…")
|
||||||
)
|
return SCRIPT_POLL_SECONDS
|
||||||
.order_by("-created_at")
|
if latest is not None and latest.status == AITask.Status.FAILED:
|
||||||
.first()
|
fail_quick_create(job, _task_public_error(latest), internal_error=latest.error_message)
|
||||||
)
|
|
||||||
if failed is not None:
|
|
||||||
fail_quick_create(job, _task_public_error(failed), internal_error=failed.error_message)
|
|
||||||
return None
|
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"))
|
started_at = _parse_iso(metadata.get("script_started_at"))
|
||||||
if started_at and timezone.now() - started_at > SCRIPT_TIMEOUT:
|
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="网络波动,正在继续生成…")
|
_save_job(job, status=QuickCreateJob.Status.RUNNING, error_message="", message="网络波动,正在继续生成…")
|
||||||
_enqueue_advance(job)
|
_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
|
return
|
||||||
if _is_finished(job):
|
if _is_finished(job):
|
||||||
return
|
return
|
||||||
@@ -983,6 +1040,7 @@ def recover_quick_create(job: QuickCreateJob) -> None:
|
|||||||
and timezone.now() - started_at > SCRIPT_STOLEN_AFTER
|
and timezone.now() - started_at > SCRIPT_STOLEN_AFTER
|
||||||
and adopted is None
|
and adopted is None
|
||||||
and not local_running
|
and not local_running
|
||||||
|
and not _script_generation_inflight(job.project)
|
||||||
)
|
)
|
||||||
if stolen:
|
if stolen:
|
||||||
metadata["script_local"] = True
|
metadata["script_local"] = True
|
||||||
@@ -990,7 +1048,7 @@ def recover_quick_create(job: QuickCreateJob) -> None:
|
|||||||
_save_job(job, metadata=metadata, message="正在继续生成分镜脚本…")
|
_save_job(job, metadata=metadata, message="正在继续生成分镜脚本…")
|
||||||
_run_quick_script_in_thread(job_id)
|
_run_quick_script_in_thread(job_id)
|
||||||
return
|
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, "脚本生成超时,请稍后重试或进入专业模式查看")
|
fail_quick_create(job, "脚本生成超时,请稍后重试或进入专业模式查看")
|
||||||
return
|
return
|
||||||
if timezone.now() - job.updated_at > STALE_AFTER and _claim_next_advance(str(job.id), SCRIPT_POLL_SECONDS):
|
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 rest_framework.test import APIClient
|
||||||
|
|
||||||
from apps.accounts.models import Team, TeamMember, User
|
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.assets.models import Asset, AssetFile
|
||||||
from apps.products.models import Product, ProductImage
|
from apps.products.models import Product, ProductImage
|
||||||
from apps.projects.models import Project, ProjectStage, QuickCreateJob, ScriptSegment, ScriptVersion, VideoSegment, VideoSegmentVersion
|
from apps.projects.models import Project, ProjectStage, QuickCreateJob, ScriptSegment, ScriptVersion, VideoSegment, VideoSegmentVersion
|
||||||
from apps.projects.serializers import ProjectListSerializer, QuickCreateJobSerializer
|
from apps.projects.serializers import ProjectListSerializer, QuickCreateJobSerializer
|
||||||
from apps.projects.services.pipeline import initialize_project_pipeline
|
from apps.projects.services.pipeline import initialize_project_pipeline
|
||||||
from apps.projects.services.quick_create import (
|
from apps.projects.services.quick_create import (
|
||||||
|
_consume_script_agent,
|
||||||
_reviews_ready,
|
_reviews_ready,
|
||||||
_start_videos,
|
_start_videos,
|
||||||
advance_quick_create,
|
advance_quick_create,
|
||||||
@@ -374,6 +375,25 @@ class QuickCreateCoordinatorTests(TestCase):
|
|||||||
)
|
)
|
||||||
initialize_project_pipeline(self.project)
|
initialize_project_pipeline(self.project)
|
||||||
self.job = QuickCreateJob.objects.create(team=self.team, created_by=self.user, project=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")
|
@patch("apps.projects.tasks.run_quick_script_task.apply_async")
|
||||||
def test_first_advance_recognizes_product_then_starts_script(self, start_script):
|
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)
|
self.assertEqual(self.job.status, QuickCreateJob.Status.RUNNING)
|
||||||
run_local.assert_called_once_with(str(self.job.id))
|
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):
|
def test_timeout_before_video_request_does_not_fail_project(self):
|
||||||
self.project.status = Project.Status.VIDEOING
|
self.project.status = Project.Status.VIDEOING
|
||||||
self.project.current_stage = ProjectStage.Stage.VIDEO
|
self.project.current_stage = ProjectStage.Stage.VIDEO
|
||||||
|
|||||||
Reference in New Issue
Block a user