You can not select more than 25 topics Topics must start with a letter or number, can include dashes ('-') and can be up to 35 characters long.
 
 
 
 
 
 

188 lines
6.4 KiB

package memory
import (
"database/sql"
"database/sql/driver"
"reflect"
"strconv"
"strings"
"testing"
"yunyan/pb"
"gorm.io/driver/mysql"
"gorm.io/gorm"
)
/*
守住 addItem 那条补写逻辑的两个前提。
背景:gorm 对带 `default:` 标签的字段,Create 时把零值换成标签里的默认值再写库。
DBMemoryItem 的 date_certain(default:true) / remind_ahead(default:5) 因此永远写不进
false / 0 —— 2026-09-18 测试库里 25 条记录全部落成 1 / 5。
⚠️ 一度以为加 `Select("*")` 能绕开,实测不行:替换发生在 ConvertToCreateValues 的
reflect.Struct 分支,只看字段值是不是零值,与 Select / Omit 无关。
TestGormCreateSwallowsZeroValues 就是把这个事实钉住,免得下次又有人去试。
*/
// ── DryRun 用的假驱动:只要能让 gorm 把 SQL 拼出来,不需要真连库 ──
type nullDriver struct{}
func (nullDriver) Open(string) (driver.Conn, error) { return nullConn{}, nil }
type nullConn struct{}
func (nullConn) Prepare(string) (driver.Stmt, error) { return nil, driver.ErrSkip }
func (nullConn) Close() error { return nil }
func (nullConn) Begin() (driver.Tx, error) { return nil, driver.ErrSkip }
func init() { sql.Register("memory_model_test_null", nullDriver{}) }
func dryRunDB(t *testing.T) *gorm.DB {
t.Helper()
conn, _ := sql.Open("memory_model_test_null", "")
db, err := gorm.Open(
mysql.New(mysql.Config{Conn: conn, SkipInitializeWithVersion: true}),
// SkipDefaultTransaction:写操作默认包事务,假驱动的 Begin 会报错,SQL 就拼不出来了
&gorm.Config{DryRun: true, DisableAutomaticPing: true, SkipDefaultTransaction: true},
)
if err != nil {
t.Fatalf("起 DryRun 连接失败: %v", err)
}
return db
}
// insertValueOf 取一条 INSERT 里某列的绑定值。
func insertValueOf(t *testing.T, st *gorm.Statement, column string) interface{} {
t.Helper()
sql := st.SQL.String()
head := sql[strings.Index(sql, "(")+1 : strings.Index(sql, ")")]
for i, c := range strings.Split(head, ",") {
if strings.Trim(strings.TrimSpace(c), "`") == column {
if i >= len(st.Vars) {
t.Fatalf("列 %s 下标 %d 超出绑定值个数 %d", column, i, len(st.Vars))
}
return st.Vars[i]
}
}
t.Fatalf("INSERT 里没有列 %s:%s", column, sql)
return nil
}
// TestGormCreateSwallowsZeroValues 记录 bug 本身:Create 会把 false/0 换成 true/5,
// 加不加 Select("*") 都一样。哪天 gorm 改了行为这条会红,那时 addItem 的补写就可以删了。
func TestGormCreateSwallowsZeroValues(t *testing.T) {
for _, tc := range []struct {
name string
make func(*gorm.DB, *pb.DBMemoryItem) *gorm.DB
}{
{"裸 Create", func(db *gorm.DB, it *pb.DBMemoryItem) *gorm.DB {
return db.Table("memory_item").Create(it)
}},
{"Select(*) 也救不了", func(db *gorm.DB, it *pb.DBMemoryItem) *gorm.DB {
return db.Table("memory_item").Select("*").Create(it)
}},
} {
t.Run(tc.name, func(t *testing.T) {
item := &pb.DBMemoryItem{
Uid: "u1", Title: "跟进合同", HappenDate: "2026-09-18",
DateCertain: false, RemindAhead: 0,
}
st := tc.make(dryRunDB(t), item).Statement
if got := insertValueOf(t, st, "date_certain"); got != true {
t.Errorf("date_certain 期望被换成 true(即 bug 仍在),实际 %v", got)
}
// ⚠️ 比数值不比类型:gorm 的默认值来自 strconv.ParseInt,塞回去是 int64 不是字段的 int32
if got := insertValueOf(t, st, "remind_ahead"); toInt64(t, got) != 5 {
t.Errorf("remind_ahead 期望被换成 5(即 bug 仍在),实际 %v", got)
}
})
}
}
// TestGormNonZeroDefaultsAreHandled 反查 pb 结构体上所有「默认值非零」的列,
// 必须与 gormZeroDefaultColumns 一字不差。
//
// proto 里给某个字段加了 `default:1` 之类而 addItem 没跟着补写,这条会红 ——
// 否则那个字段会重演同一个 bug,而且同样不报任何错。
func TestGormNonZeroDefaultsAreHandled(t *testing.T) {
got := nonZeroDefaultColumns(reflect.TypeOf(pb.DBMemoryItem{}))
want := append([]string(nil), gormZeroDefaultColumns...)
if !reflect.DeepEqual(got, sortedCopy(want)) {
t.Fatalf("DBMemoryItem 上非零默认值的列是 %v,而 addItem 只补写 %v;"+
"请同步 gormZeroDefaultColumns 与 addItem 里的补写分支", got, want)
}
}
// TestReportHasNoNonZeroDefaults 周期报告表走同一个 mysql.Insert,
// 目前它没有非零默认值的列,所以不需要补写。加了就要按 addItem 的办法处理。
func TestReportHasNoNonZeroDefaults(t *testing.T) {
if got := nonZeroDefaultColumns(reflect.TypeOf(pb.DBMemoryReport{})); len(got) > 0 {
t.Fatalf("DBMemoryReport 新增了非零默认值的列 %v,upsertReport 的 Insert 会把它们的零值吞掉", got)
}
}
// nonZeroDefaultColumns 扫结构体 tag,挑出 `gorm:"default:X"` 且 X 不是零值的列,
// 列名取 json tag(pb 生成的 json 名与建表列名一致)。返回值已排序。
func nonZeroDefaultColumns(t reflect.Type) []string {
out := make([]string, 0, 4)
for i := 0; i < t.NumField(); i++ {
f := t.Field(i)
tag := f.Tag.Get("gorm")
if tag == "" || tag == "-" {
continue
}
def, ok := "", false
for _, seg := range strings.Split(tag, ";") {
if v, found := strings.CutPrefix(strings.TrimSpace(seg), "default:"); found {
def, ok = strings.TrimSpace(v), true
}
}
if !ok || isZeroLiteral(def) {
continue
}
name := strings.Split(f.Tag.Get("json"), ",")[0]
if name == "" || name == "-" {
continue
}
out = append(out, name)
}
return sortedCopy(out)
}
// toInt64 把绑定值统一成 int64 再比,绕开 int32/int64 的类型差异。
func toInt64(t *testing.T, v interface{}) int64 {
t.Helper()
rv := reflect.ValueOf(v)
switch rv.Kind() {
case reflect.Int, reflect.Int8, reflect.Int16, reflect.Int32, reflect.Int64:
return rv.Int()
case reflect.Uint, reflect.Uint8, reflect.Uint16, reflect.Uint32, reflect.Uint64:
return int64(rv.Uint())
}
t.Fatalf("不是整数: %#v", v)
return 0
}
func isZeroLiteral(s string) bool {
switch strings.ToLower(strings.TrimSpace(s)) {
case "", "0", "false", "null", "''", `""`:
return true
}
if f, err := strconv.ParseFloat(s, 64); err == nil {
return f == 0
}
return false
}
func sortedCopy(in []string) []string {
out := append([]string(nil), in...)
for i := 1; i < len(out); i++ {
for j := i; j > 0 && out[j] < out[j-1]; j-- {
out[j], out[j-1] = out[j-1], out[j]
}
}
return out
}