拆镜复刻开发完成

This commit is contained in:
2026-06-11 17:54:40 +08:00
parent a6b6e5f822
commit 74266a126c
34 changed files with 7105 additions and 470 deletions
+3 -1
View File
@@ -10,7 +10,9 @@ try:
generation_poll_tasks,
generation_download_tasks,
generation_recovery_tasks,
hot_opening_replicate_tasks
hot_opening_replicate_tasks,
shot_replicate_tasks,
shot_replicate_flow_tasks
)
except Exception:
pass
+112 -14
View File
@@ -1,23 +1,35 @@
from __future__ import annotations
import asyncio
import os
import threading
from concurrent.futures import Future
from typing import Awaitable, TypeVar
from app.config import settings
T = TypeVar("T")
_thread_local = threading.local()
_single_loop_lock = threading.RLock()
_single_loop: asyncio.AbstractEventLoop | None = None
_single_loop_thread: threading.Thread | None = None
_single_loop_pid: int | None = None
_single_loop_ready: threading.Event | None = None
def _get_or_create_loop() -> asyncio.AbstractEventLoop:
"""
给当前进程/线程维护一个长期 event loop。
def _runner_mode() -> str:
mode = str(getattr(settings, "CELERY_ASYNC_RUNNER_MODE", "single_loop") or "single_loop").strip().lower()
if mode not in {"single_loop", "direct"}:
return "single_loop"
return mode
Linux prefork
每个 Celery 子进程通常单线程跑任务,这里相当于每个子进程一个长期 loop。
Windows -P threads
每个线程一个 loop,但注意 asyncpg pool 仍不适合跨线程共享;
Windows threads 模式建议继续用 NullPool 或只做本地调试。
def _get_or_create_thread_local_loop() -> asyncio.AbstractEventLoop:
"""兼容旧方案:当前线程持有一个长期 event loop。
仅作为降级模式使用。长期推荐 single_loop,避免 Windows threads 下
多线程 event loop 复用 asyncpg / redis.asyncio 连接对象。
"""
pid = os.getpid()
loop = getattr(_thread_local, "loop", None)
@@ -31,18 +43,104 @@ def _get_or_create_loop() -> asyncio.AbstractEventLoop:
return loop
def _single_loop_worker(loop: asyncio.AbstractEventLoop, ready: threading.Event) -> None:
asyncio.set_event_loop(loop)
ready.set()
loop.run_forever()
pending = [task for task in asyncio.all_tasks(loop) if not task.done()]
if pending:
for task in pending:
task.cancel()
loop.run_until_complete(asyncio.gather(*pending, return_exceptions=True))
loop.run_until_complete(loop.shutdown_asyncgens())
loop.close()
def _get_or_create_single_loop() -> asyncio.AbstractEventLoop:
"""获取当前 Celery 进程内唯一 async event loop。
Linux prefork:每个 Celery 子进程各自一个 loop。
Windows threads:同一 worker 进程内所有任务线程共享同一个 loop。
"""
global _single_loop, _single_loop_thread, _single_loop_pid, _single_loop_ready
pid = os.getpid()
with _single_loop_lock:
if (
_single_loop is not None
and not _single_loop.is_closed()
and _single_loop_thread is not None
and _single_loop_thread.is_alive()
and _single_loop_pid == pid
):
return _single_loop
# fork 后 pid 变化,必须丢弃父进程状态,重新创建子进程自己的 loop。
_single_loop = asyncio.new_event_loop()
_single_loop_pid = pid
_single_loop_ready = threading.Event()
_single_loop_thread = threading.Thread(
target=_single_loop_worker,
args=(_single_loop, _single_loop_ready),
name=f"celery-async-runner-{pid}",
daemon=True,
)
_single_loop_thread.start()
_single_loop_ready.wait(timeout=5)
return _single_loop
def run_async(coro: Awaitable[T]) -> T:
"""Celery 同步 task 调用异步协程的统一入口。
默认 single_loop 模式:
- 一个 Celery 子进程只有一个专用 event loop;
- 所有 asyncpg / redis.asyncio 操作都在这个 loop 内创建和使用;
- 避免 got Future attached to a different loop。
降级 direct 模式:
- 兼容旧的线程本地 loop 方案;
- 如果使用 direct,建议同时开启 CELERY_DB_USE_NULLPOOL=true。
"""
Celery 同步 task 调用异步协程的统一入口。
不使用 asyncio.run(),避免每个 task 结束时关闭 event loop。
"""
loop = _get_or_create_loop()
return loop.run_until_complete(coro)
if _runner_mode() == "direct":
loop = _get_or_create_thread_local_loop()
return loop.run_until_complete(coro)
loop = _get_or_create_single_loop()
try:
running_loop = asyncio.get_running_loop()
except RuntimeError:
running_loop = None
if running_loop is loop:
raise RuntimeError("run_async() 不能在 Celery async_runner 的事件循环内部被同步调用")
future: Future[T] = asyncio.run_coroutine_threadsafe(coro, loop)
return future.result()
def close_loop() -> None:
"""关闭当前进程内 async runner loop。"""
global _single_loop, _single_loop_thread, _single_loop_pid, _single_loop_ready
# 关闭 single_loop。
with _single_loop_lock:
loop = _single_loop
thread = _single_loop_thread
if loop is not None and not loop.is_closed() and thread is not None and thread.is_alive():
loop.call_soon_threadsafe(loop.stop)
thread.join(timeout=5)
_single_loop = None
_single_loop_thread = None
_single_loop_pid = None
_single_loop_ready = None
# 关闭 direct 降级模式的线程本地 loop。
loop = getattr(_thread_local, "loop", None)
if loop is not None and not loop.is_closed():
loop.close()
_thread_local.loop = None
_thread_local.pid = None
_thread_local.pid = None
+39 -5
View File
@@ -51,6 +51,12 @@ if broker_url:
"generation.download_generation_result_task": {"queue": "gen_result_download"},
"hot_opening.start_image_prompt_optimize": {"queue": "gen_chatapi_create"},
"hot_opening.start_video_prompt_optimize": {"queue": "gen_chatapi_create"},
"shot_replicate.analyze_original_video": {"queue": "gen_chatapi_create"},
"shot_replicate.analyze_custom_segment_video": {"queue": "gen_chatapi_create"},
"shot_replicate.split_one_segment": {"queue": "gen_result_download"},
"shot_replicate.start_image_prompt_optimize": {"queue": "gen_chatapi_create"},
"shot_replicate.start_video_prompt_optimize": {"queue": "gen_chatapi_create"},
"shot_replicate.recover_split_tasks_once": {"queue": "gen_result_download"},
"generation.recover_download_tasks_once": {"queue": "gen_result_download"},
"generation.recover_generation_tasks_once": {"queue": "gen_result_download"},
"app.tasks.cleanup.*": {"queue": "default"},
@@ -61,6 +67,18 @@ else:
celery_app = None
async def _try_acquire_startup_recovery_lock() -> bool:
"""任意 worker 启动时都可尝试抢恢复锁,避免依赖 hostname 命名。"""
from app.services.redis_registry_service import redis_acquire_lock
token = await redis_acquire_lock(
lock_key=settings.CELERY_STARTUP_RECOVERY_LOCK_KEY,
ttl_seconds=int(settings.CELERY_STARTUP_RECOVERY_LOCK_TTL_SECONDS or 120),
log_context="celery_startup_recovery",
)
return bool(token)
@worker_ready.connect
def on_worker_ready(sender=None, **kwargs):
"""Celery worker 启动时做一次容灾恢复。
@@ -68,13 +86,21 @@ def on_worker_ready(sender=None, **kwargs):
注意:
- 不启用 Celery beat。
- 不要求新增第四条启动命令。
- 只让 gen_result_download worker 投递恢复任务,避免三个 worker 同时重复扫描
- 不再依赖 worker hostname 是否包含 gen_result_download
- 所有 worker 都尝试抢 Redis 锁,只有抢到锁的 worker 投递恢复任务。
"""
if celery_app is None:
return
if not bool(getattr(settings, "CELERY_STARTUP_RECOVERY_ENABLED", True)):
logger.info("启动容灾恢复已关闭。CELERY_STARTUP_RECOVERY_ENABLED=false")
return
hostname = str(getattr(sender, "hostname", "") or "")
if "gen_result_download" not in hostname:
try:
if not run_async(_try_acquire_startup_recovery_lock()):
return
except Exception:
# Redis 不可用时不阻塞 worker 启动,避免影响稳定生成链路。
logger.exception("启动容灾恢复锁获取失败,已跳过本次自动恢复投递")
return
try:
@@ -82,6 +108,7 @@ def on_worker_ready(sender=None, **kwargs):
recover_download_tasks_once,
recover_generation_tasks_once,
)
from app.tasks.shot_replicate_tasks import recover_split_tasks_once
countdown = max(0, int(settings.DOWNLOAD_RECOVERY_STARTUP_DELAY_SECONDS or 0))
@@ -95,6 +122,13 @@ def on_worker_ready(sender=None, **kwargs):
queue="gen_result_download",
priority=settings.DOWNLOAD_TASK_PRIORITY_RECOVER,
)
recover_split_tasks_once.apply_async(
countdown=countdown + 10,
queue="gen_result_download",
priority=settings.DOWNLOAD_TASK_PRIORITY_RECOVER,
)
logger.info("启动容灾恢复任务已投递。countdown=%s", countdown)
except Exception:
logger.exception("启动容灾恢复任务投递失败")
@@ -110,14 +144,14 @@ def on_worker_process_init(**kwargs):
@worker_process_shutdown.connect
def on_worker_process_shutdown(**kwargs):
"""子进程退出前关闭连接池和 event loop。"""
"""子进程退出前关闭连接池、Redis 注册表连接和 event loop。"""
try:
run_async(engine.dispose())
except Exception:
pass
try:
from app.services.celery_download_recovery_service import close_registry_redis
from app.services.redis_registry_service import close_registry_redis
run_async(close_registry_redis())
except Exception:
@@ -11,9 +11,10 @@ from app.services.error_codes import extract_error_message
from app.services.generation_log_service import log_task_event
from app.services.generation_refund_service import mark_chat_generation_task_failed_and_refund_once
from app.services.generation_provider_service import create_provider_task
from app.services.redis_registry_service import ensure_aware_utc
from app.tasks.celery_app import celery_app
ALLOWED_GENERATION_MODES = {"chatapi_async", "hot_opening_replicate"}
ALLOWED_GENERATION_MODES = {"chatapi_async", "hot_opening_replicate", "shot_replicate"}
def _get_first_value(obj: Any, *field_names: str) -> Optional[Any]:
@@ -69,7 +70,7 @@ def _build_optimized_prompt_by_params(task: ChatGenerationTask) -> str:
# 爆款开头复刻第5步的视频生成,original_prompt 已经是视频提词 JSON schema。
# 不能再追加“时长/比例/分辨率”中文参数,否则会污染 schema。
if generation_mode == "hot_opening_replicate" and gen_type == "video":
if generation_mode in {"hot_opening_replicate", "shot_replicate"} and gen_type == "video":
stripped = base_prompt.strip()
if stripped.startswith("{") or stripped.startswith("["):
return base_prompt
@@ -129,7 +130,8 @@ async def _run(task_id: str):
if task.status != "generating":
return
if task.deadline_at and datetime.now(timezone.utc) > task.deadline_at:
deadline_at = ensure_aware_utc(task.deadline_at)
if deadline_at and datetime.now(timezone.utc) > deadline_at:
await mark_chat_generation_task_failed_and_refund_once(
db,
task=task,
@@ -242,7 +244,11 @@ async def _run(task_id: str):
else:
from app.tasks.generation_poll_tasks import poll_generation_task
poll_generation_task.delay(task.id)
poll_generation_task.apply_async(
args=[task.id],
queue="gen_provider_poll",
countdown=0,
)
except Exception as exc:
try:
@@ -21,7 +21,7 @@ from app.services.generation_refund_service import mark_chat_generation_task_fai
from app.services.resource_accounting_service import record_chat_task_generated_resource
from app.tasks.celery_app import celery_app
ALLOWED_GENERATION_MODES = {"chatapi_async", "hot_opening_replicate"}
ALLOWED_GENERATION_MODES = {"chatapi_async", "hot_opening_replicate", "shot_replicate"}
DOWNLOAD_QUEUE = "gen_result_download"
DOWNLOAD_STAGE_QUEUED = "download_queued"
+170 -57
View File
@@ -1,6 +1,7 @@
from app.tasks.async_runner import run_async
import json
from datetime import datetime, timezone
from datetime import datetime, timedelta, timezone
from typing import Any
from sqlalchemy import select
@@ -11,9 +12,21 @@ from app.services.error_codes import extract_error_message
from app.services.generation_log_service import log_task_event, log_provider_call
from app.services.generation_refund_service import mark_chat_generation_task_failed_and_refund_once
from app.services.generation_provider_service import poll_provider_task
from app.services.redis_registry_service import (
datetime_to_epoch,
ensure_aware_utc,
redis_remove_registry_item,
redis_upsert_registry_item,
utc_now,
)
from app.tasks.celery_app import celery_app
ALLOWED_GENERATION_MODES = {"chatapi_async", "hot_opening_replicate"}
ALLOWED_GENERATION_MODES = {"chatapi_async", "hot_opening_replicate", "shot_replicate"}
POLL_QUEUE = "gen_provider_poll"
def _now() -> datetime:
return datetime.now(timezone.utc)
def _is_success(status: str) -> bool:
@@ -31,6 +44,86 @@ def _engine_snapshot(task: ChatGenerationTask) -> dict:
return {}
def _deadline_expired(task: ChatGenerationTask, now: datetime | None = None) -> bool:
deadline_at = ensure_aware_utc(task.deadline_at)
return bool(deadline_at and deadline_at <= (now or _now()))
def _poll_check_at(*, delay_seconds: int | float | None = None, now: datetime | None = None) -> datetime:
current_time = now or _now()
delay = int(delay_seconds or settings.CHATAPI_ASYNC_POLL_INTERVAL_SECONDS or 30)
grace = int(settings.POLL_TASK_QUEUE_TIMEOUT_SECONDS or 120)
return current_time + timedelta(seconds=max(1, delay) + max(0, grace))
def _poll_lease_until(now: datetime | None = None) -> datetime:
current_time = now or _now()
return current_time + timedelta(seconds=int(settings.POLL_TASK_LEASE_SECONDS or 300))
def _build_poll_active_payload(
task: ChatGenerationTask,
*,
stage: str,
reason: str,
next_poll_at: datetime | None = None,
check_at: datetime | None = None,
) -> dict[str, Any]:
current_time = utc_now()
checked_next_poll_at = ensure_aware_utc(next_poll_at)
checked_check_at = ensure_aware_utc(check_at)
return {
"task_id": task.id,
"provider_task_id": task.provider_task_id,
"seedance_task_id": task.seedance_task_id,
"generation_mode": task.generation_mode,
"gen_type": task.gen_type,
"stage": stage,
"queue": POLL_QUEUE,
"poll_count": int(task.poll_count or 0),
"retry_count": int(task.retry_count or 0),
"last_poll_at": datetime_to_epoch(task.last_poll_at) if task.last_poll_at else None,
"next_poll_at": datetime_to_epoch(checked_next_poll_at) if checked_next_poll_at else None,
"deadline_at": datetime_to_epoch(task.deadline_at) if task.deadline_at else None,
"check_at": datetime_to_epoch(checked_check_at) if checked_check_at else None,
"updated_at": datetime_to_epoch(current_time),
"reason": reason,
}
async def register_poll_active(
task: ChatGenerationTask,
*,
check_at: datetime,
reason: str,
next_poll_at: datetime | None = None,
) -> None:
payload = _build_poll_active_payload(
task,
stage=task.pipeline_stage or "",
reason=reason,
next_poll_at=next_poll_at,
check_at=check_at,
)
await redis_upsert_registry_item(
hash_key=settings.POLL_ACTIVE_REDIS_HASH_KEY,
zset_key=settings.POLL_ACTIVE_REDIS_ZSET_KEY,
item_id=task.id,
payload=payload,
check_at=check_at,
log_context="poll_active",
)
async def remove_poll_active(task_id: str) -> None:
await redis_remove_registry_item(
hash_key=settings.POLL_ACTIVE_REDIS_HASH_KEY,
zset_key=settings.POLL_ACTIVE_REDIS_ZSET_KEY,
item_id=task_id,
log_context="poll_active",
)
async def _notify_finished(db, task: ChatGenerationTask) -> None:
from app.services.generation_module_hook_service import notify_chat_generation_task_finished
@@ -55,6 +148,32 @@ async def _reload_task(db, task_id: str) -> ChatGenerationTask | None:
return result.scalar_one_or_none()
async def _mark_timeout(db, task: ChatGenerationTask, *, message: str = "任务轮询超时") -> None:
await mark_chat_generation_task_failed_and_refund_once(
db,
task=task,
error_message=message,
pipeline_stage="timeout",
)
await _notify_finished(db, task)
await db.commit()
await remove_poll_active(task.id)
await log_task_event(task, event_type="TASK_TIMEOUT", to_status="failed", to_stage="timeout")
async def _mark_failed(db, task: ChatGenerationTask, *, message: str, detail: Any = None) -> None:
await mark_chat_generation_task_failed_and_refund_once(
db,
task=task,
error_message=message,
pipeline_stage="failed",
)
await _notify_finished(db, task)
await db.commit()
await remove_poll_active(task.id)
await log_task_event(task, event_type="POLL_FAILED", message=task.error_message, detail=detail)
async def _run(task_id: str):
async with async_session() as db:
result = await db.execute(select(ChatGenerationTask).where(
@@ -62,43 +181,37 @@ async def _run(task_id: str):
ChatGenerationTask.deleted_at.is_(None),
).with_for_update().limit(1))
task = result.scalar_one_or_none()
if not task or task.generation_mode not in ALLOWED_GENERATION_MODES:
if not task:
await remove_poll_active(task_id)
return
if task.generation_mode not in ALLOWED_GENERATION_MODES:
await remove_poll_active(task.id)
return
# 只处理正在生成,且处于远程等待/轮询中的任务。
if task.status != "generating" or task.pipeline_stage not in ("waiting_remote", "polling"):
await remove_poll_active(task.id)
return
if task.deadline_at and datetime.now(timezone.utc) > task.deadline_at:
await mark_chat_generation_task_failed_and_refund_once(
db,
task=task,
error_message="任务轮询超时",
pipeline_stage="timeout",
)
await _notify_finished(db, task)
await db.commit()
await log_task_event(task, event_type="TASK_TIMEOUT", to_status="failed", to_stage="timeout")
if _deadline_expired(task):
await _mark_timeout(db, task, message="任务轮询超时")
return
if not (task.seedance_task_id or task.provider_task_id):
await mark_chat_generation_task_failed_and_refund_once(
db,
task=task,
error_message="缺少外部任务ID",
pipeline_stage="failed",
)
await _notify_finished(db, task)
await db.commit()
await log_task_event(task, event_type="POLL_FAILED", message=task.error_message)
await _mark_failed(db, task, message="缺少外部任务ID")
return
# 标记本次正在轮询。
# 注意:pending 后会再改回 waiting_remote,避免任务长期卡在 polling
# 标记本次正在轮询,并登记 poll lease
# 如果 worker 在供应商接口调用过程中退出,启动恢复会在 lease 过期后重新投递
task.pipeline_stage = "polling"
task.poll_count = (task.poll_count or 0) + 1
task.last_poll_at = datetime.now(timezone.utc)
task.last_poll_at = _now()
await db.commit()
await register_poll_active(
task,
check_at=_poll_lease_until(task.last_poll_at),
reason="polling_lease",
)
try:
poll_result = await poll_provider_task(db, task)
@@ -134,20 +247,13 @@ async def _run(task_id: str):
task.provider_response_json = response_data
if not task.remote_result_url:
await mark_chat_generation_task_failed_and_refund_once(
db,
task=task,
error_message="供应商任务成功但未返回结果URL",
pipeline_stage="failed",
)
await _notify_finished(db, task)
await db.commit()
await log_task_event(task, event_type="POLL_FAILED", message=task.error_message)
await _mark_failed(db, task, message="供应商任务成功但未返回结果URL", detail=poll_result)
return
task.pipeline_stage = "result_ready"
task.retry_count = 0
await db.commit()
await remove_poll_active(task.id)
await log_task_event(task, event_type="POLL_SUCCESS", to_stage="result_ready")
@@ -158,34 +264,38 @@ async def _run(task_id: str):
if _is_failed(status):
task.provider_response_json = response_data
await mark_chat_generation_task_failed_and_refund_once(
await _mark_failed(
db,
task=task,
error_message=poll_result.get("error") or f"供应商任务失败: {status}",
pipeline_stage="failed",
task,
message=poll_result.get("error") or f"供应商任务失败: {status}",
detail=poll_result,
)
await _notify_finished(db, task)
await db.commit()
await log_task_event(task, event_type="POLL_FAILED", message=task.error_message, detail=poll_result)
return
# 关键修改 1
# 供应商仍在 pending / running 时,把阶段从 polling 改回 waiting_remote。
# 这样数据库状态表示“等待下一次轮询”,不会长期停在 polling
# 同时可以降低重复 Celery 消息形成多条轮询链的概率。
# 同时登记下一次 poll activeCelery countdown 丢失时可由恢复任务拉起
task.pipeline_stage = "waiting_remote"
task.retry_count = 0
await db.commit()
await log_task_event(task, event_type="POLL_PENDING", message=f"status={status}")
delay_seconds = int(settings.CHATAPI_ASYNC_POLL_INTERVAL_SECONDS or 30)
next_poll_at = _now() + timedelta(seconds=max(1, delay_seconds))
await register_poll_active(
task,
check_at=_poll_check_at(delay_seconds=delay_seconds),
next_poll_at=next_poll_at,
reason="poll_pending_next",
)
poll_generation_task.apply_async(
args=[task.id],
countdown=settings.CHATAPI_ASYNC_POLL_INTERVAL_SECONDS,
queue=POLL_QUEUE,
countdown=delay_seconds,
)
except Exception as exc:
# 关键修改 2
# 异常后先 rollback,再重新查询 task,不继续使用 rollback 前的旧 ORM 对象。
try:
await db.rollback()
@@ -194,30 +304,33 @@ async def _run(task_id: str):
task = await _reload_task(db, task_id)
if not task:
await remove_poll_active(task_id)
return
task.retry_count = (task.retry_count or 0) + 1
if task.retry_count > settings.CHATAPI_ASYNC_MAX_RETRIES:
error_message = extract_error_message(exc, "轮询") if callable(extract_error_message) else str(exc)
await mark_chat_generation_task_failed_and_refund_once(
db,
task=task,
error_message=error_message,
pipeline_stage="failed",
)
await _notify_finished(db, task)
await db.commit()
await log_task_event(task, event_type="POLL_FAILED", message=task.error_message)
await _mark_failed(db, task, message=error_message)
else:
# 临时轮询异常时,不让任务停在 polling。
# 回到 waiting_remote,等待下一次重试轮询。
task.pipeline_stage = "waiting_remote"
await db.commit()
delay_seconds = int(settings.CHATAPI_ASYNC_RETRY_BACKOFF_SECONDS or 30) * int(task.retry_count or 1)
next_poll_at = _now() + timedelta(seconds=max(1, delay_seconds))
await register_poll_active(
task,
check_at=_poll_check_at(delay_seconds=delay_seconds),
next_poll_at=next_poll_at,
reason="poll_exception_retry",
)
poll_generation_task.apply_async(
args=[task.id],
countdown=settings.CHATAPI_ASYNC_RETRY_BACKOFF_SECONDS * task.retry_count,
queue=POLL_QUEUE,
countdown=delay_seconds,
)
@@ -233,4 +346,4 @@ else:
def apply_async(self, *args, **kwargs):
raise RuntimeError("Celery is disabled")
poll_generation_task = _DisabledTask()
poll_generation_task = _DisabledTask()
@@ -0,0 +1,42 @@
from __future__ import annotations
from app.models.base import async_session
from app.services.shot_replicate_flow_service import run_image_prompt_optimize, run_video_prompt_optimize
from app.tasks.async_runner import run_async
from app.tasks.celery_app import celery_app
async def _run_image_prompt(project_id: str, step_id: str | None = None):
async with async_session() as db:
await run_image_prompt_optimize(db, project_id=project_id, step_id=step_id)
await db.commit()
async def _run_video_prompt(project_id: str, step_id: str | None = None):
async with async_session() as db:
await run_video_prompt_optimize(db, project_id=project_id, step_id=step_id)
await db.commit()
if celery_app:
@celery_app.task(name="shot_replicate.start_image_prompt_optimize", bind=True, max_retries=3, default_retry_delay=30)
def start_image_prompt_optimize(self, project_id: str, step_id: str | None = None):
return run_async(_run_image_prompt(project_id, step_id))
@celery_app.task(name="shot_replicate.start_video_prompt_optimize", bind=True, max_retries=3, default_retry_delay=30)
def start_video_prompt_optimize(self, project_id: str, step_id: str | None = None):
return run_async(_run_video_prompt(project_id, step_id))
else:
class _DisabledTask:
def delay(self, *args, **kwargs):
raise RuntimeError("Celery is disabled")
def apply_async(self, *args, **kwargs):
raise RuntimeError("Celery is disabled")
start_image_prompt_optimize = _DisabledTask()
start_video_prompt_optimize = _DisabledTask()
@@ -0,0 +1,505 @@
from __future__ import annotations
import logging
from datetime import datetime, timedelta, timezone
from typing import Any
from sqlalchemy import select
from app.config import settings
from app.enums.shot_replicate import (
ModuleCodeEnum,
ShotAnalysisStatusEnum,
ShotSegmentAnalysisStatusEnum,
ShotSegmentSourceModeEnum,
ShotSplitStatusEnum,
ShotTaskSetStatusEnum,
)
from app.models.base import async_session
from app.models.shot_replicate_segment import ShotReplicateSegment
from app.models.shot_replicate_task_set import ShotReplicateTaskSet
from app.services.module_generation_log_service import log_module_error, log_module_event_file, log_module_prompt_event
from app.services.redis_registry_service import redis_acquire_lock, redis_release_lock
from app.services.shot_replicate_taskset_service import refresh_task_set_split_summary
from app.services.shot_video_analysis_service import analyze_video_for_shot_split
from app.services.shot_video_split_service import split_video_segment_async
from app.services.upload_video_asset_service import validate_split_range
from app.tasks.async_runner import run_async
from app.tasks.celery_app import celery_app
logger = logging.getLogger("video_gen")
MODULE = ModuleCodeEnum.SHOT_REPLICATE.value
SPLIT_QUEUE = "gen_result_download"
ANALYSIS_QUEUE = "gen_chatapi_create"
def _now() -> datetime:
return datetime.now(timezone.utc)
def _lease_until(now: datetime | None = None) -> datetime:
return (now or _now()) + timedelta(seconds=int(settings.SHOT_SPLIT_LEASE_SECONDS or 600))
def _retry_at(attempt: int, now: datetime | None = None) -> datetime:
base = int(settings.SHOT_SPLIT_RETRY_BACKOFF_SECONDS or settings.DOWNLOAD_TASK_RETRY_BACKOFF_SECONDS or 30)
return (now or _now()) + timedelta(seconds=max(1, base * max(1, attempt)))
async def _acquire_split_semaphore(segment_id: str) -> str | None:
"""简单 Redis 并发闸门:用固定槽位锁限制 ffmpeg 同时运行数量。"""
max_concurrent = max(1, int(settings.SHOT_SPLIT_MAX_CONCURRENT or 1))
ttl = int(settings.SHOT_SPLIT_LEASE_SECONDS or 600)
for slot in range(max_concurrent):
key = f"{settings.SHOT_SPLIT_SEMAPHORE_KEY_PREFIX}:{slot}"
token = await redis_acquire_lock(lock_key=key, ttl_seconds=ttl, token=segment_id, log_context="shot_split_semaphore")
if token:
return key
return None
async def _release_split_semaphore(lock_key: str | None, segment_id: str) -> None:
if lock_key:
await redis_release_lock(lock_key=lock_key, token=segment_id, log_context="shot_split_semaphore")
async def _run_analyze_original_video(task_set_id: str) -> None:
task_set_user_id: str | None = None
video_url: str | None = None
try:
async with async_session() as db:
result = await db.execute(
select(ShotReplicateTaskSet)
.where(ShotReplicateTaskSet.id == task_set_id, ShotReplicateTaskSet.deleted_at.is_(None))
.with_for_update()
.limit(1)
)
task_set = result.scalar_one_or_none()
if not task_set:
return
if task_set.analysis_status == ShotAnalysisStatusEnum.COMPLETED.value:
return
task_set_user_id = task_set.user_id
video_url = task_set.video_url
task_set.status = ShotTaskSetStatusEnum.ANALYZING.value
task_set.analysis_status = ShotAnalysisStatusEnum.PROCESSING.value
task_set.analysis_error_message = None
await db.commit()
log_module_event_file(
module=MODULE,
event_type="SHOT_ANALYSIS_STARTED",
project_id=task_set_id,
user_id=task_set_user_id,
message="原视频拆镜分析开始",
detail={"task_set_id": task_set_id, "video_url": video_url, "analysis_mode": "full_breakdown"},
)
async with async_session() as db:
analyzed = await analyze_video_for_shot_split(db, video_url or "", user_id=task_set_user_id, mode="full_breakdown")
result = await db.execute(
select(ShotReplicateTaskSet)
.where(ShotReplicateTaskSet.id == task_set_id, ShotReplicateTaskSet.deleted_at.is_(None))
.with_for_update()
.limit(1)
)
task_set = result.scalar_one_or_none()
if not task_set:
await db.rollback()
return
result_json = analyzed.result
task_set.original_video_content = str(result_json.get("原视频内容") or "")
task_set.original_video_category = str(result_json.get("原视频分类") or "")
task_set.original_video_audience = str(result_json.get("原视频受众人群") or "")
task_set.ai_suggestion_json = result_json.get("拆镜内容剖析") or []
task_set.analysis_raw_json = analyzed.raw_response
task_set.analysis_result_json = result_json
task_set.analysis_status = ShotAnalysisStatusEnum.COMPLETED.value
task_set.status = ShotTaskSetStatusEnum.ANALYSIS_COMPLETED.value
task_set.analysis_error_message = None
await db.commit()
log_module_prompt_event(
event_type="SHOT_ANALYSIS_SUCCESS",
project_id=task_set_id,
step_id=task_set_id,
user_id=task_set_user_id or "",
module=MODULE,
prompt_type="shot_video_analysis",
request=analyzed.usage.get("log_request") if isinstance(analyzed.usage, dict) else {},
response=analyzed.result,
token_usage=analyzed.usage,
)
log_module_event_file(
module=MODULE,
event_type="SHOT_ANALYSIS_SUCCESS",
project_id=task_set_id,
user_id=task_set_user_id,
message="原视频拆镜分析成功",
detail={"suggestion_count": len(analyzed.result.get("拆镜内容剖析") or []), "token_usage": analyzed.usage},
)
except Exception as exc:
async with async_session() as db:
result = await db.execute(
select(ShotReplicateTaskSet)
.where(ShotReplicateTaskSet.id == task_set_id, ShotReplicateTaskSet.deleted_at.is_(None))
.with_for_update()
.limit(1)
)
task_set = result.scalar_one_or_none()
if task_set:
task_set_user_id = task_set_user_id or task_set.user_id
task_set.status = ShotTaskSetStatusEnum.ANALYSIS_FAILED.value
task_set.analysis_status = ShotAnalysisStatusEnum.FAILED.value
task_set.analysis_error_message = str(exc)
await db.commit()
log_module_error(
module=MODULE,
event_type="SHOT_ANALYSIS_FAILED",
project_id=task_set_id,
user_id=task_set_user_id,
message="原视频拆镜分析失败",
detail={"task_set_id": task_set_id, "video_url": video_url, "analysis_mode": "full_breakdown"},
exc=exc,
)
async def _run_analyze_custom_segment_video(segment_id: str) -> None:
user_id: str | None = None
task_set_id: str | None = None
video_url: str | None = None
try:
async with async_session() as db:
result = await db.execute(
select(ShotReplicateSegment)
.where(ShotReplicateSegment.id == segment_id, ShotReplicateSegment.deleted_at.is_(None))
.with_for_update()
.limit(1)
)
segment = result.scalar_one_or_none()
if not segment or not segment.segment_video_url:
return
if segment.analysis_status == ShotSegmentAnalysisStatusEnum.COMPLETED.value:
return
user_id = segment.user_id
task_set_id = segment.task_set_id
video_url = segment.segment_video_url
segment.analysis_status = ShotSegmentAnalysisStatusEnum.PROCESSING.value
segment.analysis_error_message = None
await db.commit()
log_module_event_file(
module=MODULE,
event_type="SHOT_SEGMENT_ANALYSIS_STARTED",
project_id=task_set_id,
step_id=segment_id,
user_id=user_id,
message="自定义拆镜片段分析开始",
detail={"segment_id": segment_id, "task_set_id": task_set_id, "video_url": video_url, "analysis_mode": "summary_only"},
)
async with async_session() as db:
analyzed = await analyze_video_for_shot_split(db, video_url or "", user_id=user_id, mode="summary_only")
result = await db.execute(
select(ShotReplicateSegment)
.where(ShotReplicateSegment.id == segment_id, ShotReplicateSegment.deleted_at.is_(None))
.with_for_update()
.limit(1)
)
segment = result.scalar_one_or_none()
if not segment:
await db.rollback()
return
result_json = analyzed.result
segment.original_video_content = str(result_json.get("原视频内容") or "")
segment.original_video_category = str(result_json.get("原视频分类") or "")
segment.original_video_audience = str(result_json.get("原视频受众人群") or "")
segment.segment_content = segment.original_video_content
segment.segment_category = segment.original_video_category
segment.segment_audience = segment.original_video_audience
segment.analysis_json = result_json
segment.analysis_status = ShotSegmentAnalysisStatusEnum.COMPLETED.value
segment.analysis_error_message = None
await db.commit()
log_module_prompt_event(
event_type="SHOT_SEGMENT_ANALYSIS_SUCCESS",
project_id=task_set_id or segment_id,
step_id=segment_id,
user_id=user_id or "",
module=MODULE,
prompt_type="shot_segment_analysis",
request=analyzed.usage.get("log_request") if isinstance(analyzed.usage, dict) else {},
response=analyzed.result,
token_usage=analyzed.usage,
)
log_module_event_file(
module=MODULE,
event_type="SHOT_SEGMENT_ANALYSIS_SUCCESS",
project_id=task_set_id,
step_id=segment_id,
user_id=user_id,
message="自定义拆镜片段分析成功",
detail={"segment_id": segment_id, "task_set_id": task_set_id, "token_usage": analyzed.usage},
)
except Exception as exc:
async with async_session() as db:
result = await db.execute(
select(ShotReplicateSegment)
.where(ShotReplicateSegment.id == segment_id, ShotReplicateSegment.deleted_at.is_(None))
.with_for_update()
.limit(1)
)
segment = result.scalar_one_or_none()
if segment:
user_id = user_id or segment.user_id
task_set_id = task_set_id or segment.task_set_id
segment.analysis_status = ShotSegmentAnalysisStatusEnum.FAILED.value
segment.analysis_error_message = str(exc)
await db.commit()
log_module_error(
module=MODULE,
event_type="SHOT_SEGMENT_ANALYSIS_FAILED",
project_id=task_set_id,
step_id=segment_id,
user_id=user_id,
message="自定义拆镜片段分析失败",
detail={"segment_id": segment_id, "task_set_id": task_set_id, "video_url": video_url, "analysis_mode": "summary_only"},
exc=exc,
)
async def _run_split_one_segment(segment_id: str) -> None:
segment_lock_key = f"{settings.SHOT_SPLIT_LOCK_KEY_PREFIX}:{segment_id}"
segment_lock_token = await redis_acquire_lock(
lock_key=segment_lock_key,
ttl_seconds=int(settings.SHOT_SPLIT_LEASE_SECONDS or 600),
log_context="shot_split_segment_lock",
)
if not segment_lock_token:
return
semaphore_key: str | None = None
user_id: str | None = None
task_set_id: str | None = None
source_path: str | None = None
try:
semaphore_key = await _acquire_split_semaphore(segment_id)
if not semaphore_key:
log_module_event_file(
module=MODULE,
event_type="SHOT_SEGMENT_SPLIT_RETRY_WAITING",
step_id=segment_id,
message="拆镜 ffmpeg 并发闸门已满,稍后重试",
detail={"segment_id": segment_id, "reason": "semaphore_full"},
)
if celery_app:
split_one_segment.apply_async(args=[segment_id], queue=SPLIT_QUEUE, countdown=10, priority=settings.DOWNLOAD_TASK_PRIORITY_NORMAL)
return
async with async_session() as db:
result = await db.execute(
select(ShotReplicateSegment)
.where(ShotReplicateSegment.id == segment_id, ShotReplicateSegment.deleted_at.is_(None))
.with_for_update()
.limit(1)
)
segment = result.scalar_one_or_none()
if not segment:
return
user_id = segment.user_id
task_set_id = segment.task_set_id
task_set_result = await db.execute(
select(ShotReplicateTaskSet)
.where(ShotReplicateTaskSet.id == segment.task_set_id, ShotReplicateTaskSet.deleted_at.is_(None))
.with_for_update()
.limit(1)
)
task_set = task_set_result.scalar_one_or_none()
if not task_set:
return
if segment.split_status == ShotSplitStatusEnum.COMPLETED.value and segment.segment_video_url:
return
validate_split_range(
start_second=segment.start_second,
end_second=segment.end_second,
video_duration_seconds=task_set.video_duration_seconds,
)
now = _now()
segment.split_status = ShotSplitStatusEnum.PROCESSING.value
segment.split_started_at = now
segment.split_lease_until = _lease_until(now)
segment.split_retry_count = int(segment.split_retry_count or 0) + 1
segment.split_next_retry_at = None
segment.split_last_error = None
task_set.status = ShotTaskSetStatusEnum.SPLITTING.value
task_set.split_status = ShotSplitStatusEnum.PROCESSING.value
await db.commit()
source_path = task_set.video_path
date_dir = (segment.created_at or now).strftime("%Y/%m/%d")
start_second = segment.start_second
end_second = segment.end_second
attempt = segment.split_retry_count
log_module_event_file(
module=MODULE,
event_type="SHOT_SEGMENT_SPLIT_STARTED",
project_id=task_set_id,
step_id=segment_id,
user_id=user_id,
message="拆镜片段 ffmpeg 切割开始",
detail={
"segment_id": segment_id,
"task_set_id": task_set_id,
"source_path": source_path,
"start_second": start_second,
"end_second": end_second,
"attempt": attempt,
},
)
split_result = await split_video_segment_async(
source_path=source_path,
segment_id=segment_id,
start_second=start_second,
end_second=end_second,
date_dir=date_dir,
)
async with async_session() as db:
result = await db.execute(
select(ShotReplicateSegment)
.where(ShotReplicateSegment.id == segment_id, ShotReplicateSegment.deleted_at.is_(None))
.with_for_update()
.limit(1)
)
segment = result.scalar_one_or_none()
if not segment:
return
segment.segment_video_url = split_result.url
segment.segment_video_path = split_result.path
segment.split_status = ShotSplitStatusEnum.COMPLETED.value
segment.split_completed_at = _now()
segment.split_lease_until = None
segment.split_next_retry_at = None
segment.split_last_error = None
await refresh_task_set_split_summary(db, segment.task_set_id)
await db.commit()
log_module_event_file(
module=MODULE,
event_type="SHOT_SEGMENT_SPLIT_SUCCESS",
project_id=segment.task_set_id,
step_id=segment.id,
user_id=segment.user_id,
message="拆镜片段 ffmpeg 切割成功",
detail={
"segment_id": segment.id,
"task_set_id": segment.task_set_id,
"segment_video_url": split_result.url,
"segment_video_path": split_result.path,
"source_mode": segment.source_mode,
},
)
if segment.source_mode == ShotSegmentSourceModeEnum.CUSTOM.value and celery_app:
analyze_custom_segment_video.apply_async(args=[segment.id], queue=ANALYSIS_QUEUE, countdown=0)
except Exception as exc:
next_retry_delay: int | None = None
final_failed = False
async with async_session() as db:
result = await db.execute(
select(ShotReplicateSegment)
.where(ShotReplicateSegment.id == segment_id, ShotReplicateSegment.deleted_at.is_(None))
.with_for_update()
.limit(1)
)
segment = result.scalar_one_or_none()
if not segment:
return
user_id = user_id or segment.user_id
task_set_id = task_set_id or segment.task_set_id
attempt = int(segment.split_retry_count or 0)
segment.split_last_error = str(exc)
segment.split_lease_until = None
if attempt >= int(settings.SHOT_SPLIT_MAX_RETRY_COUNT or 3):
segment.split_status = ShotSplitStatusEnum.FAILED.value
segment.split_next_retry_at = None
final_failed = True
else:
segment.split_status = ShotSplitStatusEnum.RETRY_WAITING.value
segment.split_next_retry_at = _retry_at(attempt)
await refresh_task_set_split_summary(db, segment.task_set_id)
await db.commit()
if segment.split_status == ShotSplitStatusEnum.RETRY_WAITING.value and celery_app:
next_retry_delay = max(1, int(((segment.split_next_retry_at or _now()) - _now()).total_seconds()))
split_one_segment.apply_async(args=[segment_id], queue=SPLIT_QUEUE, countdown=next_retry_delay, priority=settings.DOWNLOAD_TASK_PRIORITY_RECOVER)
log_module_error(
module=MODULE,
event_type="SHOT_SEGMENT_SPLIT_FAILED" if final_failed else "SHOT_SEGMENT_SPLIT_RETRY_WAITING",
project_id=task_set_id,
step_id=segment_id,
user_id=user_id,
message="拆镜片段 ffmpeg 切割失败" if final_failed else "拆镜片段 ffmpeg 切割失败,等待重试",
detail={
"segment_id": segment_id,
"task_set_id": task_set_id,
"source_path": source_path,
"next_retry_delay_seconds": next_retry_delay,
"final_failed": final_failed,
},
exc=exc,
)
finally:
await _release_split_semaphore(semaphore_key, segment_id)
await redis_release_lock(lock_key=segment_lock_key, token=segment_lock_token, log_context="shot_split_segment_lock")
async def _run_recover_split_tasks_once() -> dict[str, Any]:
from app.services.shot_replicate_recovery_service import recover_shot_split_tasks_once
async with async_session() as db:
return await recover_shot_split_tasks_once(db)
if celery_app:
@celery_app.task(name="shot_replicate.analyze_original_video")
def analyze_original_video(task_set_id: str) -> None:
return run_async(_run_analyze_original_video(task_set_id))
@celery_app.task(name="shot_replicate.split_one_segment", bind=True, max_retries=0)
def split_one_segment(self, segment_id: str) -> None:
return run_async(_run_split_one_segment(segment_id))
@celery_app.task(name="shot_replicate.analyze_custom_segment_video")
def analyze_custom_segment_video(segment_id: str) -> None:
return run_async(_run_analyze_custom_segment_video(segment_id))
@celery_app.task(name="shot_replicate.recover_split_tasks_once")
def recover_split_tasks_once() -> dict[str, Any]:
return run_async(_run_recover_split_tasks_once())
else:
class _DisabledTask:
def delay(self, *args: Any, **kwargs: Any) -> None:
raise RuntimeError("Celery is disabled")
def apply_async(self, *args: Any, **kwargs: Any) -> None:
raise RuntimeError("Celery is disabled")
analyze_original_video = _DisabledTask()
split_one_segment = _DisabledTask()
analyze_custom_segment_video = _DisabledTask()
recover_split_tasks_once = _DisabledTask()