Files
xiaoxia-saas/packages/adapters/sqlalchemy_impl/template_repository.py
T
CI Bot 19edd9aa85
CI/CD Pipeline / Validate Code Quality And Tests (pull_request) Failing after 35s
CI/CD Pipeline / Unit Tests (pull_request) Successful in 1m28s
CI/CD Pipeline / Frontend Lint (pull_request) Successful in 1m29s
CI/CD Pipeline / Build Staging API Image (pull_request) Has been skipped
CI/CD Pipeline / Build Staging Worker Image (pull_request) Has been skipped
CI/CD Pipeline / Build Staging Web Image (pull_request) Has been skipped
CI/CD Pipeline / Build Production API Image (pull_request) Has been skipped
CI/CD Pipeline / Build Production Worker Image (pull_request) Has been skipped
CI/CD Pipeline / Build Production Web Image (pull_request) Has been skipped
CI/CD Pipeline / Deploy Staging (Watchtower auto-deploy) (pull_request) Has been skipped
CI/CD Pipeline / Staging API Integration Tests (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 / Production Browser E2E (pull_request) Has been skipped
CI/CD Pipeline / Integration Tests (pull_request) Has been cancelled
feat(ci): 代码质量深度加固 - mypy/ruff/vulture 接入
- ruff 替换 flake8:增加 bugbear/pyupgrade/simplify/return 等规则集
  * 自动修复34个问题(未使用import/格式等)
  * 手动修复10个F841未使用变量 + 2个E722裸except
- mypy 类型检查:告警模式接入Validate,不阻断CI
  * 检查范围:apps/api/app + packages核心业务代码
  * 配置:ignore-missing-imports + explicit-package-bases
- vulture 死代码扫描升级:
  * 置信度阈值从80%降到70%,输出更多参考
  * 按size排序,便于人工审查高价值条目
  * 告警模式不阻断CI
- 新增 pyproject.toml:ruff 集中配置
- 修复F841未使用变量(10处)+ E722裸except(2处)
2026-07-14 17:38:13 +08:00

377 lines
13 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.
"""SQLAlchemy implementation of TemplateRepository."""
from __future__ import annotations
import uuid
from typing import List, Optional
from sqlalchemy.orm import Session
from packages.adapters.sqlalchemy_impl.models import (
EditPlanModel,
TemplateCategoryModel,
TemplateModel,
TemplateSegmentModel,
)
from packages.domain.template import Template, TemplateCategory, TemplateSegment
class SQLAlchemyTemplateRepository:
"""SQLAlchemy 剪辑计划模板仓储."""
def __init__(self, session: Session) -> None:
self.session = session
# ── Template CRUD ──
def list_by_user(
self,
user_id: str,
*,
skip: int = 0,
limit: int = 50,
category: Optional[str] = None,
tag: Optional[str] = None,
keyword: Optional[str] = None,
mode: Optional[str] = None,
) -> List[Template]:
query = self.session.query(TemplateModel).filter(
TemplateModel.user_id == user_id,
TemplateModel.is_active.is_(True),
)
if category:
query = query.filter(TemplateModel.category == category)
if mode:
query = query.filter(TemplateModel.mode == mode)
if keyword:
like_pattern = f"%{keyword}%"
query = query.filter(TemplateModel.name.like(like_pattern))
if tag:
# JSON 数组包含指定标签(MySQL JSON_CONTAINS / SQLite json_each 兼容写法用 LIKE
query = query.filter(TemplateModel.tags.like(f'%"{tag}"%'))
models = query.order_by(TemplateModel.created_at.desc()).offset(skip).limit(limit).all()
templates = [self._model_to_entity(m) for m in models]
# 批量加载所有 segments,避免 N+1 查询
if templates:
template_ids = [t.id for t in templates]
seg_models = (
self.session.query(TemplateSegmentModel)
.filter(TemplateSegmentModel.template_id.in_(template_ids))
.order_by(TemplateSegmentModel.segment_order)
.all()
)
# 按 template_id 分组
seg_map: dict[str, list] = {}
for sm in seg_models:
seg_map.setdefault(sm.template_id, []).append(
self._segment_model_to_entity(sm),
)
for t in templates:
t.segments = seg_map.get(t.id, [])
return templates
def get(self, template_id: str, user_id: str) -> Optional[Template]:
model = (
self.session.query(TemplateModel)
.filter(
TemplateModel.id == template_id,
TemplateModel.user_id == user_id,
)
.first()
)
if model is None:
return None
template = self._model_to_entity(model)
template.segments = self.list_segments(template.id)
return template
def create(self, template: Template) -> Template:
model = TemplateModel(
id=template.id,
user_id=template.user_id,
name=template.name,
mode=template.mode,
category=template.category,
tags=template.tags,
title_config=template.title_config,
subtitle_config=template.subtitle_config,
bgm_config=template.bgm_config,
estimated_duration=template.estimated_duration,
is_active=template.is_active,
)
self.session.add(model)
# flush 而非 commit,让 create + create_segments 在同一事务中提交
self.session.flush()
self.session.refresh(model)
result = self._model_to_entity(model)
result.segments = template.segments
return result
def update(self, template: Template) -> Template:
model = (
self.session.query(TemplateModel)
.filter(
TemplateModel.id == template.id,
TemplateModel.user_id == template.user_id,
)
.first()
)
if model is None:
raise ValueError(f"Template {template.id} not found")
model.name = template.name
model.mode = template.mode
model.category = template.category
model.tags = template.tags
model.title_config = template.title_config
model.subtitle_config = template.subtitle_config
model.bgm_config = template.bgm_config
model.estimated_duration = template.estimated_duration
model.is_active = template.is_active
self.session.commit()
self.session.refresh(model)
result = self._model_to_entity(model)
result.segments = template.segments
return result
def delete(self, template_id: str, user_id: str) -> bool:
model = (
self.session.query(TemplateModel)
.filter(
TemplateModel.id == template_id,
TemplateModel.user_id == user_id,
)
.first()
)
if model is None:
return False
model.is_active = False
# 级联清理关联的 segments,避免孤儿数据
self.session.query(TemplateSegmentModel).filter(
TemplateSegmentModel.template_id == template_id,
).delete(synchronize_session=False)
self.session.commit()
return True
def count_by_user(
self,
user_id: str,
*,
category: Optional[str] = None,
tag: Optional[str] = None,
keyword: Optional[str] = None,
mode: Optional[str] = None,
) -> int:
query = self.session.query(TemplateModel).filter(
TemplateModel.user_id == user_id,
TemplateModel.is_active.is_(True),
)
if category:
query = query.filter(TemplateModel.category == category)
if mode:
query = query.filter(TemplateModel.mode == mode)
if keyword:
query = query.filter(TemplateModel.name.like(f"%{keyword}%"))
if tag:
query = query.filter(TemplateModel.tags.like(f'%"{tag}"%'))
return query.count()
def copy_template(self, template_id: str, user_id: str, new_name: str) -> Template:
"""复制模板(含所有 segments)。"""
source = self.get(template_id, user_id)
if source is None:
raise ValueError(f"Template {template_id} not found")
new_id = str(uuid.uuid4())
new_template = Template(
id=new_id,
user_id=user_id,
name=new_name,
mode=source.mode,
category=source.category,
tags=list(source.tags),
title_config=dict(source.title_config),
subtitle_config=dict(source.subtitle_config),
bgm_config=dict(source.bgm_config),
estimated_duration=source.estimated_duration,
is_active=True,
)
created = self.create(new_template)
# 复制 segments
new_segments: List[TemplateSegment] = []
for seg in source.segments:
new_seg = TemplateSegment(
id=str(uuid.uuid4()),
template_id=new_id,
segment_order=seg.segment_order,
duration_min=seg.duration_min,
duration_max=seg.duration_max,
material_type=seg.material_type,
)
new_segments.append(new_seg)
model = TemplateSegmentModel(
id=new_seg.id,
template_id=new_seg.template_id,
segment_order=new_seg.segment_order,
duration_min=new_seg.duration_min,
duration_max=new_seg.duration_max,
material_type=new_seg.material_type,
)
self.session.add(model)
if new_segments:
self.session.commit()
created.segments = new_segments
return created
# ── Segments ──
def list_segments(self, template_id: str) -> List[TemplateSegment]:
models = (
self.session.query(TemplateSegmentModel)
.filter(TemplateSegmentModel.template_id == template_id)
.order_by(TemplateSegmentModel.segment_order)
.all()
)
return [self._segment_model_to_entity(m) for m in models]
def create_segments(self, segments: List[TemplateSegment]) -> List[TemplateSegment]:
for seg in segments:
model = TemplateSegmentModel(
id=seg.id,
template_id=seg.template_id,
segment_order=seg.segment_order,
duration_min=seg.duration_min,
duration_max=seg.duration_max,
material_type=seg.material_type,
)
self.session.add(model)
self.session.commit()
return segments
def delete_segments_by_template(self, template_id: str) -> int:
count = (
self.session.query(TemplateSegmentModel).filter(TemplateSegmentModel.template_id == template_id).delete()
)
self.session.commit()
return count
# ── Categories ──
def list_categories(self, user_id: str) -> List[TemplateCategory]:
models = (
self.session.query(TemplateCategoryModel)
.filter(TemplateCategoryModel.user_id == user_id)
.order_by(TemplateCategoryModel.created_at)
.all()
)
return [self._category_model_to_entity(m) for m in models]
def create_category(self, category: TemplateCategory) -> TemplateCategory:
model = TemplateCategoryModel(
id=category.id,
user_id=category.user_id,
name=category.name,
)
self.session.add(model)
self.session.commit()
self.session.refresh(model)
return self._category_model_to_entity(model)
def get_category(self, category_id: str, user_id: str) -> Optional[TemplateCategory]:
model = (
self.session.query(TemplateCategoryModel)
.filter(
TemplateCategoryModel.id == category_id,
TemplateCategoryModel.user_id == user_id,
)
.first()
)
if model is None:
return None
return self._category_model_to_entity(model)
def delete_category(self, category_id: str, user_id: str) -> bool:
model = (
self.session.query(TemplateCategoryModel)
.filter(
TemplateCategoryModel.id == category_id,
TemplateCategoryModel.user_id == user_id,
)
.first()
)
if model is None:
return False
self.session.delete(model)
self.session.commit()
return True
# ── Tags ──
def list_tags(self, user_id: str) -> List[str]:
"""获取用户所有模板的标签(去重)。"""
models = (
self.session.query(TemplateModel)
.filter(
TemplateModel.user_id == user_id,
TemplateModel.is_active.is_(True),
TemplateModel.tags.isnot(None),
)
.all()
)
tags_set: set[str] = set()
for m in models:
if m.tags:
for t in m.tags:
if t:
tags_set.add(t)
return sorted(tags_set)
# ── Usage Stats ──
def get_usage_count(self, template_id: str) -> int:
"""获取模板被使用的次数(关联的剪辑计划数量)。"""
return self.session.query(EditPlanModel).filter(EditPlanModel.template_id == template_id).count()
# ── Mapping helpers ──
@staticmethod
def _model_to_entity(model: TemplateModel) -> Template:
return Template(
id=model.id,
user_id=model.user_id,
name=model.name,
mode=model.mode,
category=model.category or "",
tags=model.tags or [],
title_config=model.title_config or {},
subtitle_config=model.subtitle_config or {},
bgm_config=model.bgm_config or {},
estimated_duration=model.estimated_duration or 0.0,
is_active=model.is_active,
created_at=model.created_at,
updated_at=model.updated_at,
)
@staticmethod
def _segment_model_to_entity(model: TemplateSegmentModel) -> TemplateSegment:
return TemplateSegment(
id=model.id,
template_id=model.template_id,
segment_order=model.segment_order,
duration_min=model.duration_min,
duration_max=model.duration_max,
material_type=model.material_type,
created_at=model.created_at,
updated_at=model.updated_at,
)
@staticmethod
def _category_model_to_entity(model: TemplateCategoryModel) -> TemplateCategory:
return TemplateCategory(
id=model.id,
user_id=model.user_id,
name=model.name,
created_at=model.created_at,
)