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, ) 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, ) 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