feat(worker): celery 队列隔离 + 孤儿任务消息作废 (#1714) #1722
@@ -0,0 +1,35 @@
|
||||
"""add celery_task_id to generation_tasks and ingest_jobs
|
||||
|
||||
Issue #1714:孤儿恢复/超时清理撤销队列消息。
|
||||
- generation_tasks.celery_task_id:入队时记录的 Celery 消息 ID,清理时 revoke
|
||||
- ingest_jobs.celery_task_id:同上(素材转码任务)
|
||||
|
||||
Revision ID: 067_celery_task_id
|
||||
Revises: 066_upload_idempotency
|
||||
Create Date: 2026-09-05
|
||||
"""
|
||||
|
||||
import sqlalchemy as sa
|
||||
|
||||
from alembic import op
|
||||
|
||||
revision = "067_celery_task_id"
|
||||
down_revision = "066_upload_idempotency"
|
||||
branch_labels = None
|
||||
depends_on = None
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
op.add_column(
|
||||
"generation_tasks",
|
||||
sa.Column("celery_task_id", sa.String(64), nullable=False, server_default=""),
|
||||
)
|
||||
op.add_column(
|
||||
"ingest_jobs",
|
||||
sa.Column("celery_task_id", sa.String(64), nullable=False, server_default=""),
|
||||
)
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
op.drop_column("ingest_jobs", "celery_task_id")
|
||||
op.drop_column("generation_tasks", "celery_task_id")
|
||||
@@ -14,6 +14,7 @@ from typing import Any
|
||||
from uuid import uuid4
|
||||
|
||||
from app.api.routes._helpers import require_project_and_library
|
||||
from app.api.routes.upload import _persist_celery_task_id
|
||||
from app.auth import AuthenticatedUser, get_current_user
|
||||
from app.core.celery_app import celery_app
|
||||
from app.core.storage import OSSStorageService, get_storage_service
|
||||
@@ -381,7 +382,8 @@ async def complete_chunked_upload(
|
||||
file_hash=request.file_hash,
|
||||
)
|
||||
)
|
||||
celery_app.send_task("worker.ingest_asset", args=[job.id])
|
||||
celery_result = celery_app.send_task("worker.ingest_asset", args=[job.id])
|
||||
_persist_celery_task_id(ingest_job_repository, job, getattr(celery_result, "id", ""))
|
||||
|
||||
# Update metadata status
|
||||
meta["status"] = "completed"
|
||||
|
||||
@@ -43,7 +43,13 @@ def submit_ingest_job(
|
||||
)
|
||||
)
|
||||
|
||||
celery_app.send_task("worker.ingest_asset", args=[job.id])
|
||||
celery_result = celery_app.send_task("worker.ingest_asset", args=[job.id])
|
||||
if getattr(celery_result, "id", ""):
|
||||
try:
|
||||
job.celery_task_id = celery_result.id
|
||||
ingest_job_repository.update(job)
|
||||
except Exception: # noqa: BLE001
|
||||
pass
|
||||
|
||||
return IngestJobResponse(
|
||||
id=job.id,
|
||||
|
||||
@@ -375,7 +375,13 @@ def retry_project_task(
|
||||
storage_key=job.storage_key,
|
||||
)
|
||||
)
|
||||
celery_app.send_task("worker.ingest_asset", args=[retried.id])
|
||||
celery_result = celery_app.send_task("worker.ingest_asset", args=[retried.id])
|
||||
if getattr(celery_result, "id", ""):
|
||||
try:
|
||||
retried.celery_task_id = celery_result.id
|
||||
ingest_job_repository.update(retried)
|
||||
except Exception: # noqa: BLE001
|
||||
pass
|
||||
return ProjectTaskResponse(
|
||||
id=f"ingest:{retried.id}",
|
||||
task_type="ingest",
|
||||
|
||||
@@ -203,6 +203,17 @@ def _create_pending_asset(
|
||||
return asset_repository.create(asset)
|
||||
|
||||
|
||||
def _persist_celery_task_id(repo: Any, job: Any, celery_task_id: str) -> None:
|
||||
"""记录 celery 消息 ID 到任务行,供孤儿清理时 revoke/清除队列消息(#1714)。"""
|
||||
if not celery_task_id:
|
||||
return
|
||||
try:
|
||||
job.celery_task_id = celery_task_id
|
||||
repo.update(job)
|
||||
except Exception: # noqa: BLE001 — 记录失败不影响主流程(执行前状态守卫兜底)
|
||||
pass
|
||||
|
||||
|
||||
def _submit_ingest_job(
|
||||
project_id: str,
|
||||
library_id: str,
|
||||
@@ -221,7 +232,8 @@ def _submit_ingest_job(
|
||||
asset_id=asset_id,
|
||||
)
|
||||
)
|
||||
celery_app.send_task("worker.ingest_asset", args=[job.id])
|
||||
celery_result = celery_app.send_task("worker.ingest_asset", args=[job.id])
|
||||
_persist_celery_task_id(ingest_job_repository, job, getattr(celery_result, "id", ""))
|
||||
return job
|
||||
|
||||
|
||||
|
||||
@@ -5,3 +5,11 @@ settings = get_settings()
|
||||
celery_app = Celery("xiaoxia-saas-api")
|
||||
celery_app.conf.broker_url = settings.CELERY_BROKER_URL
|
||||
celery_app.conf.result_backend = settings.CELERY_RESULT_BACKEND
|
||||
|
||||
# #1714 队列隔离:视频生成走 generation 队列,素材转码走 transcode 队列
|
||||
try:
|
||||
from packages.shared.celery_queues import apply_queue_settings
|
||||
|
||||
apply_queue_settings(celery_app)
|
||||
except Exception: # noqa: BLE001 — 队列配置失败不阻断 API 启动
|
||||
pass
|
||||
|
||||
@@ -279,7 +279,17 @@ def safe_enqueue_generation_task(
|
||||
|
||||
# ── 发送 Celery 任务 ──
|
||||
try:
|
||||
celery_app.send_task("worker.generate_video", args=[task.id])
|
||||
celery_result = celery_app.send_task("worker.generate_video", args=[task.id])
|
||||
# 记录 celery 消息 ID:孤儿清理/超时作废时据此 revoke + 清除队列消息(#1714)
|
||||
celery_task_id = getattr(celery_result, "id", "")
|
||||
if celery_task_id:
|
||||
try:
|
||||
task.celery_task_id = celery_task_id
|
||||
generation_task_repository.update(task)
|
||||
except Exception as persist_err: # noqa: BLE001
|
||||
logger.warning(
|
||||
"%s 持久化 celery_task_id 失败(不影响主流程): task_id=%s err=%s", log_prefix, task.id, persist_err
|
||||
)
|
||||
except Exception as e:
|
||||
logger.error(
|
||||
"%s 入队失败,标记为失败: task_id=%s error=%s",
|
||||
|
||||
@@ -6,6 +6,18 @@ celery_app = Celery(settings.worker_name)
|
||||
celery_app.conf.broker_url = settings.broker_url
|
||||
celery_app.conf.result_backend = settings.result_backend
|
||||
celery_app.conf.broker_connection_retry_on_startup = True
|
||||
|
||||
# #1714 队列隔离:generation(高优,独占 worker)/ transcode(素材转码)/ celery(默认)
|
||||
from packages.shared.celery_queues import ( # noqa: E402
|
||||
GENERATION_WORKER_PREFETCH_MULTIPLIER,
|
||||
apply_queue_settings,
|
||||
)
|
||||
|
||||
apply_queue_settings(celery_app)
|
||||
# 长渲染任务预取 1,避免任务被预取占住导致调度不均
|
||||
celery_app.conf.worker_prefetch_multiplier = GENERATION_WORKER_PREFETCH_MULTIPLIER
|
||||
celery_app.conf.task_acks_late = True # worker 崩溃时未完成任务重回队列,由执行前守卫丢弃作废消息
|
||||
|
||||
celery_app.conf.imports = (
|
||||
"worker_app.tasks.health",
|
||||
"worker_app.tasks.ingest",
|
||||
|
||||
@@ -14,7 +14,17 @@ def cleanup_stale_running_with_session(repo, timeout_minutes: int) -> int:
|
||||
Returns:
|
||||
清理的任务数量
|
||||
"""
|
||||
return repo.cleanup_stale_running(timeout_minutes)
|
||||
return len(cleanup_stale_running_with_session_ids(repo, timeout_minutes))
|
||||
|
||||
|
||||
def cleanup_stale_running_with_session_ids(repo, timeout_minutes: int) -> list[tuple[str, str]]:
|
||||
"""同 cleanup_stale_running_with_session,返回 [(task_id, celery_task_id), ...]。"""
|
||||
fn = getattr(repo, "cleanup_stale_running_with_ids", None)
|
||||
if fn is not None:
|
||||
return fn(timeout_minutes)
|
||||
# 旧仓储无 _with_ids 方法:降级为计数,无法撤销消息(执行前状态守卫兜底)
|
||||
count = repo.cleanup_stale_running(timeout_minutes)
|
||||
return [("", "") for _ in range(count)]
|
||||
|
||||
|
||||
def cleanup_stale_pending_with_session(repo, timeout_minutes: int) -> int:
|
||||
@@ -23,7 +33,44 @@ def cleanup_stale_pending_with_session(repo, timeout_minutes: int) -> int:
|
||||
Returns:
|
||||
清理的任务数量
|
||||
"""
|
||||
return repo.cleanup_stale_pending(timeout_minutes)
|
||||
return len(cleanup_stale_pending_with_session_ids(repo, timeout_minutes))
|
||||
|
||||
|
||||
def cleanup_stale_pending_with_session_ids(repo, timeout_minutes: int) -> list[tuple[str, str]]:
|
||||
"""同 cleanup_stale_pending_with_session,返回 [(task_id, celery_task_id), ...]。"""
|
||||
fn = getattr(repo, "cleanup_stale_pending_with_ids", None)
|
||||
if fn is not None:
|
||||
return fn(timeout_minutes)
|
||||
count = repo.cleanup_stale_pending(timeout_minutes)
|
||||
return [("", "") for _ in range(count)]
|
||||
|
||||
|
||||
def _revoke_and_purge_stale_messages(items: list[tuple[str, str]]) -> int:
|
||||
"""把清理掉的任务对应的 Celery 消息撤销并从 Redis 队列清除(#1714)。
|
||||
|
||||
防止「DB 已标 failed,但队列消息还在 → 重投执行 → 非法状态转换 → 半成品」。
|
||||
失败不阻断清理流程(执行前状态守卫是第二道防线)。
|
||||
"""
|
||||
biz_ids = [tid for tid, _ in items if tid]
|
||||
celery_ids = [cid for _, cid in items if cid]
|
||||
if not biz_ids and not celery_ids:
|
||||
return 0
|
||||
try:
|
||||
from worker_app.celery_app import celery_app as app
|
||||
from worker_app.core.config import get_settings
|
||||
|
||||
from packages.shared.celery_orphan_guard import revoke_and_purge
|
||||
|
||||
broker_url = get_settings().broker_url
|
||||
return revoke_and_purge(
|
||||
app,
|
||||
broker_url,
|
||||
business_task_ids=biz_ids,
|
||||
celery_task_ids=celery_ids,
|
||||
)
|
||||
except Exception as e: # noqa: BLE001
|
||||
logger.error("撤销作废任务队列消息失败(执行前守卫仍会兜底): %s", e, exc_info=True)
|
||||
return 0
|
||||
|
||||
|
||||
# 孤儿任务超时阈值:running 任务超过此时间无进度更新则视为卡死。
|
||||
@@ -31,10 +78,12 @@ def cleanup_stale_pending_with_session(repo, timeout_minutes: int) -> int:
|
||||
# 20 分钟阈值覆盖硬超时 + 重试 + 余量,绝不误杀正常任务。
|
||||
ORPHAN_TASK_TIMEOUT_MINUTES = 20
|
||||
|
||||
# Pending 任务超时阈值:任务创建后超过此时间仍未被 worker 拉取,
|
||||
# 说明 worker 已停止消费(容器异常/卡死),清掉释放限流名额。
|
||||
# 依据:满队列(20 pending)× 平均 2 分钟 / 并发 4 ≈ 10 分钟,15 分钟留余量。
|
||||
PENDING_TASK_TIMEOUT_MINUTES = 15
|
||||
# Pending 任务超时阈值:任务创建后超过此时间仍未开始执行则判死。
|
||||
# 注意区分 running 孤儿阈值(20 分钟):pending 是「排队等待」时间,
|
||||
# 队列积压(如 20+ 转码任务)时视频生成可能正常排队较久,阈值必须放宽,
|
||||
# 避免正常排队任务被误杀。队列隔离(#1714)后 generation 队列独占 worker,
|
||||
# 理论上排队极短;保留 45 分钟作为兜底,覆盖 worker 短暂停止消费的场景。
|
||||
PENDING_TASK_TIMEOUT_MINUTES = 45
|
||||
|
||||
|
||||
def cleanup_orphan_tasks(timeout_minutes: int = ORPHAN_TASK_TIMEOUT_MINUTES) -> int: # pragma: no cover
|
||||
@@ -57,11 +106,14 @@ def cleanup_orphan_tasks(timeout_minutes: int = ORPHAN_TASK_TIMEOUT_MINUTES) ->
|
||||
session = SessionLocal()
|
||||
try:
|
||||
repo = SQLAlchemyGenerationTaskRepository(session)
|
||||
count = cleanup_stale_running_with_session(repo, timeout_minutes)
|
||||
items = cleanup_stale_running_with_session_ids(repo, timeout_minutes)
|
||||
finally:
|
||||
session.close()
|
||||
count = len(items)
|
||||
if count > 0:
|
||||
logger.warning("清理了 %d 个超时的孤儿 GenerationTask(超过 %d 分钟未更新)", count, timeout_minutes)
|
||||
purged = _revoke_and_purge_stale_messages(items)
|
||||
logger.info("孤儿任务对应队列消息撤销/清除完成: %d 条", purged)
|
||||
else:
|
||||
logger.info("无孤儿 GenerationTask 需要清理")
|
||||
return count
|
||||
@@ -95,7 +147,9 @@ def cleanup_stale_jobs(timeout_minutes: int = ORPHAN_TASK_TIMEOUT_MINUTES) -> in
|
||||
.all()
|
||||
)
|
||||
count = 0
|
||||
stale_items: list[tuple[str, str]] = []
|
||||
for model in stale_jobs:
|
||||
stale_items.append((model.id, getattr(model, "celery_task_id", "") or ""))
|
||||
model.status = JobStatus.FAILED.value
|
||||
model.error_message = f"任务执行中断(超过 {timeout_minutes} 分钟未更新)"
|
||||
count += 1
|
||||
@@ -105,6 +159,8 @@ def cleanup_stale_jobs(timeout_minutes: int = ORPHAN_TASK_TIMEOUT_MINUTES) -> in
|
||||
else:
|
||||
logger.info("无孤儿 Job 需要清理")
|
||||
session.close()
|
||||
if count > 0:
|
||||
_revoke_and_purge_generation(stale_items)
|
||||
return count
|
||||
except Exception as e:
|
||||
logger.error("清理孤儿 Job 失败: %s", e, exc_info=True)
|
||||
@@ -130,9 +186,12 @@ def cleanup_stale_pending_tasks(timeout_minutes: int = PENDING_TASK_TIMEOUT_MINU
|
||||
session = SessionLocal()
|
||||
try:
|
||||
repo = SQLAlchemyGenerationTaskRepository(session)
|
||||
count = cleanup_stale_pending_with_session(repo, timeout_minutes)
|
||||
items = cleanup_stale_pending_with_session_ids(repo, timeout_minutes)
|
||||
count = len(items)
|
||||
if count > 0:
|
||||
logger.warning("清理了 %d 个超时的 pending GenerationTask(超过 %d 分钟未处理)", count, timeout_minutes)
|
||||
purged = _revoke_and_purge_stale_messages(items)
|
||||
logger.info("超时 pending 任务对应队列消息撤销/清除完成: %d 条", purged)
|
||||
else:
|
||||
logger.info("无超时 pending GenerationTask 需要清理")
|
||||
return count
|
||||
@@ -143,6 +202,31 @@ def cleanup_stale_pending_tasks(timeout_minutes: int = PENDING_TASK_TIMEOUT_MINU
|
||||
session.close()
|
||||
|
||||
|
||||
def _revoke_and_purge_generation(items: list[tuple[str, str]]) -> int:
|
||||
"""撤销 Job 表孤儿任务(TTS/配音等)的队列消息,队列覆盖全部已知队列。"""
|
||||
biz_ids = [tid for tid, _ in items if tid]
|
||||
celery_ids = [cid for _, cid in items if cid]
|
||||
if not biz_ids and not celery_ids:
|
||||
return 0
|
||||
try:
|
||||
from worker_app.celery_app import celery_app as app
|
||||
from worker_app.core.config import get_settings
|
||||
|
||||
from packages.shared.celery_orphan_guard import revoke_and_purge
|
||||
|
||||
broker_url = get_settings().broker_url
|
||||
return revoke_and_purge(
|
||||
app,
|
||||
broker_url,
|
||||
business_task_ids=biz_ids,
|
||||
celery_task_ids=celery_ids,
|
||||
queue_names=("generation", "transcode", "celery"),
|
||||
)
|
||||
except Exception as e: # noqa: BLE001
|
||||
logger.error("撤销 Job 队列消息失败: %s", e, exc_info=True)
|
||||
return 0
|
||||
|
||||
|
||||
def cleanup_all_stale_tasks(timeout_minutes: int = ORPHAN_TASK_TIMEOUT_MINUTES) -> dict: # pragma: no cover
|
||||
"""统一清理所有超时的孤儿任务。
|
||||
|
||||
|
||||
@@ -25,10 +25,12 @@ def scheduled_cleanup_stale_pending(timeout_minutes: int = PENDING_TASK_TIMEOUT_
|
||||
|
||||
每 5 分钟执行一次(由 celery_app.py 的 beat_schedule 配置),
|
||||
查找所有 status='pending' 且 created_at < NOW() - timeout_minutes
|
||||
的 generation_tasks,批量更新为 failed,释放限流名额。
|
||||
的 generation_tasks,批量更新为 failed,释放限流名额;同时 revoke 并清除
|
||||
Redis 队列中对应的 Celery 消息,杜绝作废消息重投执行(#1714)。
|
||||
|
||||
Args:
|
||||
timeout_minutes: 超时时间(分钟),默认 15 分钟
|
||||
timeout_minutes: 超时时间(分钟),默认 45 分钟(pending 排队阈值放宽,
|
||||
与 running 孤儿 20 分钟区分,避免正常排队任务被误杀)
|
||||
|
||||
Returns:
|
||||
{"cleaned": int}
|
||||
|
||||
@@ -22,6 +22,8 @@ from worker_app.celery_app import celery_app
|
||||
from worker_app.db import SessionLocal
|
||||
from worker_app.tasks.generation_plan_builder import build_error_info as _build_error_info
|
||||
|
||||
from packages.shared.celery_orphan_guard import TERMINAL_STATUS_VALUES
|
||||
|
||||
OUTPUT_WIDTH = 1280
|
||||
OUTPUT_HEIGHT = 720
|
||||
OUTPUT_FPS = 25.0
|
||||
@@ -658,6 +660,34 @@ def generate_video(self, task_id: str) -> dict:
|
||||
finally:
|
||||
_session.close()
|
||||
|
||||
# ── 0. 执行前状态守卫(#1714):任务已被超时清理/孤儿恢复标记为终态时,
|
||||
# 这是作废消息(worker 崩溃前未 ack 的旧消息重投/重复投递),直接丢弃,
|
||||
# 不进入渲染,杜绝 failed→running 非法转换后继续跑产出半成品。
|
||||
if gen_task is not None and gen_task.status.value in TERMINAL_STATUS_VALUES:
|
||||
logger.warning(
|
||||
"[task_id=%s] 任务状态已为 %s,丢弃作废消息,不执行渲染",
|
||||
task_id,
|
||||
gen_task.status.value,
|
||||
)
|
||||
return {
|
||||
"status": "discarded",
|
||||
"task_id": task_id,
|
||||
"reason": f"task already terminal: {gen_task.status.value}",
|
||||
}
|
||||
|
||||
# 标记任务为 running —— 必须成功:状态机非法转换(如 failed→running)说明
|
||||
# 任务已被作废,安全中止,禁止继续执行。
|
||||
if not _update_task_status(task_id, "mark_processing"):
|
||||
logger.error(
|
||||
"[task_id=%s] 标记 running 失败(任务可能已被作废/取消),安全中止,不执行渲染",
|
||||
task_id,
|
||||
)
|
||||
return {
|
||||
"status": "discarded",
|
||||
"task_id": task_id,
|
||||
"reason": "claim failed (invalid state transition)",
|
||||
}
|
||||
|
||||
# 记录接收任务日志
|
||||
if gen_task:
|
||||
gen_task.append_log(
|
||||
@@ -669,8 +699,6 @@ def generate_video(self, task_id: str) -> dict:
|
||||
)
|
||||
_flush_logs(task_id, gen_task)
|
||||
|
||||
# 标记任务为 running
|
||||
_update_task_status(task_id, "mark_processing")
|
||||
_update_task_progress(task_id, 10, "任务启动")
|
||||
|
||||
try:
|
||||
|
||||
@@ -428,6 +428,19 @@ def ingest_asset(job_id: str) -> dict:
|
||||
if job is None:
|
||||
return {"status": "failed", "error": "job not found"}
|
||||
|
||||
# ── 执行前状态守卫(#1714):job 已终态(失败/完成)说明这是作废消息
|
||||
# (超时清理标记 failed 后旧消息重投、或重复投递),直接丢弃不执行,
|
||||
# 避免重复转码、重复回写。processing 是本任务自己第一次置位前的旧消息
|
||||
# 极少见,保守起见也丢弃(processing 的占位由恢复流程处理)。
|
||||
current_status = job.status.value if hasattr(job.status, "value") else str(job.status)
|
||||
if current_status in ("failed", "completed"):
|
||||
logger.warning(
|
||||
"[ingest job_id=%s] 任务状态已为 %s,丢弃作废消息,不执行转码",
|
||||
job_id,
|
||||
current_status,
|
||||
)
|
||||
return {"status": "discarded", "job_id": job_id, "reason": f"job already terminal: {current_status}"}
|
||||
|
||||
# 记录原始 storage_key:HEVC 转码成功后 job.storage_key 会改写为 *_h264,
|
||||
# 而 complete 阶段的占位 asset 始终以原始 key 创建,关联回写必须保留它。
|
||||
original_storage_key = job.storage_key
|
||||
|
||||
@@ -115,6 +115,8 @@ services:
|
||||
APP_VERSION: ${APP_VERSION:-unknown}
|
||||
WORKER_CONCURRENCY: ${WORKER_CONCURRENCY:-4}
|
||||
WORKER_MAX_TASKS_PER_CHILD: ${WORKER_MAX_TASKS_PER_CHILD:-100}
|
||||
# #1714 队列隔离:generation 队列独占 worker(默认并发 2),其余并发给转码
|
||||
GENERATION_CONCURRENCY: ${GENERATION_CONCURRENCY:-2}
|
||||
GENERATED_FILES_DIR: /app/generated
|
||||
GENERATED_FILES_URL_PREFIX: /generated-files
|
||||
PUBLIC_API_BASE_URL: ${PUBLIC_API_BASE_URL:-https://api.xiaoxiajianji.com}
|
||||
@@ -128,11 +130,11 @@ services:
|
||||
# 健康检查配置
|
||||
# 注:celery inspect ping 依赖 broker 连接,在容器内不可靠,改用进程检查
|
||||
healthcheck:
|
||||
test: ["CMD-SHELL", "grep -q celery /proc/1/cmdline || exit 1"]
|
||||
test: ["CMD-SHELL", "pgrep -f 'celery.*worker' | head -n1 >/dev/null 2>&1 || grep -q celery /proc/1/cmdline || exit 1"]
|
||||
interval: 30s
|
||||
timeout: 10s
|
||||
retries: 3
|
||||
start_period: 30s
|
||||
start_period: 40s
|
||||
|
||||
logging: *default-logging
|
||||
|
||||
@@ -140,7 +142,8 @@ services:
|
||||
# 资源限制建议(生产环境建议启用)
|
||||
# =========================================
|
||||
# 注意: Worker 需要处理视频,建议分配更多资源
|
||||
# 并发 4 时需要 4C8G 以上,确保视频渲染不 OOM
|
||||
# #1714 队列隔离后容器内运行 generation + transcode 两个 worker 进程,
|
||||
# 总并发 = WORKER_CONCURRENCY(默认 4),4C8G 以上确保视频渲染不 OOM
|
||||
deploy:
|
||||
resources:
|
||||
limits:
|
||||
|
||||
@@ -146,6 +146,7 @@ docker run -d \
|
||||
-e APP_ENV=production \
|
||||
-e APP_VERSION="$IMAGE_TAG" \
|
||||
-e WORKER_CONCURRENCY="${WORKER_CONCURRENCY:-4}" \
|
||||
-e GENERATION_CONCURRENCY="${GENERATION_CONCURRENCY:-2}" \
|
||||
-e WORKER_MAX_TASKS_PER_CHILD=100 \
|
||||
-e GENERATED_FILES_DIR=/app/generated \
|
||||
-e GENERATED_FILES_URL_PREFIX=/generated-files \
|
||||
@@ -154,7 +155,7 @@ docker run -d \
|
||||
--restart unless-stopped \
|
||||
--cpus 2 \
|
||||
--memory 2g \
|
||||
--health-cmd "sh -c \"grep -q celery /proc/1/cmdline || exit 1\"" \
|
||||
--health-cmd "sh -c \"pgrep -f 'celery.*worker' >/dev/null 2>&1 || grep -q celery /proc/1/cmdline || exit 1\"" \
|
||||
--health-interval 30s \
|
||||
--health-timeout 10s \
|
||||
--health-retries 3 \
|
||||
|
||||
@@ -109,13 +109,14 @@ docker run -d \
|
||||
-e APP_ENV=staging \
|
||||
-e APP_VERSION="$IMAGE_TAG" \
|
||||
-e WORKER_CONCURRENCY="${WORKER_CONCURRENCY:-4}" \
|
||||
-e GENERATION_CONCURRENCY="${GENERATION_CONCURRENCY:-2}" \
|
||||
-e WORKER_MAX_TASKS_PER_CHILD=100 \
|
||||
-e GENERATED_FILES_DIR=/app/generated \
|
||||
-e GENERATED_FILES_URL_PREFIX=/generated-files \
|
||||
-v "$GENERATED_DIR:/app/generated" \
|
||||
--restart unless-stopped \
|
||||
--label com.centurylinklabs.watchtower.enable=true \
|
||||
--health-cmd "sh -c \"grep -q celery /proc/1/cmdline || exit 1\"" \
|
||||
--health-cmd "sh -c \"pgrep -f 'celery.*worker' >/dev/null 2>&1 || grep -q celery /proc/1/cmdline || exit 1\"" \
|
||||
--health-interval 30s \
|
||||
--health-timeout 10s \
|
||||
--health-retries 3 \
|
||||
|
||||
@@ -1,18 +1,63 @@
|
||||
#!/bin/bash
|
||||
# Worker 启动脚本 — 支持 WORKER_CONCURRENCY 环境变量
|
||||
# 未设置时默认 2(保持向后兼容)
|
||||
# Worker 启动脚本 — #1714 队列隔离
|
||||
#
|
||||
# 部署约束:worker 容器单实例(replicas=1),容器内启动两个 celery 进程:
|
||||
# 1. generation-worker:独占消费 generation 队列(用户视频生成,高优先级),
|
||||
# 内嵌 celery beat(-B),定时清理任务只在一个进程里跑,避免重复执行;
|
||||
# 2. transcode-worker:消费 transcode + celery 默认队列(素材转码/分类/查重/
|
||||
# 配音/下载等后台任务)。
|
||||
# 转码队列积压时,generation 队列仍有独立 worker 立即领取视频生成任务。
|
||||
#
|
||||
# 环境变量:
|
||||
# WORKER_CONCURRENCY 总并发槽参考(默认 4);生成 worker 并发默认 2,
|
||||
# 可用 GENERATION_CONCURRENCY 覆盖
|
||||
# GENERATION_CONCURRENCY generation worker 并发(默认 2)
|
||||
# TRANSCODE_CONCURRENCY transcode worker 并发(默认 = WORKER_CONCURRENCY - 2,最小 1)
|
||||
# WORKER_MAX_TASKS_PER_CHILD 每个子进程最大任务数(默认 100)
|
||||
|
||||
set -e
|
||||
|
||||
CONCURRENCY="${WORKER_CONCURRENCY:-2}"
|
||||
CONCURRENCY="${WORKER_CONCURRENCY:-4}"
|
||||
MAX_TASKS="${WORKER_MAX_TASKS_PER_CHILD:-100}"
|
||||
|
||||
# ⚠️ 部署约束:此 Worker 必须且只能运行单实例(replicas=1)
|
||||
# -B 标志嵌入 celery beat,beat 负责定期触发 pending 超时清理等定时任务
|
||||
# 多实例部署会导致每个 Worker 独立运行 Beat,造成定时任务重复执行
|
||||
# 若需横向扩展 Worker,必须将 Beat 拆分为独立服务(celery beat -A worker_app.celery_app)
|
||||
exec celery \
|
||||
GEN_CONCURRENCY="${GENERATION_CONCURRENCY:-2}"
|
||||
if [ -z "$TRANSCODE_CONCURRENCY" ]; then
|
||||
TRANS_CONCURRENCY=$((CONCURRENCY - GEN_CONCURRENCY))
|
||||
if [ "$TRANS_CONCURRENCY" -lt 1 ]; then
|
||||
TRANS_CONCURRENCY=1
|
||||
fi
|
||||
else
|
||||
TRANS_CONCURRENCY="$TRANSCODE_CONCURRENCY"
|
||||
fi
|
||||
|
||||
echo "Starting generation worker (queue=generation, concurrency=$GEN_CONCURRENCY, beat embedded)"
|
||||
celery \
|
||||
-A worker_app.celery_app \
|
||||
worker \
|
||||
--loglevel=info \
|
||||
"-B" \
|
||||
"--concurrency=${CONCURRENCY}"
|
||||
-Q generation \
|
||||
"--concurrency=${GEN_CONCURRENCY}" \
|
||||
"--max-tasks-per-child=${MAX_TASKS}" \
|
||||
-n generation@%h &
|
||||
GEN_PID=$!
|
||||
|
||||
echo "Starting transcode worker (queues=transcode,celery, concurrency=$TRANS_CONCURRENCY)"
|
||||
celery \
|
||||
-A worker_app.celery_app \
|
||||
worker \
|
||||
--loglevel=info \
|
||||
-Q transcode,celery \
|
||||
"--concurrency=${TRANS_CONCURRENCY}" \
|
||||
"--max-tasks-per-child=${MAX_TASKS}" \
|
||||
-n transcode@%h &
|
||||
TRANS_PID=$!
|
||||
|
||||
# 任一进程退出则终止另一个,让容器整体重启(restart: unless-stopped)
|
||||
trap 'echo "Shutting down workers..."; kill -TERM $GEN_PID $TRANS_PID 2>/dev/null || true' TERM INT
|
||||
|
||||
wait -n $GEN_PID $TRANS_PID
|
||||
EXIT_CODE=$?
|
||||
echo "One worker exited (code=$EXIT_CODE), stopping the other..."
|
||||
kill -TERM $GEN_PID $TRANS_PID 2>/dev/null || true
|
||||
exit $EXIT_CODE
|
||||
|
||||
@@ -38,6 +38,7 @@ def _to_domain(model: GenerationTaskModel) -> GenerationTask:
|
||||
bgm_config=dict(getattr(model, "bgm_config", {}) or {}),
|
||||
is_preview=bool(getattr(model, "is_preview", False)),
|
||||
source_task_id=getattr(model, "source_task_id", "") or "",
|
||||
celery_task_id=getattr(model, "celery_task_id", "") or "",
|
||||
output_width=getattr(model, "output_width", 1280) or 1280,
|
||||
output_height=getattr(model, "output_height", 720) or 720,
|
||||
cover_url=getattr(model, "cover_url", "") or "",
|
||||
@@ -82,6 +83,7 @@ class SQLAlchemyGenerationTaskRepository:
|
||||
bgm_config=task.bgm_config or {},
|
||||
is_preview=task.is_preview or False,
|
||||
source_task_id=task.source_task_id or "",
|
||||
celery_task_id=getattr(task, "celery_task_id", "") or "",
|
||||
output_width=task.output_width,
|
||||
output_height=task.output_height,
|
||||
cover_url=task.cover_url or "",
|
||||
@@ -315,6 +317,7 @@ class SQLAlchemyGenerationTaskRepository:
|
||||
if hasattr(model, "is_preview"):
|
||||
model.is_preview = task.is_preview or False
|
||||
model.source_task_id = task.source_task_id or ""
|
||||
model.celery_task_id = getattr(task, "celery_task_id", "") or model.celery_task_id or ""
|
||||
model.output_width = task.output_width
|
||||
model.output_height = task.output_height
|
||||
model.cover_url = task.cover_url or ""
|
||||
@@ -326,12 +329,14 @@ class SQLAlchemyGenerationTaskRepository:
|
||||
def cleanup_stale_running(self, timeout_minutes: int = 10) -> int:
|
||||
"""清理超时未更新的 running 任务(孤儿任务)。
|
||||
|
||||
将 status=running 且 updated_at 超过 timeout_minutes 分钟未更新的任务
|
||||
标记为 failed,error_message 标记为任务执行中断。
|
||||
|
||||
Returns:
|
||||
清理的任务数量
|
||||
清理的任务数量(仅计数,保持旧签名兼容)
|
||||
"""
|
||||
items = self.cleanup_stale_running_with_ids(timeout_minutes)
|
||||
return len(items)
|
||||
|
||||
def cleanup_stale_running_with_ids(self, timeout_minutes: int = 10) -> list[tuple[str, str]]:
|
||||
"""同 cleanup_stale_running,但返回 [(task_id, celery_task_id), ...] 供撤销队列消息。"""
|
||||
from datetime import timedelta
|
||||
|
||||
cutoff = datetime.now(timezone.utc) - timedelta(minutes=timeout_minutes)
|
||||
@@ -344,8 +349,10 @@ class SQLAlchemyGenerationTaskRepository:
|
||||
.all()
|
||||
)
|
||||
if not models:
|
||||
return 0
|
||||
return []
|
||||
result: list[tuple[str, str]] = []
|
||||
for model in models:
|
||||
result.append((model.id, getattr(model, "celery_task_id", "") or ""))
|
||||
model.status = GenerationTaskStatus.FAILED.value
|
||||
model.error_message = "任务执行中断(worker重启/超时)"
|
||||
model.error_info = {
|
||||
@@ -355,43 +362,43 @@ class SQLAlchemyGenerationTaskRepository:
|
||||
}
|
||||
model.completed_at = datetime.now(timezone.utc)
|
||||
self.session.commit()
|
||||
return len(models)
|
||||
return result
|
||||
|
||||
def cleanup_stale_pending(self, timeout_minutes: int = 30) -> int:
|
||||
"""清理超时的 pending 任务(未被 Worker 拉取的任务)。
|
||||
|
||||
全局任务队列有 pending 数量上限,长期卡在 pending 的任务会占满队列,
|
||||
导致新用户无法创建任务。将超时的 pending 任务标记为 failed。
|
||||
|
||||
Args:
|
||||
timeout_minutes: 超时时间(分钟),默认 30 分钟
|
||||
|
||||
Returns:
|
||||
清理的任务数量
|
||||
清理的任务数量(仅计数,保持旧签名兼容)
|
||||
"""
|
||||
items = self.cleanup_stale_pending_with_ids(timeout_minutes)
|
||||
return len(items)
|
||||
|
||||
def cleanup_stale_pending_with_ids(self, timeout_minutes: int = 30) -> list[tuple[str, str]]:
|
||||
"""同 cleanup_stale_pending,但返回 [(task_id, celery_task_id), ...] 供撤销队列消息。"""
|
||||
from datetime import timedelta
|
||||
|
||||
cutoff = datetime.now(timezone.utc) - timedelta(minutes=timeout_minutes)
|
||||
error_info = {
|
||||
"error_type": "PendingTimeout",
|
||||
"message": f"任务在 pending 状态停留超过 {timeout_minutes} 分钟,自动清理",
|
||||
"failed_at": datetime.now(timezone.utc).isoformat(),
|
||||
}
|
||||
count = (
|
||||
models = (
|
||||
self.session.query(GenerationTaskModel)
|
||||
.filter(
|
||||
GenerationTaskModel.status == GenerationTaskStatus.PENDING.value,
|
||||
GenerationTaskModel.created_at < cutoff,
|
||||
)
|
||||
.update(
|
||||
{
|
||||
GenerationTaskModel.status: GenerationTaskStatus.FAILED.value,
|
||||
GenerationTaskModel.error_message: "pending timeout: auto cleanup",
|
||||
GenerationTaskModel.error_info: error_info,
|
||||
GenerationTaskModel.completed_at: datetime.now(timezone.utc),
|
||||
},
|
||||
synchronize_session=False,
|
||||
)
|
||||
.all()
|
||||
)
|
||||
if not models:
|
||||
return []
|
||||
error_info = {
|
||||
"error_type": "PendingTimeout",
|
||||
"message": f"任务在 pending 状态停留超过 {timeout_minutes} 分钟,自动清理",
|
||||
"failed_at": datetime.now(timezone.utc).isoformat(),
|
||||
}
|
||||
result: list[tuple[str, str]] = []
|
||||
for model in models:
|
||||
result.append((model.id, getattr(model, "celery_task_id", "") or ""))
|
||||
model.status = GenerationTaskStatus.FAILED.value
|
||||
model.error_message = "pending timeout: auto cleanup"
|
||||
model.error_info = error_info
|
||||
model.completed_at = datetime.now(timezone.utc)
|
||||
self.session.commit()
|
||||
return count
|
||||
return result
|
||||
|
||||
@@ -19,6 +19,7 @@ class SQLAlchemyIngestJobRepository:
|
||||
result_asset_id=job.result_asset_id,
|
||||
file_hash=job.file_hash,
|
||||
asset_id=job.asset_id or "",
|
||||
celery_task_id=getattr(job, "celery_task_id", "") or "",
|
||||
created_at=job.created_at,
|
||||
updated_at=job.updated_at,
|
||||
)
|
||||
@@ -40,6 +41,7 @@ class SQLAlchemyIngestJobRepository:
|
||||
result_asset_id=model.result_asset_id,
|
||||
file_hash=model.file_hash or "",
|
||||
asset_id=getattr(model, "asset_id", "") or "",
|
||||
celery_task_id=getattr(model, "celery_task_id", "") or "",
|
||||
created_at=model.created_at,
|
||||
updated_at=model.updated_at,
|
||||
)
|
||||
@@ -59,6 +61,9 @@ class SQLAlchemyIngestJobRepository:
|
||||
model.storage_key = job.storage_key
|
||||
if job.asset_id:
|
||||
model.asset_id = job.asset_id
|
||||
celery_tid = getattr(job, "celery_task_id", "")
|
||||
if celery_tid:
|
||||
model.celery_task_id = celery_tid
|
||||
model.updated_at = job.updated_at
|
||||
self.session.commit()
|
||||
return job
|
||||
|
||||
@@ -246,6 +246,7 @@ class IngestJobModel(Base):
|
||||
result_asset_id = Column(String(36), nullable=False, default="")
|
||||
file_hash = Column(String(64), nullable=True, index=True)
|
||||
asset_id = Column(String(36), nullable=False, default="", index=True)
|
||||
celery_task_id = Column(String(64), nullable=False, default="", server_default="")
|
||||
created_at = Column(DateTime, nullable=False, default=lambda: datetime.now(timezone.utc))
|
||||
updated_at = Column(DateTime, nullable=False, default=lambda: datetime.now(timezone.utc))
|
||||
|
||||
@@ -298,6 +299,7 @@ class GenerationTaskModel(Base):
|
||||
resolution = Column(String(20), nullable=False, default="")
|
||||
is_preview = Column(Boolean, nullable=False, default=False, index=True)
|
||||
source_task_id = Column(String(32), nullable=False, default="", index=True)
|
||||
celery_task_id = Column(String(64), nullable=False, default="", server_default="")
|
||||
output_width = Column(Integer, nullable=False, default=1280)
|
||||
output_height = Column(Integer, nullable=False, default=720)
|
||||
cover_url = Column(String(1000), nullable=False, default="")
|
||||
|
||||
@@ -13,6 +13,7 @@ class SubmitIngestJobCommand:
|
||||
storage_key: str
|
||||
file_hash: str = ""
|
||||
asset_id: str = ""
|
||||
celery_task_id: str = ""
|
||||
|
||||
|
||||
class SubmitIngestJobUseCase:
|
||||
@@ -26,5 +27,6 @@ class SubmitIngestJobUseCase:
|
||||
storage_key=command.storage_key,
|
||||
file_hash=command.file_hash,
|
||||
asset_id=command.asset_id,
|
||||
celery_task_id=command.celery_task_id,
|
||||
)
|
||||
return self.ingest_job_repository.create(job)
|
||||
|
||||
@@ -270,6 +270,7 @@ class IngestJob:
|
||||
result_asset_id: str = ""
|
||||
file_hash: str = ""
|
||||
asset_id: str = ""
|
||||
celery_task_id: str = ""
|
||||
created_at: datetime = field(default_factory=lambda: datetime.now(timezone.utc))
|
||||
updated_at: datetime = field(default_factory=lambda: datetime.now(timezone.utc))
|
||||
|
||||
@@ -281,6 +282,7 @@ class IngestJob:
|
||||
storage_key: str,
|
||||
file_hash: str = "",
|
||||
asset_id: str = "",
|
||||
celery_task_id: str = "",
|
||||
) -> "IngestJob":
|
||||
if not project_id.strip():
|
||||
raise ValueError("project_id 不能为空")
|
||||
@@ -295,4 +297,5 @@ class IngestJob:
|
||||
storage_key=storage_key.strip(),
|
||||
file_hash=file_hash.strip(),
|
||||
asset_id=asset_id.strip(),
|
||||
celery_task_id=celery_task_id.strip(),
|
||||
)
|
||||
|
||||
@@ -117,6 +117,7 @@ class GenerationTask:
|
||||
bgm_config: dict = field(default_factory=dict)
|
||||
is_preview: bool = False
|
||||
source_task_id: str = ""
|
||||
celery_task_id: str = ""
|
||||
output_width: int = 1280
|
||||
output_height: int = 720
|
||||
cover_url: str = ""
|
||||
|
||||
@@ -0,0 +1,222 @@
|
||||
"""孤儿任务消息撤销与执行前状态守卫(API / Worker 共享)。
|
||||
|
||||
#1714 / #1710 缺陷修复:超时清理/孤儿恢复把 DB 任务标记为 failed/cancelled
|
||||
后,Redis 队列里对应的 Celery 消息仍然存在;worker 重启或重新拉取时该消息
|
||||
被再次执行,状态机抛「非法状态转换: failed → running」,旧实现打印 ERROR 后
|
||||
继续跑,最终产出半成品。
|
||||
|
||||
防御两道:
|
||||
1. 清理任务标 failed 时,调用 revoke_and_purge() 撤销(celery revoke 广播,
|
||||
通知在线 worker 丢弃)并直接扫描 Redis 队列移除消息体(worker 下线期间
|
||||
队列中的消息 revoke 广播收不到,必须物理移除);
|
||||
2. 任务真正开始业务逻辑前,调用 ensure_task_claimable() 校验 DB 状态,
|
||||
非 pending 的消息直接丢弃(抛 StaleTaskDiscarded,task 捕获后安全返回,
|
||||
不进入渲染/转码,不产出半成品)。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import base64
|
||||
import json
|
||||
import logging
|
||||
from collections.abc import Callable, Iterable
|
||||
from typing import Any
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
class StaleTaskDiscarded(Exception):
|
||||
"""任务消息已作废(DB 中任务已是终态),应安全中止、丢弃消息。"""
|
||||
|
||||
def __init__(self, task_id: str, status: str):
|
||||
self.task_id = task_id
|
||||
self.status = status
|
||||
super().__init__(f"任务 {task_id} 已是终态 {status},丢弃重复/作废消息")
|
||||
|
||||
|
||||
# 终态状态值集合:处于这些状态的任务消息一律不执行
|
||||
TERMINAL_STATUS_VALUES = frozenset({"failed", "cancelled", "completed"})
|
||||
|
||||
|
||||
def ensure_task_claimable(
|
||||
task_id: str,
|
||||
get_status: Callable[[str], str | None],
|
||||
*,
|
||||
task_label: str = "任务",
|
||||
) -> str:
|
||||
"""执行前守卫:任务必须处于可领取状态(pending)。
|
||||
|
||||
Args:
|
||||
task_id: 业务任务 ID
|
||||
get_status: 回调,返回 DB 中任务当前状态字符串;返回 None 表示任务不存在
|
||||
task_label: 日志用任务类型名
|
||||
|
||||
Returns:
|
||||
当前状态字符串(pending);任务不存在时返回空串(由调用方处理 not found)
|
||||
|
||||
Raises:
|
||||
StaleTaskDiscarded: 任务已是终态(failed/cancelled/completed),消息必须丢弃
|
||||
"""
|
||||
status = get_status(task_id)
|
||||
if status is None:
|
||||
return ""
|
||||
if status in TERMINAL_STATUS_VALUES:
|
||||
logger.warning("[%s] task_id=%s 状态已为 %s,消息作废,丢弃不执行", task_label, task_id, status)
|
||||
raise StaleTaskDiscarded(task_id, status)
|
||||
return status
|
||||
|
||||
|
||||
def _extract_business_ids(raw: bytes) -> tuple[str | None, str | None]:
|
||||
"""从 Redis 中的 Celery 消息提取 (celery 消息 ID, 业务任务 ID)。
|
||||
|
||||
Redis transport 存储格式为 JSON 信封:
|
||||
{"body": base64(json), "headers": {"id": <celery id>, "task": <name>, ...}, ...}
|
||||
body 解码后 Celery task 协议为 [args, kwargs, embed];
|
||||
generate_video / ingest_asset 均以 args=[业务任务ID] 投递。
|
||||
|
||||
无法解析时返回 (None, None)(保守保留该消息,绝不误删)。
|
||||
"""
|
||||
try:
|
||||
envelope = json.loads(raw)
|
||||
celery_id = None
|
||||
headers = envelope.get("headers") or {}
|
||||
if isinstance(headers, dict):
|
||||
celery_id = headers.get("id")
|
||||
body = envelope.get("body")
|
||||
if not body:
|
||||
return celery_id, None
|
||||
decoded = base64.b64decode(body)
|
||||
payload = json.loads(decoded)
|
||||
# 两种 body 形态:
|
||||
# 1. 标准 Celery task 消息:[args, kwargs, embed] 三元组 → 业务 ID 在 payload[0][0]
|
||||
# 2. 裸 producer 发布:body 即 args 数组 ["biz-id"] → 业务 ID 在 payload[0]
|
||||
args = None
|
||||
if isinstance(payload, dict):
|
||||
args = payload.get("args")
|
||||
elif isinstance(payload, (list, tuple)) and payload:
|
||||
first = payload[0]
|
||||
if isinstance(first, (list, tuple)):
|
||||
args = first # 三元组:[args, kwargs, embed]
|
||||
else:
|
||||
args = payload # body 本身就是 args
|
||||
if isinstance(args, (list, tuple)) and args and args[0] is not None:
|
||||
return celery_id, str(args[0])
|
||||
return celery_id, None
|
||||
except Exception:
|
||||
return None, None
|
||||
|
||||
|
||||
def purge_stale_messages_from_queues(
|
||||
broker_url: str,
|
||||
queue_names: Iterable[str],
|
||||
business_task_ids: Iterable[str] = (),
|
||||
celery_task_ids: Iterable[str] = (),
|
||||
) -> int:
|
||||
"""扫描 Redis 队列,移除作废任务的待消费消息。
|
||||
|
||||
同时按业务任务 ID(消息 args[0])和 celery 消息 ID(headers.id)匹配,
|
||||
任一命中即移除。未命中或无法解析的消息原样保留(保持相对顺序)。
|
||||
|
||||
Returns:
|
||||
实际移除的消息条数
|
||||
"""
|
||||
biz_ids = {bid for bid in business_task_ids if bid}
|
||||
msg_ids = {mid for mid in celery_task_ids if mid}
|
||||
if not biz_ids and not msg_ids:
|
||||
return 0
|
||||
|
||||
try:
|
||||
import redis
|
||||
except ImportError:
|
||||
logger.warning("redis-py 不可用,跳过队列消息清理")
|
||||
return 0
|
||||
|
||||
try:
|
||||
client = redis.Redis.from_url(broker_url)
|
||||
client.ping()
|
||||
except Exception as e:
|
||||
logger.warning("连接 Redis 清理作废消息失败: %s", e)
|
||||
return 0
|
||||
|
||||
removed_total = 0
|
||||
try:
|
||||
for queue in queue_names:
|
||||
removed_total += _purge_one_queue(client, queue, biz_ids, msg_ids)
|
||||
finally:
|
||||
try:
|
||||
client.close()
|
||||
except Exception:
|
||||
pass
|
||||
if removed_total:
|
||||
logger.info(
|
||||
"从 Redis 队列移除 %d 条作废消息(biz=%s, celery=%s)",
|
||||
removed_total,
|
||||
sorted(biz_ids),
|
||||
sorted(msg_ids),
|
||||
)
|
||||
return removed_total
|
||||
|
||||
|
||||
def _purge_one_queue(client: Any, queue_name: str, biz_ids: set[str], msg_ids: set[str]) -> int:
|
||||
try:
|
||||
raw_messages = client.lrange(queue_name, 0, -1)
|
||||
except Exception as e:
|
||||
logger.warning("读取队列 %s 失败: %s", queue_name, e)
|
||||
return 0
|
||||
if not raw_messages:
|
||||
return 0
|
||||
|
||||
keep: list[bytes] = []
|
||||
removed = 0
|
||||
for raw in raw_messages:
|
||||
celery_id, biz_id = _extract_business_ids(raw)
|
||||
hit = (biz_id is not None and biz_id in biz_ids) or (celery_id is not None and celery_id in msg_ids)
|
||||
if hit:
|
||||
removed += 1
|
||||
continue
|
||||
keep.append(raw)
|
||||
|
||||
if removed:
|
||||
try:
|
||||
pipe = client.pipeline()
|
||||
pipe.delete(queue_name)
|
||||
if keep:
|
||||
pipe.rpush(queue_name, *keep)
|
||||
pipe.execute()
|
||||
except Exception as e:
|
||||
logger.warning("重写队列 %s 失败: %s", queue_name, e)
|
||||
return 0
|
||||
return removed
|
||||
|
||||
|
||||
def revoke_and_purge(
|
||||
celery_app: Any,
|
||||
broker_url: str,
|
||||
business_task_ids: Iterable[str] = (),
|
||||
celery_task_ids: Iterable[str] = (),
|
||||
*,
|
||||
queue_names: Iterable[str] = ("generation", "transcode", "celery"),
|
||||
) -> int:
|
||||
"""撤销作废任务:revoke 广播(在线 worker)+ 物理清理 Redis 队列消息。
|
||||
|
||||
Args:
|
||||
celery_app: Celery app 实例(worker 端 worker_app.celery_app.celery_app)
|
||||
broker_url: Redis broker URL
|
||||
business_task_ids: 业务任务 ID(generation_tasks.id / ingest_jobs.id)
|
||||
celery_task_ids: 入队时记录的 celery 消息 ID
|
||||
queue_names: 需要扫描清理的队列名
|
||||
|
||||
Returns:
|
||||
从队列中实际移除的消息条数
|
||||
"""
|
||||
for tid in celery_task_ids:
|
||||
if not tid:
|
||||
continue
|
||||
try:
|
||||
celery_app.control.revoke(tid)
|
||||
except Exception as e:
|
||||
logger.warning("revoke celery 消息 %s 失败: %s", tid, e)
|
||||
|
||||
return purge_stale_messages_from_queues(
|
||||
broker_url, queue_names, business_task_ids=business_task_ids, celery_task_ids=celery_task_ids
|
||||
)
|
||||
@@ -0,0 +1,58 @@
|
||||
"""Celery 队列定义与路由配置(API / Worker 共享)。
|
||||
|
||||
#1714 队列隔离:用户等待的视频生成任务路由到高优先级 `generation` 队列,
|
||||
由专用 worker 进程独占消费;素材入库/转码等后台批量任务路由到 `transcode`
|
||||
队列;其余杂项任务走默认 `celery` 队列。转码队列积压时,视频生成任务
|
||||
仍能被 generation worker 立即领取执行,不会排队。
|
||||
|
||||
队列说明:
|
||||
- generation: 用户提交的视频生成/预览渲染(延迟敏感,资源消耗大)
|
||||
- transcode: 素材入库(HEVC 转码)、AI 分类、素材查重(批量、可排队)
|
||||
- celery(默认): 配音、语音、下载缩略图、定时清理等杂项
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from kombu import Queue
|
||||
|
||||
# ── 队列名常量(生产端与消费端共用,禁止拼写漂移) ──
|
||||
QUEUE_GENERATION = "generation"
|
||||
QUEUE_TRANSCODE = "transcode"
|
||||
QUEUE_DEFAULT = "celery"
|
||||
|
||||
# Worker 消费的队列列表(顺序即优先级:高优队列排在前面)
|
||||
WORKER_QUEUES = (QUEUE_GENERATION, QUEUE_TRANSCODE, QUEUE_DEFAULT)
|
||||
|
||||
# 队列声明:持久化队列,broker 重启不丢消息
|
||||
task_queues = (
|
||||
Queue(QUEUE_GENERATION, routing_key=QUEUE_GENERATION, durable=True),
|
||||
Queue(QUEUE_TRANSCODE, routing_key=QUEUE_TRANSCODE, durable=True),
|
||||
Queue(QUEUE_DEFAULT, routing_key=QUEUE_DEFAULT, durable=True),
|
||||
)
|
||||
|
||||
# ── 任务路由表:task name → 队列 ──
|
||||
# 键支持 celery 标准通配符。
|
||||
task_routes = {
|
||||
# 高优先级:用户等待的视频生成
|
||||
"worker.generate_video": {"queue": QUEUE_GENERATION},
|
||||
# 后台批量:素材入库/转码 + AI 分类 + 素材查重,积压不影响生成
|
||||
"worker.ingest_asset": {"queue": QUEUE_TRANSCODE},
|
||||
"worker.classify_asset": {"queue": QUEUE_TRANSCODE},
|
||||
"worker.process_duplication_check": {"queue": QUEUE_TRANSCODE},
|
||||
"worker.check_duplicate": {"queue": QUEUE_TRANSCODE},
|
||||
}
|
||||
|
||||
# 生成任务的预取数:渲染是长任务,预取 1 避免任务被某个 worker 占住不调度
|
||||
GENERATION_WORKER_PREFETCH_MULTIPLIER = 1
|
||||
|
||||
|
||||
def apply_queue_settings(app) -> None:
|
||||
"""把队列隔离配置应用到 Celery app(API 生产端与 Worker 消费端都要调用)。
|
||||
|
||||
配置 task_queues / task_routes / task_default_queue。生产端靠 task_routes
|
||||
把消息投递到对应队列;消费端靠 task_queues 声明自己消费哪些队列
|
||||
(实际消费集由启动参数 -Q 控制)。
|
||||
"""
|
||||
app.conf.task_queues = task_queues
|
||||
app.conf.task_routes = task_routes
|
||||
app.conf.task_default_queue = QUEUE_DEFAULT
|
||||
+1
-1
@@ -14,4 +14,4 @@ Write-Host "`n启动 Celery Worker..." -ForegroundColor Yellow
|
||||
Write-Host "监听任务队列: Redis (47.98.113.167:6379)" -ForegroundColor Cyan
|
||||
Write-Host "`n按 Ctrl+C 停止服务`n" -ForegroundColor Gray
|
||||
|
||||
celery -A celery_app worker --loglevel=info --pool=solo
|
||||
celery -A celery_app worker --loglevel=info --pool=solo -Q generation,transcode,celery
|
||||
|
||||
@@ -0,0 +1,176 @@
|
||||
"""#1714 队列隔离 + 作废消息清除 单元测试。
|
||||
|
||||
覆盖:
|
||||
1. task_routes:generate_video → generation,ingest_asset/classify/duplication → transcode
|
||||
2. purge_stale_messages_from_queues:Redis 队列中作废任务消息被物理移除,未命中保留
|
||||
3. revoke_and_purge:revoke 广播 + 队列清理同时生效
|
||||
4. ensure_task_claimable:终态任务抛 StaleTaskDiscarded,pending 放行
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from unittest.mock import MagicMock
|
||||
|
||||
import pytest
|
||||
from celery import Celery
|
||||
|
||||
from packages.shared.celery_orphan_guard import (
|
||||
StaleTaskDiscarded,
|
||||
_extract_business_ids,
|
||||
ensure_task_claimable,
|
||||
purge_stale_messages_from_queues,
|
||||
revoke_and_purge,
|
||||
)
|
||||
from packages.shared.celery_queues import (
|
||||
QUEUE_GENERATION,
|
||||
QUEUE_TRANSCODE,
|
||||
apply_queue_settings,
|
||||
task_routes,
|
||||
)
|
||||
|
||||
BROKER_URL = "redis://localhost:6379/15"
|
||||
TEST_QUEUES = ("_test_gen_q", "_test_transcode_q")
|
||||
|
||||
|
||||
# ── 1. 路由表 ──────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
def test_routes_send_generation_to_generation_queue():
|
||||
assert task_routes["worker.generate_video"]["queue"] == QUEUE_GENERATION
|
||||
|
||||
|
||||
def test_routes_send_ingest_to_transcode_queue():
|
||||
assert task_routes["worker.ingest_asset"]["queue"] == QUEUE_TRANSCODE
|
||||
assert task_routes["worker.classify_asset"]["queue"] == QUEUE_TRANSCODE
|
||||
assert task_routes["worker.process_duplication_check"]["queue"] == QUEUE_TRANSCODE
|
||||
assert task_routes["worker.check_duplicate"]["queue"] == QUEUE_TRANSCODE
|
||||
|
||||
|
||||
def test_apply_queue_settings_configures_celery_app():
|
||||
app = Celery("test-routes")
|
||||
apply_queue_settings(app)
|
||||
queue_names = {q.name for q in app.conf.task_queues}
|
||||
assert queue_names == {"generation", "transcode", "celery"}
|
||||
assert app.conf.task_default_queue == "celery"
|
||||
|
||||
|
||||
# ── Redis 队列消息清理(需要本地 redis;不可用时 skip) ─────────────────
|
||||
|
||||
|
||||
def _redis_available() -> bool:
|
||||
try:
|
||||
import redis
|
||||
|
||||
return bool(redis.Redis.from_url(BROKER_URL).ping())
|
||||
except Exception:
|
||||
return False
|
||||
|
||||
|
||||
@pytest.fixture()
|
||||
def redis_client():
|
||||
import redis
|
||||
|
||||
client = redis.Redis.from_url(BROKER_URL)
|
||||
for q in TEST_QUEUES:
|
||||
client.delete(q)
|
||||
yield client
|
||||
for q in TEST_QUEUES:
|
||||
client.delete(q)
|
||||
|
||||
|
||||
def _publish(app: Celery, queue: str, celery_id: str, business_id: str) -> None:
|
||||
from kombu import Queue
|
||||
from kombu.pools import producers
|
||||
|
||||
with app.connection_for_write() as conn:
|
||||
with producers[conn].acquire(block=True) as prod:
|
||||
prod.publish(
|
||||
(business_id,),
|
||||
exchange="",
|
||||
routing_key=queue,
|
||||
serializer="json",
|
||||
headers={"id": celery_id, "task": "worker.generate_video"},
|
||||
retry=False,
|
||||
delivery_mode=1,
|
||||
declare=[Queue(queue, routing_key=queue, durable=False)],
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.skipif(not _redis_available(), reason="本地 redis 不可用")
|
||||
def test_purge_removes_stale_business_message_and_keeps_others(redis_client):
|
||||
app = Celery("test-purge")
|
||||
app.conf.broker_url = BROKER_URL
|
||||
_publish(app, TEST_QUEUES[0], "celery-1", "task-KEEP-A")
|
||||
_publish(app, TEST_QUEUES[0], "celery-2", "task-STALE-B")
|
||||
_publish(app, TEST_QUEUES[0], "celery-3", "task-KEEP-C")
|
||||
_publish(app, TEST_QUEUES[1], "celery-4", "task-STALE-B") # 同一业务任务在转码队列?不应出现但验证全队列扫描
|
||||
|
||||
removed = purge_stale_messages_from_queues(BROKER_URL, TEST_QUEUES, business_task_ids={"task-STALE-B"})
|
||||
assert removed == 2
|
||||
|
||||
remaining = []
|
||||
for raw in redis_client.lrange(TEST_QUEUES[0], 0, -1):
|
||||
_celery_id, biz_id = _extract_business_ids(raw)
|
||||
remaining.append(biz_id)
|
||||
assert set(remaining) == {"task-KEEP-A", "task-KEEP-C"}
|
||||
assert redis_client.llen(TEST_QUEUES[1]) == 0
|
||||
|
||||
|
||||
@pytest.mark.skipif(not _redis_available(), reason="本地 redis 不可用")
|
||||
def test_purge_matches_by_celery_message_id(redis_client):
|
||||
app = Celery("test-purge-msg-id")
|
||||
app.conf.broker_url = BROKER_URL
|
||||
_publish(app, TEST_QUEUES[0], "celery-stale-id", "task-X")
|
||||
_publish(app, TEST_QUEUES[0], "celery-good-id", "task-Y")
|
||||
|
||||
removed = purge_stale_messages_from_queues(BROKER_URL, TEST_QUEUES, celery_task_ids={"celery-stale-id"})
|
||||
assert removed == 1
|
||||
assert redis_client.llen(TEST_QUEUES[0]) == 1
|
||||
|
||||
|
||||
@pytest.mark.skipif(not _redis_available(), reason="本地 redis 不可用")
|
||||
def test_revoke_and_purge_calls_control_revoke(redis_client):
|
||||
app = Celery("test-revoke")
|
||||
app.conf.broker_url = BROKER_URL
|
||||
app.control = MagicMock()
|
||||
_publish(app, TEST_QUEUES[0], "celery-revoke-1", "task-R")
|
||||
|
||||
removed = revoke_and_purge(
|
||||
app,
|
||||
BROKER_URL,
|
||||
business_task_ids={"task-R"},
|
||||
celery_task_ids={"celery-revoke-1"},
|
||||
queue_names=TEST_QUEUES,
|
||||
)
|
||||
assert removed == 1
|
||||
app.control.revoke.assert_called_once_with("celery-revoke-1")
|
||||
|
||||
|
||||
def test_purge_empty_ids_is_noop():
|
||||
assert purge_stale_messages_from_queues(BROKER_URL, TEST_QUEUES) == 0
|
||||
|
||||
|
||||
# ── 2. 执行前状态守卫 ──────────────────────────────────────────────────
|
||||
|
||||
|
||||
def test_guard_allows_pending():
|
||||
status = ensure_task_claimable("t1", lambda _id: "pending", task_label="generation")
|
||||
assert status == "pending"
|
||||
|
||||
|
||||
def test_guard_rejects_failed():
|
||||
with pytest.raises(StaleTaskDiscarded) as exc:
|
||||
ensure_task_claimable("t2", lambda _id: "failed", task_label="generation")
|
||||
assert exc.value.task_id == "t2"
|
||||
assert exc.value.status == "failed"
|
||||
|
||||
|
||||
def test_guard_rejects_cancelled_and_completed():
|
||||
with pytest.raises(StaleTaskDiscarded):
|
||||
ensure_task_claimable("t3", lambda _id: "cancelled")
|
||||
with pytest.raises(StaleTaskDiscarded):
|
||||
ensure_task_claimable("t4", lambda _id: "completed")
|
||||
|
||||
|
||||
def test_guard_missing_task_returns_empty():
|
||||
assert ensure_task_claimable("t5", lambda _id: None) == ""
|
||||
@@ -0,0 +1,57 @@
|
||||
"""#1714:入队成功后 celery 消息 ID 必须持久化到任务行(供清理时 revoke)。"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import sys
|
||||
from pathlib import Path
|
||||
from unittest.mock import MagicMock
|
||||
|
||||
sys.path.insert(0, str(Path(__file__).resolve().parents[2] / "apps" / "api"))
|
||||
|
||||
from app.core import task_enqueue # noqa: E402
|
||||
|
||||
|
||||
class _FakeTask:
|
||||
def __init__(self):
|
||||
self.id = "task-enqueue-1"
|
||||
self.status = "pending"
|
||||
self.celery_task_id = ""
|
||||
|
||||
def mark_failed(self, msg): # noqa: ARG002
|
||||
self.status = "failed"
|
||||
|
||||
|
||||
class _FakeRepo:
|
||||
def __init__(self):
|
||||
self.updated = None
|
||||
|
||||
def count_pending_total(self):
|
||||
return 0
|
||||
|
||||
def count_pending_by_user(self, user_id): # noqa: ARG002
|
||||
return 0
|
||||
|
||||
def update(self, task):
|
||||
self.updated = task
|
||||
return task
|
||||
|
||||
|
||||
def test_safe_enqueue_persists_celery_message_id(monkeypatch):
|
||||
fake_result = MagicMock()
|
||||
fake_result.id = "celery-msg-id-enqueue-999"
|
||||
mock_celery = MagicMock()
|
||||
mock_celery.send_task.return_value = fake_result
|
||||
monkeypatch.setattr(task_enqueue, "celery_app", mock_celery)
|
||||
|
||||
task = _FakeTask()
|
||||
repo = _FakeRepo()
|
||||
|
||||
ok = task_enqueue.safe_enqueue_generation_task(task, repo, user_id="u1")
|
||||
assert ok is True
|
||||
# celery_task_id 已持久化
|
||||
assert task.celery_task_id == "celery-msg-id-enqueue-999"
|
||||
assert repo.updated is task
|
||||
mock_celery.send_task.assert_called_once()
|
||||
args, kwargs = mock_celery.send_task.call_args
|
||||
assert args[0] == "worker.generate_video"
|
||||
assert kwargs.get("args") == [task.id]
|
||||
@@ -0,0 +1,360 @@
|
||||
"""#1714:孤儿消息撤销/清理逻辑测试(mock redis,CI 无真实 redis 时也产生覆盖)。
|
||||
|
||||
覆盖 packages/shared/celery_orphan_guard.py:
|
||||
- _extract_business_ids:三元组 body / 裸 args body / dict args / headers 提取 /
|
||||
无 body / 坏 JSON / 坏 base64 / 空 args
|
||||
- _purge_one_queue:biz id 命中、celery id 命中、未命中保序(重写 rpush)、
|
||||
lrange 异常、重写异常、空队列
|
||||
- purge_stale_messages_from_queues:空 ids 早退、redis 未安装、连接失败、
|
||||
正常清理并 close
|
||||
- revoke_and_purge:revoke 逐消息调用、revoke 异常不阻断、空 id 跳过
|
||||
- ensure_task_claimable:任务不存在返回空串、终态抛错、pending 放行
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import base64
|
||||
import json
|
||||
import sys
|
||||
import types
|
||||
from pathlib import Path
|
||||
from unittest.mock import MagicMock
|
||||
|
||||
import pytest
|
||||
|
||||
sys.path.insert(0, str(Path(__file__).resolve().parents[2] / "apps" / "api"))
|
||||
|
||||
from packages.shared import celery_orphan_guard as guard # noqa: E402
|
||||
|
||||
|
||||
def _envelope(celery_id: str | None, body_payload) -> bytes:
|
||||
"""构造 Redis transport 存储的 celery 消息(JSON 信封)。"""
|
||||
if body_payload is None:
|
||||
body = None
|
||||
else:
|
||||
body = base64.b64encode(json.dumps(body_payload).encode()).decode()
|
||||
envelope = {"body": body, "headers": {"id": celery_id, "task": "worker.generate_video"}}
|
||||
return json.dumps(envelope).encode()
|
||||
|
||||
|
||||
# ── _extract_business_ids ───────────────────────────────────────────────
|
||||
|
||||
|
||||
def test_extract_ids_standard_tuple_body():
|
||||
raw = _envelope("celery-1", [["biz-task-1"], {}, {"callbacks": None}])
|
||||
assert guard._extract_business_ids(raw) == ("celery-1", "biz-task-1")
|
||||
|
||||
|
||||
def test_extract_ids_bare_args_body():
|
||||
raw = _envelope("celery-2", ["biz-task-2"])
|
||||
assert guard._extract_business_ids(raw) == ("celery-2", "biz-task-2")
|
||||
|
||||
|
||||
def test_extract_ids_dict_body_with_args():
|
||||
raw = _envelope("celery-3", {"args": ["biz-task-3"], "kwargs": {}})
|
||||
assert guard._extract_business_ids(raw) == ("celery-3", "biz-task-3")
|
||||
|
||||
|
||||
def test_extract_ids_non_dict_headers_returns_celery_id_none():
|
||||
raw = json.dumps({"body": base64.b64encode(json.dumps([["biz-4"]]).encode()).decode(), "headers": "x"}).encode()
|
||||
celery_id, biz_id = guard._extract_business_ids(raw)
|
||||
assert celery_id is None
|
||||
assert biz_id == "biz-4"
|
||||
|
||||
|
||||
def test_extract_ids_no_body_returns_celery_id_only():
|
||||
raw = json.dumps({"headers": {"id": "celery-5"}}).encode()
|
||||
assert guard._extract_business_ids(raw) == ("celery-5", None)
|
||||
|
||||
|
||||
def test_extract_ids_empty_args_returns_no_biz_id():
|
||||
raw = _envelope("celery-6", [[], {}, {}])
|
||||
assert guard._extract_business_ids(raw) == ("celery-6", None)
|
||||
|
||||
|
||||
def test_extract_ids_args_first_none_returns_no_biz_id():
|
||||
raw = _envelope("celery-7", [[None], {}, {}])
|
||||
assert guard._extract_business_ids(raw) == ("celery-7", None)
|
||||
|
||||
|
||||
def test_extract_ids_bad_json_returns_none_none():
|
||||
assert guard._extract_business_ids(b"not-json{") == (None, None)
|
||||
|
||||
|
||||
def test_extract_ids_bad_base64_returns_none_none():
|
||||
raw = json.dumps({"body": "!!!not-base64!!!", "headers": {"id": "c"}}).encode()
|
||||
assert guard._extract_business_ids(raw) == (None, None)
|
||||
|
||||
|
||||
def test_extract_ids_int_arg_coerced_to_str():
|
||||
raw = _envelope("celery-9", [[12345], {}, {}])
|
||||
celery_id, biz_id = guard._extract_business_ids(raw)
|
||||
assert celery_id == "celery-9"
|
||||
assert biz_id == "12345"
|
||||
|
||||
|
||||
# ── _purge_one_queue ────────────────────────────────────────────────────
|
||||
|
||||
|
||||
def _queue_with_messages(*payloads: bytes):
|
||||
"""返回 list-backed mock redis client(记录当前队列内容)。"""
|
||||
client = MagicMock()
|
||||
store: dict[str, list[bytes]] = {"q": list(payloads)}
|
||||
|
||||
def lrange(name, start, end): # noqa: ARG001
|
||||
return list(store.get(name, []))
|
||||
|
||||
client.lrange.side_effect = lrange
|
||||
|
||||
pipe = MagicMock()
|
||||
pipe.delete.side_effect = lambda name: store.pop(name, None)
|
||||
pipe.rpush.side_effect = lambda name, *items: store.setdefault(name, []).extend(items)
|
||||
client.pipeline.return_value = pipe
|
||||
return client, store, pipe
|
||||
|
||||
|
||||
def test_purge_one_queue_removes_by_biz_id_and_keeps_order():
|
||||
stale = _envelope("c-stale", [["biz-stale"], {}, {}])
|
||||
keep1 = _envelope("c-keep-1", [["biz-keep-1"], {}, {}])
|
||||
keep2 = _envelope("c-keep-2", [["biz-keep-2"], {}, {}])
|
||||
client, store, pipe = _queue_with_messages(keep1, stale, keep2)
|
||||
|
||||
removed = guard._purge_one_queue(client, "q", {"biz-stale"}, set())
|
||||
assert removed == 1
|
||||
# 队列被 delete + rpush 重写,未命中消息保持相对顺序
|
||||
pipe.delete.assert_called_once_with("q")
|
||||
pipe.rpush.assert_called_once()
|
||||
args, _ = pipe.rpush.call_args
|
||||
assert args[0] == "q"
|
||||
assert list(args[1:]) == [keep1, keep2]
|
||||
pipe.execute.assert_called_once()
|
||||
|
||||
|
||||
def test_purge_one_queue_removes_by_celery_message_id():
|
||||
stale = _envelope("celery-xyz", [["biz-whatever"], {}, {}])
|
||||
keep = _envelope("celery-aaa", [["biz-keep"], {}, {}])
|
||||
client, store, pipe = _queue_with_messages(stale, keep)
|
||||
|
||||
removed = guard._purge_one_queue(client, "q", set(), {"celery-xyz"})
|
||||
assert removed == 1
|
||||
args, _ = pipe.rpush.call_args
|
||||
assert list(args[1:]) == [keep]
|
||||
|
||||
|
||||
def test_purge_one_queue_no_hit_no_rewrite():
|
||||
msg1 = _envelope("c1", [["b1"], {}, {}])
|
||||
msg2 = _envelope("c2", [["b2"], {}, {}])
|
||||
client, store, pipe = _queue_with_messages(msg1, msg2)
|
||||
|
||||
removed = guard._purge_one_queue(client, "q", {"other"}, {"other-c"})
|
||||
assert removed == 0
|
||||
# 没有命中:不重写队列
|
||||
pipe.delete.assert_not_called()
|
||||
pipe.rpush.assert_not_called()
|
||||
|
||||
|
||||
def test_purge_one_queue_all_removed_deletes_without_rpush():
|
||||
stale = _envelope("c-stale", [["biz-stale"], {}, {}])
|
||||
client, store, pipe = _queue_with_messages(stale)
|
||||
|
||||
removed = guard._purge_one_queue(client, "q", {"biz-stale"}, set())
|
||||
assert removed == 1
|
||||
pipe.delete.assert_called_once_with("q")
|
||||
pipe.rpush.assert_not_called()
|
||||
|
||||
|
||||
def test_purge_one_queue_lrange_exception_returns_zero():
|
||||
client = MagicMock()
|
||||
client.lrange.side_effect = RuntimeError("redis down")
|
||||
assert guard._purge_one_queue(client, "q", {"b"}, set()) == 0
|
||||
|
||||
|
||||
def test_purge_one_queue_empty_queue_returns_zero():
|
||||
client = MagicMock()
|
||||
client.lrange.return_value = []
|
||||
assert guard._purge_one_queue(client, "q", {"b"}, set()) == 0
|
||||
client.pipeline.assert_not_called()
|
||||
|
||||
|
||||
def test_purge_one_queue_rewrite_exception_returns_zero():
|
||||
stale = _envelope("c-stale", [["biz-stale"], {}, {}])
|
||||
client, store, pipe = _queue_with_messages(stale)
|
||||
pipe.execute.side_effect = RuntimeError("write fail")
|
||||
|
||||
removed = guard._purge_one_queue(client, "q", {"biz-stale"}, set())
|
||||
assert removed == 0
|
||||
|
||||
|
||||
def test_purge_one_queue_unparseable_message_conservatively_kept():
|
||||
stale = _envelope("c-stale", [["biz-stale"], {}, {}])
|
||||
garbage = b"garbage-not-a-message"
|
||||
client, store, pipe = _queue_with_messages(garbage, stale)
|
||||
|
||||
removed = guard._purge_one_queue(client, "q", {"biz-stale"}, set())
|
||||
assert removed == 1
|
||||
args, _ = pipe.rpush.call_args
|
||||
# 无法解析的消息保守保留,绝不误删
|
||||
assert list(args[1:]) == [garbage]
|
||||
|
||||
|
||||
# ── purge_stale_messages_from_queues ────────────────────────────────────
|
||||
|
||||
|
||||
def test_purge_queues_no_ids_returns_zero_without_connecting():
|
||||
assert guard.purge_stale_messages_from_queues("redis://x", ("q",)) == 0
|
||||
|
||||
|
||||
def test_purge_queues_blank_ids_filtered_out():
|
||||
assert guard.purge_stale_messages_from_queues("redis://x", ("q",), business_task_ids=["", None]) == 0
|
||||
|
||||
|
||||
def test_purge_queues_redis_not_installed(monkeypatch):
|
||||
"""redis-py 不可用(ImportError)时安全返回 0。"""
|
||||
import builtins
|
||||
|
||||
real_import = builtins.__import__
|
||||
|
||||
def fake_import(name, *args, **kwargs):
|
||||
if name == "redis":
|
||||
raise ImportError("no redis")
|
||||
return real_import(name, *args, **kwargs)
|
||||
|
||||
monkeypatch.setattr(builtins, "__import__", fake_import)
|
||||
assert guard.purge_stale_messages_from_queues("redis://x", ("q",), business_task_ids=["b1"]) == 0
|
||||
|
||||
|
||||
def test_purge_queues_connection_failure_returns_zero():
|
||||
fake_redis = types.ModuleType("redis")
|
||||
|
||||
class _FakeRedis:
|
||||
@classmethod
|
||||
def from_url(cls, url): # noqa: ARG003
|
||||
client = MagicMock()
|
||||
client.ping.side_effect = ConnectionError("connect refused")
|
||||
return client
|
||||
|
||||
fake_redis.Redis = _FakeRedis
|
||||
sys.modules["redis"] = fake_redis
|
||||
try:
|
||||
assert guard.purge_stale_messages_from_queues("redis://x", ("q",), celery_task_ids=["c1"]) == 0
|
||||
finally:
|
||||
sys.modules.pop("redis", None)
|
||||
|
||||
|
||||
def test_purge_queues_happy_path_closes_client():
|
||||
stale = _envelope("c-stale", [["biz-stale"], {}, {}])
|
||||
fake_redis = types.ModuleType("redis")
|
||||
|
||||
client = MagicMock()
|
||||
client.lrange.return_value = [stale]
|
||||
pipe = MagicMock()
|
||||
client.pipeline.return_value = pipe
|
||||
|
||||
class _FakeRedis:
|
||||
@classmethod
|
||||
def from_url(cls, url): # noqa: ARG003
|
||||
return client
|
||||
|
||||
fake_redis.Redis = _FakeRedis
|
||||
sys.modules["redis"] = fake_redis
|
||||
try:
|
||||
removed = guard.purge_stale_messages_from_queues(
|
||||
"redis://x", ("generation", "transcode"), business_task_ids=["biz-stale"]
|
||||
)
|
||||
finally:
|
||||
sys.modules.pop("redis", None)
|
||||
|
||||
# mock client 对两个队列都返回同一条作废消息 → 各移除 1 条
|
||||
assert removed == 2
|
||||
client.ping.assert_called_once()
|
||||
client.close.assert_called_once()
|
||||
# 两个队列都扫描
|
||||
assert client.lrange.call_count == 2
|
||||
|
||||
|
||||
def test_purge_queues_close_exception_swallowed():
|
||||
fake_redis = types.ModuleType("redis")
|
||||
|
||||
client = MagicMock()
|
||||
client.lrange.return_value = []
|
||||
client.close.side_effect = RuntimeError("close fail")
|
||||
|
||||
class _FakeRedis:
|
||||
@classmethod
|
||||
def from_url(cls, url): # noqa: ARG003
|
||||
return client
|
||||
|
||||
fake_redis.Redis = _FakeRedis
|
||||
sys.modules["redis"] = fake_redis
|
||||
try:
|
||||
removed = guard.purge_stale_messages_from_queues("redis://x", ("q",), celery_task_ids=["c1"])
|
||||
finally:
|
||||
sys.modules.pop("redis", None)
|
||||
assert removed == 0
|
||||
|
||||
|
||||
# ── revoke_and_purge ────────────────────────────────────────────────────
|
||||
|
||||
|
||||
def test_revoke_and_purge_revokes_each_message(monkeypatch):
|
||||
fake_app = MagicMock()
|
||||
purge_mock = MagicMock(return_value=2)
|
||||
monkeypatch.setattr(guard, "purge_stale_messages_from_queues", purge_mock)
|
||||
|
||||
removed = guard.revoke_and_purge(
|
||||
fake_app,
|
||||
"redis://x",
|
||||
business_task_ids=["b1"],
|
||||
celery_task_ids=["c1", "c2"],
|
||||
queue_names=("generation",),
|
||||
)
|
||||
assert removed == 2
|
||||
assert fake_app.control.revoke.call_count == 2
|
||||
fake_app.control.revoke.assert_any_call("c1")
|
||||
fake_app.control.revoke.assert_any_call("c2")
|
||||
purge_mock.assert_called_once_with(
|
||||
"redis://x", ("generation",), business_task_ids=["b1"], celery_task_ids=["c1", "c2"]
|
||||
)
|
||||
|
||||
|
||||
def test_revoke_and_purge_revoke_exception_does_not_block(monkeypatch):
|
||||
fake_app = MagicMock()
|
||||
fake_app.control.revoke.side_effect = RuntimeError("broadcast fail")
|
||||
purge_mock = MagicMock(return_value=0)
|
||||
monkeypatch.setattr(guard, "purge_stale_messages_from_queues", purge_mock)
|
||||
|
||||
removed = guard.revoke_and_purge(fake_app, "redis://x", celery_task_ids=["c1"])
|
||||
assert removed == 0
|
||||
purge_mock.assert_called_once()
|
||||
|
||||
|
||||
def test_revoke_and_purge_skips_blank_ids(monkeypatch):
|
||||
fake_app = MagicMock()
|
||||
purge_mock = MagicMock(return_value=0)
|
||||
monkeypatch.setattr(guard, "purge_stale_messages_from_queues", purge_mock)
|
||||
|
||||
guard.revoke_and_purge(fake_app, "redis://x", celery_task_ids=["", None])
|
||||
fake_app.control.revoke.assert_not_called()
|
||||
|
||||
|
||||
# ── ensure_task_claimable ───────────────────────────────────────────────
|
||||
|
||||
|
||||
def test_ensure_claimable_missing_task_returns_empty():
|
||||
assert guard.ensure_task_claimable("t1", lambda _tid: None) == ""
|
||||
|
||||
|
||||
def test_ensure_claimable_terminal_raises():
|
||||
with pytest.raises(guard.StaleTaskDiscarded) as exc_info:
|
||||
guard.ensure_task_claimable("t1", lambda _tid: "failed")
|
||||
assert exc_info.value.task_id == "t1"
|
||||
assert exc_info.value.status == "failed"
|
||||
|
||||
|
||||
def test_ensure_claimable_cancelled_raises():
|
||||
with pytest.raises(guard.StaleTaskDiscarded):
|
||||
guard.ensure_task_claimable("t1", lambda _tid: "cancelled")
|
||||
|
||||
|
||||
def test_ensure_claimable_pending_passes():
|
||||
assert guard.ensure_task_claimable("t1", lambda _tid: "pending") == "pending"
|
||||
@@ -0,0 +1,259 @@
|
||||
"""#1714:入队后 celery_task_id 持久化路径覆盖(routes / enqueue / celery_app / 仓储)。
|
||||
|
||||
CI 无 redis、不走完整 HTTP 流程,这些 try/except 与早退分支此前覆盖率为 0。
|
||||
用真实 SQLite 仓储 + monkeypatch celery_app.send_task 直接驱动路由函数:
|
||||
- routes/ingest_jobs.submit_ingest_job:正常持久化 + 持久化异常吞掉不影响响应
|
||||
- routes/task_center.retry_project_task(ingest 分支):重试后持久化 + 异常吞掉
|
||||
- routes/upload._persist_celery_task_id:空 id 早退 + 异常吞掉
|
||||
- core/task_enqueue.safe_enqueue_generation_task:持久化失败仅 warning,入队仍 True
|
||||
- core/celery_app:apply_queue_settings 抛异常时 API 启动不炸
|
||||
- adapters/ingest_job_repository.update:写 celery_task_id 分支落库
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import importlib
|
||||
import importlib.util
|
||||
import os
|
||||
import sys
|
||||
from pathlib import Path
|
||||
from types import SimpleNamespace
|
||||
from unittest.mock import MagicMock
|
||||
|
||||
os.environ.setdefault("JWT_SECRET_KEY", "unit-test-secret-key-for-testing")
|
||||
os.environ.setdefault("DATABASE_URL", "sqlite:///test.db")
|
||||
|
||||
API_PATH = str(Path(__file__).resolve().parents[2] / "apps" / "api")
|
||||
if API_PATH not in sys.path:
|
||||
sys.path.insert(0, API_PATH)
|
||||
|
||||
import pytest # noqa: E402
|
||||
from app.api.routes import ingest_jobs as ingest_jobs_route # noqa: E402
|
||||
from app.api.routes import task_center as task_center_route # noqa: E402
|
||||
from app.api.routes import upload as upload_route # noqa: E402
|
||||
from app.schemas.ingest_job import SubmitIngestJobRequest # noqa: E402
|
||||
from sqlalchemy import create_engine # noqa: E402
|
||||
from sqlalchemy.orm import sessionmaker # noqa: E402
|
||||
|
||||
from packages.adapters.sqlalchemy_impl.ingest_job_repository import ( # noqa: E402
|
||||
SQLAlchemyIngestJobRepository,
|
||||
)
|
||||
from packages.adapters.sqlalchemy_impl.models import Base # noqa: E402
|
||||
from packages.domain import IngestJob, IngestJobStatus # noqa: E402
|
||||
|
||||
|
||||
def _ingest_repo():
|
||||
engine = create_engine("sqlite:///:memory:")
|
||||
Base.metadata.create_all(engine)
|
||||
session = sessionmaker(bind=engine)()
|
||||
return SQLAlchemyIngestJobRepository(session), session
|
||||
|
||||
|
||||
def _fake_celery_result(task_id: str = "celery-route-msg-1"):
|
||||
result = MagicMock()
|
||||
result.id = task_id
|
||||
return result
|
||||
|
||||
|
||||
# ── routes/ingest_jobs.submit_ingest_job ────────────────────────────────
|
||||
|
||||
|
||||
def test_submit_ingest_job_persists_celery_task_id(monkeypatch):
|
||||
repo, session = _ingest_repo()
|
||||
monkeypatch.setattr(ingest_jobs_route.celery_app, "send_task", MagicMock(return_value=_fake_celery_result()))
|
||||
|
||||
request = SubmitIngestJobRequest(project_id="proj-1", library_id="lib-1", storage_key="uploads/x.mov")
|
||||
response = ingest_jobs_route.submit_ingest_job(request, ingest_job_repository=repo)
|
||||
|
||||
assert response.status == "pending"
|
||||
saved = repo.get(response.id)
|
||||
assert saved.celery_task_id == "celery-route-msg-1"
|
||||
|
||||
|
||||
def test_submit_ingest_job_persist_failure_swallowed(monkeypatch):
|
||||
repo, _ = _ingest_repo()
|
||||
|
||||
class _BoomRepo:
|
||||
def __init__(self, inner):
|
||||
self.inner = inner
|
||||
|
||||
def create(self, job):
|
||||
return self.inner.create(job)
|
||||
|
||||
def get(self, job_id):
|
||||
return self.inner.get(job_id)
|
||||
|
||||
def update(self, job): # noqa: ARG002
|
||||
raise RuntimeError("db write fail")
|
||||
|
||||
boom_repo = _BoomRepo(repo)
|
||||
monkeypatch.setattr(ingest_jobs_route.celery_app, "send_task", MagicMock(return_value=_fake_celery_result()))
|
||||
|
||||
request = SubmitIngestJobRequest(project_id="proj-1", library_id="lib-1", storage_key="uploads/y.mov")
|
||||
# 持久化异常被吞掉,主流程(响应)不受影响
|
||||
response = ingest_jobs_route.submit_ingest_job(request, ingest_job_repository=boom_repo)
|
||||
assert response.id
|
||||
assert response.status == "pending"
|
||||
|
||||
|
||||
# ── routes/task_center.retry_project_task(ingest 分支) ────────────────
|
||||
|
||||
|
||||
def _auth_user():
|
||||
user = SimpleNamespace(id="user-1")
|
||||
return SimpleNamespace(user=user, session_id=None, token_type=None)
|
||||
|
||||
|
||||
def test_retry_ingest_job_persists_celery_task_id(monkeypatch):
|
||||
repo, session = _ingest_repo()
|
||||
# 造一条 failed 的 ingest job
|
||||
job = IngestJob.create(project_id="proj-1", library_id="lib-1", storage_key="uploads/z.mov")
|
||||
job.status = IngestJobStatus.FAILED
|
||||
repo.create(job)
|
||||
|
||||
monkeypatch.setattr(
|
||||
task_center_route.celery_app, "send_task", MagicMock(return_value=_fake_celery_result("celery-retry-1"))
|
||||
)
|
||||
|
||||
response = task_center_route.retry_project_task(
|
||||
"ingest", job.id, authenticated_user=_auth_user(), ingest_job_repository=repo
|
||||
)
|
||||
assert response.task_type == "ingest"
|
||||
new_id = response.id.split("ingest:")[1]
|
||||
retried = repo.get(new_id)
|
||||
assert retried is not None
|
||||
assert retried.celery_task_id == "celery-retry-1"
|
||||
|
||||
|
||||
def test_retry_ingest_job_persist_failure_swallowed(monkeypatch):
|
||||
repo, _ = _ingest_repo()
|
||||
job = IngestJob.create(project_id="proj-1", library_id="lib-1", storage_key="uploads/w.mov")
|
||||
job.status = IngestJobStatus.FAILED
|
||||
repo.create(job)
|
||||
|
||||
real_update = repo.update
|
||||
|
||||
def _update_that_booms(entity):
|
||||
# 仅在写 celery_task_id 的那次 update 抛错(新建 job 后路由内的持久化)
|
||||
if getattr(entity, "celery_task_id", ""):
|
||||
raise RuntimeError("db write fail")
|
||||
return real_update(entity)
|
||||
|
||||
repo.update = _update_that_booms # type: ignore[method-assign]
|
||||
monkeypatch.setattr(task_center_route.celery_app, "send_task", MagicMock(return_value=_fake_celery_result()))
|
||||
|
||||
# 持久化异常吞掉,重试接口仍正常返回
|
||||
response = task_center_route.retry_project_task(
|
||||
"ingest", job.id, authenticated_user=_auth_user(), ingest_job_repository=repo
|
||||
)
|
||||
assert response.task_type == "ingest"
|
||||
|
||||
|
||||
# ── routes/upload._persist_celery_task_id ───────────────────────────────
|
||||
|
||||
|
||||
def test_upload_persist_helper_empty_id_early_return():
|
||||
repo = MagicMock()
|
||||
job = MagicMock()
|
||||
upload_route._persist_celery_task_id(repo, job, "")
|
||||
repo.update.assert_not_called()
|
||||
upload_route._persist_celery_task_id(repo, job, None) # type: ignore[arg-type]
|
||||
repo.update.assert_not_called()
|
||||
|
||||
|
||||
def test_upload_persist_helper_exception_swallowed():
|
||||
repo = MagicMock()
|
||||
repo.update.side_effect = RuntimeError("db fail")
|
||||
job = MagicMock()
|
||||
# 不抛异常
|
||||
upload_route._persist_celery_task_id(repo, job, "celery-upload-1")
|
||||
repo.update.assert_called_once()
|
||||
assert job.celery_task_id == "celery-upload-1"
|
||||
|
||||
|
||||
# ── core/task_enqueue:持久化失败仅 warning ─────────────────────────────
|
||||
|
||||
|
||||
def test_safe_enqueue_persist_failure_still_returns_true(monkeypatch):
|
||||
from app.core import task_enqueue
|
||||
|
||||
class _FakeTask:
|
||||
def __init__(self):
|
||||
self.id = "task-enqueue-persist-fail"
|
||||
self.status = "pending"
|
||||
self.celery_task_id = ""
|
||||
|
||||
def mark_failed(self, msg): # noqa: ARG002
|
||||
self.status = "failed"
|
||||
|
||||
class _FakeRepo:
|
||||
def count_pending_total(self):
|
||||
return 0
|
||||
|
||||
def count_pending_by_user(self, user_id): # noqa: ARG002
|
||||
return 0
|
||||
|
||||
def update(self, task): # noqa: ARG002
|
||||
raise RuntimeError("persist celery_task_id failed")
|
||||
|
||||
fake_result = MagicMock()
|
||||
fake_result.id = "celery-enqueue-fail-1"
|
||||
mock_celery = MagicMock()
|
||||
mock_celery.send_task.return_value = fake_result
|
||||
monkeypatch.setattr(task_enqueue, "celery_app", mock_celery)
|
||||
|
||||
task = _FakeTask()
|
||||
repo = _FakeRepo()
|
||||
ok = task_enqueue.safe_enqueue_generation_task(task, repo, user_id="u1")
|
||||
# 持久化失败不影响入队结果
|
||||
assert ok is True
|
||||
mock_celery.send_task.assert_called_once()
|
||||
|
||||
|
||||
# ── core/celery_app:队列配置失败不阻断 API 启动 ────────────────────────
|
||||
|
||||
|
||||
def test_api_celery_app_survives_queue_settings_failure(monkeypatch):
|
||||
"""apply_queue_settings 抛异常时 API 启动不炸(core/celery_app.py 的 try/except 分支)。
|
||||
|
||||
通过让 `from packages.shared.celery_queues import apply_queue_settings` 本身
|
||||
抛异常来触发 except 分支;用全新模块名 reload,不替换已被其他模块持有的
|
||||
app.core.celery_app 模块对象,避免污染 task_enqueue 等导入方。
|
||||
"""
|
||||
import builtins
|
||||
|
||||
real_import = builtins.__import__
|
||||
|
||||
def _failing_import(name, globals=None, locals=None, fromlist=(), level=0): # noqa: A002
|
||||
if name == "packages.shared.celery_queues" and "apply_queue_settings" in (fromlist or ()):
|
||||
raise RuntimeError("config boom")
|
||||
return real_import(name, globals, locals, fromlist, level)
|
||||
|
||||
monkeypatch.setattr(builtins, "__import__", _failing_import)
|
||||
|
||||
spec = importlib.util.find_spec("app.core.celery_app")
|
||||
fresh_mod = importlib.util.module_from_spec(spec)
|
||||
spec.loader.exec_module(fresh_mod) # 异常在模块内被 try/except 吞掉
|
||||
assert fresh_mod.celery_app is not None
|
||||
assert fresh_mod.celery_app.main == "xiaoxia-saas-api"
|
||||
|
||||
# 已加载的原模块对象不受影响(无 reload 污染)
|
||||
import app.core.celery_app as api_celery_mod
|
||||
|
||||
assert api_celery_mod.celery_app is not None
|
||||
|
||||
|
||||
# ── 仓储:update 写 celery_task_id 落库 ─────────────────────────────────
|
||||
|
||||
|
||||
def test_ingest_repo_update_persists_celery_task_id():
|
||||
repo, session = _ingest_repo()
|
||||
job = IngestJob.create(project_id="proj-1", library_id="lib-1", storage_key="uploads/repo.mov")
|
||||
repo.create(job)
|
||||
|
||||
job.celery_task_id = "celery-repo-update-1"
|
||||
repo.update(job)
|
||||
|
||||
session.expire_all()
|
||||
saved = repo.get(job.id)
|
||||
assert saved.celery_task_id == "celery-repo-update-1"
|
||||
@@ -0,0 +1,204 @@
|
||||
"""Issue #1714:孤儿/超时清理标记 failed 时必须撤销并清除 Redis 队列消息。
|
||||
|
||||
覆盖:
|
||||
- cleanup_stale_pending_with_session_ids:超时 pending 标记 failed 并返回
|
||||
(task_id, celery_task_id),worker 清理流程据此 revoke + purge 队列消息
|
||||
- 队列中对应业务任务的 celery 消息被物理移除(作废消息不会重投执行)
|
||||
- 旧仓储(无 _with_ids 方法)降级为计数模式,不抛异常
|
||||
- cleanup_stale_running_with_ids 同样返回 id 列表
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import sys
|
||||
from datetime import datetime, timedelta, timezone
|
||||
from pathlib import Path
|
||||
from unittest.mock import MagicMock
|
||||
|
||||
import pytest
|
||||
from sqlalchemy import create_engine, text
|
||||
from sqlalchemy.orm import sessionmaker
|
||||
|
||||
sys.path.insert(0, str(Path(__file__).resolve().parents[2] / "apps" / "api"))
|
||||
|
||||
from packages.adapters.sqlalchemy_impl.generation_task_repository import ( # noqa: E402
|
||||
SQLAlchemyGenerationTaskRepository,
|
||||
)
|
||||
from packages.adapters.sqlalchemy_impl.models import Base # noqa: E402
|
||||
from packages.domain import GenerationTask, GenerationTaskStatus # noqa: E402
|
||||
|
||||
BROKER_URL = "redis://localhost:6379/15"
|
||||
TEST_QUEUE = "_test_revoke_q"
|
||||
|
||||
|
||||
def _repository():
|
||||
engine = create_engine("sqlite:///:memory:")
|
||||
Base.metadata.create_all(engine)
|
||||
session = sessionmaker(bind=engine)()
|
||||
return SQLAlchemyGenerationTaskRepository(session), session, engine
|
||||
|
||||
|
||||
def _make_task(**kwargs) -> GenerationTask:
|
||||
defaults = dict(project_id="proj-1", asset_library_id="lib-1", created_by_user_id="user-1")
|
||||
defaults.update(kwargs)
|
||||
return GenerationTask.create(**defaults)
|
||||
|
||||
|
||||
def _redis_available() -> bool:
|
||||
try:
|
||||
import redis
|
||||
|
||||
return bool(redis.Redis.from_url(BROKER_URL).ping())
|
||||
except Exception:
|
||||
return False
|
||||
|
||||
|
||||
# ── 仓储层:返回 ids ────────────────────────────────────────────────────
|
||||
|
||||
|
||||
def test_cleanup_stale_pending_returns_ids_with_celery_task_id():
|
||||
repo, _, engine = _repository()
|
||||
task = _make_task()
|
||||
task.celery_task_id = "celery-msg-id-001"
|
||||
repo.create(task)
|
||||
# created_at 改到 60 分钟前
|
||||
with engine.connect() as conn:
|
||||
conn.execute(
|
||||
text("UPDATE generation_tasks SET created_at = :ts WHERE id = :id"),
|
||||
{"ts": datetime.now(timezone.utc) - timedelta(minutes=60), "id": task.id},
|
||||
)
|
||||
conn.commit()
|
||||
|
||||
items = repo.cleanup_stale_pending_with_ids(timeout_minutes=45)
|
||||
assert len(items) == 1
|
||||
biz_id, celery_id = items[0]
|
||||
assert biz_id == task.id
|
||||
assert celery_id == "celery-msg-id-001"
|
||||
|
||||
saved = repo.get(task.id)
|
||||
assert saved.status == GenerationTaskStatus.FAILED
|
||||
|
||||
|
||||
def test_cleanup_stale_running_returns_ids():
|
||||
repo, _, engine = _repository()
|
||||
task = _make_task()
|
||||
repo.create(task)
|
||||
task.mark_processing()
|
||||
task.celery_task_id = "celery-msg-id-002"
|
||||
repo.update(task)
|
||||
with engine.connect() as conn:
|
||||
conn.execute(
|
||||
text("UPDATE generation_tasks SET updated_at = :ts WHERE id = :id"),
|
||||
{"ts": datetime.now(timezone.utc) - timedelta(minutes=60), "id": task.id},
|
||||
)
|
||||
conn.commit()
|
||||
|
||||
items = repo.cleanup_stale_running_with_ids(timeout_minutes=20)
|
||||
assert len(items) == 1
|
||||
assert items[0][0] == task.id
|
||||
assert items[0][1] == "celery-msg-id-002"
|
||||
assert repo.get(task.id).status == GenerationTaskStatus.FAILED
|
||||
|
||||
|
||||
def test_legacy_repo_without_with_ids_falls_back_to_count():
|
||||
"""旧仓储只有 cleanup_stale_pending(返回 int)时降级可用,不抛异常。"""
|
||||
# worker 模块加载(标准 mock 模式)
|
||||
saved = set(sys.modules.keys())
|
||||
mock_db = MagicMock()
|
||||
mock_db.SessionLocal = MagicMock()
|
||||
sys.modules["worker_app.db"] = mock_db
|
||||
sys.modules["worker_app.core.config"] = MagicMock()
|
||||
mock_celery = MagicMock()
|
||||
mock_celery.celery_app.task = MagicMock(
|
||||
side_effect=(lambda *a, **k: (a[0] if a and callable(a[0]) else (lambda f: f)))
|
||||
)
|
||||
sys.modules["worker_app.celery_app"] = mock_celery
|
||||
worker_path = str(Path(__file__).resolve().parents[2] / "apps" / "worker")
|
||||
if worker_path not in sys.path:
|
||||
sys.path.insert(0, worker_path)
|
||||
|
||||
from worker_app.tasks import _startup # noqa: E402
|
||||
|
||||
class LegacyRepo:
|
||||
def cleanup_stale_pending(self, timeout_minutes): # noqa: ARG002
|
||||
return 3
|
||||
|
||||
def cleanup_stale_running(self, timeout_minutes): # noqa: ARG002
|
||||
return 2
|
||||
|
||||
items_p = _startup.cleanup_stale_pending_with_session_ids(LegacyRepo(), 45)
|
||||
items_r = _startup.cleanup_stale_running_with_session_ids(LegacyRepo(), 20)
|
||||
assert len(items_p) == 3
|
||||
assert len(items_r) == 2
|
||||
|
||||
for key in list(sys.modules.keys()):
|
||||
if key not in saved and not key.startswith("video_processing"):
|
||||
del sys.modules[key]
|
||||
|
||||
|
||||
# ── 端到端:清理 → 队列消息被移除(作废消息不重投) ────────────────────
|
||||
|
||||
|
||||
@pytest.mark.skipif(not _redis_available(), reason="本地 redis 不可用")
|
||||
def test_stale_pending_cleanup_purges_redis_message():
|
||||
"""任务标 failed 后,其在 Redis 队列里的 celery 消息被清除,不会被重投。"""
|
||||
import redis
|
||||
from celery import Celery
|
||||
from kombu import Queue
|
||||
from kombu.pools import producers
|
||||
|
||||
from packages.shared.celery_orphan_guard import purge_stale_messages_from_queues
|
||||
|
||||
repo, _, engine = _repository()
|
||||
task = _make_task()
|
||||
task.celery_task_id = "celery-stale-xyz"
|
||||
repo.create(task)
|
||||
with engine.connect() as conn:
|
||||
conn.execute(
|
||||
text("UPDATE generation_tasks SET created_at = :ts WHERE id = :id"),
|
||||
{"ts": datetime.now(timezone.utc) - timedelta(minutes=60), "id": task.id},
|
||||
)
|
||||
conn.commit()
|
||||
|
||||
# 模拟该任务的 celery 消息仍在 generation 队列里(worker 下线期间未消费)
|
||||
client = redis.Redis.from_url(BROKER_URL)
|
||||
client.delete(TEST_QUEUE)
|
||||
app = Celery("test-e2e-revoke")
|
||||
app.conf.broker_url = BROKER_URL
|
||||
with app.connection_for_write() as conn:
|
||||
with producers[conn].acquire(block=True) as prod:
|
||||
# 作废任务消息
|
||||
prod.publish(
|
||||
(task.id,),
|
||||
exchange="",
|
||||
routing_key=TEST_QUEUE,
|
||||
serializer="json",
|
||||
headers={"id": "celery-stale-xyz", "task": "worker.generate_video"},
|
||||
retry=False,
|
||||
delivery_mode=1,
|
||||
declare=[Queue(TEST_QUEUE, routing_key=TEST_QUEUE, durable=False)],
|
||||
)
|
||||
# 另一条正常任务消息(必须保留)
|
||||
prod.publish(
|
||||
("other-task-id",),
|
||||
exchange="",
|
||||
routing_key=TEST_QUEUE,
|
||||
serializer="json",
|
||||
headers={"id": "celery-keep", "task": "worker.generate_video"},
|
||||
retry=False,
|
||||
delivery_mode=1,
|
||||
)
|
||||
|
||||
assert client.llen(TEST_QUEUE) == 2
|
||||
|
||||
# 执行清理(与 worker beat 相同流程:标 failed → 拿 ids → purge)
|
||||
items = repo.cleanup_stale_pending_with_ids(timeout_minutes=45)
|
||||
biz_ids = [bid for bid, _ in items]
|
||||
celery_ids = [cid for _, cid in items if cid]
|
||||
removed = purge_stale_messages_from_queues(
|
||||
BROKER_URL, (TEST_QUEUE,), business_task_ids=biz_ids, celery_task_ids=celery_ids
|
||||
)
|
||||
|
||||
assert removed == 1
|
||||
assert client.llen(TEST_QUEUE) == 1 # 正常任务消息保留
|
||||
client.delete(TEST_QUEUE)
|
||||
@@ -0,0 +1,243 @@
|
||||
"""Issue #1714:任务执行前状态守卫 — 已作废消息必须丢弃,禁止非法转换后继续跑。
|
||||
|
||||
覆盖:
|
||||
- ingest_asset:job 已 failed/completed 时直接返回 discarded,不下载、不转码、不回写
|
||||
- generate_video:GenerationTask 已 failed 时返回 discarded,不进入渲染
|
||||
- generate_video:pending → running 标记失败(非法转换)时安全中止
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import sys
|
||||
from pathlib import Path
|
||||
from types import SimpleNamespace
|
||||
from unittest.mock import MagicMock
|
||||
|
||||
# ── worker 模块标准加载方式 ──
|
||||
# 显式保存将要覆盖的注入键旧值:全量收集时更早的测试文件(如
|
||||
# test_ingest_validation.py)可能已向 sys.modules 注入 worker_app.* mock,
|
||||
# 导入完成后必须精确恢复旧值,否则本文件的 bind 感知透传装饰器会残留,
|
||||
# 污染后续懒加载短路径 worker_app.celery_app 的 worker 测试。
|
||||
_INJECTED_KEYS = ("worker_app.db", "worker_app.core.config", "worker_app.celery_app")
|
||||
_SAVED_MODULE_VALUES = {k: sys.modules.get(k) for k in _INJECTED_KEYS}
|
||||
_SAVED_MODULES_KEYS = set(sys.modules.keys())
|
||||
|
||||
_mock_db_module = MagicMock()
|
||||
_mock_db_module.SessionLocal = MagicMock()
|
||||
sys.modules["worker_app.db"] = _mock_db_module
|
||||
sys.modules["worker_app.core.config"] = MagicMock()
|
||||
|
||||
_mock_celery_module = MagicMock()
|
||||
|
||||
|
||||
def _passthrough_decorator(*args, **kwargs):
|
||||
if len(args) == 1 and callable(args[0]):
|
||||
return args[0]
|
||||
bind = kwargs.get("bind", False)
|
||||
|
||||
def _wrap(f):
|
||||
if bind:
|
||||
# 模拟 celery bind=True:task(task_id) 调用时注入 self(MagicMock)
|
||||
return lambda *a, **kw: f(MagicMock(), *a, **kw)
|
||||
return f
|
||||
|
||||
return _wrap
|
||||
|
||||
|
||||
_mock_celery_module.celery_app.task = MagicMock(side_effect=_passthrough_decorator)
|
||||
sys.modules["worker_app.celery_app"] = _mock_celery_module
|
||||
|
||||
_WORKER_PATH = str(Path(__file__).resolve().parents[2] / "apps" / "worker")
|
||||
sys.path.insert(0, _WORKER_PATH)
|
||||
|
||||
import pytest # noqa: E402
|
||||
from worker_app.tasks import ingest as ingest_mod # noqa: E402
|
||||
|
||||
# video_processing 相关 mock(generation 模块导入链)
|
||||
for _mod_name in [
|
||||
"video_processing",
|
||||
"video_processing.ffmpeg_utils",
|
||||
"video_processing.oss_helpers",
|
||||
]:
|
||||
sys.modules.setdefault(_mod_name, MagicMock())
|
||||
|
||||
from worker_app.tasks import generation as gen_mod # noqa: E402
|
||||
|
||||
from packages.domain import IngestJobStatus # noqa: E402
|
||||
|
||||
# 模块导入完成后立即清理:删除本次 import 新引入的模块缓存(本模块已通过名字绑定
|
||||
# 持有 ingest_mod/gen_mod/IngestJobStatus,删除缓存不影响调用),再把三个注入键
|
||||
# 精确恢复为注入前的旧值(旧值不存在则移除),杜绝 mock 残留污染其他 worker 测试。
|
||||
for _key in list(sys.modules.keys()):
|
||||
if _key not in _SAVED_MODULES_KEYS and not _key.startswith("video_processing"):
|
||||
del sys.modules[_key]
|
||||
for _k, _v in _SAVED_MODULE_VALUES.items():
|
||||
if _v is None:
|
||||
sys.modules.pop(_k, None)
|
||||
else:
|
||||
sys.modules[_k] = _v
|
||||
del _SAVED_MODULES_KEYS, _SAVED_MODULE_VALUES
|
||||
|
||||
|
||||
# ── ingest 守卫 ────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
class _FakeJobRepo:
|
||||
def __init__(self, job):
|
||||
self.job = job
|
||||
|
||||
def get(self, job_id):
|
||||
return self.job
|
||||
|
||||
|
||||
def _make_ingest_job(status):
|
||||
job = MagicMock()
|
||||
job.id = "job-stale-1"
|
||||
job.storage_key = "uploads/proj/stale.mov"
|
||||
job.status = status
|
||||
job.file_hash = "h"
|
||||
job.asset_id = ""
|
||||
return job
|
||||
|
||||
|
||||
def test_ingest_discards_failed_job_message():
|
||||
"""job 已 failed:消息丢弃,不进入下载/转码/回写。"""
|
||||
job = _make_ingest_job(IngestJobStatus.FAILED)
|
||||
fake_session = MagicMock()
|
||||
_mock_db_module.SessionLocal = MagicMock(return_value=fake_session)
|
||||
|
||||
# SQLAlchemy 仓储构造返回 fake
|
||||
fake_job_repo = _FakeJobRepo(job)
|
||||
fake_asset_repo = MagicMock()
|
||||
|
||||
orig_job_repo = ingest_mod.SQLAlchemyIngestJobRepository
|
||||
orig_asset_repo = ingest_mod.SQLAlchemyAssetRepository
|
||||
ingest_mod.SQLAlchemyIngestJobRepository = MagicMock(return_value=fake_job_repo)
|
||||
ingest_mod.SQLAlchemyAssetRepository = MagicMock(return_value=fake_asset_repo)
|
||||
try:
|
||||
result = ingest_mod.ingest_asset("job-stale-1")
|
||||
finally:
|
||||
ingest_mod.SQLAlchemyIngestJobRepository = orig_job_repo
|
||||
ingest_mod.SQLAlchemyAssetRepository = orig_asset_repo
|
||||
|
||||
assert result["status"] == "discarded"
|
||||
# 没有任何 update / commit / 下载动作
|
||||
fake_session.commit.assert_not_called()
|
||||
fake_asset_repo.create.assert_not_called()
|
||||
|
||||
|
||||
def test_ingest_discards_completed_job_message():
|
||||
job = _make_ingest_job(IngestJobStatus.COMPLETED)
|
||||
fake_session = MagicMock()
|
||||
_mock_db_module.SessionLocal = MagicMock(return_value=fake_session)
|
||||
fake_job_repo = _FakeJobRepo(job)
|
||||
|
||||
orig = ingest_mod.SQLAlchemyIngestJobRepository
|
||||
ingest_mod.SQLAlchemyIngestJobRepository = MagicMock(return_value=fake_job_repo)
|
||||
ingest_mod.SQLAlchemyAssetRepository = MagicMock(return_value=MagicMock())
|
||||
try:
|
||||
result = ingest_mod.ingest_asset("job-stale-1")
|
||||
finally:
|
||||
ingest_mod.SQLAlchemyIngestJobRepository = orig
|
||||
|
||||
assert result["status"] == "discarded"
|
||||
|
||||
|
||||
# ── generation 守卫 ────────────────────────────────────────────────────
|
||||
|
||||
|
||||
def _make_gen_task(status_value: str):
|
||||
from packages.domain import GenerationTask
|
||||
|
||||
task = GenerationTask.create(project_id="p", asset_library_id="l", created_by_user_id="u")
|
||||
task.status = type(task.status)(status_value)
|
||||
return task
|
||||
|
||||
|
||||
def test_generate_video_discards_failed_task(monkeypatch):
|
||||
"""GenerationTask 已 failed:直接 discarded,不加载渲染数据。"""
|
||||
failed_task = _make_gen_task("failed")
|
||||
|
||||
fake_repo = MagicMock()
|
||||
fake_repo.get.return_value = failed_task
|
||||
|
||||
fake_session = MagicMock()
|
||||
_mock_db_module.SessionLocal = MagicMock(return_value=fake_session)
|
||||
|
||||
import packages.adapters.sqlalchemy_impl.generation_task_repository as gen_repo_mod
|
||||
|
||||
orig = gen_repo_mod.SQLAlchemyGenerationTaskRepository
|
||||
gen_repo_mod.SQLAlchemyGenerationTaskRepository = MagicMock(return_value=fake_repo)
|
||||
|
||||
update_status_mock = MagicMock(return_value=False)
|
||||
monkeypatch.setattr(gen_mod, "_update_task_status", update_status_mock)
|
||||
monkeypatch.setattr(
|
||||
gen_mod,
|
||||
"_load_task_info",
|
||||
lambda task_id: {
|
||||
"project_id": "p",
|
||||
"template_id": "",
|
||||
"task_asset_ids": [],
|
||||
"batch_id": "",
|
||||
"user_id": "u",
|
||||
"mode": "one_take",
|
||||
},
|
||||
)
|
||||
monkeypatch.setattr(gen_mod, "_flush_logs", lambda *a, **k: None)
|
||||
|
||||
task_fn = gen_mod.generate_video
|
||||
if hasattr(task_fn, "__wrapped__"):
|
||||
task_fn = task_fn.__wrapped__
|
||||
try:
|
||||
result = task_fn("task-stale-1")
|
||||
finally:
|
||||
gen_repo_mod.SQLAlchemyGenerationTaskRepository = orig
|
||||
|
||||
assert result["status"] == "discarded"
|
||||
# 状态守卫命中终态,根本不应尝试 mark_processing
|
||||
update_status_mock.assert_not_called()
|
||||
|
||||
|
||||
def test_generate_video_aborts_when_claim_fails(monkeypatch):
|
||||
"""pending 但 mark_processing 返回 False(状态机非法转换)时安全中止。"""
|
||||
pending_task = _make_gen_task("pending")
|
||||
|
||||
fake_repo = MagicMock()
|
||||
fake_repo.get.return_value = pending_task
|
||||
fake_session = MagicMock()
|
||||
_mock_db_module.SessionLocal = MagicMock(return_value=fake_session)
|
||||
|
||||
import packages.adapters.sqlalchemy_impl.generation_task_repository as gen_repo_mod
|
||||
|
||||
orig = gen_repo_mod.SQLAlchemyGenerationTaskRepository
|
||||
gen_repo_mod.SQLAlchemyGenerationTaskRepository = MagicMock(return_value=fake_repo)
|
||||
|
||||
monkeypatch.setattr(
|
||||
gen_mod,
|
||||
"_load_task_info",
|
||||
lambda task_id: {
|
||||
"project_id": "p",
|
||||
"template_id": "",
|
||||
"task_asset_ids": [],
|
||||
"batch_id": "",
|
||||
"user_id": "u",
|
||||
"mode": "one_take",
|
||||
},
|
||||
)
|
||||
monkeypatch.setattr(gen_mod, "_flush_logs", lambda *a, **k: None)
|
||||
# 模拟 mark_processing 失败(failed→running 非法转换被 _update_task_status 吞掉返回 False)
|
||||
update_status_mock = MagicMock(return_value=False)
|
||||
monkeypatch.setattr(gen_mod, "_update_task_status", update_status_mock)
|
||||
render_mock = MagicMock(side_effect=AssertionError("must not render"))
|
||||
monkeypatch.setattr(gen_mod, "_render_from_edit_plan", render_mock)
|
||||
|
||||
task_fn = gen_mod.generate_video
|
||||
if hasattr(task_fn, "__wrapped__"):
|
||||
task_fn = task_fn.__wrapped__
|
||||
try:
|
||||
result = task_fn("task-claim-fail")
|
||||
finally:
|
||||
gen_repo_mod.SQLAlchemyGenerationTaskRepository = orig
|
||||
|
||||
assert result["status"] == "discarded"
|
||||
render_mock.assert_not_called()
|
||||
@@ -164,7 +164,9 @@ class TestSafeEnqueueWithLimits:
|
||||
result = safe_enqueue_generation_task(task, repo, user_id="user-1")
|
||||
assert result is True
|
||||
mock_celery.assert_called_once_with("worker.generate_video", args=["task-1"])
|
||||
assert len(repo.updated_tasks) == 0 # 成功不需要更新状态
|
||||
# 成功入队后持久化 celery 消息 ID(#1714:孤儿清理据此 revoke/清队列)
|
||||
assert len(repo.updated_tasks) == 1
|
||||
assert task.celery_task_id
|
||||
|
||||
def test_user_limit_rejected_with_failed_status(self, mock_celery):
|
||||
"""用户超限:任务标记为 failed,抛 UserPendingLimitExceeded。"""
|
||||
@@ -335,7 +337,9 @@ class TestPostEnqueueFinalCheck:
|
||||
assert result is True
|
||||
mock_celery.assert_called_once()
|
||||
assert task.status == "pending" # 状态没变
|
||||
assert len(repo.updated_tasks) == 0 # 没更新 DB
|
||||
# 入队成功后持久化 celery_task_id(#1714),业务状态不变
|
||||
assert len(repo.updated_tasks) == 1
|
||||
assert task.celery_task_id
|
||||
|
||||
def test_post_enqueue_no_user_id_skips_user_check(self, mock_celery):
|
||||
"""不传 user_id 时,入队后校验也跳过用户级,只查全局。"""
|
||||
|
||||
Reference in New Issue
Block a user