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 }