diff --git a/apps/worker/celery_app.py b/apps/worker/celery_app.py deleted file mode 100644 index 5db9732fa..000000000 --- a/apps/worker/celery_app.py +++ /dev/null @@ -1,12 +0,0 @@ -import os -from celery import Celery - - -def create_celery_app() -> Celery: - app = Celery("xiaoxia_saas_worker") - app.conf.broker_url = os.getenv("CELERY_BROKER_URL", "redis://redis:6379/0") - app.conf.result_backend = os.getenv("CELERY_RESULT_BACKEND", "redis://redis:6379/1") - return app - - -celery_app = create_celery_app() diff --git a/apps/worker/video_processing/dedup.py b/apps/worker/video_processing/dedup.py index 068d0f858..7468e0164 100644 --- a/apps/worker/video_processing/dedup.py +++ b/apps/worker/video_processing/dedup.py @@ -14,16 +14,12 @@ from celery import Task from sqlalchemy.orm import Session from packages.adapters.sqlalchemy_impl.generated_video_repository import SQLAlchemyGeneratedVideoRepository -from packages.adapters.sqlalchemy_impl.session import SessionLocal, build_session_factory -from packages.shared.config import get_shared_settings from packages.shared.storage import get_storage_service -from apps.worker.celery_app import celery_app +from worker_app.celery_app import celery_app +from worker_app.db import SessionLocal logger = logging.getLogger(__name__) -settings = get_shared_settings() -if SessionLocal is None: - build_session_factory(settings.database_url) def compute_phash(image: np.ndarray, hash_size: int = 8) -> str: diff --git a/apps/worker/worker_app/celery_app.py b/apps/worker/worker_app/celery_app.py index a14dd2e29..5dd182a0e 100644 --- a/apps/worker/worker_app/celery_app.py +++ b/apps/worker/worker_app/celery_app.py @@ -5,9 +5,12 @@ settings = get_settings() 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 celery_app.conf.imports = ( "worker_app.tasks.health", "worker_app.tasks.ingest", "worker_app.tasks.classification", "worker_app.tasks.generation", + "worker_app.tasks.voice_extraction", + "apps.worker.video_processing.dedup", ) diff --git a/apps/worker/worker_app/tasks/classification.py b/apps/worker/worker_app/tasks/classification.py index 88a409c5b..091d1d65c 100755 --- a/apps/worker/worker_app/tasks/classification.py +++ b/apps/worker/worker_app/tasks/classification.py @@ -13,12 +13,14 @@ from packages.domain import ( ClassificationStatus, ) +from worker_app.celery_app import celery_app from .asset_analyzer import classify_asset_real logger = get_task_logger(__name__) -def classify_asset(job_id: str) -> dict: +@celery_app.task(bind=True, name="worker.classify_asset", max_retries=2) +def classify_asset(self, job_id: str) -> dict: """ Classify asset task. diff --git a/apps/worker/worker_app/tasks/generation.py b/apps/worker/worker_app/tasks/generation.py index 57c7242af..f8cd1a5a4 100755 --- a/apps/worker/worker_app/tasks/generation.py +++ b/apps/worker/worker_app/tasks/generation.py @@ -14,6 +14,8 @@ from typing import Optional import oss2 +from worker_app.celery_app import celery_app + OUTPUT_WIDTH = 1280 OUTPUT_HEIGHT = 720 OUTPUT_FPS = 25.0 @@ -214,29 +216,37 @@ def _process_with_editing_mode( ) -def generate_video( - task_id: str, - workspace_id: str, - project_id: str, - asset_library_id: str, - voice_library_id: str = "", - mode: str = "one_take", -) -> dict: +@celery_app.task(bind=True, name="worker.generate_video", max_retries=2) +def generate_video(self, task_id: str) -> dict: """ 生成视频任务 Args: - task_id: 任务 ID - workspace_id: 工作空间 ID - project_id: 项目 ID - asset_library_id: 素材库 ID - voice_library_id: 配音库 ID(可选) - mode: 剪辑模式,默认 one_take + task_id: 任务 ID(从数据库加载完整任务信息) Returns: 生成结果字典 """ from packages.domain import GeneratedVideo, GenerationMode, GenerationTaskStatus + from packages.adapters.sqlalchemy_impl.generation_task_repository import ( + SQLAlchemyGenerationTaskRepository, + ) + from worker_app.db import SessionLocal + + # 从数据库加载任务信息 + session = SessionLocal() + try: + task_repo = SQLAlchemyGenerationTaskRepository(session) + gen_task = task_repo.get(task_id) + if gen_task is None: + return {"status": "failed", "error": f"generation task {task_id} not found"} + workspace_id = gen_task.workspace_id + project_id = gen_task.project_id + asset_library_id = gen_task.asset_library_id + voice_library_id = gen_task.voice_library_id or "" + mode = gen_task.strategy_id or "one_take" + finally: + session.close() try: editing_mode = GenerationMode(mode) diff --git a/apps/worker/worker_app/tasks/voice_extraction.py b/apps/worker/worker_app/tasks/voice_extraction.py index b3cae3ec9..38641e66b 100644 --- a/apps/worker/worker_app/tasks/voice_extraction.py +++ b/apps/worker/worker_app/tasks/voice_extraction.py @@ -10,16 +10,12 @@ from celery import Task from sqlalchemy.orm import Session from packages.adapters.sqlalchemy_impl.asset_repository import SQLAlchemyAssetRepository -from packages.adapters.sqlalchemy_impl.session import SessionLocal, build_session_factory -from packages.shared.config import get_shared_settings from packages.shared.storage import get_storage_service -from .celery_app import celery_app +from worker_app.celery_app import celery_app +from worker_app.db import SessionLocal logger = logging.getLogger(__name__) -settings = get_shared_settings() -if SessionLocal is None: - build_session_factory(settings.database_url) class VoiceExtractor: diff --git a/infra/docker/worker.Dockerfile b/infra/docker/worker.Dockerfile index 80b167dab..e76503555 100644 --- a/infra/docker/worker.Dockerfile +++ b/infra/docker/worker.Dockerfile @@ -36,6 +36,13 @@ COPY migrations/ /app/migrations/ ENV PYTHONPATH=/app ENV PYTHONUNBUFFERED=1 +# 创建非 root 用户运行 Worker +RUN groupadd -r celery && useradd -r -g celery -d /app -s /sbin/nologin celery \ + && chown -R celery:celery /app +RUN mkdir -p /app/generated && chown celery:celery /app/generated + +USER celery + # Worker 入口点 WORKDIR /app/apps/worker -CMD ["celery", "-A", "celery_app", "worker", "--loglevel=info", "--concurrency=2"] +CMD ["celery", "-A", "worker_app.celery_app", "worker", "--loglevel=info", "--concurrency=2"]