Files

438 lines
18 KiB
Python

from __future__ import annotations
import json
from collections import defaultdict
from typing import Any, Iterable
from sqlalchemy import case, func, or_, select
from sqlalchemy.ext.asyncio import AsyncSession
from sqlalchemy.sql import Select
from app.enums.home_material import (
HomeMaterialAssetStatus,
HomeMaterialMediaType,
HomeMaterialPublicResponseMode,
)
from app.models.home_material import HomeMaterialAsset, HomeMaterialCategory, HomeMaterialWatermark
from app.schemas.home_material import (
HomeMaterialAssetOut,
HomeMaterialMediaReference,
HomeMaterialCategoryOut,
HomeMaterialPublicAssetOut,
HomeMaterialPublicCategoryGroupOut,
HomeMaterialPublicCategoryOut,
HomeMaterialPublicFlatItemOut,
HomeMaterialWatermarkOut,
)
def _unique(values: Iterable[str | None]) -> list[str]:
return list({v for v in values if v})
def _load_json(value: str | None) -> dict[str, Any] | None:
if not value:
return None
try:
data = json.loads(value)
return data if isinstance(data, dict) else None
except Exception:
return None
def _load_media_references(value: str | None) -> list[HomeMaterialMediaReference]:
if not value:
return []
try:
data = json.loads(value)
if not isinstance(data, list):
return []
refs: list[HomeMaterialMediaReference] = []
for item in data:
try:
refs.append(HomeMaterialMediaReference(**item))
except Exception:
continue
return refs
except Exception:
return []
class HomeMaterialQueryService:
"""首页素材高性能查询组装层:列表查询 → ID 去重 → 批量查询 → map 组装。"""
async def batch_categories_map(self, db: AsyncSession, category_ids: Iterable[str | None]) -> dict[str, HomeMaterialCategory]:
ids = _unique(category_ids)
if not ids:
return {}
result = await db.execute(
select(HomeMaterialCategory).where(HomeMaterialCategory.id.in_(ids), HomeMaterialCategory.deleted_at.is_(None))
)
return {row.id: row for row in result.scalars().all()}
async def batch_watermarks_map(self, db: AsyncSession, watermark_ids: Iterable[str | None]) -> dict[str, HomeMaterialWatermark]:
ids = _unique(watermark_ids)
if not ids:
return {}
result = await db.execute(
select(HomeMaterialWatermark).where(HomeMaterialWatermark.id.in_(ids), HomeMaterialWatermark.deleted_at.is_(None))
)
return {row.id: row for row in result.scalars().all()}
async def asset_counts_map(
self,
db: AsyncSession,
category_ids: Iterable[str] | None = None,
*,
media_type: HomeMaterialMediaType | str | None = None,
public_only: bool = False,
) -> dict[str, dict[str, int]]:
stmt = select(
HomeMaterialAsset.category_id,
func.count(HomeMaterialAsset.id).label("asset_count"),
func.coalesce(func.sum(case((HomeMaterialAsset.media_type == HomeMaterialMediaType.IMAGE.value, 1), else_=0)), 0).label("image_count"),
func.coalesce(func.sum(case((HomeMaterialAsset.media_type == HomeMaterialMediaType.VIDEO.value, 1), else_=0)), 0).label("video_count"),
).where(HomeMaterialAsset.deleted_at.is_(None))
ids = _unique(category_ids or [])
if ids:
stmt = stmt.where(HomeMaterialAsset.category_id.in_(ids))
if media_type:
stmt = stmt.where(HomeMaterialAsset.media_type == HomeMaterialMediaType(media_type).value)
if public_only:
stmt = stmt.where(
HomeMaterialAsset.is_active.is_(True),
HomeMaterialAsset.status == HomeMaterialAssetStatus.SUCCESS.value,
HomeMaterialAsset.watermarked_url.is_not(None),
)
stmt = stmt.group_by(HomeMaterialAsset.category_id)
result = await db.execute(stmt)
out: dict[str, dict[str, int]] = {}
for row in result.all():
out[row.category_id] = {
"asset_count": int(row.asset_count or 0),
"image_count": int(row.image_count or 0),
"video_count": int(row.video_count or 0),
}
return out
def category_to_out(self, category: HomeMaterialCategory, counts: dict[str, int] | None = None) -> HomeMaterialCategoryOut:
c = counts or {}
return HomeMaterialCategoryOut(
id=category.id,
name=category.name,
key=category.key,
description=category.description,
icon=category.icon,
is_active=category.is_active,
sort_order=category.sort_order,
asset_count=int(c.get("asset_count", 0)),
image_count=int(c.get("image_count", 0)),
video_count=int(c.get("video_count", 0)),
created_at=category.created_at,
updated_at=category.updated_at,
)
def watermark_to_out(self, watermark: HomeMaterialWatermark) -> HomeMaterialWatermarkOut:
return HomeMaterialWatermarkOut(
id=watermark.id,
name=watermark.name,
file_url=watermark.file_url,
file_name=watermark.file_name,
file_size_bytes=watermark.file_size_bytes,
width=watermark.width,
height=watermark.height,
is_default=watermark.is_default,
is_active=watermark.is_active,
created_at=watermark.created_at,
updated_at=watermark.updated_at,
)
def asset_to_out(
self,
asset: HomeMaterialAsset,
*,
category_map: dict[str, HomeMaterialCategory] | None = None,
watermark_map: dict[str, HomeMaterialWatermark] | None = None,
include_original: bool = True,
) -> HomeMaterialAssetOut:
category = (category_map or {}).get(asset.category_id)
watermark = (watermark_map or {}).get(asset.watermark_id or "")
return HomeMaterialAssetOut(
id=asset.id,
category_id=asset.category_id,
category_name=category.name if category else None,
category_key=category.key if category else None,
title=asset.title,
media_type=HomeMaterialMediaType(asset.media_type),
status=HomeMaterialAssetStatus(asset.status),
original_url=asset.original_url if include_original else None,
watermarked_url=asset.watermarked_url,
cover_url=asset.cover_url,
watermark_id=asset.watermark_id,
watermark_name=watermark.name if watermark else None,
watermark_config=_load_json(asset.watermark_config_json),
generation_prompt=asset.generation_prompt,
media_references=_load_media_references(asset.media_references_json),
width=asset.width,
height=asset.height,
duration_seconds=asset.duration_seconds,
file_size_bytes=asset.file_size_bytes,
watermarked_file_size_bytes=asset.watermarked_file_size_bytes,
is_active=asset.is_active,
sort_order=asset.sort_order,
error_message=asset.error_message,
processed_at=asset.processed_at,
created_at=asset.created_at,
updated_at=asset.updated_at,
)
async def list_admin_assets(
self,
db: AsyncSession,
*,
page: int,
page_size: int,
category_id: str | None = None,
media_type: HomeMaterialMediaType | str | None = None,
status: HomeMaterialAssetStatus | str | None = None,
is_active: bool | None = None,
keyword: str | None = None,
include_original: bool = True,
) -> tuple[list[HomeMaterialAssetOut], int]:
stmt = select(HomeMaterialAsset).where(HomeMaterialAsset.deleted_at.is_(None))
count_stmt = select(func.count(HomeMaterialAsset.id)).where(HomeMaterialAsset.deleted_at.is_(None))
conditions = []
if category_id:
conditions.append(HomeMaterialAsset.category_id == category_id)
if media_type:
conditions.append(HomeMaterialAsset.media_type == HomeMaterialMediaType(media_type).value)
if status:
conditions.append(HomeMaterialAsset.status == HomeMaterialAssetStatus(status).value)
if is_active is not None:
conditions.append(HomeMaterialAsset.is_active.is_(is_active))
if keyword:
conditions.append(HomeMaterialAsset.title.ilike(f"%{keyword}%"))
for condition in conditions:
stmt = stmt.where(condition)
count_stmt = count_stmt.where(condition)
total = int((await db.execute(count_stmt)).scalar_one() or 0)
result = await db.execute(
stmt.order_by(HomeMaterialAsset.sort_order.asc(), HomeMaterialAsset.created_at.desc())
.offset((page - 1) * page_size)
.limit(page_size)
)
assets = result.scalars().all()
category_map = await self.batch_categories_map(db, [a.category_id for a in assets])
watermark_map = await self.batch_watermarks_map(db, [a.watermark_id for a in assets])
return [self.asset_to_out(a, category_map=category_map, watermark_map=watermark_map, include_original=include_original) for a in assets], total
async def resolve_category_ids(
self,
db: AsyncSession,
*,
category_id: str | None = None,
category_key: str | None = None,
category_ids: str | None = None,
category_keys: str | None = None,
active_only: bool = True,
) -> list[str] | None:
"""解析前台行业筛选参数,优先级:category_id > category_key > category_ids > category_keys。None 表示不限制。"""
if category_id:
return [category_id]
stmt = select(HomeMaterialCategory.id).where(HomeMaterialCategory.deleted_at.is_(None))
if active_only:
stmt = stmt.where(HomeMaterialCategory.is_active.is_(True))
if category_key:
result = await db.execute(stmt.where(HomeMaterialCategory.key == category_key).limit(1))
cid = result.scalar_one_or_none()
return [cid] if cid else []
if category_ids:
return [v.strip() for v in category_ids.split(",") if v.strip()]
if category_keys:
keys = [v.strip() for v in category_keys.split(",") if v.strip()]
if not keys:
return []
result = await db.execute(stmt.where(HomeMaterialCategory.key.in_(keys)))
return list(result.scalars().all())
return None
async def list_public_categories(
self,
db: AsyncSession,
*,
with_asset_count: bool = True,
media_type: HomeMaterialMediaType | str | None = None,
only_has_assets: bool = True,
) -> list[HomeMaterialPublicCategoryOut]:
result = await db.execute(
select(HomeMaterialCategory)
.where(HomeMaterialCategory.deleted_at.is_(None), HomeMaterialCategory.is_active.is_(True))
.order_by(HomeMaterialCategory.sort_order.asc(), HomeMaterialCategory.created_at.desc())
)
categories = result.scalars().all()
counts = await self.asset_counts_map(db, [c.id for c in categories], media_type=media_type, public_only=True) if with_asset_count or only_has_assets else {}
items: list[HomeMaterialPublicCategoryOut] = []
for category in categories:
c = counts.get(category.id, {})
if only_has_assets and int(c.get("asset_count", 0)) <= 0:
continue
items.append(
HomeMaterialPublicCategoryOut(
id=category.id,
name=category.name,
key=category.key,
icon=category.icon,
description=category.description,
sort_order=category.sort_order,
asset_count=int(c.get("asset_count", 0)),
image_count=int(c.get("image_count", 0)),
video_count=int(c.get("video_count", 0)),
)
)
return items
async def list_public_grouped(
self,
db: AsyncSession,
*,
category_ids: list[str] | None,
media_type: HomeMaterialMediaType | str | None,
limit_per_category: int,
include_empty_categories: bool,
) -> list[HomeMaterialPublicCategoryGroupOut]:
cat_stmt = select(HomeMaterialCategory).where(HomeMaterialCategory.deleted_at.is_(None), HomeMaterialCategory.is_active.is_(True))
if category_ids is not None:
if not category_ids:
return []
cat_stmt = cat_stmt.where(HomeMaterialCategory.id.in_(category_ids))
cat_result = await db.execute(cat_stmt.order_by(HomeMaterialCategory.sort_order.asc(), HomeMaterialCategory.created_at.desc()))
categories = cat_result.scalars().all()
ids = [c.id for c in categories]
if not ids:
return []
asset_base = select(
HomeMaterialAsset.id.label("id"),
func.row_number()
.over(
partition_by=HomeMaterialAsset.category_id,
order_by=(HomeMaterialAsset.sort_order.asc(), HomeMaterialAsset.created_at.desc()),
)
.label("rn"),
).where(
HomeMaterialAsset.deleted_at.is_(None),
HomeMaterialAsset.is_active.is_(True),
HomeMaterialAsset.status == HomeMaterialAssetStatus.SUCCESS.value,
HomeMaterialAsset.watermarked_url.is_not(None),
HomeMaterialAsset.category_id.in_(ids),
)
if media_type:
asset_base = asset_base.where(HomeMaterialAsset.media_type == HomeMaterialMediaType(media_type).value)
ranked = asset_base.subquery()
asset_result = await db.execute(
select(HomeMaterialAsset)
.where(HomeMaterialAsset.id.in_(select(ranked.c.id).where(ranked.c.rn <= limit_per_category)))
.order_by(HomeMaterialAsset.category_id.asc(), HomeMaterialAsset.sort_order.asc(), HomeMaterialAsset.created_at.desc())
)
grouped: dict[str, list[HomeMaterialAsset]] = defaultdict(list)
for asset in asset_result.scalars().all():
grouped[asset.category_id].append(asset)
out: list[HomeMaterialPublicCategoryGroupOut] = []
for category in categories:
assets = grouped.get(category.id, [])
if not include_empty_categories and not assets:
continue
out.append(
HomeMaterialPublicCategoryGroupOut(
id=category.id,
name=category.name,
key=category.key,
icon=category.icon,
description=category.description,
sort_order=category.sort_order,
assets=[
HomeMaterialPublicAssetOut(
id=a.id,
title=a.title,
media_type=HomeMaterialMediaType(a.media_type),
url=a.watermarked_url or "",
cover_url=a.cover_url,
width=a.width,
height=a.height,
duration_seconds=a.duration_seconds,
generation_prompt=a.generation_prompt,
media_references=_load_media_references(a.media_references_json),
sort_order=a.sort_order,
)
for a in assets
],
)
)
return out
async def list_public_flat(
self,
db: AsyncSession,
*,
category_ids: list[str] | None,
media_type: HomeMaterialMediaType | str | None,
page: int,
page_size: int,
) -> tuple[list[HomeMaterialPublicFlatItemOut], int]:
stmt = select(HomeMaterialAsset).where(
HomeMaterialAsset.deleted_at.is_(None),
HomeMaterialAsset.is_active.is_(True),
HomeMaterialAsset.status == HomeMaterialAssetStatus.SUCCESS.value,
HomeMaterialAsset.watermarked_url.is_not(None),
)
count_stmt = select(func.count(HomeMaterialAsset.id)).where(
HomeMaterialAsset.deleted_at.is_(None),
HomeMaterialAsset.is_active.is_(True),
HomeMaterialAsset.status == HomeMaterialAssetStatus.SUCCESS.value,
HomeMaterialAsset.watermarked_url.is_not(None),
)
if category_ids is not None:
if not category_ids:
return [], 0
stmt = stmt.where(HomeMaterialAsset.category_id.in_(category_ids))
count_stmt = count_stmt.where(HomeMaterialAsset.category_id.in_(category_ids))
if media_type:
stmt = stmt.where(HomeMaterialAsset.media_type == HomeMaterialMediaType(media_type).value)
count_stmt = count_stmt.where(HomeMaterialAsset.media_type == HomeMaterialMediaType(media_type).value)
total = int((await db.execute(count_stmt)).scalar_one() or 0)
result = await db.execute(
stmt.order_by(HomeMaterialAsset.sort_order.asc(), HomeMaterialAsset.created_at.desc())
.offset((page - 1) * page_size)
.limit(page_size)
)
assets = result.scalars().all()
category_map = await self.batch_categories_map(db, [a.category_id for a in assets])
items: list[HomeMaterialPublicFlatItemOut] = []
for asset in assets:
category = category_map.get(asset.category_id)
if not category:
continue
items.append(
HomeMaterialPublicFlatItemOut(
id=asset.id,
category_id=category.id,
category_name=category.name,
category_key=category.key,
title=asset.title,
media_type=HomeMaterialMediaType(asset.media_type),
url=asset.watermarked_url or "",
cover_url=asset.cover_url,
width=asset.width,
height=asset.height,
duration_seconds=asset.duration_seconds,
generation_prompt=asset.generation_prompt,
media_references=_load_media_references(asset.media_references_json),
sort_order=asset.sort_order,
)
)
return items, total
query_service = HomeMaterialQueryService()