解决极速成片问题
This commit is contained in:
@@ -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.
|
||||
|
||||
@@ -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():
|
||||
|
||||
@@ -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"))
|
||||
|
||||
@@ -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"):
|
||||
|
||||
@@ -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")
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user