Files
xiaoxia-saas/packages/adapters/sqlalchemy_impl/asset_repository.py
T
灵应 4fee87c5e8
CI/CD Pipeline / Deploy Staging (push) Has been skipped
CI/CD Pipeline / Staging E2E Tests (push) Has been skipped
CI/CD Pipeline / Build Production Runtime Images (push) Has been skipped
CI/CD Pipeline / Deploy Production (push) Has been skipped
CI/CD Pipeline / Production Browser E2E (push) Has been skipped
CI/CD Pipeline / Frontend Lint (push) Failing after 45h56m35s
CI/CD Pipeline / Validate Code Quality And Tests (push) Failing after 45h56m48s
feat: 素材重复上传检测 + 批量生成视频
任务1: 素材重复上传检测
- 上传接口支持 file_hash 参数,通过 MD5+素材库ID 去重
- 命中去重直接返回已有 asset_id,不重复存 OSS
- file_hash 透传: API → IngestJob → Asset 全链路
- 三条上传路径(表单/直传/分片)均支持去重
- Alembic 031: assets + ingest_jobs 加 file_hash 列+索引
- 6 个单元测试覆盖去重命中/未命中/空hash/透传

任务3: 批量生成视频
- POST /generations 支持 count 参数,一次创建多条生成任务
- 每条任务独立状态跟踪,响应返回 task_ids 列表

Co-Authored-By: Claude Fable 5 <noreply@anthropic.com>
2026-07-07 14:57:37 +08:00

293 lines
11 KiB
Python
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.
import json
from datetime import datetime, timezone
from sqlalchemy.orm import Session
from packages.adapters.sqlalchemy_impl.models import AssetModel, AssetTagModel
from packages.domain import Asset, AssetStatus, ClassificationStatus
class SQLAlchemyAssetRepository:
def __init__(self, session: Session):
self.session = session
def find_by_library(
self,
library_id: str,
skip: int = 0,
limit: int = 100,
) -> list[Asset]:
models = (
self.session.query(AssetModel)
.filter(AssetModel.asset_library_id == library_id)
.offset(skip)
.limit(limit)
.all()
)
return [self._to_domain(model) for model in models]
def find_by_project(
self,
project_id: str,
skip: int = 0,
limit: int = 100,
) -> list[Asset]:
models = (
self.session.query(AssetModel).filter(AssetModel.project_id == project_id).offset(skip).limit(limit).all()
)
return [self._to_domain(model) for model in models]
def find_by_library_and_file_type(
self,
library_id: str,
file_type: str,
skip: int = 0,
limit: int = 100,
) -> list[Asset]:
models = (
self.session.query(AssetModel)
.filter(AssetModel.asset_library_id == library_id, AssetModel.file_type == file_type)
.offset(skip)
.limit(limit)
.all()
)
return [self._to_domain(model) for model in models]
def find_by_id(self, asset_id: str) -> Asset | None:
model = self.session.query(AssetModel).filter(AssetModel.id == asset_id).first()
if model is None:
return None
return self._to_domain(model)
def get(self, asset_id: str) -> Asset | None:
return self.find_by_id(asset_id)
def create(self, asset: Asset) -> Asset:
now = datetime.now(timezone.utc)
model = AssetModel(
id=asset.id,
project_id=asset.project_id,
asset_library_id=asset.library_id,
name=asset.name,
file_type=(asset.mime_type.split("/")[0] if "/" in asset.mime_type else asset.mime_type),
file_size=asset.file_size,
file_url=asset.storage_key,
thumbnail_url=asset.thumbnail_url,
duration=asset.duration,
width=asset.width,
height=asset.height,
fps=asset.fps,
codec=asset.codec,
status=asset.status.value,
classification_status=asset.classification_status.value,
classification_result=(json.dumps(asset.metadata) if asset.metadata else None),
quality_score=asset.quality_score,
uploaded_by_user_id=asset.uploaded_by_user_id or "system",
file_hash=asset.file_hash or None,
created_at=asset.created_at,
updated_at=now,
)
self.session.add(model)
self.session.flush()
self._sync_asset_tags(asset.id, asset.tag_ids)
self.session.commit()
return asset
def update(self, asset: Asset) -> Asset:
model = self.session.query(AssetModel).filter(AssetModel.id == asset.id).first()
if model is None:
raise ValueError(f"Asset {asset.id} not found")
model.name = asset.name
model.file_size = asset.file_size
model.file_url = asset.storage_key
model.thumbnail_url = asset.thumbnail_url
model.duration = asset.duration
model.width = asset.width
model.height = asset.height
model.fps = asset.fps
model.codec = asset.codec
model.status = asset.status.value
model.classification_status = asset.classification_status.value
model.classification_result = json.dumps(asset.metadata) if asset.metadata else None
model.quality_score = asset.quality_score
model.uploaded_by_user_id = asset.uploaded_by_user_id or model.uploaded_by_user_id
model.file_hash = asset.file_hash or model.file_hash
model.updated_at = datetime.now(timezone.utc)
self.session.flush()
self._sync_asset_tags(asset.id, asset.tag_ids)
self.session.commit()
return asset
def delete(self, asset_id: str) -> bool:
model = self.session.query(AssetModel).filter(AssetModel.id == asset_id).first()
if model:
self.session.delete(model)
self.session.commit()
return True
return False
def batch_delete(self, asset_ids: list[str]) -> int:
"""批量删除素材,返回实际删除数量。"""
if not asset_ids:
return 0
count = self.session.query(AssetModel).filter(AssetModel.id.in_(asset_ids)).delete(synchronize_session=False)
self.session.commit()
return count
def count_by_project(self, project_id: str) -> int:
return self.session.query(AssetModel).filter(AssetModel.project_id == project_id).count()
def count_by_project_ids(self, project_ids: list[str]) -> int:
if not project_ids:
return 0
return self.session.query(AssetModel).filter(AssetModel.project_id.in_(project_ids)).count()
def sum_storage_by_project_ids(self, project_ids: list[str]) -> int:
if not project_ids:
return 0
from sqlalchemy import func
result = (
self.session.query(func.coalesce(func.sum(AssetModel.file_size), 0))
.filter(AssetModel.project_id.in_(project_ids))
.scalar()
)
return int(result or 0)
def search_candidates(
self,
project_id: str,
*,
file_type: str | None = None,
min_quality_score: float | None = None,
min_duration: float | None = None,
max_duration: float | None = None,
classification_category: str | None = None,
tags: list[str] | None = None,
status: str | None = None,
limit: int = 50,
) -> list[Asset]:
"""按筛选条件搜索候选素材,按质量分降序排列。"""
query = self.session.query(AssetModel).filter(
AssetModel.project_id == project_id,
)
if file_type is not None:
query = query.filter(AssetModel.file_type == file_type)
if min_quality_score is not None:
query = query.filter(AssetModel.quality_score >= min_quality_score)
if min_duration is not None:
query = query.filter(AssetModel.duration >= min_duration)
if max_duration is not None:
query = query.filter(AssetModel.duration <= max_duration)
if status is not None:
query = query.filter(AssetModel.status == status)
if classification_category is not None:
# classification_result 是 JSON Text,用 LIKE 匹配 category 字段
query = query.filter(AssetModel.classification_result.like(f'%"{classification_category}"%'))
query = query.order_by(AssetModel.quality_score.desc().nullslast())
if limit > 0:
query = query.limit(limit)
models = query.all()
candidates = [self._to_domain(m) for m in models]
# 内存中过滤 tagstags 存在 metadata 中)
if tags:
tag_set = set(tags)
candidates = [a for a in candidates if tag_set.issubset(set(a.metadata.get("tags", [])))]
return candidates
def _to_domain(self, model: AssetModel) -> Asset:
metadata = {}
if model.classification_result:
try:
metadata = json.loads(model.classification_result)
except Exception:
metadata = {}
mime_type = model.file_type
if "/" not in mime_type:
mime_type = {
"video": "video/mp4",
"audio": "audio/mpeg",
"image": "image/jpeg",
}.get(mime_type, mime_type)
# 查询关联的 tag_ids
tag_ids = [
row.tag_id
for row in self.session.query(AssetTagModel.tag_id).filter(AssetTagModel.asset_id == model.id).all()
]
return Asset(
id=model.id,
project_id=model.project_id,
library_id=model.asset_library_id,
name=model.name,
storage_key=model.file_url,
mime_type=mime_type,
file_size=int(model.file_size or 0),
thumbnail_url=model.thumbnail_url,
duration=model.duration,
width=int(model.width) if model.width is not None else None,
height=int(model.height) if model.height is not None else None,
fps=model.fps,
codec=model.codec,
status=AssetStatus(model.status),
classification_status=ClassificationStatus(model.classification_status),
quality_score=model.quality_score,
uploaded_by_user_id=model.uploaded_by_user_id,
file_hash=model.file_hash or "",
metadata=metadata,
tag_ids=tag_ids,
created_at=model.created_at,
updated_at=model.updated_at,
)
def _sync_asset_tags(self, asset_id: str, tag_ids: list[str]) -> None:
"""同步素材-标签关联表(全量替换)。"""
self.session.query(AssetTagModel).filter(AssetTagModel.asset_id == asset_id).delete(synchronize_session=False)
for tag_id in tag_ids:
self.session.add(AssetTagModel(asset_id=asset_id, tag_id=tag_id))
def find_by_tag_ids(
self,
tag_ids: list[str],
skip: int = 0,
limit: int = 100,
) -> list[Asset]:
"""查找包含所有指定标签的素材。"""
if not tag_ids:
return []
from sqlalchemy import func
# 找出同时拥有所有指定 tag_id 的 asset_id
tag_set = set(tag_ids)
asset_ids = (
self.session.query(AssetTagModel.asset_id)
.filter(AssetTagModel.tag_id.in_(tag_set))
.group_by(AssetTagModel.asset_id)
.having(func.count(AssetTagModel.tag_id) == len(tag_set))
.all()
)
ids = [row[0] for row in asset_ids]
if not ids:
return []
models = self.session.query(AssetModel).filter(AssetModel.id.in_(ids)).offset(skip).limit(limit).all()
return [self._to_domain(m) for m in models]
def find_by_library_and_file_hash(
self,
library_id: str,
file_hash: str,
) -> Asset | None:
"""按素材库 + 文件哈希查找已有素材(去重检测)。"""
if not file_hash:
return None
model = (
self.session.query(AssetModel)
.filter(
AssetModel.asset_library_id == library_id,
AssetModel.file_hash == file_hash,
)
.first()
)
if model is None:
return None
return self._to_domain(model)