Files
backend_v2/internal/repository/article_repository.go
toom1996 30c9f21da4 update
2026-08-26 10:45:21 +08:00

209 lines
7.0 KiB
Go
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.

package repository
import (
"context"
"errors"
"fashionapi/internal/dto"
"fashionapi/internal/model"
"gorm.io/gorm"
)
// ArticleRepository 走秀档案(文章)数据访问接口。
type ArticleRepository interface {
// List 按筛选条件分页查询文章,同时返回符合条件的总数。
List(ctx context.Context, q dto.ArticleQuery) ([]model.RunwayRow, int64, error)
// FindByID 查询单篇文章(含 JOIN 出的品牌名)。不存在时返回 ErrNotFound。
FindByID(ctx context.Context, id string) (*model.RunwayRow, error)
// ListImages 查询某篇文章的全部图片,按排序值升序。
ListImages(ctx context.Context, runwayID string) ([]model.BrandRunwayImage, error)
// ImagesByRunwayIDs 批量查询多篇文章的图片,单次 IN 查询避免 N+1。
ImagesByRunwayIDs(ctx context.Context, ids []uint32) (map[uint32][]model.BrandRunwayImage, error)
// IDs 返回全部未删除文章的 id(升序),供 SSG 的 getStaticPaths 枚举路径使用。
// 只回主键,不带回封面/描述等大字段,比拉 size=500 再 .map(id) 省得多。
IDs(ctx context.Context) ([]uint32, error)
}
type articleRepository struct {
db *gorm.DB
}
// NewArticleRepository 创建文章仓储。
func NewArticleRepository(db *gorm.DB) ArticleRepository {
return &articleRepository{db: db}
}
// selectColumns 列表查询的投影列:只取对外需要的字段,不做 SELECT *。
const articleListColumns = `brand_runway.id, brand_runway.brand_id, brand_runway.title,
brand_runway.title_en, brand_runway.title_cn,
brand_runway.description, brand_runway.description_en, brand_runway.description_cn,
brand_runway.cover, brand_runway.year, brand_runway.image_count,
brand_runway.created_at, brand_runway.collection_type, brand_runway.season,
brand_runway.season_code, b.name AS brand_name, b.name_en AS brand_name_en, b.name_cn AS brand_name_cn`
const articleDetailColumns = `brand_runway.id, brand_runway.title, brand_runway.title_en, brand_runway.title_cn,
brand_runway.description, brand_runway.description_en, brand_runway.description_cn,
brand_runway.cover, brand_runway.year, brand_runway.image_count, brand_runway.created_at,
brand_runway.source_url, b.name AS brand_name, b.name_en AS brand_name_en, b.name_cn AS brand_name_cn`
// brandJoin 关联品牌表取品牌名;LEFT JOIN 保证品牌被软删时文章依然可见。
const brandJoin = "LEFT JOIN brand b ON b.id = brand_runway.brand_id AND b.is_deleted = 0"
// filterScope 把查询条件编译为 GORM Scope。
//
// 用 Scope 而非复用同一个 *gorm.DB:GORM v2 中在 Count 等终结方法之后复用同一实例
// 会带上残留的 Statement 状态,Scope 每次作用于全新查询,杜绝这类隐患。
func filterScope(q dto.ArticleQuery) func(*gorm.DB) *gorm.DB {
return func(db *gorm.DB) *gorm.DB {
db = db.Where("brand_runway.is_deleted = 0")
if q.Keyword != "" {
kw := "%" + q.Keyword + "%"
db = db.Where("brand_runway.title_en LIKE ? OR brand_runway.title_cn LIKE ?", kw, kw)
}
if q.BrandID != "" && q.BrandID != "0" {
db = db.Where("brand_runway.brand_id = ?", q.BrandID)
}
// 多值筛选统一策略:单值用 = (命中索引更精准),多值用 IN,空集合不加条件。
db = whereMulti(db, "brand_runway.brand_id", toAnySlice(q.BrandIDs))
db = whereMulti(db, "brand_runway.collection_type", toAnySlice(q.CollectionTypes))
db = whereMulti(db, "brand_runway.season", toAnySlice(q.Seasons))
db = whereMulti(db, "brand_runway.year", toAnySlice(q.Years))
if q.SeasonCode != "" {
db = db.Where("brand_runway.season_code = ?", q.SeasonCode)
}
return db
}
}
// whereMulti 按值数量选择 = 或 IN。
func whereMulti(db *gorm.DB, column string, values []any) *gorm.DB {
switch len(values) {
case 0:
return db
case 1:
return db.Where(column+" = ?", values[0])
default:
return db.Where(column+" IN ?", values)
}
}
func toAnySlice[T any](in []T) []any {
if len(in) == 0 {
return nil
}
out := make([]any, 0, len(in))
for _, v := range in {
out = append(out, v)
}
return out
}
// orderBy 把 sort 参数映射为 ORDER BY 子句。
//
// 白名单映射而非直接拼接用户输入,从根上排除 SQL 注入;
// 每个分支都以 created_at + id 兜底,保证分页结果稳定不跳行。
func orderBy(sort string) string {
switch sort {
case "year_desc":
return "brand_runway.year DESC, brand_runway.created_at DESC, brand_runway.id DESC"
case "year_asc":
return "brand_runway.year ASC, brand_runway.created_at DESC, brand_runway.id DESC"
case "image_count":
return "brand_runway.image_count DESC, brand_runway.created_at DESC, brand_runway.id DESC"
default: // newest
return "brand_runway.created_at DESC, brand_runway.id DESC"
}
}
func (r *articleRepository) List(ctx context.Context, q dto.ArticleQuery) ([]model.RunwayRow, int64, error) {
scope := filterScope(q)
var total int64
if err := r.db.WithContext(ctx).
Model(&model.BrandRunway{}).
Scopes(scope).
Count(&total).Error; err != nil {
return nil, 0, err
}
if total == 0 {
return []model.RunwayRow{}, 0, nil
}
var rows []model.RunwayRow
if err := r.db.WithContext(ctx).
Model(&model.BrandRunway{}).
Scopes(scope).
Select(articleListColumns).
Joins(brandJoin).
Order(orderBy(q.Sort)).
Offset(q.Offset()).
Limit(q.Size).
Scan(&rows).Error; err != nil {
return nil, 0, err
}
return rows, total, nil
}
func (r *articleRepository) FindByID(ctx context.Context, id string) (*model.RunwayRow, error) {
var row model.RunwayRow
err := r.db.WithContext(ctx).
Model(&model.BrandRunway{}).
Select(articleDetailColumns).
Joins(brandJoin).
Where("brand_runway.id = ? AND brand_runway.is_deleted = 0", id).
First(&row).Error
if err != nil {
if errors.Is(err, gorm.ErrRecordNotFound) {
return nil, ErrNotFound
}
return nil, err
}
return &row, nil
}
func (r *articleRepository) ListImages(ctx context.Context, runwayID string) ([]model.BrandRunwayImage, error) {
var imgs []model.BrandRunwayImage
err := r.db.WithContext(ctx).
Where("runway_id = ? AND is_deleted = 0", runwayID).
Order("sort_order ASC, id ASC").
Find(&imgs).Error
return imgs, err
}
func (r *articleRepository) ImagesByRunwayIDs(ctx context.Context, ids []uint32) (map[uint32][]model.BrandRunwayImage, error) {
result := make(map[uint32][]model.BrandRunwayImage, len(ids))
if len(ids) == 0 {
return result, nil
}
var imgs []model.BrandRunwayImage
err := r.db.WithContext(ctx).
Select("runway_id, image, name, sort_order").
Where("runway_id IN ? AND is_deleted = 0", ids).
Order("sort_order ASC, id ASC").
Find(&imgs).Error
if err != nil {
return nil, err
}
for _, im := range imgs {
result[im.RunwayID] = append(result[im.RunwayID], im)
}
return result, nil
}
// IDs 返回全部未删除文章的 id(升序),供 SSG 构建期枚举详情页路径。
func (r *articleRepository) IDs(ctx context.Context) ([]uint32, error) {
var ids []uint32
if err := r.db.WithContext(ctx).
Model(&model.BrandRunway{}).
Where("is_deleted = 0").
Order("id ASC").
Pluck("id", &ids).Error; err != nil {
return nil, err
}
return ids, nil
}