Files
xiaoxia-saas/packages/adapters/sqlalchemy_impl/asset_repository.py
CI Bot 9e97473eec
CI/CD Pipeline / Build Staging Web Image (pull_request) Has been skipped
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 / Deploy Staging (Watchtower auto-deploy) (pull_request) Has been skipped
CI/CD Pipeline / ACR Image Cleanup (pull_request) Has been skipped
CI/CD Pipeline / Staging E2E Tests (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 58s
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 API Image (pull_request) Failing after 0s
CI/CD Pipeline / PR Build Web Image (pull_request) Has been skipped
CI/CD Pipeline / PR Build Worker Image (pull_request) Failing after 0s
PR Automation / Auto Merge on CI Green + Approved (pull_request) Successful in 1m46s
CI/CD Pipeline / Validate - Migration (alembic) (pull_request) Successful in 2m18s
AI Code Review / AI Code Review (pull_request) Failing after 2m19s
CI/CD Pipeline / Validate - Type Check (mypy) (pull_request) Successful in 2m26s
Preview Deploy / Deploy Preview Environment (pull_request) Successful in 2m52s
PR Automation / Auto Approve on CI Green (pull_request) Successful in 3m45s
CI/CD Pipeline / Unit Tests (pull_request) Failing after 3m50s
CI/CD Pipeline / Validate - Code Quality (pull_request) Successful in 6m39s
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
CI/CD Pipeline / Integration Tests (pull_request) Has been cancelled
CI/CD Pipeline / CI Gate (pull_request) Has been cancelled
fix: batch query assets to resolve N+1 in _build_asset_url_map
- Add find_by_ids() to SQLAlchemyAssetRepository (single SQL IN query)
- Replace per-id find_by_id loop with single batch call
- Fixes AI Code Review blocking performance issue in PR #1404
2026-08-17 16:28:49 +08:00

448 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 find_by_ids(self, asset_ids: list[str]) -> list[Asset]:
"""批量查询素材(单次 SQL IN 查询,避免 N+1)。"""
if not asset_ids:
return []
models = self.session.query(AssetModel).filter(AssetModel.id.in_(asset_ids)).all()
return [self._to_domain(m) for m in models]
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)