feat(worker): celery 队列隔离 + 孤儿任务消息作废 (#1714) #1722

Merged
xiaoxia merged 2 commits from feature/celery-queue-isolation-1714 into develop 2026-09-05 20:10:13 +08:00
32 changed files with 1934 additions and 63 deletions
@@ -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")
+3 -1
View File
@@ -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"
+7 -1
View File
@@ -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,
+7 -1
View File
@@ -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",
+13 -1
View File
@@ -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
+8
View File
@@ -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
+11 -1
View File
@@ -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",
+12
View File
@@ -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",
+92 -8
View File
@@ -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
"""统一清理所有超时的孤儿任务。
+4 -2
View File
@@ -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}
+30 -2
View File
@@ -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:
+13
View File
@@ -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_keyHEVC 转码成功后 job.storage_key 会改写为 *_h264
# 而 complete 阶段的占位 asset 始终以原始 key 创建,关联回写必须保留它。
original_storage_key = job.storage_key
+6 -3
View File
@@ -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:
+2 -1
View File
@@ -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 \
+2 -1
View File
@@ -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 \
+54 -9
View File
@@ -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 beatbeat 负责定期触发 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 分钟未更新的任务
标记为 failederror_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="")
+2
View File
@@ -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)
+3
View File
@@ -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(),
)
+1
View File
@@ -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 = ""
+222
View File
@@ -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 的消息直接丢弃(抛 StaleTaskDiscardedtask 捕获后安全返回,
不进入渲染/转码,不产出半成品)。
"""
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 消息 IDheaders.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: 业务任务 IDgeneration_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
)
+58
View File
@@ -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 appAPI 生产端与 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
View File
@@ -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_routesgenerate_video → generationingest_asset/classify/duplication → transcode
2. purge_stale_messages_from_queuesRedis 队列中作废任务消息被物理移除,未命中保留
3. revoke_and_purgerevoke 广播 + 队列清理同时生效
4. ensure_task_claimable:终态任务抛 StaleTaskDiscardedpending 放行
"""
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_queuebiz id 命中、celery id 命中、未命中保序(重写 rpush)、
lrange 异常、重写异常、空队列
- purge_stale_messages_from_queues:空 ids 早退、redis 未安装、连接失败、
正常清理并 close
- revoke_and_purgerevoke 逐消息调用、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_taskingest 分支):重试后持久化 + 异常吞掉
- routes/upload._persist_celery_task_id:空 id 早退 + 异常吞掉
- core/task_enqueue.safe_enqueue_generation_task:持久化失败仅 warning,入队仍 True
- core/celery_appapply_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_taskingest 分支) ────────────────
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"
+204
View File
@@ -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)
+243
View File
@@ -0,0 +1,243 @@
"""Issue #1714:任务执行前状态守卫 — 已作废消息必须丢弃,禁止非法转换后继续跑。
覆盖:
- ingest_assetjob 已 failed/completed 时直接返回 discarded,不下载、不转码、不回写
- generate_videoGenerationTask 已 failed 时返回 discarded,不进入渲染
- generate_videopending → 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=Truetask(task_id) 调用时注入 selfMagicMock
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 相关 mockgeneration 模块导入链)
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()
+6 -2
View File
@@ -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 时,入队后校验也跳过用户级,只查全局。"""