"""爆款视频任务 SQLAlchemy 仓储实现。""" from datetime import datetime, timezone from sqlalchemy.orm import Session from packages.adapters.sqlalchemy_impl.models import ( ViralVideoJobModel, ViralVideoPromptTemplateModel, ViralVideoStyleTemplateModel, ) from packages.domain.viral_video import ViralVideoJob, ViralVideoStatus def _to_domain(model: ViralVideoJobModel) -> ViralVideoJob: """ORM → 领域实体。""" return ViralVideoJob( id=model.id, user_id=model.user_id, images=list(model.images or []), industry=model.industry or "", target_customer=model.target_customer or "", persona_id=model.persona_id or "", viral_structure=model.viral_structure or "", marketing_purpose=model.marketing_purpose or "", bgm_preference=model.bgm_preference or "", duration=model.duration or 15, user_copy_text=model.user_copy_text or "", fusion_level=model.fusion_level or "ai_polish", reference_audio_path=model.reference_audio_path or "", reference_video_url=getattr(model, "reference_video_url", "") or "", style_strength=getattr(model, "style_strength", "medium") or "medium", style_guide=dict(model.style_guide) if model.style_guide else None, style_template_id=getattr(model, "style_template_id", "") or "", voice_id=getattr(model, "voice_id", "") or "", voice_source=getattr(model, "voice_source", "") or "", video_ratio=getattr(model, "video_ratio", "9:16") or "9:16", video_model=getattr(model, "video_model", "") or "", status=ViralVideoStatus(model.status) if model.status else ViralVideoStatus.PENDING, intent_result=dict(model.intent_result) if model.intent_result else None, image_analysis=dict(model.image_analysis) if getattr(model, "image_analysis", None) else None, storyboard=list(model.storyboard) if getattr(model, "storyboard", None) else None, generated_copy_text=getattr(model, "generated_copy_text", "") or "", copy_result=dict(model.copy_result) if getattr(model, "copy_result", None) else None, result_video_url=model.result_video_url or "", credits_cost=model.credits_cost or 0, error_msg=model.error_msg or "", retry_count=model.retry_count or 0, started_at=model.started_at, completed_at=model.completed_at, created_at=model.created_at, updated_at=model.updated_at, ) class SQLAlchemyViralVideoJobRepository: """爆款视频任务仓储。""" def __init__(self, session: Session): self.session = session def save(self, job: ViralVideoJob) -> ViralVideoJob: model = ViralVideoJobModel( id=job.id, user_id=job.user_id, images=job.images, industry=job.industry, target_customer=job.target_customer, persona_id=job.persona_id, viral_structure=job.viral_structure, marketing_purpose=job.marketing_purpose, bgm_preference=job.bgm_preference, duration=job.duration, user_copy_text=job.user_copy_text, fusion_level=job.fusion_level, reference_audio_path=job.reference_audio_path, reference_video_url=job.reference_video_url, style_strength=job.style_strength, style_guide=job.style_guide, style_template_id=job.style_template_id, voice_id=job.voice_id, voice_source=job.voice_source, video_ratio=job.video_ratio, video_model=job.video_model, status=job.status, intent_result=job.intent_result, image_analysis=job.image_analysis, storyboard=job.storyboard, generated_copy_text=job.generated_copy_text, copy_result=job.copy_result, result_video_url=job.result_video_url, credits_cost=job.credits_cost, error_msg=job.error_msg, retry_count=job.retry_count, started_at=job.started_at, completed_at=job.completed_at, created_at=job.created_at, updated_at=job.updated_at, ) self.session.add(model) self.session.commit() return job def update(self, job: ViralVideoJob) -> None: model = self.session.query(ViralVideoJobModel).filter(ViralVideoJobModel.id == job.id).first() if model is None: raise ValueError(f"ViralVideoJob {job.id} not found") model.status = job.status model.intent_result = job.intent_result model.image_analysis = job.image_analysis model.storyboard = job.storyboard model.generated_copy_text = job.generated_copy_text or "" model.copy_result = job.copy_result model.result_video_url = job.result_video_url model.credits_cost = job.credits_cost model.error_msg = job.error_msg model.retry_count = job.retry_count model.started_at = job.started_at model.completed_at = job.completed_at model.style_guide = job.style_guide # v1.5 three-stage: persist user-editable params so resume uses latest values model.user_copy_text = job.user_copy_text model.industry = job.industry model.target_customer = job.target_customer model.persona_id = job.persona_id model.viral_structure = job.viral_structure model.marketing_purpose = job.marketing_purpose model.bgm_preference = job.bgm_preference model.duration = job.duration model.fusion_level = job.fusion_level model.reference_audio_path = job.reference_audio_path model.reference_video_url = job.reference_video_url model.style_strength = job.style_strength model.style_template_id = job.style_template_id model.voice_id = job.voice_id or "" model.voice_source = job.voice_source or "" model.video_ratio = job.video_ratio or "9:16" model.video_model = job.video_model or "" model.updated_at = datetime.now(timezone.utc) self.session.commit() def get(self, job_id: str) -> ViralVideoJob | None: model = self.session.query(ViralVideoJobModel).filter(ViralVideoJobModel.id == job_id).first() if model is None: return None return _to_domain(model) def list_by_user(self, user_id: str, limit: int = 50, offset: int = 0) -> list[ViralVideoJob]: models = ( self.session.query(ViralVideoJobModel) .filter(ViralVideoJobModel.user_id == user_id) .order_by(ViralVideoJobModel.created_at.desc()) .offset(offset) .limit(limit) .all() ) return [_to_domain(m) for m in models] def count_pending_by_user(self, user_id: str) -> int: return ( self.session.query(ViralVideoJobModel) .filter( ViralVideoJobModel.user_id == user_id, ViralVideoJobModel.status.in_( ["pending", "running", "wait_user_confirm", "image_analyzed", "copy_generated"] ), ) .count() ) class SQLAlchemyViralVideoStyleTemplateRepository: """风格模板仓储。""" def __init__(self, session: Session): self.session = session def list_all(self) -> list[dict]: models = ( self.session.query(ViralVideoStyleTemplateModel) .order_by(ViralVideoStyleTemplateModel.sort_order.asc()) .all() ) return [ { "id": m.id, "name": m.name, "description": m.description or "", "thumbnail_url": m.thumbnail_url or "", "style_config": dict(m.style_config) if m.style_config else {}, "is_system": m.is_system, } for m in models ] def get(self, template_id: str) -> dict | None: model = ( self.session.query(ViralVideoStyleTemplateModel) .filter(ViralVideoStyleTemplateModel.id == template_id) .first() ) if model is None: return None return { "id": model.id, "name": model.name, "description": model.description or "", "thumbnail_url": model.thumbnail_url or "", "style_config": dict(model.style_config) if model.style_config else {}, "is_system": model.is_system, } class SQLAlchemyViralVideoPromptTemplateRepository: """Prompt 模板仓储(由 #2040 seed,这里只读取)。""" def __init__(self, session: Session): self.session = session def get_active_by_type(self, prompt_type: str) -> dict | None: model = ( self.session.query(ViralVideoPromptTemplateModel) .filter( ViralVideoPromptTemplateModel.prompt_type == prompt_type, ViralVideoPromptTemplateModel.is_active.is_(True), ) .order_by(ViralVideoPromptTemplateModel.version.desc()) .first() ) if model is None: return None return { "id": model.id, "prompt_type": model.prompt_type, "name": model.name, "content": model.content, "variables": list(model.variables or []), "version": model.version, }