Files
xiaoxia-saas/packages/adapters/sqlalchemy_impl/asset_repository.py
CI Bot e09fdda74c
CI/CD Pipeline / Build Staging Worker Image (pull_request) Has been skipped
CI/CD Pipeline / Build Staging API Image (pull_request) Has been skipped
CI/CD Pipeline / Build Staging Web Image (pull_request) Has been skipped
CI/CD Pipeline / Deploy Staging (Watchtower auto-deploy) (pull_request) Has been skipped
CI/CD Pipeline / Staging E2E Tests (pull_request) Has been skipped
CI/CD Pipeline / ACR Image Cleanup (pull_request) Has been skipped
CI/CD Pipeline / Staging API Integration Tests (pull_request) Has been skipped
CI/CD Pipeline / Check if frontend-only change (pull_request) Successful in 21s
CI/CD Pipeline / Frontend Lint (pull_request) Has been skipped
CI/CD Pipeline / Frontend Unit Tests (pull_request) Has been skipped
CI/CD Pipeline / PR Build Web Image (pull_request) Has been skipped
Preview Deploy / Deploy Preview Environment (pull_request) Successful in 1m27s
PR Automation / Auto Merge on CI Green + Approved (pull_request) Successful in 1m51s
CI/CD Pipeline / Validate - Migration (alembic) (pull_request) Successful in 1m53s
CI/CD Pipeline / Validate - Type Check (mypy) (pull_request) Successful in 1m59s
CI/CD Pipeline / PR Build API Image (pull_request) Successful in 2m15s
CI/CD Pipeline / PR Build Worker Image (pull_request) Successful in 2m17s
PR Automation / Auto Approve on CI Green (pull_request) Successful in 3m24s
CI/CD Pipeline / Unit Tests (pull_request) Failing after 3m29s
CI/CD Pipeline / Validate - Code Quality (pull_request) Successful in 6m8s
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 Production (pull_request) Has been skipped
CI/CD Pipeline / Canary Release to Production (pull_request) Has been skipped
CI/CD Pipeline / Production Browser E2E (pull_request) Has been skipped
AI Code Review / AI Code Review (pull_request) Successful in 6m43s
CI/CD Pipeline / Integration Tests (pull_request) Successful in 2m31s
CI/CD Pipeline / CI Gate (pull_request) Failing after 8s
fix: 确认生成兜底增强 — project_id 为空时通过 user_id 查找素材
根因:模板编辑器直接创建的草稿 plan 没有 project_id 和 config.asset_ids,
导致 _auto_fallback_auto_material_mode 跳过素材分配,
can_generate 报 "没有可渲染的就绪片段"。

修复:
- asset_repository 新增 find_ready_videos_by_user 方法
- _auto_fallback_auto_material_mode 增加策略2:project_id 为空时
  通过 uploaded_by_user_id 查找用户上传的就绪视频
- generation.py 传递 current_user.user.id 给兜底函数
2026-08-10 20:19:12 +08:00

441 lines
17 KiB
Python
Executable File
Raw Permalink 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,
status: list[str] | None = None,
) -> list[Asset]:
query = self.session.query(AssetModel).filter(AssetModel.asset_library_id == library_id)
if status:
query = query.filter(AssetModel.status.in_(status))
models = query.order_by(AssetModel.created_at.desc()).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,
status: list[str] | None = None,
) -> list[Asset]:
query = self.session.query(AssetModel).filter(AssetModel.project_id == project_id)
if status:
query = query.filter(AssetModel.status.in_(status))
models = query.order_by(AssetModel.created_at.desc()).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,
status: list[str] | None = None,
) -> list[Asset]:
query = self.session.query(AssetModel).filter(
AssetModel.asset_library_id == library_id, AssetModel.file_type == file_type
)
if status:
query = query.filter(AssetModel.status.in_(status))
models = query.order_by(AssetModel.created_at.desc()).offset(skip).limit(limit).all()
return [self._to_domain(model) for model in models]
def count_by_library_and_file_type(
self,
library_id: str,
file_type: str,
status: list[str] | None = None,
) -> int:
query = self.session.query(AssetModel).filter(
AssetModel.asset_library_id == library_id, AssetModel.file_type == file_type
)
if status:
query = query.filter(AssetModel.status.in_(status))
return query.count()
def find_by_project_and_file_type(
self,
project_id: str,
file_type: str,
skip: int = 0,
limit: int = 100,
status: list[str] | None = None,
) -> list[Asset]:
query = self.session.query(AssetModel).filter(
AssetModel.project_id == project_id, AssetModel.file_type == file_type
)
if status:
query = query.filter(AssetModel.status.in_(status))
models = query.order_by(AssetModel.created_at.desc()).offset(skip).limit(limit).all()
return [self._to_domain(model) for model in models]
def count_by_project_and_file_type(
self,
project_id: str,
file_type: str,
status: list[str] | None = None,
) -> int:
query = self.session.query(AssetModel).filter(
AssetModel.project_id == project_id, AssetModel.file_type == file_type
)
if status:
query = query.filter(AssetModel.status.in_(status))
return query.count()
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,
storage_key=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.storage_key = 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:
"""批量删除素材(软删除,标记 status=deleted),返回实际影响数量。"""
if not asset_ids:
return 0
from datetime import datetime, timezone
now = datetime.now(timezone.utc)
count = (
self.session.query(AssetModel)
.filter(AssetModel.id.in_(asset_ids), AssetModel.status != "deleted")
.update({AssetModel.status: "deleted", AssetModel.updated_at: now}, synchronize_session=False)
)
self.session.commit()
return count
def batch_update_metadata(self, asset_ids: list[str], metadata_patch: dict[str, object]) -> int:
"""批量更新素材 metadata(合并 patch),返回实际影响数量。"""
if not asset_ids:
return 0
from datetime import datetime, timezone
now = datetime.now(timezone.utc)
# 逐条读取 + 合并 + 更新,保证 JSON 合并正确
models = self.session.query(AssetModel).filter(AssetModel.id.in_(asset_ids)).all()
count = 0
for model in models:
existing = {}
if model.classification_result:
try:
existing = json.loads(model.classification_result)
except Exception:
existing = {}
merged = {**existing, **metadata_patch}
model.classification_result = json.dumps(merged, ensure_ascii=False)
model.updated_at = now
count += 1
self.session.commit()
return count
def batch_add_tags(self, asset_ids: list[str], tag_ids: list[str]) -> int:
"""批量给素材添加标签(合并去重),返回实际影响数量。"""
if not asset_ids or not tag_ids:
return 0
from datetime import datetime, timezone
now = datetime.now(timezone.utc)
clean_tag_ids = list(set(tag_ids))
count = 0
for aid in asset_ids:
# 查询现有标签
existing = {
row.tag_id
for row in self.session.query(AssetTagModel.tag_id).filter(AssetTagModel.asset_id == aid).all()
}
new_tags = [t for t in clean_tag_ids if t not in existing]
if new_tags:
for tid in new_tags:
self.session.add(AssetTagModel(asset_id=aid, tag_id=tid))
# 更新 updated_at
self.session.query(AssetModel).filter(AssetModel.id == aid).update(
{AssetModel.updated_at: now}, synchronize_session=False
)
count += 1
self.session.commit()
return count
def batch_replace_tags(self, asset_ids: list[str], tag_ids: list[str]) -> int:
"""批量替换素材标签(全量覆盖),返回实际影响数量。"""
if not asset_ids:
return 0
from datetime import datetime, timezone
now = datetime.now(timezone.utc)
clean_tag_ids = list(set(tag_ids))
count = 0
for aid in asset_ids:
# 先删再加
self.session.query(AssetTagModel).filter(AssetTagModel.asset_id == aid).delete(synchronize_session=False)
for tid in clean_tag_ids:
self.session.add(AssetTagModel(asset_id=aid, tag_id=tid))
# 更新 updated_at
self.session.query(AssetModel).filter(AssetModel.id == aid).update(
{AssetModel.updated_at: now}, synchronize_session=False
)
count += 1
self.session.commit()
return count
def count_by_project(self, project_id: str, status: list[str] | None = None) -> int:
query = self.session.query(AssetModel).filter(AssetModel.project_id == project_id)
if status:
query = query.filter(AssetModel.status.in_(status))
return query.count()
def count_by_project_ids(self, project_ids: list[str], status: list[str] | None = None) -> int:
if not project_ids:
return 0
query = self.session.query(AssetModel).filter(AssetModel.project_id.in_(project_ids))
if status:
query = query.filter(AssetModel.status.in_(status))
return query.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 find_ready_videos_by_user(
self,
user_id: str,
*,
limit: int = 50,
) -> list[Asset]:
"""查找用户上传的所有就绪视频素材。"""
query = self.session.query(AssetModel).filter(
AssetModel.uploaded_by_user_id == user_id,
AssetModel.status == "ready",
AssetModel.file_type == "video",
)
query = query.order_by(AssetModel.created_at.desc())
if limit > 0:
query = query.limit(limit)
models = query.all()
return [self._to_domain(m) for m in models]
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.storage_key or 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)