提交 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 { ...@@ -68,12 +68,12 @@ type resolveSearchTag struct {
Type string Type string
Column string Column string
Table string Table string
On []string On string
Join string Join string
} }
// makeTag 解析search的tag标签 // parseTag 解析search的tag标签
func makeTag(tag string) *resolveSearchTag { func parseTag(tag string) *resolveSearchTag {
r := &resolveSearchTag{} r := &resolveSearchTag{}
tags := strings.Split(tag, ";") tags := strings.Split(tag, ";")
var ts []string var ts []string
...@@ -97,7 +97,7 @@ func makeTag(tag string) *resolveSearchTag { ...@@ -97,7 +97,7 @@ func makeTag(tag string) *resolveSearchTag {
} }
case "on": case "on":
if len(ts) > 1 { if len(ts) > 1 {
r.On = ts[1:] r.On = ts[1]
} }
case "join": case "join":
if len(ts) > 1 { 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 ( ...@@ -4,6 +4,7 @@ import (
"context" "context"
"database/sql" "database/sql"
"fmt" "fmt"
"strings"
"gitlab.wanzhuangkj.com/tush/xpkg/database" "gitlab.wanzhuangkj.com/tush/xpkg/database"
"gitlab.wanzhuangkj.com/tush/xpkg/goredis" "gitlab.wanzhuangkj.com/tush/xpkg/goredis"
...@@ -195,14 +196,14 @@ func (x *ODao[T]) UpdateByIDs(ctx context.Context, ids []xsf.ID, tb *T) error { ...@@ -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 { 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 必传 // id 必传
// upd 可以是struct、map // upd 可以是struct、map
// Updates 方法支持 struct 和 map[string]interface{} 参数。当使用 struct 更新时,默认情况下GORM 只会更新非零值的字段 // 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 { if id <= 0 {
return nil return nil
} }
...@@ -216,24 +217,20 @@ func (x *ODao[T]) UpdByIDTx(ctx context.Context, tx *gorm.DB, id xsf.ID, upd any ...@@ -216,24 +217,20 @@ func (x *ODao[T]) UpdByIDTx(ctx context.Context, tx *gorm.DB, id xsf.ID, upd any
return nil 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 { 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 必传 // ids 必传
// upd 可以是struct、map // upd 可以是struct、map
// Updates 方法支持 struct 和 map[string]interface{} 参数。当使用 struct 更新时,默认情况下GORM 只会更新非零值的字段 // 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 { if len(ids) == 0 {
return nil return nil
} }
var table T 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()) return xerror.New(err.Error())
} }
_ = x.deleteCaches(ctx, ids) _ = x.deleteCaches(ctx, ids)
...@@ -282,6 +279,24 @@ func (x *ODao[T]) Page(ctx context.Context, where IParams) (rs []*T, total int64 ...@@ -282,6 +279,24 @@ func (x *ODao[T]) Page(ctx context.Context, where IParams) (rs []*T, total int64
return rs, total, nil 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) { func (x *ODao[T]) GetPageByColumns(ctx context.Context, params *query.Params) (rs []*T, total int64, err error) {
queryStr, args, err := params.ConvertToGormConditions() queryStr, args, err := params.ConvertToGormConditions()
if err != nil { if err != nil {
...@@ -531,7 +546,9 @@ func (x *ODao[T]) MakeCondition(q Query) func(db *gorm.DB) *gorm.DB { ...@@ -531,7 +546,9 @@ func (x *ODao[T]) MakeCondition(q Query) func(db *gorm.DB) *gorm.DB {
GormPublic: GormPublic{}, GormPublic: GormPublic{},
Join: make([]*GormJoin, 0), Join: make([]*GormJoin, 0),
} }
ResolveSearchQuery("mysql", q, condition, q.TableName()) ResolveSearchQuery("mysql", q, condition, q.TableName())
for _, join := range condition.Join { for _, join := range condition.Join {
if join == nil { if join == nil {
continue continue
......
...@@ -91,14 +91,15 @@ func ResolveSearchQuery(driver string, q any, condition Condition, pTName string ...@@ -91,14 +91,15 @@ func ResolveSearchQuery(driver string, q any, condition Condition, pTName string
ResolveSearchQuery(driver, qValue.Field(i).Interface(), condition, tname) ResolveSearchQuery(driver, qValue.Field(i).Interface(), condition, tname)
continue continue
} }
switch tag { if tag == "-" {
case "-":
continue continue
} }
if qValue.Field(i).IsZero() { t = parseTag(tag)
continue if t.Join == "" {
if qValue.Field(i).IsZero() {
continue
}
} }
t = makeTag(tag)
if t.Column == "" { if t.Column == "" {
t.Column = snakeCase(qType.Field(i).Name, false) t.Column = snakeCase(qType.Field(i).Name, false)
} }
...@@ -272,14 +273,16 @@ func otherSql(driver string, t *resolveSearchTag, condition Condition, qValue re ...@@ -272,14 +273,16 @@ func otherSql(driver string, t *resolveSearchTag, condition Condition, qValue re
return return
case JOIN: 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( join := condition.SetJoinOn(t.Type, fmt.Sprintf(
"left join `%s` on `%s`.`%s` = `%s`.`%s`", "%s",
t.Join,
t.Join, 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) ResolveSearchQuery(driver, qValue.Field(i).Interface(), join, tname)
return return
default: default:
......
Markdown 格式
0%
您添加了 0 到此讨论。请谨慎行事。
请先完成此评论的编辑!
注册 或者 后发表评论