提交 e06545fc authored 作者: mooncake9527's avatar mooncake9527

support multi tables join query pagination

上级 389c9a37
package main
import (
"fmt"
"log"
"os"
"reflect"
"strings"
"time"
"gitlab.wanzhuangkj.com/tush/xpkg/xcommon/odao"
"gorm.io/driver/mysql"
"gorm.io/gorm"
"gorm.io/gorm/logger"
)
// select f.*,fa.* from friend f left join friend_apply fa on f.id = fa.f_id where f.name = ? and fa.no = ? order by f.id asc offset 10 limit 20;
type Bike struct {
ID int64 `gorm:"column:id"`
ModelID int64 `gorm:"column:model_id"`
SiteID int64 `gorm:"column:site_id"`
Nickname string `gorm:"column:nickname"`
}
func (x Bike) TableName() string {
return "bike"
}
type BikeModel struct {
ID int64 `gorm:"column:id"`
FlatImage string `gorm:"column:flat_image"`
}
func (x BikeModel) TableName() string {
return "bike_model"
}
type Site struct {
ID int64 `gorm:"column:id"`
Address string `gorm:"column:address"`
}
func (x Site) TableName() string {
return "site"
}
type Record struct {
Bike Bike `gorm:"embedded"`
Site Site `gorm:"embedded"`
BikeModel BikeModel `gorm:"embedded"`
UID int `gorm:"column:bike.uid"`
}
var (
db *gorm.DB
)
type BikeReq struct {
}
type Req struct {
BikeNickname string `query:"table:bike;column:nickname;type:eq"`
SiteID int64 `query:"table:site;column:id"`
Join int64 `query:"table:site;column:id;type:join;join:inner join bike on site.id = bike.site_id"`
BikeModelID int64 `query:"table:bike_model;column:id;type:join;join:left join bike_model on bike.model_id = bike_model.id"`
Name string `query:"table:site;column:name;type:eq"`
}
func (Req) TableName() string {
return "site"
}
func main1() {
newLogger := logger.New(
log.New(os.Stdout, "\r\n", log.LstdFlags),
logger.Config{
SlowThreshold: time.Second,
LogLevel: logger.Info, // 设置为Info级别打印所有SQL
Colorful: true,
},
)
req := Req{
BikeNickname: "研发车辆",
SiteID: 1320976731124793344,
Name: "优盟店",
}
dsn := "qitu_test:By1rGembg6@tcp(rm-bp15831972934dz2c.mysql.rds.aliyuncs.com:3306)/qitu_athena_test?charset=utf8mb4&parseTime=True&loc=Local"
// 方法2:全局DryRun配置
dryRunDB, _ := gorm.Open(mysql.Open(dsn), &gorm.Config{Logger: newLogger})
db = dryRunDB
rs, err := query(req, 10, 0)
if err != nil {
fmt.Println(err.Error())
return
}
fmt.Println(rs[0].Site.ID)
fmt.Println(rs[0].BikeModel.ID)
fmt.Println(rs[0].Bike.ID)
}
func query(where odao.Query, limit, offset int) ([]Record, error) {
var rs []Record
if err := db.Table(where.TableName()).Scopes(MakeCondition(where)).Select("bike.id,site.id,bike_model.id").Limit(limit).Offset(offset).Scan(&rs).Error; err != nil {
return nil, err
}
return rs, nil
}
func MakeCondition(q odao.Query) func(db *gorm.DB) *gorm.DB {
return func(db *gorm.DB) *gorm.DB {
condition := &odao.GormCondition{
GormPublic: odao.GormPublic{},
Join: make([]*odao.GormJoin, 0),
}
odao.ResolveSearchQuery("mysql", q, condition, q.TableName())
for _, join := range condition.Join {
if join == nil {
continue
}
db = db.Joins(join.JoinOn)
for k, v := range join.Where {
db = db.Where(k, v...)
}
for k, v := range join.Or {
db = db.Or(k, v...)
}
for _, o := range join.Order {
db = db.Order(o)
}
}
for k, v := range condition.Where {
db = db.Where(k, v...)
}
for k, v := range condition.Or {
db = db.Or(k, v...)
}
for _, o := range condition.Order {
db = db.Order(o)
}
return db
}
}
func GetGormTags(obj interface{}) []string {
var tags []string
val := reflect.ValueOf(obj)
if val.Kind() == reflect.Ptr {
val = val.Elem()
}
switch val.Kind() {
case reflect.Struct:
return processStruct(val)
case reflect.Slice, reflect.Array:
if val.Len() > 0 {
return processStruct(val.Index(0))
}
}
return tags
}
func processStruct(val reflect.Value) []string {
var tags []string
for i := 0; i < val.NumField(); i++ {
fieldVal := val.Field(i)
fieldType := val.Type().Field(i)
switch fieldVal.Kind() {
case reflect.Struct:
if prefix := getTablePrefix(fieldVal); prefix != "" {
tags = append(tags, extractFieldTags(fieldVal, prefix)...)
}
case reflect.Ptr:
if fieldVal.Elem().Kind() == reflect.Struct {
if prefix := getTablePrefix(fieldVal); prefix != "" {
tags = append(tags, extractFieldTags(fieldVal, prefix)...)
}
}
default:
tag := fieldType.Tag.Get("gorm")
if column := extractColumn(tag); column != "" {
tags = append(tags, column)
}
}
}
return tags
}
func getTablePrefix(val reflect.Value) string {
if val.CanInterface() {
if namer, ok := val.Interface().(TableNamer); ok {
return namer.TableName()
}
}
return ""
}
func extractFieldTags(val reflect.Value, prefix string) []string {
var tags []string
typ := val.Type()
for i := 0; i < typ.NumField(); i++ {
if column := extractColumn(typ.Field(i).Tag.Get("gorm")); column != "" {
tags = append(tags, fmt.Sprintf("%s.%s", prefix, column))
}
}
return tags
}
func extractColumn(tag string) string {
for _, part := range strings.Split(tag, ";") {
if strings.HasPrefix(part, "column:") {
return strings.Trim(strings.TrimPrefix(part, "column:"), `"`)
}
}
return ""
}
type TableNamer interface {
TableName() string
}
func main() {
record := Record{
Bike: Bike{ID: 1},
Site: Site{ID: 1},
}
record2 := Record{
Bike: Bike{ID: 2},
Site: Site{ID: 2},
}
records := make([]Record, 0, 2)
records = append(records, record, record2)
fmt.Println(GetGormTags(record))
fmt.Println(GetGormTags(records))
}
......@@ -68,12 +68,12 @@ type resolveSearchTag struct {
Type string
Column string
Table string
On []string
On string
Join string
}
// makeTag 解析search的tag标签
func makeTag(tag string) *resolveSearchTag {
// parseTag 解析search的tag标签
func parseTag(tag string) *resolveSearchTag {
r := &resolveSearchTag{}
tags := strings.Split(tag, ";")
var ts []string
......@@ -97,7 +97,7 @@ func makeTag(tag string) *resolveSearchTag {
}
case "on":
if len(ts) > 1 {
r.On = ts[1:]
r.On = ts[1]
}
case "join":
if len(ts) > 1 {
......
package odao
import (
"fmt"
"reflect"
"strings"
)
func GetGormTags(obj interface{}) []string {
var tags []string
val := reflect.ValueOf(obj)
if val.Kind() == reflect.Ptr {
val = val.Elem()
}
switch val.Kind() {
case reflect.Struct:
return processStruct(val)
case reflect.Slice, reflect.Array:
if val.Len() > 0 {
return processStruct(val.Index(0))
}
}
return tags
}
func processStruct(val reflect.Value) []string {
var tags []string
for i := 0; i < val.NumField(); i++ {
fieldVal := val.Field(i)
fieldType := val.Type().Field(i)
switch fieldVal.Kind() {
case reflect.Struct:
if prefix := getTablePrefix(fieldVal); prefix != "" {
tags = append(tags, extractFieldTags(fieldVal, prefix)...)
}
case reflect.Ptr:
if fieldVal.Elem().Kind() == reflect.Struct {
if prefix := getTablePrefix(fieldVal); prefix != "" {
tags = append(tags, extractFieldTags(fieldVal, prefix)...)
}
}
default:
tag := fieldType.Tag.Get("gorm")
if tag == "" {
tableName := toSnakeCase(val.Type().Name())
fieldName := toSnakeCase(fieldType.Name)
tags = append(tags, tableName+"."+fieldName)
}
if tag != "" {
if column := extractColumn(tag); column != "" {
tags = append(tags, column)
}
}
}
}
return tags
}
// 辅助函数:字段名转下划线格式
func toSnakeCase(name string) string {
var b strings.Builder
for i, r := range name {
if i == 0 {
b.WriteRune(tolower(r))
} else {
if isUpper(r) {
b.WriteByte('_')
b.WriteRune(tolower(r))
} else {
b.WriteRune(r)
}
}
}
return b.String()
}
// 判断单个字符是否为大写字母
func isUpper(r rune) bool {
return 'A' <= r && r <= 'Z'
}
func tolower(r rune) rune {
if 'A' <= r && r <= 'Z' {
return r - ('A' - 'a')
}
return r
}
func getTablePrefix(val reflect.Value) string {
if val.CanInterface() {
if namer, ok := val.Interface().(TableNamer); ok {
return namer.TableName()
}
}
return ""
}
func extractFieldTags(val reflect.Value, prefix string) []string {
var tags []string
typ := val.Type()
for i := 0; i < typ.NumField(); i++ {
if column := extractColumn(typ.Field(i).Tag.Get("gorm")); column != "" {
tags = append(tags, fmt.Sprintf("%s.%s", prefix, column))
}
}
return tags
}
func extractColumn(tag string) string {
for _, part := range strings.Split(tag, ";") {
if strings.HasPrefix(part, "column:") {
return strings.Trim(strings.TrimPrefix(part, "column:"), `"`)
}
}
return ""
}
type TableNamer interface {
TableName() string
}
......@@ -4,6 +4,7 @@ import (
"context"
"database/sql"
"fmt"
"strings"
"gitlab.wanzhuangkj.com/tush/xpkg/database"
"gitlab.wanzhuangkj.com/tush/xpkg/goredis"
......@@ -195,14 +196,14 @@ func (x *ODao[T]) UpdateByIDs(ctx context.Context, ids []xsf.ID, tb *T) error {
}
func (x *ODao[T]) UpdateByIDTx(ctx context.Context, tx *gorm.DB, id xsf.ID, tb *T) error {
return x.UpdByIDTx(ctx, tx, id, tb)
return x.updByIDTx(ctx, tx, id, tb)
}
// UpdByIDTx
// updByIDTx
// id 必传
// upd 可以是struct、map
// Updates 方法支持 struct 和 map[string]interface{} 参数。当使用 struct 更新时,默认情况下GORM 只会更新非零值的字段
func (x *ODao[T]) UpdByIDTx(ctx context.Context, tx *gorm.DB, id xsf.ID, upd any) error {
func (x *ODao[T]) updByIDTx(ctx context.Context, tx *gorm.DB, id xsf.ID, upd any) error {
if id <= 0 {
return nil
}
......@@ -216,24 +217,20 @@ func (x *ODao[T]) UpdByIDTx(ctx context.Context, tx *gorm.DB, id xsf.ID, upd any
return nil
}
func (x *ODao[T]) UpdByID(ctx context.Context, id xsf.ID, upd any) error {
return x.UpdByIDTx(ctx, x.DB(), id, upd)
}
func (x *ODao[T]) UpdateByIDsTx(ctx context.Context, tx *gorm.DB, ids []xsf.ID, tb *T) error {
return x.UpdByIDsTx(ctx, tx, ids, tb)
return x.updByIDsTx(ctx, tx, ids, tb)
}
// UpdByIDsTx
// updByIDsTx
// ids 必传
// upd 可以是struct、map
// Updates 方法支持 struct 和 map[string]interface{} 参数。当使用 struct 更新时,默认情况下GORM 只会更新非零值的字段
func (x *ODao[T]) UpdByIDsTx(ctx context.Context, tx *gorm.DB, ids []xsf.ID, upd any) error {
func (x *ODao[T]) updByIDsTx(ctx context.Context, tx *gorm.DB, ids []xsf.ID, upd any) error {
if len(ids) == 0 {
return nil
}
var table T
if err := tx.WithContext(ctx).Model(&table).Where("id in (?)", ids).Updates(upd).Error; err != nil {
if err := tx.WithContext(ctx).Model(&table).Where("id IN (?)", ids).Updates(upd).Error; err != nil {
return xerror.New(err.Error())
}
_ = x.deleteCaches(ctx, ids)
......@@ -282,6 +279,24 @@ func (x *ODao[T]) Page(ctx context.Context, where IParams) (rs []*T, total int64
return rs, total, nil
}
func (x *ODao[T]) JPage(ctx context.Context, where IParams, result any, resultType any) (total int64, err error) {
condition := x.DB().WithContext(ctx).Model(new(T)).Scopes(x.MakeCondition(where))
if where.Unscoped() {
condition = condition.Unscoped()
}
if err = condition.Limit(-1).Offset(-1).Count(&total).Error; err != nil {
return total, xerror.New(err.Error())
}
if total == 0 {
return total, nil
}
fields := GetGormTags(resultType)
if err = condition.Select(strings.Join(fields, ",")).Limit(where.GetLimit()).Offset(where.GetOffset()).Scan(result).Error; err != nil {
return total, xerror.New(err.Error())
}
return total, nil
}
func (x *ODao[T]) GetPageByColumns(ctx context.Context, params *query.Params) (rs []*T, total int64, err error) {
queryStr, args, err := params.ConvertToGormConditions()
if err != nil {
......@@ -531,7 +546,9 @@ func (x *ODao[T]) MakeCondition(q Query) func(db *gorm.DB) *gorm.DB {
GormPublic: GormPublic{},
Join: make([]*GormJoin, 0),
}
ResolveSearchQuery("mysql", q, condition, q.TableName())
for _, join := range condition.Join {
if join == nil {
continue
......
......@@ -91,14 +91,15 @@ func ResolveSearchQuery(driver string, q any, condition Condition, pTName string
ResolveSearchQuery(driver, qValue.Field(i).Interface(), condition, tname)
continue
}
switch tag {
case "-":
if tag == "-" {
continue
}
t = parseTag(tag)
if t.Join == "" {
if qValue.Field(i).IsZero() {
continue
}
t = makeTag(tag)
}
if t.Column == "" {
t.Column = snakeCase(qType.Field(i).Name, false)
}
......@@ -272,14 +273,16 @@ func otherSql(driver string, t *resolveSearchTag, condition Condition, qValue re
return
case JOIN:
//左关联
// join := condition.SetJoinOn(t.Type, fmt.Sprintf(
// "left join `%s` on %s",
// t.Join,
// t.On,
// ))
join := condition.SetJoinOn(t.Type, fmt.Sprintf(
"left join `%s` on `%s`.`%s` = `%s`.`%s`",
t.Join,
"%s",
t.Join,
t.On[0],
t.Table,
t.On[1],
))
// "left join `%s` on `%s`.`%s` = `%s`.`%s`",
ResolveSearchQuery(driver, qValue.Field(i).Interface(), join, tname)
return
default:
......
Markdown 格式
0%
您添加了 0 到此讨论。请谨慎行事。
请先完成此评论的编辑!
注册 或者 后发表评论