diff --git a/core/backend/README.md b/core/backend/README.md index 8279de1..4b36644 100644 --- a/core/backend/README.md +++ b/core/backend/README.md @@ -21,7 +21,7 @@ Start workers in separate terminals: ```bash cd /Users/maidong/Desktop/zyc/qiyuan_gitea/AirShelf/core/backend source .venv/bin/activate -celery -A airshelf worker -l info -P threads -c 4 # 线程池 4 并发(出图/视频可并行;勿用 -P solo 串行) +celery -A airshelf worker -l info -P threads -c 4 -Q celery,airshelf.quick # 必须带 airshelf.quick,否则极速成片会一直停在「等待开始」 ``` `ffmpeg` must be available on `PATH` for Stage5 export jobs. diff --git a/core/backend/apps/common/celery_health.py b/core/backend/apps/common/celery_health.py index 5aa821f..5e2b256 100644 --- a/core/backend/apps/common/celery_health.py +++ b/core/backend/apps/common/celery_health.py @@ -29,6 +29,7 @@ _FAIL_TTL = 5.0 _cache = {"ok": False, "expires": 0.0} # True / False / None(inspect 没答上来,不当成「未部署」) _registered_task_cache: dict[str, tuple[bool | None, float]] = {} +_queue_cache: dict[str, tuple[bool | None, float]] = {} def _task_name_listed(task_name: str, task_names) -> bool: @@ -77,6 +78,42 @@ def celery_worker_available() -> bool: return ok +def _inspect_queue_consumed(queue_name: str) -> bool | None: + """True 有 worker 在听该队列 / False 在线 worker 都不听 / None 问不到。""" + try: + from airshelf.celery import app as celery_app + + active_queues = celery_app.control.inspect(timeout=2.5).active_queues() + except Exception: # noqa: BLE001 — inspect 失败不能当成「没人听」 + return None + if not active_queues: + return None + found_any_queue = False + for queues in active_queues.values(): + for item in queues or []: + found_any_queue = True + name = item.get("name") if isinstance(item, dict) else str(item) + if name == queue_name: + return True + if found_any_queue: + return False + return None + + +def worker_consumes_queue(queue_name: str) -> bool | None: + """线上 worker 默认只听 celery;极速成片发到 airshelf.quick 时用来决定要不要本机兜底。""" + if getattr(settings, "CELERY_TASK_ALWAYS_EAGER", False): + return True + now = time.monotonic() + cached = _queue_cache.get(queue_name) + if cached and now < cached[1]: + return cached[0] + available = _inspect_queue_consumed(queue_name) + ttl = _OK_TTL if available else _FAIL_TTL + _queue_cache[queue_name] = (available, now + ttl) + return available + + def require_worker() -> None: """生成类提交入口的前置检查:无 worker 直接 503,不让任务出门。""" if not celery_worker_available(): diff --git a/core/backend/apps/common/test_celery_health.py b/core/backend/apps/common/test_celery_health.py index 44f0b9b..fb250d7 100644 --- a/core/backend/apps/common/test_celery_health.py +++ b/core/backend/apps/common/test_celery_health.py @@ -4,16 +4,18 @@ from django.test import SimpleTestCase, override_settings from rest_framework.exceptions import APIException from apps.common import celery_health -from apps.common.celery_health import require_worker_task +from apps.common.celery_health import require_worker_task, worker_consumes_queue class RequireWorkerTaskTests(SimpleTestCase): def setUp(self): celery_health._registered_task_cache.clear() + celery_health._queue_cache.clear() celery_health._cache["expires"] = 0.0 def tearDown(self): celery_health._registered_task_cache.clear() + celery_health._queue_cache.clear() celery_health._cache["expires"] = 0.0 def test_registered_name_accepts_short_name_and_rate_suffix(self): @@ -39,3 +41,13 @@ class RequireWorkerTaskTests(SimpleTestCase): self.assertEqual(raised.exception.status_code, 503) self.assertIn("重启", str(raised.exception.detail)) self.assertNotIn("正在升级", str(raised.exception.detail)) + + @override_settings(CELERY_TASK_ALWAYS_EAGER=False) + @patch("apps.common.celery_health._inspect_queue_consumed", return_value=False) + def test_missing_quick_queue_is_detected(self, _inspect): + self.assertIs(worker_consumes_queue("airshelf.quick"), False) + + @override_settings(CELERY_TASK_ALWAYS_EAGER=False) + @patch("apps.common.celery_health._inspect_queue_consumed", return_value=None) + def test_unknown_queue_inspect_does_not_pretend_missing(self, _inspect): + self.assertIsNone(worker_consumes_queue("airshelf.quick")) diff --git a/core/backend/apps/projects/services/quick_create.py b/core/backend/apps/projects/services/quick_create.py index 67a9286..c8b4f29 100644 --- a/core/backend/apps/projects/services/quick_create.py +++ b/core/backend/apps/projects/services/quick_create.py @@ -423,15 +423,23 @@ def _advance_script(job: QuickCreateJob) -> int | None: if not metadata.get("script_started"): from apps.projects.tasks import run_quick_script_task + from apps.common.celery_health import worker_consumes_queue + metadata["script_started"] = True metadata["script_started_at"] = timezone.now().isoformat() + listens = worker_consumes_queue("airshelf.quick") + if listens is False: + metadata["script_local"] = True _save_job( job, metadata=metadata, progress=28, message="正在根据商品名称与图片生成带货脚本", ) - run_quick_script_task.apply_async(args=[str(job.id)], queue="airshelf.quick") + if listens is False: + _run_quick_script_in_thread(str(job.id)) + else: + run_quick_script_task.apply_async(args=[str(job.id)], queue="airshelf.quick") return SCRIPT_POLL_SECONDS latest = _latest_script_task(job.project) @@ -946,12 +954,26 @@ def _run_quick_script_in_thread(job_id: str) -> None: def _enqueue_advance(job: QuickCreateJob) -> None: + from apps.common.celery_health import worker_consumes_queue from apps.projects.tasks import advance_quick_create_task job_id = str(job.id) + if worker_consumes_queue("airshelf.quick") is False: + _advance_without_quick_queue(job_id) + return try: advance_quick_create_task.apply_async(args=[job_id], queue="airshelf.quick") except Exception: # noqa: BLE001 — 队列不可用时就地推进一步 + _advance_without_quick_queue(job_id) + + +def _advance_without_quick_queue(job_id: str) -> None: + """部署后的 worker 若只听 celery 队列,airshelf.quick 里的任务会永远没人拿。""" + delay = advance_quick_create(job_id) + job = QuickCreateJob.objects.filter(id=job_id).first() + if job is None or delay is None: + return + if job.phase == QuickCreateJob.Phase.SCRIPT and not (job.metadata or {}).get("script_started"): advance_quick_create(job_id) @@ -1028,6 +1050,13 @@ def recover_quick_create(job: QuickCreateJob) -> None: return if _is_finished(job): return + if job.status == QuickCreateJob.Status.QUEUED or job.phase == QuickCreateJob.Phase.PRODUCT: + if timezone.now() - job.updated_at > STALE_AFTER: + try: + _advance_without_quick_queue(str(job.id)) + except Exception: # noqa: BLE001 — 恢复失败不能把进度接口打成 500 + logger.exception("quick create recover advance failed for job %s", job.id) + return job_id = str(job.id) metadata = dict(job.metadata or {}) if job.phase == QuickCreateJob.Phase.SCRIPT and metadata.get("script_started"): diff --git a/core/backend/apps/projects/test_quick_create.py b/core/backend/apps/projects/test_quick_create.py index 9fb5b06..bf353e0 100644 --- a/core/backend/apps/projects/test_quick_create.py +++ b/core/backend/apps/projects/test_quick_create.py @@ -48,7 +48,7 @@ class QuickCreateApiTests(TestCase): @patch("apps.projects.services.quick_create.get_quick_script_model", return_value=object()) @patch("apps.projects.views.get_default_model") - @patch("apps.projects.views.advance_quick_create_task.apply_async") + @patch("apps.projects.tasks.advance_quick_create_task.apply_async") @patch("apps.projects.views.require_worker_task") @patch("apps.projects.views._store_uploaded_asset") def test_submit_creates_product_project_and_persistent_job(self, store_asset, require_worker_task, enqueue, get_model, get_quick_model): @@ -106,7 +106,7 @@ class QuickCreateApiTests(TestCase): @patch("apps.projects.services.quick_create.get_quick_script_model", return_value=object()) @patch("apps.projects.views.get_default_model", return_value=object()) - @patch("apps.projects.views.advance_quick_create_task.apply_async") + @patch("apps.projects.tasks.advance_quick_create_task.apply_async") @patch("apps.projects.views.require_worker_task") @patch("apps.projects.views._store_uploaded_asset") def test_submit_reuses_source_product_images(self, store_asset, require_worker_task, enqueue, get_model, get_quick_model): @@ -447,6 +447,32 @@ class QuickCreateCoordinatorTests(TestCase): self.assertEqual(self.job.status, QuickCreateJob.Status.RUNNING) run_local.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( + status=QuickCreateJob.Status.QUEUED, + phase=QuickCreateJob.Phase.PRODUCT, + message="等待开始极速成片", + 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: + recover_quick_create(self.job) + advance.assert_called_once_with(str(self.job.id)) + + @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): + self.job.status = QuickCreateJob.Status.RUNNING + self.job.phase = QuickCreateJob.Phase.SCRIPT + 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, 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)) + @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") diff --git a/core/backend/apps/projects/views.py b/core/backend/apps/projects/views.py index 9a41898..a081f29 100644 --- a/core/backend/apps/projects/views.py +++ b/core/backend/apps/projects/views.py @@ -90,7 +90,7 @@ from .services.pipeline import ( ) from .services.script_import import ScriptFileError, extract_script_text from .services.templates import build_template_fields, coerce_persona, coerce_template_combo, render_outline_text -from .tasks import advance_quick_create_task, poll_video_segment_task +from .tasks import poll_video_segment_task logger = logging.getLogger(__name__) @@ -661,7 +661,9 @@ class ProjectViewSet(TeamScopedViewSetMixin, ModelViewSet): ) try: - advance_quick_create_task.apply_async(args=[str(job.id)], queue="airshelf.quick") + from .services.quick_create import _enqueue_advance + + _enqueue_advance(job) except Exception as exc: # noqa: BLE001 — broker 极小窗口失败也必须给任务落终态 from .services.quick_create import fail_quick_create diff --git a/k8s/core/worker-deployment.yaml b/k8s/core/worker-deployment.yaml index 24da11f..4151abe 100644 --- a/k8s/core/worker-deployment.yaml +++ b/k8s/core/worker-deployment.yaml @@ -23,7 +23,7 @@ spec: # Celery worker connects to the external (Volcano managed) Redis broker # configured via the airshelf-core-env secret. Uses `args` (not `command`) # so the image ENTRYPOINT still runs but skips migrate/collectstatic ($1=celery). - args: ["celery", "-A", "airshelf.celery:app", "worker", "-l", "info", "--concurrency", "4"] + args: ["celery", "-A", "airshelf.celery:app", "worker", "-l", "info", "-Q", "celery,airshelf.quick", "--concurrency", "4"] envFrom: - secretRef: name: airshelf-core-env