115 lines
2.7 KiB
Go
115 lines
2.7 KiB
Go
package dao
|
||
|
||
import (
|
||
"errors"
|
||
"fmt"
|
||
"twin-api/app/common/model"
|
||
"twin-api/base/global"
|
||
|
||
"gorm.io/gorm"
|
||
)
|
||
|
||
type UserPkg struct {
|
||
db *gorm.DB
|
||
}
|
||
|
||
func NewUserPkg() *UserPkg {
|
||
return &UserPkg{
|
||
db: global.GetDB(),
|
||
}
|
||
}
|
||
|
||
func (p *UserPkg) TableName() string {
|
||
return "tw_user_package"
|
||
}
|
||
|
||
func (p *UserPkg) Create(m *model.UserPkg) error {
|
||
return p.db.Table(p.TableName()).Create(m).Error
|
||
}
|
||
|
||
// CreateWithinLimit 在不超过 limit 的前提下为用户落库包记录。
|
||
// 仅当该包(小版本)当前人数 < limit 时才创建;返回 true 表示成功,false 表示已满。
|
||
// limit <= 0 表示不限制。
|
||
func (p *UserPkg) CreateWithinLimit(m *model.UserPkg, limit int64) (bool, error) {
|
||
if limit > 0 {
|
||
count, err := p.CountByPackageId(m.PackageId)
|
||
if err != nil {
|
||
return false, err
|
||
}
|
||
if count >= limit {
|
||
return false, nil
|
||
}
|
||
}
|
||
if err := p.Create(m); err != nil {
|
||
return false, err
|
||
}
|
||
return true, nil
|
||
}
|
||
|
||
// CountByPackageId 统计某个包(小版本)当前的用户数
|
||
func (p *UserPkg) CountByPackageId(packageId int64) (int64, error) {
|
||
var count int64
|
||
err := p.db.Table(p.TableName()).Where("package_id = ?", packageId).Count(&count).Error
|
||
return count, err
|
||
}
|
||
|
||
func (p *UserPkg) Update(m *model.UserPkg) error {
|
||
return p.db.Table(p.TableName()).Save(m).Error
|
||
}
|
||
|
||
func (p *UserPkg) Get(id, userId int64) (*model.UserPkg, error) {
|
||
var res model.UserPkg
|
||
if err := p.db.Table(p.TableName()).Where("id = ? AND user_id = ?", id, userId).First(&res).Error; err != nil {
|
||
if errors.Is(err, gorm.ErrRecordNotFound) {
|
||
return nil, nil
|
||
}
|
||
return nil, err
|
||
}
|
||
|
||
return &res, nil
|
||
}
|
||
|
||
func (p *UserPkg) GetLatest(userId int64) (*model.UserPkgInfo, error) {
|
||
var res model.UserPkgInfo
|
||
tx := p.db.Table(p.TableName() + " u ")
|
||
tx = tx.Joins(fmt.Sprintf(" JOIN %s p ON u.package_id = p.id", NewPackage().TableName()))
|
||
tx = tx.Where("u.user_id = ?", userId)
|
||
tx = tx.Order("u.id desc")
|
||
tx = tx.Where("p.status = ?", model.PkgNormal)
|
||
if err := tx.Select("u.*,p.*").Limit(1).Find(&res).Error; err != nil {
|
||
if errors.Is(err, gorm.ErrRecordNotFound) {
|
||
return nil, nil
|
||
}
|
||
|
||
return nil, err
|
||
}
|
||
return &res, nil
|
||
}
|
||
|
||
func (p *UserPkg) CountByPackageIds(packageIds []int64) (map[int64]int64, error) {
|
||
res := make(map[int64]int64)
|
||
if len(packageIds) == 0 {
|
||
return res, nil
|
||
}
|
||
type countResult struct {
|
||
PackageId int64 `gorm:"column:package_id"`
|
||
Count int64 `gorm:"column:count"`
|
||
}
|
||
var rows []countResult
|
||
if err := p.db.Table(p.TableName()).
|
||
Select("package_id, count(*) as count").
|
||
Where("package_id IN ?", packageIds).
|
||
Group("package_id").
|
||
Find(&rows).Error; err != nil {
|
||
return nil, err
|
||
}
|
||
for _, row := range rows {
|
||
res[row.PackageId] = row.Count
|
||
}
|
||
return res, nil
|
||
}
|
||
|
||
func (p *UserPkg) Query(...any) (any, error) {
|
||
panic("implement me")
|
||
}
|