From 43d0b088e406ad37be4c3613b2488b3f8217151b Mon Sep 17 00:00:00 2001 From: "Azmat@qq.com" Date: Wed, 26 Aug 2026 09:00:00 +0800 Subject: [PATCH] =?UTF-8?q?=E4=BC=98=E5=8C=96=E6=9E=81=E9=80=9F=E6=88=90?= =?UTF-8?q?=E5=93=81?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- .../apps/projects/services/quick_create.py | 88 +++++++++++++++---- .../apps/projects/test_quick_create.py | 85 +++++++++++++++++- 2 files changed, 157 insertions(+), 16 deletions(-) diff --git a/core/backend/apps/projects/services/quick_create.py b/core/backend/apps/projects/services/quick_create.py index 116e27c..67a9286 100644 --- a/core/backend/apps/projects/services/quick_create.py +++ b/core/backend/apps/projects/services/quick_create.py @@ -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): diff --git a/core/backend/apps/projects/test_quick_create.py b/core/backend/apps/projects/test_quick_create.py index 48975e7..9fb5b06 100644 --- a/core/backend/apps/projects/test_quick_create.py +++ b/core/backend/apps/projects/test_quick_create.py @@ -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