Files
xiaoxia-saas/packages/adapters/sqlalchemy_impl/generation_task_repository.py
T
xiaoxia 48c8d1fe3e
CI/CD Pipeline / Check if frontend-only change (push) Has been skipped
CI Build & Deploy Pipeline / Build Staging Web Image (push) Successful in 27s
CI Build & Deploy Pipeline / Build Production API Image (push) Has been skipped
CI Build & Deploy Pipeline / Build Production Web Image (push) Has been skipped
CI Build & Deploy Pipeline / Build Production Worker Image (push) Has been skipped
CI Build & Deploy Pipeline / Deploy Production (push) Has been skipped
CI Build & Deploy Pipeline / Production Browser E2E (push) Has been skipped
CI/CD Pipeline / Frontend Lint (push) Successful in 1m8s
CI/CD Pipeline / Validate Code Quality And Tests (push) Successful in 3m2s
CI/CD Pipeline / Unit Tests (push) Successful in 3m18s
CI/CD Pipeline / Integration Tests (push) Successful in 1m25s
CI Build & Deploy Pipeline / Build Staging API Image (push) Successful in 6m28s
CI Build & Deploy Pipeline / Build Staging Worker Image (push) Successful in 7m41s
CI Build & Deploy Pipeline / Deploy Staging (Watchtower auto-deploy) (push) Successful in 50s
CI Build & Deploy Pipeline / Staging E2E Tests (push) Failing after 3m17s
CI Build & Deploy Pipeline / Staging API Integration Tests (push) Successful in 3m57s
fix(P0+P1): 孤儿任务清理 + 时间时区统一 (#535)
Co-authored-by: xiaoxia <dev@xiaoxiajianji.com>
Co-committed-by: xiaoxia <dev@xiaoxiajianji.com>
2026-07-18 19:58:08 +08:00

266 lines
10 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 "",
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 "",
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_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 ""
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 分钟未更新的任务
标记为 failederror_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)