Files
xiaoxia-saas/packages/adapters/sqlalchemy_impl/generation_task_repository.py
T
xiaoxia 1712c1693a
CI/CD Pipeline / Check if frontend-only change (pull_request) Successful in 42s
CI/CD Pipeline / PR Build Worker Image (pull_request) Successful in 47s
CI/CD Pipeline / Validate - Type Check (mypy) (pull_request) Successful in 1m40s
PR Automation / Auto Merge on CI Green + Approved (pull_request) Successful in 1m44s
CI/CD Pipeline / Validate - Migration (alembic) (pull_request) Successful in 1m50s
AI Code Review / AI Code Review (pull_request) Failing after 2m8s
Preview Deploy / Deploy Preview Environment (pull_request) Successful in 2m16s
CI/CD Pipeline / PR Build API Image (pull_request) Successful in 1m41s
CI/CD Pipeline / Unit Tests (pull_request) Successful in 2m32s
PR Automation / Auto Approve on CI Green (pull_request) Successful in 3m28s
CI/CD Pipeline / Validate - Code Quality (pull_request) Has been cancelled
CI/CD Pipeline / Integration Tests (pull_request) Has been cancelled
CI/CD Pipeline / Build Production API Image (pull_request) Has been cancelled
CI/CD Pipeline / Build Production Web Image (pull_request) Has been cancelled
CI/CD Pipeline / Build Production Worker Image (pull_request) Has been cancelled
CI/CD Pipeline / Deploy Production (pull_request) Has been cancelled
CI/CD Pipeline / Production Browser E2E (pull_request) Has been cancelled
CI/CD Pipeline / Canary Release to Production (pull_request) Has been cancelled
CI/CD Pipeline / CI Gate (pull_request) Has been cancelled
CI/CD Pipeline / PR Build Web Image (pull_request) Failing after 578h45m36s
CI/CD Pipeline / Frontend Unit Tests (pull_request) Failing after 578h45m37s
CI/CD Pipeline / ACR Image Cleanup (pull_request) Failing after 578h46m16s
CI/CD Pipeline / Staging API Integration Tests (pull_request) Failing after 578h46m16s
CI/CD Pipeline / Deploy Staging (Watchtower auto-deploy) (pull_request) Failing after 578h46m18s
CI/CD Pipeline / Build Staging Worker Image (pull_request) Failing after 578h46m20s
CI/CD Pipeline / Build Staging API Image (pull_request) Failing after 578h46m20s
CI/CD Pipeline / Frontend Lint (pull_request) Failing after 579h19m25s
CI/CD Pipeline / Staging E2E Tests (pull_request) Failing after 579h20m3s
CI/CD Pipeline / Build Staging Web Image (pull_request) Failing after 579h20m7s
fix: address AI review - session leak & bulk update for pending cleanup
- _startup.py: move session.close() to finally block to prevent DB connection leak
- generation_task_repository.py: replace .all()+loop with bulk .update() to avoid OOM
  when large number of pending tasks accumulate
2026-08-23 12:39:49 +08:00

352 lines
14 KiB
Python
Executable File
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
from datetime import datetime, timezone
from sqlalchemy.orm import Session
from packages.adapters.sqlalchemy_impl.models import GenerationTaskModel
from packages.domain import GenerationTask
from packages.domain.generation_task import GenerationTaskStatus
def _to_domain(model: GenerationTaskModel) -> GenerationTask:
"""Convert ORM model to domain entity."""
return GenerationTask(
id=model.id,
project_id=model.project_id,
strategy_id=model.strategy_id,
asset_library_id=model.asset_library_id,
voice_library_id=model.voice_library_id,
template_id=model.template_id,
asset_ids=list(model.asset_ids or []),
title_ids=list(model.title_ids or []),
voice_ids=list(model.voice_ids or []),
status=GenerationTaskStatus(model.status) if model.status else GenerationTaskStatus.PENDING,
progress=model.progress,
result_count=int(model.result_count or 0),
error_message=model.error_message,
error_info=dict(model.error_info) if model.error_info else {},
retry_count=model.retry_count or 0,
auto_retry_enabled=bool(model.auto_retry_enabled),
auto_retry_max=model.auto_retry_max or 0,
started_at=model.started_at,
completed_at=model.completed_at,
created_by_user_id=model.created_by_user_id,
source_edit_plan_id=model.source_edit_plan_id or "",
asset_select_mode=model.asset_select_mode or "",
batch_id=model.batch_id or "",
video_title=getattr(model, "video_title", "") or "",
resolution=getattr(model, "resolution", "") or "",
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 "",
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 "",
custom_title=getattr(model, "custom_title", "") or "",
logs=model.logs or "[]",
created_at=model.created_at,
updated_at=model.updated_at,
)
class SQLAlchemyGenerationTaskRepository:
def __init__(self, session: Session):
self.session = session
def create(self, task: GenerationTask) -> GenerationTask:
model = GenerationTaskModel(
id=task.id,
project_id=task.project_id,
strategy_id=task.strategy_id,
asset_library_id=task.asset_library_id,
voice_library_id=task.voice_library_id,
template_id=task.template_id,
asset_ids=task.asset_ids,
title_ids=task.title_ids,
voice_ids=task.voice_ids,
status=task.status,
progress=task.progress,
result_count=task.result_count,
error_message=task.error_message,
error_info=task.error_info or None,
retry_count=task.retry_count or 0,
auto_retry_enabled=task.auto_retry_enabled,
auto_retry_max=task.auto_retry_max or 0,
started_at=task.started_at,
completed_at=task.completed_at,
created_by_user_id=task.created_by_user_id,
source_edit_plan_id=task.source_edit_plan_id or None,
asset_select_mode=task.asset_select_mode or "",
batch_id=task.batch_id or "",
video_title=task.video_title or "",
resolution=task.resolution or "",
bgm_config=task.bgm_config or {},
is_preview=task.is_preview or False,
source_task_id=task.source_task_id or "",
output_width=task.output_width,
output_height=task.output_height,
cover_url=task.cover_url or "",
custom_title=task.custom_title or "",
logs=task.logs,
created_at=task.created_at,
updated_at=task.updated_at,
)
self.session.add(model)
self.session.commit()
return task
def get(self, task_id: str) -> GenerationTask | None:
model = self.session.query(GenerationTaskModel).filter(GenerationTaskModel.id == task_id).first()
if model is None:
return None
return _to_domain(model)
def list_by_project(self, project_id: str) -> list[GenerationTask]:
models = (
self.session.query(GenerationTaskModel)
.filter(GenerationTaskModel.project_id == project_id)
.order_by(GenerationTaskModel.created_at.desc())
.all()
)
return [_to_domain(m) for m in models]
def list_by_user(self, user_id: str) -> list[GenerationTask]:
models = (
self.session.query(GenerationTaskModel)
.filter(GenerationTaskModel.created_by_user_id == user_id)
.order_by(GenerationTaskModel.created_at.desc())
.all()
)
return [_to_domain(m) for m in models]
def count_by_user(self, user_id: str) -> int:
return self.session.query(GenerationTaskModel).filter(GenerationTaskModel.created_by_user_id == user_id).count()
def count_pending_by_user(self, user_id: str) -> int:
return (
self.session.query(GenerationTaskModel)
.filter(
GenerationTaskModel.created_by_user_id == user_id,
GenerationTaskModel.status == GenerationTaskStatus.PENDING.value,
)
.count()
)
def count_pending_total(self) -> int:
return (
self.session.query(GenerationTaskModel)
.filter(GenerationTaskModel.status == GenerationTaskStatus.PENDING.value)
.count()
)
def list_recent_by_user(self, user_id: str, limit: int = 5) -> list[GenerationTask]:
models = (
self.session.query(GenerationTaskModel)
.filter(GenerationTaskModel.created_by_user_id == user_id)
.order_by(GenerationTaskModel.created_at.desc())
.limit(limit)
.all()
)
return [_to_domain(m) for m in models]
def list_latest_completed_preview(self, user_id: str, template_id: str, limit: int = 1) -> list[GenerationTask]:
"""按用户+模板查找最近已完成的预览任务。"""
models = (
self.session.query(GenerationTaskModel)
.filter(
GenerationTaskModel.created_by_user_id == user_id,
GenerationTaskModel.template_id == template_id,
GenerationTaskModel.is_preview,
GenerationTaskModel.status == "completed",
)
.order_by(GenerationTaskModel.created_at.desc())
.limit(limit)
.all()
)
return [_to_domain(m) for m in models]
def list_by_source_edit_plan(self, plan_id: str) -> list[GenerationTask]:
models = (
self.session.query(GenerationTaskModel)
.filter(GenerationTaskModel.source_edit_plan_id == plan_id)
.order_by(GenerationTaskModel.created_at.desc())
.all()
)
return [_to_domain(m) for m in models]
def list_by_user_filtered(
self,
user_id: str,
*,
status: str | None = None,
limit: int | None = None,
offset: int = 0,
) -> list[GenerationTask]:
"""按用户+状态筛选任务列表。"""
query = self.session.query(GenerationTaskModel).filter(GenerationTaskModel.created_by_user_id == user_id)
if status:
query = query.filter(GenerationTaskModel.status == status)
query = query.order_by(GenerationTaskModel.created_at.desc())
if offset:
query = query.offset(offset)
if limit:
query = query.limit(limit)
return [_to_domain(m) for m in query.all()]
def count_by_user_filtered(
self,
user_id: str,
*,
status: str | None = None,
) -> int:
"""按用户+状态筛选计数。"""
query = self.session.query(GenerationTaskModel).filter(GenerationTaskModel.created_by_user_id == user_id)
if status:
query = query.filter(GenerationTaskModel.status == status)
return query.count()
def list_by_project_filtered(
self,
project_id: str,
*,
status: str | None = None,
limit: int | None = None,
offset: int = 0,
) -> list[GenerationTask]:
"""按项目+状态筛选任务列表。"""
query = self.session.query(GenerationTaskModel).filter(GenerationTaskModel.project_id == project_id)
if status:
query = query.filter(GenerationTaskModel.status == status)
query = query.order_by(GenerationTaskModel.created_at.desc())
if offset:
query = query.offset(offset)
if limit:
query = query.limit(limit)
return [_to_domain(m) for m in query.all()]
def count_by_project_filtered(
self,
project_id: str,
*,
status: str | None = None,
) -> int:
"""按项目+状态筛选计数。"""
query = self.session.query(GenerationTaskModel).filter(GenerationTaskModel.project_id == project_id)
if status:
query = query.filter(GenerationTaskModel.status == status)
return query.count()
def update(self, task: GenerationTask) -> GenerationTask:
model = self.session.query(GenerationTaskModel).filter(GenerationTaskModel.id == task.id).first()
if model is None:
raise ValueError(f"GenerationTask {task.id} not found")
model.project_id = task.project_id
model.asset_library_id = task.asset_library_id
model.strategy_id = task.strategy_id
model.voice_library_id = task.voice_library_id
model.template_id = task.template_id
model.asset_ids = task.asset_ids
model.title_ids = task.title_ids
model.voice_ids = task.voice_ids
model.status = task.status
model.progress = task.progress
model.result_count = task.result_count
model.error_message = task.error_message
model.error_info = task.error_info or None
model.retry_count = task.retry_count or 0
model.auto_retry_enabled = task.auto_retry_enabled
model.auto_retry_max = task.auto_retry_max or 0
model.started_at = task.started_at
model.completed_at = task.completed_at
model.source_edit_plan_id = task.source_edit_plan_id or None
model.asset_select_mode = task.asset_select_mode or ""
model.batch_id = task.batch_id or ""
if hasattr(model, "video_title"):
model.video_title = task.video_title or ""
if hasattr(model, "resolution"):
model.resolution = task.resolution or ""
if hasattr(model, "bgm_config"):
model.bgm_config = task.bgm_config or {}
if hasattr(model, "is_preview"):
model.is_preview = task.is_preview or False
model.source_task_id = task.source_task_id or ""
model.output_width = task.output_width
model.output_height = task.output_height
model.cover_url = task.cover_url or ""
model.custom_title = task.custom_title or ""
model.logs = task.logs
self.session.commit()
return task
def cleanup_stale_running(self, timeout_minutes: int = 10) -> int:
"""清理超时未更新的 running 任务(孤儿任务)。
将 status=running 且 updated_at 超过 timeout_minutes 分钟未更新的任务
标记为 failed,error_message 标记为任务执行中断。
Returns:
清理的任务数量
"""
from datetime import timedelta
cutoff = datetime.now(timezone.utc) - timedelta(minutes=timeout_minutes)
models = (
self.session.query(GenerationTaskModel)
.filter(
GenerationTaskModel.status == GenerationTaskStatus.RUNNING.value,
GenerationTaskModel.updated_at < cutoff,
)
.all()
)
if not models:
return 0
for model in models:
model.status = GenerationTaskStatus.FAILED.value
model.error_message = "任务执行中断(worker重启/超时)"
model.error_info = {
"error_type": "WorkerInterrupted",
"message": "任务在运行中中断,可能因 worker 重启或超时",
"failed_at": datetime.now(timezone.utc).isoformat(),
}
model.completed_at = datetime.now(timezone.utc)
self.session.commit()
return len(models)
def cleanup_stale_pending(self, timeout_minutes: int = 30) -> int:
"""清理超时的 pending 任务(未被 Worker 拉取的任务)。
全局任务队列有 pending 数量上限,长期卡在 pending 的任务会占满队列,
导致新用户无法创建任务。将超时的 pending 任务标记为 failed。
Args:
timeout_minutes: 超时时间(分钟),默认 30 分钟
Returns:
清理的任务数量
"""
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 = (
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="fetch",
)
)
self.session.commit()
return count