Files
xiaoxia-saas/packages/adapters/sqlalchemy_impl/generation_task_repository.py
T
CI Bot a9896e1507
CI/CD Pipeline / Check if frontend-only change (pull_request) Successful in 21s
CI/CD Pipeline / Frontend Lint (pull_request) Successful in 37s
CI/CD Pipeline / PR Build Web Image (pull_request) Successful in 25s
CI/CD Pipeline / Validate - Migration (alembic) (pull_request) Successful in 1m21s
CI/CD Pipeline / Build Staging API Image (pull_request) Has been skipped
CI/CD Pipeline / Validate - Type Check (mypy) (pull_request) Successful in 1m23s
CI/CD Pipeline / Build Staging Web Image (pull_request) Has been skipped
CI/CD Pipeline / Build Staging Worker Image (pull_request) Has been skipped
CI/CD Pipeline / Validate - Code Quality (pull_request) Failing after 1m55s
Preview Deploy / Deploy Preview Environment (pull_request) Failing after 23s
PR Automation / Auto Approve on CI Green (pull_request) Successful in 59s
PR Automation / Auto Merge on CI Green + Approved (pull_request) Successful in 53s
CI/CD Pipeline / Frontend Unit Tests (pull_request) Has been skipped
CI/CD Pipeline / Build Production API Image (pull_request) Has been skipped
CI/CD Pipeline / Build Production Web Image (pull_request) Has been skipped
CI/CD Pipeline / Build Production Worker Image (pull_request) Has been skipped
CI/CD Pipeline / Deploy Staging (Watchtower auto-deploy) (pull_request) Has been skipped
CI/CD Pipeline / Deploy Production (pull_request) Has been skipped
CI/CD Pipeline / Staging E2E Tests (pull_request) Has been skipped
CI/CD Pipeline / Staging API Integration Tests (pull_request) Has been skipped
CI/CD Pipeline / ACR Image Cleanup (pull_request) Has been skipped
CI/CD Pipeline / Production Browser E2E (pull_request) Has been skipped
CI/CD Pipeline / PR Build API Image (pull_request) Successful in 4m34s
CI/CD Pipeline / Integration Tests (pull_request) Successful in 2m40s
AI Code Review / AI Code Review (pull_request) Successful in 4m4s
CI/CD Pipeline / Unit Tests (pull_request) Successful in 5m49s
CI/CD Pipeline / PR Build Worker Image (pull_request) Successful in 10m3s
Preview Cleanup / Cleanup Preview Environment (pull_request) Successful in 19s
feat(#642): 一键生成支持自定义BGM
- 领域模型 GenerationTask 新增 bgm_config 字段
- ORM/Repository/UseCase/API 全链路透传 bgm_config
- Worker 端新增 BGM 配置合并逻辑(用户配置 > 模板配置)
- enabled 字段特殊处理:用户显式传才覆盖模板状态
- 合并逻辑抽至 packages/domain/bgm_utils.py 纯函数
- 14个BGM合并单测 + 2个领域单测,全量4246通过
2026-07-23 18:25:06 +08:00

278 lines
11 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 {}),
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 {},
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 ""
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 {}
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)