package dataway import ( "errors" "fmt" "strconv" "time" "xorm.io/xorm" cache "xps/cmd/cache/redis" "xps/cmd/constant" "xps/pkg/base" ) type DatabaseIO struct { engine *xorm.Engine } func NewDatabaseIO(engine *xorm.Engine) *DatabaseIO { return &DatabaseIO{ engine: engine, } } func (d *DatabaseIO) BuildQueryParam(params []base.QueryParam) string { sqlWhere := " WHERE 1 = 1 " for _, param := range params { paramType := param.Type if paramType == constant.DATA_TYPE_DATE || paramType == constant.DATA_TYPE_DATETIME || paramType == constant.DATA_TYPE_VARCHAR { sqlWhere += fmt.Sprintf(" %s %v %s '%v'", param.LogicalOperator, param.Name, param.CompareOperator, param.Value) } else { sqlWhere += fmt.Sprintf(" %s %v %s %v", param.LogicalOperator, param.Name, param.CompareOperator, param.Value) } } return sqlWhere } func (d *DatabaseIO) GetPage(sqlStatement, sqlCount string, params []base.QueryParam, page int, limit int) (*base.PageResult, error) { sqlWhere := d.BuildQueryParam(params) res, err1 := d.engine.Query(sqlCount + sqlWhere) if err1 != nil { return nil, err1 } total := int64(0) for _, v := range res[0] { total, _ = strconv.ParseInt(string(v), 10, 64) } pageStart := (page - 1) * limit sqlLimit := fmt.Sprintf(" LIMIT %d OFFSET %d ", limit, pageStart) data, err2 := d.engine.QueryInterface(sqlStatement + sqlWhere + sqlLimit) if err2 != nil { return nil, err2 } pageResult := &base.PageResult{ Total: total, PageSize: limit, Page: page, Data: data, } return pageResult, nil } func (d *DatabaseIO) GetList(sqlStatement string, params []base.QueryParam) ([]map[string]interface{}, error) { sqlWhere := d.BuildQueryParam(params) results, err := d.engine.QueryInterface(sqlStatement + sqlWhere) if err != nil { return nil, err } return results, nil } func (d *DatabaseIO) GetById(sqlStatement string, params []base.QueryParam) (map[string]interface{}, error) { sqlWhere := d.BuildQueryParam(params) results, err := d.engine.QueryInterface(sqlStatement + sqlWhere) if err != nil { return nil, err } if len(results) > 1 { return nil, errors.New(" 主键不唯一") } return results[0], nil } func (d *DatabaseIO) Create(objectCode string, values map[string]interface{}) error { object, err1 := cache.GetModelObject(objectCode) if err1 != nil { return err1 } tableName := object.ObjectCode fields, err2 := cache.GetModelObjectAttrs(objectCode) if err2 != nil { return err2 } sql := "INSERT INTO" + fmt.Sprintf(" %s ", tableName) sqlFields := "( " sqlValues := "VALUES ( " for idx, field := range fields { fieldName := field.AttrCode fieldType := field.AttrType fieldValue, ok := values[fieldName] if !ok { return nil } if idx > 0 { sqlFields += "," sqlValues += "," } sqlFields += fieldName + "" if fieldType == constant.DATA_TYPE_DATE || fieldType == constant.DATA_TYPE_DATETIME || fieldType == constant.DATA_TYPE_VARCHAR { sqlValues += fmt.Sprintf("'%v'", fieldValue) + "" } else { sqlValues += fmt.Sprintf("%v", fieldValue) + "" } } sqlFields += ") " sqlValues += ") " _, err3 := d.engine.Exec(sql + sqlFields + sqlValues) return err3 } func (d *DatabaseIO) Update(objectCode string, values map[string]interface{}) error { object, err1 := cache.GetModelObject(objectCode) if err1 != nil { return err1 } tableName := object.ObjectCode fields, err2 := cache.GetModelObjectAttrs(objectCode) if err2 != nil { return err2 } sql := fmt.Sprintf("UPDATE %s ", tableName) sqlFieldSets := " SET " sqlWhere := " WHERE 1 = 1 " bFirstSetValue := true for _, field := range fields { fieldName := field.AttrCode fieldType := field.AttrType fieldValue, ok := values[fieldName] if !ok { return nil } if !bFirstSetValue { sqlFieldSets += "," } if field.IsPKey <= 0 { bFirstSetValue = false if fieldType == constant.DATA_TYPE_DATE || fieldType == constant.DATA_TYPE_DATETIME || fieldType == constant.DATA_TYPE_VARCHAR { sqlFieldSets += fmt.Sprintf(" %s = '%v' ", fieldName, fieldValue) + "" } else { sqlFieldSets += fmt.Sprintf(" %s = %v ", fieldName, fieldValue) + "" } } else { if fieldType == constant.DATA_TYPE_DATE || fieldType == constant.DATA_TYPE_DATETIME || fieldType == constant.DATA_TYPE_VARCHAR { sqlWhere += fmt.Sprintf(" AND %s = '%v' ", fieldName, fieldValue) + "" } else { sqlWhere += fmt.Sprintf(" AND %s = %v ", fieldName, fieldValue) + "" } } } _, err := d.engine.Exec(sql + sqlFieldSets + sqlWhere) return err } func (d *DatabaseIO) Patch(objectCode string, values map[string]interface{}) error { object, err1 := cache.GetModelObject(objectCode) if err1 != nil { return err1 } tableName := object.ObjectCode fields, err2 := cache.GetModelObjectAttrs(objectCode) if err2 != nil { return err2 } sql := fmt.Sprintf("UPDATE table %s ", tableName) sqlFieldSets := "SET " sqlWhere := "WHERE 1 = 1" for idx, field := range fields { fieldName := field.AttrCode fieldType := field.AttrType fieldValue, ok := values[fieldName] if !ok { return nil } if idx > 0 { sqlFieldSets += "," } if field.IsPKey <= 0 { if fieldType == constant.DATA_TYPE_DATE || fieldType == constant.DATA_TYPE_DATETIME || fieldType == constant.DATA_TYPE_VARCHAR { sqlFieldSets += fmt.Sprintf(" %s = '%v' ", fieldName, fieldValue) + "" } else { sqlFieldSets += fmt.Sprintf(" %s = %v ", fieldName, fieldValue) + "" } } else { if fieldType == constant.DATA_TYPE_DATE || fieldType == constant.DATA_TYPE_DATETIME || fieldType == constant.DATA_TYPE_VARCHAR { sqlWhere += fmt.Sprintf(" AND %s = '%v' ", fieldName, fieldValue) + "" } else { sqlWhere += fmt.Sprintf(" AND %s = %v ", fieldName, fieldValue) + "" } } } _, err := d.engine.Exec(sql + sqlFieldSets + sqlWhere) return err } func (d *DatabaseIO) Delete(objectCode string, values map[string]interface{}) error { object, err1 := cache.GetModelObject(objectCode) if err1 == nil { return err1 } tableName := object.ObjectCode fields, err2 := cache.GetModelObjectAttrs(objectCode) if err2 != nil { return err2 } deletedBy := values["deletedBy"].(int64) sql := fmt.Sprintf("UPDATE table %s ", tableName) sqlFieldSets := fmt.Sprintf("SET deleted_flag = 1 AND deleted_by = %d AND deleted_at = %s", deletedBy, time.Now().String()) sqlWhere := "WHERE 1 = 1" for idx, field := range fields { fieldName := field.AttrCode fieldType := field.AttrType fieldValue, ok := values[fieldName] if !ok { return nil } if idx > 0 { sqlFieldSets += "," } if field.IsPKey >= 1 { if fieldType == constant.DATA_TYPE_DATE || fieldType == constant.DATA_TYPE_DATETIME || fieldType == constant.DATA_TYPE_VARCHAR { sqlFieldSets += fmt.Sprintf(" %s = '%v' ", fieldName, fieldValue) + "" } else { sqlFieldSets += fmt.Sprintf(" %s = %v ", fieldName, fieldValue) + "" } } } _, err := d.engine.Exec(sql + sqlFieldSets + sqlWhere) return err }