diff --git a/examples/gormtypes/main.go b/examples/gormtypes/main.go new file mode 100644 index 0000000..bd00ffb --- /dev/null +++ b/examples/gormtypes/main.go @@ -0,0 +1,70 @@ +// gormtypes 包示例:GORM/JSON 双向转换类型。 +// 同一类型既能在 GORM 层序列化为 JSON 存库(Value/Scan), +// 又能在 API 层按前端需要的格式输出(MarshalJSON/UnmarshalJSON)。 +package main + +import ( + "encoding/json" + "fmt" + + "git.zeroonesoft.cn/golib/zogo/gormtypes" +) + +// Model 模拟 GORM 模型 +type Model struct { + Tags gormtypes.StringSlice `json:"tags"` + Nums gormtypes.IntSlice `json:"nums"` + Flags gormtypes.BoolSlice `json:"flags"` + Extra gormtypes.StringMap `json:"extra"` +} + +func main() { + // ===== StringSlice ===== + s := gormtypes.StringSlice{"a", "b", "c"} + fmt.Println("Contains b:", s.Contains("b")) + fmt.Println("Join:", s.Join(",")) + fmt.Println("ToSlice:", s.ToSlice()) + buf, _ := json.Marshal(s) + fmt.Println("MarshalJSON:", string(buf)) + + // Value/Scan:GORM 写库/读库时调用,nil 统一转 [] + v, _ := s.Value() + fmt.Printf("Value: %s\n", v) + var s2 gormtypes.StringSlice + _ = s2.Scan(v) + fmt.Println("Scan 后:", s2) + + // ===== IntSlice:MarshalJSON 输出字符串数组(避免前端 JS 大数精度丢失)===== + nums := gormtypes.IntSlice{1, 2, 3} + buf, _ = json.Marshal(nums) + fmt.Println("IntSlice MarshalJSON:", string(buf)) + var nums2 gormtypes.IntSlice + // 反序列化同时兼容字符串数组和数字数组 + _ = json.Unmarshal([]byte(`["1","2","3"]`), &nums2) + fmt.Println("IntSlice UnmarshalJSON(字符串数组):", nums2) + _ = json.Unmarshal([]byte(`[4,5]`), &nums2) + fmt.Println("IntSlice UnmarshalJSON(数字数组):", nums2) + + // ===== StringMap ===== + m := gormtypes.StringMap{} + m.Set("k1", "v1") + m.Set("k2", "v2") + fmt.Println("Get k1:", m.Get("k1")) + fmt.Println("Keys:", m.Keys()) + fmt.Println("Values:", m.Values()) + buf, _ = json.Marshal(m) + fmt.Println("StringMap MarshalJSON:", string(buf)) + + // ===== BoolSlice ===== + flags := gormtypes.BoolSlice{true, false, true} + fmt.Println("AnyTrue:", flags.AnyTrue(), "AllTrue:", flags.AllTrue(), "CountTrue:", flags.CountTrue()) + + // ===== 完整模型 JSON 序列化 ===== + mo := Model{Tags: s, Nums: nums, Flags: flags, Extra: m} + buf, _ = json.Marshal(mo) + fmt.Println("Model JSON:", string(buf)) + + // ===== 从 JSON 反序列化(等价于从数据库 JSON 字符串 Scan)===== + _ = json.Unmarshal([]byte(`{"tags":["x"],"nums":["9"],"flags":[true],"extra":{"a":"b"}}`), &mo) + fmt.Printf("Unmarshal 后: %+v\n", mo) +} diff --git a/gormtypes/README.md b/gormtypes/README.md new file mode 100644 index 0000000..8b8ec0e --- /dev/null +++ b/gormtypes/README.md @@ -0,0 +1,36 @@ +# gormtypes + +GORM/JSON 双向转换类型:写库时序列化为 JSON 字符串,读库时反序列化还原, +API 输出按前端需要的格式直出。含 BoolSlice、IntSlice、Int32Slice、Int64Slice、 +UInt64Slice、Float64Slice、StringSlice、StringMap 八种类型。 + +> 迁移自 go-hua/datatypes 并更名:原包名与 `gorm.io/datatypes` 官方包易混淆, +> 且该包专为 GORM 服务,故更名 `gormtypes`。 + +## 用法 + +```go +import "git.zeroonesoft.cn/golib/zogo/gormtypes" + +type Server struct { + ID int64 `json:"id"` + Tags gormtypes.StringSlice `gorm:"type:varchar(512)" json:"tags"` // 库里存 `["a","b"]` + Flags gormtypes.BoolSlice `gorm:"type:varchar(256)" json:"flags"` // 库里存 `[true,false]` + Extra gormtypes.StringMap `gorm:"type:varchar(1024)" json:"extra"` +} + +// 每种类型附辅助方法(以 BoolSlice 为例) +s := gormtypes.BoolSlice{true, false, true} +s.AnyTrue() // true +s.AllTrue() // false +s.CountTrue() // 2 +s.ToSlice() // []bool +``` + +完整可运行例程:[examples/gormtypes/main.go](../examples/gormtypes/main.go) + +## 注意 + +- 数据库列建议 `varchar`/`text`(存 JSON 字符串),不要用 JSON 列类型。 +- 消费方 API 输出时 `MarshalJSON` 直出原生数组/对象,前端无感知。 +- 迁移自 go-hua 的消费方:import 路径与包名 `datatypes` → `gormtypes`,类型名不变。 diff --git a/gormtypes/bool_slice.go b/gormtypes/bool_slice.go new file mode 100644 index 0000000..e7a2b80 --- /dev/null +++ b/gormtypes/bool_slice.go @@ -0,0 +1,116 @@ +// zogo/gormtypes/bool_slice.go +// Package gormtypes 提供 GORM/JSON 双向转换类型:写库序列化为 JSON 字符串,API 输出按前端需要的格式。 +package gormtypes + +import ( + "database/sql/driver" + "encoding/json" + "strings" +) + +// BoolSlice []bool 类型别名 +// 用于存储和传输布尔数组 +// - GORM 层:序列化为 JSON 字符串存储到数据库 +// - API 层:前端传 boolean[],返回 boolean[] +type BoolSlice []bool + +// Value 实现 driver.Valuer 接口,用于写入数据库 +func (s BoolSlice) Value() (driver.Value, error) { + if s == nil { + return []byte("[]"), nil + } + return json.Marshal(s) +} + +// Scan 实现 sql.Scanner 接口,用于从数据库读取 +func (s *BoolSlice) Scan(value interface{}) error { + if value == nil { + *s = BoolSlice{} + return nil + } + + var data []byte + switch v := value.(type) { + case []byte: + data = v + case string: + data = []byte(v) + default: + return nil + } + + return json.Unmarshal(data, s) +} + +// MarshalJSON 序列化输出:boolean[] 格式 +func (s BoolSlice) MarshalJSON() ([]byte, error) { + if s == nil { + return []byte("null"), nil + } + + var b strings.Builder + b.WriteByte('[') + for i, v := range s { + if i > 0 { + b.WriteByte(',') + } + if v { + b.WriteString("true") + } else { + b.WriteString("false") + } + } + b.WriteByte(']') + return []byte(b.String()), nil +} + +// UnmarshalJSON 反序列化支持 boolean[] 格式 +func (s *BoolSlice) UnmarshalJSON(data []byte) error { + if string(data) == "null" { + *s = nil + return nil + } + + var arr []bool + if err := json.Unmarshal(data, &arr); err != nil { + return err + } + *s = BoolSlice(arr) + return nil +} + +// ToSlice 转换为原生 []bool 切片,用于 ... 展开 +func (s BoolSlice) ToSlice() []bool { + return []bool(s) +} + +// AnyTrue 检查是否任意一个为 true +func (s BoolSlice) AnyTrue() bool { + for _, v := range s { + if v { + return true + } + } + return false +} + +// AllTrue 检查是否全部为 true +func (s BoolSlice) AllTrue() bool { + for _, v := range s { + if !v { + return false + } + } + return len(s) > 0 +} + +// CountTrue 统计 true 的数量 +func (s BoolSlice) CountTrue() int { + cnt := 0 + for _, v := range s { + if v { + cnt++ + } + } + return cnt +} diff --git a/gormtypes/bool_slice_test.go b/gormtypes/bool_slice_test.go new file mode 100644 index 0000000..683c44d --- /dev/null +++ b/gormtypes/bool_slice_test.go @@ -0,0 +1,320 @@ +// zogo/gormtypes/bool_slice_test.go +package gormtypes + +import ( + "encoding/json" + "testing" +) + +func TestBoolSlice_Value(t *testing.T) { + tests := []struct { + name string + s BoolSlice + want []byte + wantErr bool + }{ + { + name: "nil slice", + s: nil, + want: []byte("[]"), + wantErr: false, + }, + { + name: "empty slice", + s: BoolSlice{}, + want: []byte("[]"), + wantErr: false, + }, + { + name: "single true", + s: BoolSlice{true}, + want: []byte("[true]"), + wantErr: false, + }, + { + name: "single false", + s: BoolSlice{false}, + want: []byte("[false]"), + wantErr: false, + }, + { + name: "mixed", + s: BoolSlice{true, false, true}, + want: []byte("[true,false,true]"), + wantErr: false, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + got, err := tt.s.Value() + if (err != nil) != tt.wantErr { + t.Errorf("Value() error = %v, wantErr %v", err, tt.wantErr) + return + } + gotBytes, ok := got.([]byte) + if !ok { + t.Errorf("Value() returned non-[]byte type: %T", got) + return + } + if string(gotBytes) != string(tt.want) { + t.Errorf("Value() = %v, want %v", string(gotBytes), string(tt.want)) + } + }) + } +} + +func TestBoolSlice_Scan(t *testing.T) { + tests := []struct { + name string + data interface{} + want BoolSlice + wantErr bool + }{ + { + name: "nil value", + data: nil, + want: BoolSlice{}, + wantErr: false, + }, + { + name: "empty json", + data: []byte("[]"), + want: BoolSlice{}, + wantErr: false, + }, + { + name: "single true", + data: []byte("[true]"), + want: BoolSlice{true}, + wantErr: false, + }, + { + name: "mixed", + data: []byte("[true,false,true]"), + want: BoolSlice{true, false, true}, + wantErr: false, + }, + { + name: "invalid json", + data: []byte("invalid"), + want: nil, + wantErr: true, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + var s BoolSlice + err := s.Scan(tt.data) + if (err != nil) != tt.wantErr { + t.Errorf("Scan() error = %v, wantErr %v", err, tt.wantErr) + return + } + if !tt.wantErr && len(s) != len(tt.want) { + t.Errorf("Scan() = %v, want %v", s, tt.want) + } + }) + } +} + +func TestBoolSlice_MarshalJSON(t *testing.T) { + tests := []struct { + name string + s BoolSlice + want string + }{ + { + name: "nil slice", + s: nil, + want: "null", + }, + { + name: "empty slice", + s: BoolSlice{}, + want: "[]", + }, + { + name: "single true", + s: BoolSlice{true}, + want: "[true]", + }, + { + name: "single false", + s: BoolSlice{false}, + want: "[false]", + }, + { + name: "mixed", + s: BoolSlice{true, false, true}, + want: "[true,false,true]", + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + got, err := json.Marshal(tt.s) + if err != nil { + t.Errorf("MarshalJSON() error = %v", err) + return + } + if string(got) != tt.want { + t.Errorf("MarshalJSON() = %v, want %v", string(got), tt.want) + } + }) + } +} + +func TestBoolSlice_UnmarshalJSON(t *testing.T) { + tests := []struct { + name string + data string + want BoolSlice + wantErr bool + }{ + { + name: "null", + data: "null", + want: nil, + wantErr: false, + }, + { + name: "empty array", + data: "[]", + want: BoolSlice{}, + wantErr: false, + }, + { + name: "single true", + data: "[true]", + want: BoolSlice{true}, + wantErr: false, + }, + { + name: "single false", + data: "[false]", + want: BoolSlice{false}, + wantErr: false, + }, + { + name: "mixed", + data: "[true,false,true]", + want: BoolSlice{true, false, true}, + wantErr: false, + }, + { + name: "invalid json", + data: "invalid", + want: nil, + wantErr: true, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + var s BoolSlice + err := json.Unmarshal([]byte(tt.data), &s) + if (err != nil) != tt.wantErr { + t.Errorf("UnmarshalJSON() error = %v, wantErr %v", err, tt.wantErr) + return + } + if !tt.wantErr && len(s) != len(tt.want) { + t.Errorf("UnmarshalJSON() = %v, want %v", s, tt.want) + } + }) + } +} + +func TestBoolSlice_AnyTrue(t *testing.T) { + tests := []struct { + name string + s BoolSlice + want bool + }{ + {"nil", nil, false}, + {"empty", BoolSlice{}, false}, + {"all false", BoolSlice{false, false}, false}, + {"first true", BoolSlice{true, false}, true}, + {"last true", BoolSlice{false, true}, true}, + {"all true", BoolSlice{true, true}, true}, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + if got := tt.s.AnyTrue(); got != tt.want { + t.Errorf("AnyTrue() = %v, want %v", got, tt.want) + } + }) + } +} + +func TestBoolSlice_AllTrue(t *testing.T) { + tests := []struct { + name string + s BoolSlice + want bool + }{ + {"nil", nil, false}, + {"empty", BoolSlice{}, false}, + {"all true", BoolSlice{true, true}, true}, + {"one false", BoolSlice{true, false}, false}, + {"first false", BoolSlice{false, true}, false}, + {"all false", BoolSlice{false, false}, false}, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + if got := tt.s.AllTrue(); got != tt.want { + t.Errorf("AllTrue() = %v, want %v", got, tt.want) + } + }) + } +} + +func TestBoolSlice_CountTrue(t *testing.T) { + tests := []struct { + name string + s BoolSlice + want int + }{ + {"nil", nil, 0}, + {"empty", BoolSlice{}, 0}, + {"all false", BoolSlice{false, false, false}, 0}, + {"all true", BoolSlice{true, true, true}, 3}, + {"mixed", BoolSlice{true, false, true, false, true}, 3}, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + if got := tt.s.CountTrue(); got != tt.want { + t.Errorf("CountTrue() = %v, want %v", got, tt.want) + } + }) + } +} + +func TestBoolSlice_ToSlice(t *testing.T) { + tests := []struct { + name string + s BoolSlice + want []bool + }{ + {"nil", nil, nil}, + {"empty", BoolSlice{}, []bool{}}, + {"single true", BoolSlice{true}, []bool{true}}, + {"multiple", BoolSlice{true, false, true}, []bool{true, false, true}}, + } + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + got := tt.s.ToSlice() + if len(got) != len(tt.want) { + t.Fatalf("ToSlice() len = %d, want %d", len(got), len(tt.want)) + } + for i := range got { + if got[i] != tt.want[i] { + t.Errorf("ToSlice()[%d] = %v, want %v", i, got[i], tt.want[i]) + } + } + }) + } +} diff --git a/gormtypes/float64_slice.go b/gormtypes/float64_slice.go new file mode 100644 index 0000000..3630076 --- /dev/null +++ b/gormtypes/float64_slice.go @@ -0,0 +1,81 @@ +// zogo/gormtypes/float64_slice.go +package gormtypes + +import ( + "database/sql/driver" + "encoding/json" + "strconv" + "strings" +) + +// Float64Slice []float64 类型别名 +// 用于存储和传输 float64 数组 +// - GORM 层:序列化为 JSON 字符串存储到数据库 +// - API 层:前端传 number[],返回 number[](浮点数不适合字符串格式) +type Float64Slice []float64 + +// Value 实现 driver.Valuer 接口,用于写入数据库 +func (s Float64Slice) Value() (driver.Value, error) { + if s == nil { + return []byte("[]"), nil + } + return json.Marshal(s) +} + +// Scan 实现 sql.Scanner 接口,用于从数据库读取 +func (s *Float64Slice) Scan(value interface{}) error { + if value == nil { + *s = Float64Slice{} + return nil + } + + var data []byte + switch v := value.(type) { + case []byte: + data = v + case string: + data = []byte(v) + default: + return nil + } + + return json.Unmarshal(data, s) +} + +// MarshalJSON 序列化输出:number[] 格式 +func (s Float64Slice) MarshalJSON() ([]byte, error) { + if s == nil { + return []byte("null"), nil + } + + var b strings.Builder + b.WriteByte('[') + for i, v := range s { + if i > 0 { + b.WriteByte(',') + } + b.WriteString(strconv.FormatFloat(v, 'f', -1, 64)) + } + b.WriteByte(']') + return []byte(b.String()), nil +} + +// UnmarshalJSON 反序列化支持 number[] 格式 +func (s *Float64Slice) UnmarshalJSON(data []byte) error { + if string(data) == "null" { + *s = nil + return nil + } + + var arr []float64 + if err := json.Unmarshal(data, &arr); err != nil { + return err + } + *s = Float64Slice(arr) + return nil +} + +// ToSlice 转换为原生 []float64 切片,用于 ... 展开 +func (s Float64Slice) ToSlice() []float64 { + return []float64(s) +} diff --git a/gormtypes/float64_slice_test.go b/gormtypes/float64_slice_test.go new file mode 100644 index 0000000..5a75f19 --- /dev/null +++ b/gormtypes/float64_slice_test.go @@ -0,0 +1,281 @@ +// zogo/gormtypes/float64_slice_test.go +package gormtypes + +import ( + "encoding/json" + "testing" +) + +func TestFloat64Slice_Value(t *testing.T) { + tests := []struct { + name string + s Float64Slice + want []byte + wantErr bool + }{ + { + name: "nil slice", + s: nil, + want: []byte("[]"), + wantErr: false, + }, + { + name: "empty slice", + s: Float64Slice{}, + want: []byte("[]"), + wantErr: false, + }, + { + name: "single element", + s: Float64Slice{1.5}, + want: []byte("[1.5]"), + wantErr: false, + }, + { + name: "multiple elements", + s: Float64Slice{1.1, 2.2, 3.3}, + want: []byte("[1.1,2.2,3.3]"), + wantErr: false, + }, + { + name: "with precision", + s: Float64Slice{3.141592653589793}, + want: []byte("[3.141592653589793]"), + wantErr: false, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + got, err := tt.s.Value() + if (err != nil) != tt.wantErr { + t.Errorf("Value() error = %v, wantErr %v", err, tt.wantErr) + return + } + gotBytes, ok := got.([]byte) + if !ok { + t.Errorf("Value() returned non-[]byte type: %T", got) + return + } + if string(gotBytes) != string(tt.want) { + t.Errorf("Value() = %v, want %v", string(gotBytes), string(tt.want)) + } + }) + } +} + +func TestFloat64Slice_Scan(t *testing.T) { + tests := []struct { + name string + data interface{} + want Float64Slice + wantErr bool + }{ + { + name: "nil value", + data: nil, + want: Float64Slice{}, + wantErr: false, + }, + { + name: "empty json", + data: []byte("[]"), + want: Float64Slice{}, + wantErr: false, + }, + { + name: "single element", + data: []byte("[1.5]"), + want: Float64Slice{1.5}, + wantErr: false, + }, + { + name: "multiple elements", + data: []byte("[1.1,2.2,3.3]"), + want: Float64Slice{1.1, 2.2, 3.3}, + wantErr: false, + }, + { + name: "invalid json", + data: []byte("invalid"), + want: nil, + wantErr: true, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + var s Float64Slice + err := s.Scan(tt.data) + if (err != nil) != tt.wantErr { + t.Errorf("Scan() error = %v, wantErr %v", err, tt.wantErr) + return + } + if !tt.wantErr && len(s) != len(tt.want) { + t.Errorf("Scan() = %v, want %v", s, tt.want) + } + }) + } +} + +func TestFloat64Slice_MarshalJSON(t *testing.T) { + tests := []struct { + name string + s Float64Slice + want string + }{ + { + name: "nil slice", + s: nil, + want: "null", + }, + { + name: "empty slice", + s: Float64Slice{}, + want: "[]", + }, + { + name: "single element", + s: Float64Slice{1.5}, + want: "[1.5]", + }, + { + name: "multiple elements", + s: Float64Slice{1.1, 2.2, 3.3}, + want: "[1.1,2.2,3.3]", + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + got, err := json.Marshal(tt.s) + if err != nil { + t.Errorf("MarshalJSON() error = %v", err) + return + } + if string(got) != tt.want { + t.Errorf("MarshalJSON() = %v, want %v", string(got), tt.want) + } + }) + } +} + +func TestFloat64Slice_UnmarshalJSON(t *testing.T) { + tests := []struct { + name string + data string + want Float64Slice + wantErr bool + }{ + { + name: "null", + data: "null", + want: nil, + wantErr: false, + }, + { + name: "empty array", + data: "[]", + want: Float64Slice{}, + wantErr: false, + }, + { + name: "single element", + data: "[1.5]", + want: Float64Slice{1.5}, + wantErr: false, + }, + { + name: "multiple elements", + data: "[1.1,2.2,3.3]", + want: Float64Slice{1.1, 2.2, 3.3}, + wantErr: false, + }, + { + name: "integer as float", + data: "[1,2,3]", + want: Float64Slice{1, 2, 3}, + wantErr: false, + }, + { + name: "invalid json", + data: "invalid", + want: nil, + wantErr: true, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + var s Float64Slice + err := json.Unmarshal([]byte(tt.data), &s) + if (err != nil) != tt.wantErr { + t.Errorf("UnmarshalJSON() error = %v, wantErr %v", err, tt.wantErr) + return + } + if !tt.wantErr && len(s) != len(tt.want) { + t.Errorf("UnmarshalJSON() = %v, want %v", s, tt.want) + } + }) + } +} + +func TestFloat64Slice_RoundTrip(t *testing.T) { + original := Float64Slice{1.1, 2.2, 3.3} + + // GORM Value -> Scan + val, err := original.Value() + if err != nil { + t.Fatalf("Value() error = %v", err) + } + + var scanned Float64Slice + if err := scanned.Scan(val); err != nil { + t.Fatalf("Scan() error = %v", err) + } + + if len(scanned) != len(original) { + t.Errorf("GORM round-trip failed: got %v, want %v", scanned, original) + } + + // JSON Marshal -> Unmarshal + jsonData, err := json.Marshal(original) + if err != nil { + t.Fatalf("MarshalJSON() error = %v", err) + } + + var unmarshaled Float64Slice + if err := json.Unmarshal(jsonData, &unmarshaled); err != nil { + t.Fatalf("UnmarshalJSON() error = %v", err) + } + + if len(unmarshaled) != len(original) { + t.Errorf("JSON round-trip failed: got %v, want %v", unmarshaled, original) + } +} + +func TestFloat64Slice_ToSlice(t *testing.T) { + tests := []struct { + name string + s Float64Slice + want []float64 + }{ + {"nil", nil, nil}, + {"empty", Float64Slice{}, []float64{}}, + {"single", Float64Slice{1.5}, []float64{1.5}}, + {"multiple", Float64Slice{1.1, 2.2, 3.3}, []float64{1.1, 2.2, 3.3}}, + } + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + got := tt.s.ToSlice() + if len(got) != len(tt.want) { + t.Fatalf("ToSlice() len = %d, want %d", len(got), len(tt.want)) + } + for i := range got { + if got[i] != tt.want[i] { + t.Errorf("ToSlice()[%d] = %f, want %f", i, got[i], tt.want[i]) + } + } + }) + } +} diff --git a/gormtypes/int32_slice.go b/gormtypes/int32_slice.go new file mode 100644 index 0000000..5dc539d --- /dev/null +++ b/gormtypes/int32_slice.go @@ -0,0 +1,119 @@ +// zogo/gormtypes/int32_slice.go +package gormtypes + +import ( + "database/sql/driver" + "encoding/json" + "strconv" + "strings" +) + +// Int32Slice []int32 类型别名 +// 用于存储和传输 int32 数组 +// - GORM 层:序列化为 JSON 字符串存储到数据库 +// - API 层:前端传 string[],返回 string[](避免 JavaScript 大数字精度问题) +type Int32Slice []int32 + +// +// ========== GORM / SQL 层 ========== +// + +// Value 实现 driver.Valuer 接口,用于写入数据库 +// 数据库存储格式:JSON 数组字符串,如 "[1,2,3]" +func (s Int32Slice) Value() (driver.Value, error) { + if s == nil { + return []byte("[]"), nil + } + // 使用类型转换避免调用 MarshalJSON 方法 + // 直接序列化为数字数组格式存储到数据库 + return json.Marshal([]int32(s)) +} + +// Scan 实现 sql.Scanner 接口,用于从数据库读取 +// 支持读取数据库中的 JSON 数组字符串 +func (s *Int32Slice) Scan(value interface{}) error { + if value == nil { + *s = Int32Slice{} + return nil + } + + var data []byte + + switch v := value.(type) { + case []byte: + data = v + case string: + data = []byte(v) + default: + return nil + } + + return json.Unmarshal(data, s) +} + +// +// ========== API / JSON 层 ========== +// + +// MarshalJSON 实现 json.Marshaler 接口 +// 序列化输出:string[] 格式,如 ["1","2","3"] +// 使用字符串格式是为了避免 JavaScript 前端大数字精度丢失问题 +func (s Int32Slice) MarshalJSON() ([]byte, error) { + if s == nil { + return []byte("null"), nil + } + + var b strings.Builder + b.WriteByte('[') + + for i, v := range s { + if i > 0 { + b.WriteByte(',') + } + b.WriteByte('"') + b.WriteString(strconv.FormatInt(int64(v), 10)) + b.WriteByte('"') + } + + b.WriteByte(']') + return []byte(b.String()), nil +} + +// UnmarshalJSON 实现 json.Unmarshaler 接口 +// 反序列化支持两种输入格式: +// - string[]:["1","2","3"](前端标准传参格式) +// - number[]:[1,2,3](兼容其他场景) +func (s *Int32Slice) UnmarshalJSON(data []byte) error { + if string(data) == "null" { + *s = nil + return nil + } + + // 优先尝试解析为字符串数组 ["1","2","3"] + var strArr []string + if err := json.Unmarshal(data, &strArr); err == nil { + res := make(Int32Slice, 0, len(strArr)) + for _, v := range strArr { + i, err := strconv.ParseInt(v, 10, 32) + if err != nil { + return err + } + res = append(res, int32(i)) + } + *s = res + return nil + } + + // 兼容数字数组格式 [1,2,3] + var numArr []int32 + if err := json.Unmarshal(data, &numArr); err != nil { + return err + } + *s = numArr + return nil +} + +// ToSlice 转换为原生 []int32 切片,用于 ... 展开 +func (s Int32Slice) ToSlice() []int32 { + return []int32(s) +} diff --git a/gormtypes/int32_slice_test.go b/gormtypes/int32_slice_test.go new file mode 100644 index 0000000..0318555 --- /dev/null +++ b/gormtypes/int32_slice_test.go @@ -0,0 +1,223 @@ +// zogo/gormtypes/int32_slice_test.go +package gormtypes + +import ( + "testing" +) + +func TestInt32SliceValue(t *testing.T) { + tests := []struct { + name string + slice Int32Slice + want string + wantErr bool + }{ + { + name: "nil slice", + slice: nil, + want: "[]", + }, + { + name: "empty slice", + slice: Int32Slice{}, + want: "[]", + }, + { + name: "single element", + slice: Int32Slice{1}, + want: "[1]", + }, + { + name: "multiple elements", + slice: Int32Slice{1, 2, 3}, + want: "[1,2,3]", + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + got, err := tt.slice.Value() + if (err != nil) != tt.wantErr { + t.Errorf("Value() error = %v, wantErr %v", err, tt.wantErr) + return + } + if got != nil && string(got.([]byte)) != tt.want { + t.Errorf("Value() = %v, want %v", string(got.([]byte)), tt.want) + } + }) + } +} + +func TestInt32SliceScan(t *testing.T) { + tests := []struct { + name string + value interface{} + want Int32Slice + wantErr bool + }{ + { + name: "nil value", + value: nil, + want: Int32Slice{}, + }, + { + name: "byte array", + value: []byte("[1,2,3]"), + want: Int32Slice{1, 2, 3}, + }, + { + name: "string", + value: "[4,5,6]", + want: Int32Slice{4, 5, 6}, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + var got Int32Slice + err := got.Scan(tt.value) + if (err != nil) != tt.wantErr { + t.Errorf("Scan() error = %v, wantErr %v", err, tt.wantErr) + return + } + if len(got) != len(tt.want) { + t.Errorf("Scan() length = %v, want %v", len(got), len(tt.want)) + return + } + for i := range got { + if got[i] != tt.want[i] { + t.Errorf("Scan()[%d] = %v, want %v", i, got[i], tt.want[i]) + } + } + }) + } +} + +func TestInt32SliceMarshalJSON(t *testing.T) { + tests := []struct { + name string + slice Int32Slice + want string + wantErr bool + }{ + { + name: "nil slice", + slice: nil, + want: "null", + }, + { + name: "empty slice", + slice: Int32Slice{}, + want: "[]", + }, + { + name: "single element", + slice: Int32Slice{1}, + want: `["1"]`, + }, + { + name: "multiple elements", + slice: Int32Slice{1, 2, 3}, + want: `["1","2","3"]`, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + got, err := tt.slice.MarshalJSON() + if (err != nil) != tt.wantErr { + t.Errorf("MarshalJSON() error = %v, wantErr %v", err, tt.wantErr) + return + } + if string(got) != tt.want { + t.Errorf("MarshalJSON() = %v, want %v", string(got), tt.want) + } + }) + } +} + +func TestInt32SliceToSlice(t *testing.T) { + tests := []struct { + name string + s Int32Slice + want []int32 + }{ + {"nil", nil, nil}, + {"empty", Int32Slice{}, []int32{}}, + {"single", Int32Slice{1}, []int32{1}}, + {"multiple", Int32Slice{1, 2, 3}, []int32{1, 2, 3}}, + } + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + got := tt.s.ToSlice() + if len(got) != len(tt.want) { + t.Fatalf("ToSlice() len = %d, want %d", len(got), len(tt.want)) + } + for i := range got { + if got[i] != tt.want[i] { + t.Errorf("ToSlice()[%d] = %d, want %d", i, got[i], tt.want[i]) + } + } + }) + } +} + +func TestInt32SliceUnmarshalJSON(t *testing.T) { + tests := []struct { + name string + data string + want Int32Slice + wantErr bool + }{ + { + name: "null", + data: "null", + want: nil, + }, + { + name: "empty array", + data: "[]", + want: Int32Slice{}, + }, + { + name: "string array", + data: `["1","2","3"]`, + want: Int32Slice{1, 2, 3}, + }, + { + name: "number array", + data: "[4,5,6]", + want: Int32Slice{4, 5, 6}, + }, + { + name: "single string element", + data: `["1"]`, + want: Int32Slice{1}, + }, + { + name: "single number element", + data: "[1]", + want: Int32Slice{1}, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + var got Int32Slice + err := got.UnmarshalJSON([]byte(tt.data)) + if (err != nil) != tt.wantErr { + t.Errorf("UnmarshalJSON() error = %v, wantErr %v", err, tt.wantErr) + return + } + if len(got) != len(tt.want) { + t.Errorf("UnmarshalJSON() length = %v, want %v", len(got), len(tt.want)) + return + } + for i := range got { + if got[i] != tt.want[i] { + t.Errorf("UnmarshalJSON()[%d] = %v, want %v", i, got[i], tt.want[i]) + } + } + }) + } +} diff --git a/gormtypes/int64_slice.go b/gormtypes/int64_slice.go new file mode 100644 index 0000000..1fcdfaa --- /dev/null +++ b/gormtypes/int64_slice.go @@ -0,0 +1,117 @@ +// zogo/gormtypes/int64_slice.go +package gormtypes + +import ( + "database/sql/driver" + "encoding/json" + "strconv" + "strings" +) + +// Int64Slice []int64 类型别名 +// 用于存储和传输 int64 数组 +// - GORM 层:序列化为 JSON 字符串存储到数据库 +// - API 层:前端传 string[],返回 string[](避免 JavaScript 大数字精度问题) +type Int64Slice []int64 + +// +// ========== GORM / SQL 层 ========== +// + +// Value 实现 driver.Valuer 接口,用于写入数据库 +// 数据库存储格式:JSON 数组字符串,如 "[1,2,3]" +func (s Int64Slice) Value() (driver.Value, error) { + if s == nil { + return []byte("[]"), nil + } + return json.Marshal(s) +} + +// Scan 实现 sql.Scanner 接口,用于从数据库读取 +// 支持读取数据库中的 JSON 数组字符串 +func (s *Int64Slice) Scan(value interface{}) error { + if value == nil { + *s = Int64Slice{} + return nil + } + + var data []byte + + switch v := value.(type) { + case []byte: + data = v + case string: + data = []byte(v) + default: + return nil + } + + return json.Unmarshal(data, s) +} + +// +// ========== API / JSON 层 ========== +// + +// MarshalJSON 实现 json.Marshaler 接口 +// 序列化输出:string[] 格式,如 ["1","2","3"] +// 使用字符串格式是为了避免 JavaScript 前端大数字精度丢失问题 +func (s Int64Slice) MarshalJSON() ([]byte, error) { + if s == nil { + return []byte("null"), nil + } + + var b strings.Builder + b.WriteByte('[') + + for i, v := range s { + if i > 0 { + b.WriteByte(',') + } + b.WriteByte('"') + b.WriteString(strconv.FormatInt(v, 10)) + b.WriteByte('"') + } + + b.WriteByte(']') + return []byte(b.String()), nil +} + +// UnmarshalJSON 实现 json.Unmarshaler 接口 +// 反序列化支持两种输入格式: +// - string[]:["1","2","3"](前端标准传参格式) +// - number[]:[1,2,3](兼容其他场景) +func (s *Int64Slice) UnmarshalJSON(data []byte) error { + if string(data) == "null" { + *s = nil + return nil + } + + // 优先尝试解析为字符串数组 ["1","2","3"] + var strArr []string + if err := json.Unmarshal(data, &strArr); err == nil { + res := make(Int64Slice, 0, len(strArr)) + for _, v := range strArr { + i, err := strconv.ParseInt(v, 10, 64) + if err != nil { + return err + } + res = append(res, i) + } + *s = res + return nil + } + + // 兼容数字数组格式 [1,2,3] + var numArr []int64 + if err := json.Unmarshal(data, &numArr); err != nil { + return err + } + *s = numArr + return nil +} + +// ToSlice 转换为原生 []int64 切片,用于 ... 展开 +func (s Int64Slice) ToSlice() []int64 { + return []int64(s) +} diff --git a/gormtypes/int64_slice_test.go b/gormtypes/int64_slice_test.go new file mode 100644 index 0000000..3e09a55 --- /dev/null +++ b/gormtypes/int64_slice_test.go @@ -0,0 +1,326 @@ +// zogo/gormtypes/int64_slice_test.go +package gormtypes + +import ( + "encoding/json" + "testing" +) + +func TestInt64Slice_Value(t *testing.T) { + tests := []struct { + name string + s Int64Slice + want []byte + wantErr bool + }{ + { + name: "nil slice", + s: nil, + want: []byte("[]"), + wantErr: false, + }, + { + name: "empty slice", + s: Int64Slice{}, + want: []byte("[]"), + wantErr: false, + }, + { + name: "single element", + s: Int64Slice{1}, + want: []byte("[\"1\"]"), + wantErr: false, + }, + { + name: "multiple elements", + s: Int64Slice{1, 2, 3}, + want: []byte("[\"1\",\"2\",\"3\"]"), + wantErr: false, + }, + { + name: "large number", + s: Int64Slice{9223372036854775807}, + want: []byte("[\"9223372036854775807\"]"), + wantErr: false, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + got, err := tt.s.Value() + if (err != nil) != tt.wantErr { + t.Errorf("Value() error = %v, wantErr %v", err, tt.wantErr) + return + } + gotBytes, ok := got.([]byte) + if !ok { + t.Errorf("Value() returned non-[]byte type: %T", got) + return + } + if string(gotBytes) != string(tt.want) { + t.Errorf("Value() = %v, want %v", string(gotBytes), string(tt.want)) + } + }) + } +} + +func TestInt64Slice_Scan(t *testing.T) { + tests := []struct { + name string + data interface{} + want Int64Slice + wantErr bool + }{ + { + name: "nil value", + data: nil, + want: Int64Slice{}, + wantErr: false, + }, + { + name: "empty json", + data: []byte("[]"), + want: Int64Slice{}, + wantErr: false, + }, + { + name: "single element", + data: []byte("[1]"), + want: Int64Slice{1}, + wantErr: false, + }, + { + name: "multiple elements", + data: []byte("[1,2,3]"), + want: Int64Slice{1, 2, 3}, + wantErr: false, + }, + { + name: "string data", + data: "[1,2,3]", + want: Int64Slice{1, 2, 3}, + wantErr: false, + }, + { + name: "invalid json", + data: []byte("invalid"), + want: nil, + wantErr: true, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + var s Int64Slice + err := s.Scan(tt.data) + if (err != nil) != tt.wantErr { + t.Errorf("Scan() error = %v, wantErr %v", err, tt.wantErr) + return + } + if !tt.wantErr && len(s) != len(tt.want) { + t.Errorf("Scan() = %v, want %v", s, tt.want) + } + }) + } +} + +func TestInt64Slice_MarshalJSON(t *testing.T) { + tests := []struct { + name string + s Int64Slice + want string + }{ + { + name: "nil slice", + s: nil, + want: "null", + }, + { + name: "empty slice", + s: Int64Slice{}, + want: "[]", + }, + { + name: "single element", + s: Int64Slice{1}, + want: "[\"1\"]", + }, + { + name: "multiple elements", + s: Int64Slice{1, 2, 3}, + want: "[\"1\",\"2\",\"3\"]", + }, + { + name: "large number", + s: Int64Slice{9223372036854775807}, + want: "[\"9223372036854775807\"]", + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + got, err := json.Marshal(tt.s) + if err != nil { + t.Errorf("MarshalJSON() error = %v", err) + return + } + if string(got) != tt.want { + t.Errorf("MarshalJSON() = %v, want %v", string(got), tt.want) + } + }) + } +} + +func TestInt64Slice_UnmarshalJSON(t *testing.T) { + tests := []struct { + name string + data string + want Int64Slice + wantErr bool + }{ + { + name: "null", + data: "null", + want: nil, + wantErr: false, + }, + { + name: "empty array", + data: "[]", + want: Int64Slice{}, + wantErr: false, + }, + { + name: "string array single", + data: "[\"1\"]", + want: Int64Slice{1}, + wantErr: false, + }, + { + name: "string array multiple", + data: "[\"1\",\"2\",\"3\"]", + want: Int64Slice{1, 2, 3}, + wantErr: false, + }, + { + name: "number array single", + data: "[1]", + want: Int64Slice{1}, + wantErr: false, + }, + { + name: "number array multiple", + data: "[1,2,3]", + want: Int64Slice{1, 2, 3}, + wantErr: false, + }, + { + name: "invalid string", + data: "[\"invalid\"]", + want: nil, + wantErr: true, + }, + { + name: "invalid json", + data: "invalid", + want: nil, + wantErr: true, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + var s Int64Slice + err := json.Unmarshal([]byte(tt.data), &s) + if (err != nil) != tt.wantErr { + t.Errorf("UnmarshalJSON() error = %v, wantErr %v", err, tt.wantErr) + return + } + if !tt.wantErr && len(s) != len(tt.want) { + t.Errorf("UnmarshalJSON() = %v, want %v", s, tt.want) + } + }) + } +} + +func TestInt64Slice_Equal(t *testing.T) { + a := Int64Slice{1, 2, 3} + b := Int64Slice{1, 2, 3} + c := Int64Slice{1, 2, 4} + d := Int64Slice{1, 2} + var nilSlice Int64Slice + + // 手动比较 + if len(a) != len(b) || (len(a) > 0 && (a[0] != b[0] || a[1] != b[1] || a[2] != b[2])) { + t.Error("Equal slices should return true") + } + if len(a) == len(c) && (a[0] == c[0] && a[1] == c[1] && a[2] == c[2]) { + t.Error("Different slices should return false") + } + if len(a) == len(d) && (a[0] == d[0] && a[1] == d[1]) { + t.Error("Different length slices should return false") + } + if len(a) == len(nilSlice) { + t.Error("Non-empty vs empty slice should return false") + } +} + +func TestInt64Slice_RoundTrip(t *testing.T) { + original := Int64Slice{1, 2, 3, 9223372036854775807} + + // GORM Value -> Scan + val, err := original.Value() + if err != nil { + t.Fatalf("Value() error = %v", err) + } + + var scanned Int64Slice + if err := scanned.Scan(val); err != nil { + t.Fatalf("Scan() error = %v", err) + } + + if len(scanned) != len(original) { + t.Errorf("GORM round-trip failed: got %v, want %v", scanned, original) + } + + // JSON Marshal -> Unmarshal + jsonData, err := json.Marshal(original) + if err != nil { + t.Fatalf("MarshalJSON() error = %v", err) + } + + var unmarshaled Int64Slice + if err := json.Unmarshal(jsonData, &unmarshaled); err != nil { + t.Fatalf("UnmarshalJSON() error = %v", err) + } + + if len(unmarshaled) != len(original) { + t.Errorf("JSON round-trip failed: got %v, want %v", unmarshaled, original) + } +} + +func TestInt64Slice_ToSlice(t *testing.T) { + tests := []struct { + name string + s Int64Slice + want []int64 + }{ + {"nil", nil, nil}, + {"empty", Int64Slice{}, []int64{}}, + {"single", Int64Slice{1}, []int64{1}}, + {"multiple", Int64Slice{1, 2, 3}, []int64{1, 2, 3}}, + } + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + got := tt.s.ToSlice() + if len(got) != len(tt.want) { + t.Fatalf("ToSlice() len = %d, want %d", len(got), len(tt.want)) + } + for i := range got { + if got[i] != tt.want[i] { + t.Errorf("ToSlice()[%d] = %d, want %d", i, got[i], tt.want[i]) + } + } + }) + } +} diff --git a/gormtypes/int_slice.go b/gormtypes/int_slice.go new file mode 100644 index 0000000..99af587 --- /dev/null +++ b/gormtypes/int_slice.go @@ -0,0 +1,99 @@ +// zogo/gormtypes/int_slice.go +package gormtypes + +import ( + "database/sql/driver" + "encoding/json" + "strconv" + "strings" +) + +// IntSlice []int 类型别名 +// 用于存储和传输 int 数组 +// - GORM 层:序列化为 JSON 字符串存储到数据库 +// - API 层:前端传 string[],返回 string[](避免 JavaScript 大数字精度问题) +type IntSlice []int + +// Value 实现 driver.Valuer 接口,用于写入数据库 +func (s IntSlice) Value() (driver.Value, error) { + if s == nil { + return []byte("[]"), nil + } + return json.Marshal(s) +} + +// Scan 实现 sql.Scanner 接口,用于从数据库读取 +func (s *IntSlice) Scan(value interface{}) error { + if value == nil { + *s = IntSlice{} + return nil + } + + var data []byte + switch v := value.(type) { + case []byte: + data = v + case string: + data = []byte(v) + default: + return nil + } + + return json.Unmarshal(data, s) +} + +// MarshalJSON 序列化输出:string[] 格式 +func (s IntSlice) MarshalJSON() ([]byte, error) { + if s == nil { + return []byte("null"), nil + } + + var b strings.Builder + b.WriteByte('[') + for i, v := range s { + if i > 0 { + b.WriteByte(',') + } + b.WriteByte('"') + b.WriteString(strconv.Itoa(v)) + b.WriteByte('"') + } + b.WriteByte(']') + return []byte(b.String()), nil +} + +// UnmarshalJSON 反序列化支持 string[] 和 number[] 两种格式 +func (s *IntSlice) UnmarshalJSON(data []byte) error { + if string(data) == "null" { + *s = nil + return nil + } + + // 优先尝试解析为字符串数组 + var strArr []string + if err := json.Unmarshal(data, &strArr); err == nil { + res := make(IntSlice, 0, len(strArr)) + for _, v := range strArr { + i, err := strconv.Atoi(v) + if err != nil { + return err + } + res = append(res, i) + } + *s = res + return nil + } + + // 兼容数字数组格式 + var numArr []int + if err := json.Unmarshal(data, &numArr); err != nil { + return err + } + *s = numArr + return nil +} + +// ToSlice 转换为原生 []int 切片,用于 ... 展开 +func (s IntSlice) ToSlice() []int { + return []int(s) +} diff --git a/gormtypes/int_slice_test.go b/gormtypes/int_slice_test.go new file mode 100644 index 0000000..585d06c --- /dev/null +++ b/gormtypes/int_slice_test.go @@ -0,0 +1,150 @@ +// zogo/gormtypes/int_slice_test.go +package gormtypes + +import ( + "encoding/json" + "testing" +) + +func TestIntSlice_Value(t *testing.T) { + tests := []struct { + name string + s IntSlice + want []byte + wantErr bool + }{ + {name: "nil slice", s: nil, want: []byte("[]"), wantErr: false}, + {name: "empty slice", s: IntSlice{}, want: []byte("[]"), wantErr: false}, + {name: "single element", s: IntSlice{1}, want: []byte("[\"1\"]"), wantErr: false}, + {name: "multiple elements", s: IntSlice{1, 2, 3}, want: []byte("[\"1\",\"2\",\"3\"]"), wantErr: false}, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + got, err := tt.s.Value() + if (err != nil) != tt.wantErr { + t.Errorf("Value() error = %v, wantErr %v", err, tt.wantErr) + return + } + gotBytes, ok := got.([]byte) + if !ok { + t.Errorf("Value() returned non-[]byte type: %T", got) + return + } + if string(gotBytes) != string(tt.want) { + t.Errorf("Value() = %v, want %v", string(gotBytes), string(tt.want)) + } + }) + } +} + +func TestIntSlice_ToSlice(t *testing.T) { + tests := []struct { + name string + s IntSlice + want []int + }{ + {"nil", nil, nil}, + {"empty", IntSlice{}, []int{}}, + {"single", IntSlice{1}, []int{1}}, + {"multiple", IntSlice{1, 2, 3}, []int{1, 2, 3}}, + } + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + got := tt.s.ToSlice() + if len(got) != len(tt.want) { + t.Fatalf("ToSlice() len = %d, want %d", len(got), len(tt.want)) + } + for i := range got { + if got[i] != tt.want[i] { + t.Errorf("ToSlice()[%d] = %d, want %d", i, got[i], tt.want[i]) + } + } + }) + } +} + +func TestIntSlice_Scan(t *testing.T) { + tests := []struct { + name string + data interface{} + want IntSlice + wantErr bool + }{ + {name: "nil value", data: nil, want: IntSlice{}, wantErr: false}, + {name: "empty json", data: []byte("[]"), want: IntSlice{}, wantErr: false}, + {name: "single element", data: []byte("[1]"), want: IntSlice{1}, wantErr: false}, + {name: "multiple elements", data: []byte("[1,2,3]"), want: IntSlice{1, 2, 3}, wantErr: false}, + {name: "invalid json", data: []byte("invalid"), want: nil, wantErr: true}, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + var s IntSlice + err := s.Scan(tt.data) + if (err != nil) != tt.wantErr { + t.Errorf("Scan() error = %v, wantErr %v", err, tt.wantErr) + return + } + if !tt.wantErr && len(s) != len(tt.want) { + t.Errorf("Scan() = %v, want %v", s, tt.want) + } + }) + } +} + +func TestIntSlice_MarshalJSON(t *testing.T) { + tests := []struct { + name string + s IntSlice + want string + }{ + {name: "nil slice", s: nil, want: "null"}, + {name: "empty slice", s: IntSlice{}, want: "[]"}, + {name: "single element", s: IntSlice{1}, want: "[\"1\"]"}, + {name: "multiple elements", s: IntSlice{1, 2, 3}, want: "[\"1\",\"2\",\"3\"]"}, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + got, err := json.Marshal(tt.s) + if err != nil { + t.Errorf("MarshalJSON() error = %v", err) + return + } + if string(got) != tt.want { + t.Errorf("MarshalJSON() = %v, want %v", string(got), tt.want) + } + }) + } +} + +func TestIntSlice_UnmarshalJSON(t *testing.T) { + tests := []struct { + name string + data string + want IntSlice + wantErr bool + }{ + {name: "null", data: "null", want: nil, wantErr: false}, + {name: "empty array", data: "[]", want: IntSlice{}, wantErr: false}, + {name: "string array", data: "[\"1\",\"2\",\"3\"]", want: IntSlice{1, 2, 3}, wantErr: false}, + {name: "number array", data: "[1,2,3]", want: IntSlice{1, 2, 3}, wantErr: false}, + {name: "invalid string", data: "[\"invalid\"]", want: nil, wantErr: true}, + {name: "invalid json", data: "invalid", want: nil, wantErr: true}, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + var s IntSlice + err := json.Unmarshal([]byte(tt.data), &s) + if (err != nil) != tt.wantErr { + t.Errorf("UnmarshalJSON() error = %v, wantErr %v", err, tt.wantErr) + return + } + if !tt.wantErr && len(s) != len(tt.want) { + t.Errorf("UnmarshalJSON() = %v, want %v", s, tt.want) + } + }) + } +} diff --git a/gormtypes/string_map.go b/gormtypes/string_map.go new file mode 100644 index 0000000..e63748a --- /dev/null +++ b/gormtypes/string_map.go @@ -0,0 +1,112 @@ +// zogo/gormtypes/string_map.go +package gormtypes + +import ( + "database/sql/driver" + "encoding/json" +) + +// StringMap map[string]string 类型别名 +// 用于存储和传输字符串键值对 +// - GORM 层:序列化为 JSON 对象存储到数据库 +// - API 层:前端传 object,返回 object +type StringMap map[string]string + +// Value 实现 driver.Valuer 接口,用于写入数据库 +func (m StringMap) Value() (driver.Value, error) { + if m == nil { + return []byte("{}"), nil + } + return json.Marshal(m) +} + +// Scan 实现 sql.Scanner 接口,用于从数据库读取 +func (m *StringMap) Scan(value interface{}) error { + if value == nil { + *m = StringMap{} + return nil + } + + var data []byte + switch v := value.(type) { + case []byte: + data = v + case string: + data = []byte(v) + default: + return nil + } + + return json.Unmarshal(data, m) +} + +// MarshalJSON 序列化输出:object 格式 +func (m StringMap) MarshalJSON() ([]byte, error) { + if m == nil { + return []byte("null"), nil + } + return json.Marshal(map[string]string(m)) +} + +// UnmarshalJSON 反序列化支持 object 格式 +func (m *StringMap) UnmarshalJSON(data []byte) error { + if string(data) == "null" { + *m = nil + return nil + } + + var mp map[string]string + if err := json.Unmarshal(data, &mp); err != nil { + return err + } + *m = StringMap(mp) + return nil +} + +// Get 获取指定 key 的值,不存在返回空字符串 +func (m StringMap) Get(key string) string { + if m == nil { + return "" + } + return m[key] +} + +// Set 设置指定 key 的值 +func (m *StringMap) Set(key, value string) { + if *m == nil { + *m = StringMap{} + } + (*m)[key] = value +} + +// Delete 删除指定 key +func (m StringMap) Delete(key string) { + if m == nil { + return + } + delete(m, key) +} + +// Keys 返回所有 key +func (m StringMap) Keys() []string { + if m == nil { + return []string{} + } + keys := make([]string, 0, len(m)) + for k := range m { + keys = append(keys, k) + } + return keys +} + +// Values 返回所有 value +func (m StringMap) Values() []string { + if m == nil { + return []string{} + } + values := make([]string, 0, len(m)) + for _, v := range m { + values = append(values, v) + } + return values +} diff --git a/gormtypes/string_map_test.go b/gormtypes/string_map_test.go new file mode 100644 index 0000000..1b330cb --- /dev/null +++ b/gormtypes/string_map_test.go @@ -0,0 +1,338 @@ +// zogo/gormtypes/string_map_test.go +package gormtypes + +import ( + "encoding/json" + "testing" +) + +func TestStringMap_Value(t *testing.T) { + tests := []struct { + name string + m StringMap + want []byte + wantErr bool + }{ + { + name: "nil map", + m: nil, + want: []byte("{}"), + wantErr: false, + }, + { + name: "empty map", + m: StringMap{}, + want: []byte("{}"), + wantErr: false, + }, + { + name: "single key", + m: StringMap{"key": "value"}, + want: []byte("{\"key\":\"value\"}"), + wantErr: false, + }, + { + name: "multiple keys", + m: StringMap{"a": "1", "b": "2"}, + want: []byte("{\"a\":\"1\",\"b\":\"2\"}"), + wantErr: false, + }, + { + name: "with special chars", + m: StringMap{"key": "value\"with\"quotes"}, + want: []byte("{\"key\":\"value\\\"with\\\"quotes\"}"), + wantErr: false, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + got, err := tt.m.Value() + if (err != nil) != tt.wantErr { + t.Errorf("Value() error = %v, wantErr %v", err, tt.wantErr) + return + } + gotBytes, ok := got.([]byte) + if !ok { + t.Errorf("Value() returned non-[]byte type: %T", got) + return + } + if string(gotBytes) != string(tt.want) { + t.Errorf("Value() = %v, want %v", string(gotBytes), string(tt.want)) + } + }) + } +} + +func TestStringMap_Scan(t *testing.T) { + tests := []struct { + name string + data interface{} + want StringMap + wantErr bool + }{ + { + name: "nil value", + data: nil, + want: StringMap{}, + wantErr: false, + }, + { + name: "empty json", + data: []byte("{}"), + want: StringMap{}, + wantErr: false, + }, + { + name: "single key", + data: []byte("{\"key\":\"value\"}"), + want: StringMap{"key": "value"}, + wantErr: false, + }, + { + name: "multiple keys", + data: []byte("{\"a\":\"1\",\"b\":\"2\"}"), + want: StringMap{"a": "1", "b": "2"}, + wantErr: false, + }, + { + name: "invalid json", + data: []byte("invalid"), + want: nil, + wantErr: true, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + var m StringMap + err := m.Scan(tt.data) + if (err != nil) != tt.wantErr { + t.Errorf("Scan() error = %v, wantErr %v", err, tt.wantErr) + return + } + if !tt.wantErr && len(m) != len(tt.want) { + t.Errorf("Scan() = %v, want %v", m, tt.want) + } + }) + } +} + +func TestStringMap_MarshalJSON(t *testing.T) { + tests := []struct { + name string + m StringMap + want string + }{ + { + name: "nil map", + m: nil, + want: "null", + }, + { + name: "empty map", + m: StringMap{}, + want: "{}", + }, + { + name: "single key", + m: StringMap{"key": "value"}, + want: "{\"key\":\"value\"}", + }, + { + name: "multiple keys", + m: StringMap{"a": "1", "b": "2"}, + want: "{\"a\":\"1\",\"b\":\"2\"}", + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + got, err := json.Marshal(tt.m) + if err != nil { + t.Errorf("MarshalJSON() error = %v", err) + return + } + if string(got) != tt.want { + t.Errorf("MarshalJSON() = %v, want %v", string(got), tt.want) + } + }) + } +} + +func TestStringMap_UnmarshalJSON(t *testing.T) { + tests := []struct { + name string + data string + want StringMap + wantErr bool + }{ + { + name: "null", + data: "null", + want: nil, + wantErr: false, + }, + { + name: "empty object", + data: "{}", + want: StringMap{}, + wantErr: false, + }, + { + name: "single key", + data: "{\"key\":\"value\"}", + want: StringMap{"key": "value"}, + wantErr: false, + }, + { + name: "multiple keys", + data: "{\"a\":\"1\",\"b\":\"2\"}", + want: StringMap{"a": "1", "b": "2"}, + wantErr: false, + }, + { + name: "invalid json", + data: "invalid", + want: nil, + wantErr: true, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + var m StringMap + err := json.Unmarshal([]byte(tt.data), &m) + if (err != nil) != tt.wantErr { + t.Errorf("UnmarshalJSON() error = %v, wantErr %v", err, tt.wantErr) + return + } + if !tt.wantErr && len(m) != len(tt.want) { + t.Errorf("UnmarshalJSON() = %v, want %v", m, tt.want) + } + }) + } +} + +func TestStringMap_Get(t *testing.T) { + m := StringMap{"key": "value", "empty": ""} + + if m.Get("key") != "value" { + t.Error("Get('key') should return 'value'") + } + if m.Get("empty") != "" { + t.Error("Get('empty') should return ''") + } + if m.Get("nonexistent") != "" { + t.Error("Get('nonexistent') should return ''") + } + if StringMap(nil).Get("key") != "" { + t.Error("nil Get() should return ''") + } +} + +func TestStringMap_Set(t *testing.T) { + var m StringMap + m.Set("key", "value") + + if m["key"] != "value" { + t.Error("Set() should set value correctly") + } + + var nilMap StringMap + nilMap.Set("key", "value") + if nilMap["key"] != "value" { + t.Error("Set() on nil map should initialize and set value") + } +} + +func TestStringMap_Delete(t *testing.T) { + m := StringMap{"key": "value", "keep": "this"} + + m.Delete("key") + if _, ok := m["key"]; ok { + t.Error("Delete() should remove key") + } + if m["keep"] != "this" { + t.Error("Delete() should not remove other keys") + } + + var nilMap StringMap + nilMap.Delete("key") // should not panic +} + +func TestStringMap_Keys(t *testing.T) { + m := StringMap{"a": "1", "b": "2", "c": "3"} + keys := m.Keys() + + if len(keys) != 3 { + t.Errorf("Keys() length = %d, want 3", len(keys)) + } + + keyMap := make(map[string]bool) + for _, k := range keys { + keyMap[k] = true + } + if !keyMap["a"] || !keyMap["b"] || !keyMap["c"] { + t.Error("Keys() should contain all keys") + } + + if len(StringMap(nil).Keys()) != 0 { + t.Error("nil Keys() should return empty slice") + } +} + +func TestStringMap_Values(t *testing.T) { + m := StringMap{"a": "1", "b": "2", "c": "3"} + values := m.Values() + + if len(values) != 3 { + t.Errorf("Values() length = %d, want 3", len(values)) + } + + valueMap := make(map[string]bool) + for _, v := range values { + valueMap[v] = true + } + if !valueMap["1"] || !valueMap["2"] || !valueMap["3"] { + t.Error("Values() should contain all values") + } + + if len(StringMap(nil).Values()) != 0 { + t.Error("nil Values() should return empty slice") + } +} + +func TestStringMap_RoundTrip(t *testing.T) { + original := StringMap{"key1": "value1", "key2": "value2"} + + // GORM Value -> Scan + val, err := original.Value() + if err != nil { + t.Fatalf("Value() error = %v", err) + } + + var scanned StringMap + if err := scanned.Scan(val); err != nil { + t.Fatalf("Scan() error = %v", err) + } + + if len(scanned) != len(original) { + t.Errorf("GORM round-trip failed: got %v, want %v", scanned, original) + } + + // JSON Marshal -> Unmarshal + jsonData, err := json.Marshal(original) + if err != nil { + t.Fatalf("MarshalJSON() error = %v", err) + } + + var unmarshaled StringMap + if err := json.Unmarshal(jsonData, &unmarshaled); err != nil { + t.Fatalf("UnmarshalJSON() error = %v", err) + } + + if len(unmarshaled) != len(original) { + t.Errorf("JSON round-trip failed: got %v, want %v", unmarshaled, original) + } +} diff --git a/gormtypes/string_slice.go b/gormtypes/string_slice.go new file mode 100644 index 0000000..9a8a7b7 --- /dev/null +++ b/gormtypes/string_slice.go @@ -0,0 +1,110 @@ +// zogo/gormtypes/string_slice.go +package gormtypes + +import ( + "database/sql/driver" + "encoding/json" + "strings" +) + +// StringSlice []string 类型别名 +// 用于存储和传输字符串数组 +// - GORM 层:序列化为 JSON 字符串存储到数据库 +// - API 层:返回 string[],前端传 string[] +type StringSlice []string + +// Value 实现 driver.Valuer 接口,用于写入数据库 +func (s StringSlice) Value() (driver.Value, error) { + if s == nil { + return []byte("[]"), nil + } + return json.Marshal(s) +} + +// Scan 实现 sql.Scanner 接口,用于从数据库读取 +func (s *StringSlice) Scan(value interface{}) error { + if value == nil { + *s = StringSlice{} + return nil + } + + var data []byte + switch v := value.(type) { + case []byte: + data = v + case string: + data = []byte(v) + default: + *s = StringSlice{} + return nil + } + + // 空字符串或 "null" 统一转成空切片 [] + if len(data) == 0 || strings.TrimSpace(string(data)) == "null" { + *s = StringSlice{} + return nil + } + + return json.Unmarshal(data, s) +} + +// MarshalJSON 序列化输出:string[] 格式 +// nil 统一输出为 [](空数组),符合"null/nil 都转成 []"的预期 +func (s StringSlice) MarshalJSON() ([]byte, error) { + if s == nil { + return []byte("[]"), nil + } + return json.Marshal([]string(s)) +} + +// UnmarshalJSON 反序列化支持 string[] 格式 +// null 统一解码为空切片 [],符合"null 转成 []"的预期 +func (s *StringSlice) UnmarshalJSON(data []byte) error { + if string(data) == "null" { + *s = StringSlice{} + return nil + } + + // 标准 JSON 数组解析 + var arr []string + if err := json.Unmarshal(data, &arr); err != nil { + // 尝试兼容单字符串情况 + var single string + if err := json.Unmarshal(data, &single); err == nil { + *s = StringSlice{single} + return nil + } + return err + } + *s = StringSlice(arr) + return nil +} + +// ToSlice 转换为原生 []string 切片,用于 ... 展开 +func (s StringSlice) ToSlice() []string { + return []string(s) +} + +// Contains 检查是否包含指定字符串 +func (s StringSlice) Contains(val string) bool { + for _, v := range s { + if v == val { + return true + } + } + return false +} + +// ToMap 转换为 map,用于去重和快速查找 +func (s StringSlice) ToMap() map[string]struct{} { + m := make(map[string]struct{}, len(s)) + for _, v := range s { + m[v] = struct{}{} + } + return m +} + +// Join 使用分隔符连接所有元素 +func (s StringSlice) Join(sep string) string { + return strings.Join([]string(s), sep) +} diff --git a/gormtypes/string_slice_test.go b/gormtypes/string_slice_test.go new file mode 100644 index 0000000..06fec25 --- /dev/null +++ b/gormtypes/string_slice_test.go @@ -0,0 +1,309 @@ +// zogo/gormtypes/string_slice_test.go +package gormtypes + +import ( + "encoding/json" + "testing" +) + +func TestStringSlice_Value(t *testing.T) { + tests := []struct { + name string + s StringSlice + want []byte + wantErr bool + }{ + { + name: "nil slice", + s: nil, + want: []byte("[]"), + wantErr: false, + }, + { + name: "empty slice", + s: StringSlice{}, + want: []byte("[]"), + wantErr: false, + }, + { + name: "single element", + s: StringSlice{"a"}, + want: []byte("[\"a\"]"), + wantErr: false, + }, + { + name: "multiple elements", + s: StringSlice{"a", "b", "c"}, + want: []byte("[\"a\",\"b\",\"c\"]"), + wantErr: false, + }, + { + name: "with special chars", + s: StringSlice{"hello\"world", "test\nvalue"}, + want: []byte("[\"hello\\\"world\",\"test\\nvalue\"]"), + wantErr: false, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + got, err := tt.s.Value() + if (err != nil) != tt.wantErr { + t.Errorf("Value() error = %v, wantErr %v", err, tt.wantErr) + return + } + gotBytes, ok := got.([]byte) + if !ok { + t.Errorf("Value() returned non-[]byte type: %T", got) + return + } + if string(gotBytes) != string(tt.want) { + t.Errorf("Value() = %v, want %v", string(gotBytes), string(tt.want)) + } + }) + } +} + +func TestStringSlice_Scan(t *testing.T) { + tests := []struct { + name string + data interface{} + want StringSlice + wantErr bool + }{ + { + name: "nil value", + data: nil, + want: StringSlice{}, + wantErr: false, + }, + { + name: "empty json", + data: []byte("[]"), + want: StringSlice{}, + wantErr: false, + }, + { + name: "single element", + data: []byte("[\"hello\"]"), + want: StringSlice{"hello"}, + wantErr: false, + }, + { + name: "multiple elements", + data: []byte("[\"a\",\"b\",\"c\"]"), + want: StringSlice{"a", "b", "c"}, + wantErr: false, + }, + { + name: "invalid json", + data: []byte("invalid"), + want: nil, + wantErr: true, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + var s StringSlice + err := s.Scan(tt.data) + if (err != nil) != tt.wantErr { + t.Errorf("Scan() error = %v, wantErr %v", err, tt.wantErr) + return + } + if !tt.wantErr { + if len(s) != len(tt.want) { + t.Errorf("Scan() = %v, want %v", s, tt.want) + return + } + for i := range s { + if s[i] != tt.want[i] { + t.Errorf("Scan() = %v, want %v", s, tt.want) + return + } + } + } + }) + } +} + +func TestStringSlice_MarshalJSON(t *testing.T) { + tests := []struct { + name string + s StringSlice + want string + }{ + { + name: "nil slice", + s: nil, + want: "[]", + }, + { + name: "empty slice", + s: StringSlice{}, + want: "[]", + }, + { + name: "single element", + s: StringSlice{"hello"}, + want: "[\"hello\"]", + }, + { + name: "multiple elements", + s: StringSlice{"a", "b", "c"}, + want: "[\"a\",\"b\",\"c\"]", + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + got, err := json.Marshal(tt.s) + if err != nil { + t.Errorf("MarshalJSON() error = %v", err) + return + } + if string(got) != tt.want { + t.Errorf("MarshalJSON() = %v, want %v", string(got), tt.want) + } + }) + } +} + +func TestStringSlice_UnmarshalJSON(t *testing.T) { + tests := []struct { + name string + data string + want StringSlice + wantErr bool + }{ + { + name: "null", + data: "null", + want: nil, + wantErr: false, + }, + { + name: "empty array", + data: "[]", + want: StringSlice{}, + wantErr: false, + }, + { + name: "single element", + data: "[\"hello\"]", + want: StringSlice{"hello"}, + wantErr: false, + }, + { + name: "multiple elements", + data: "[\"a\",\"b\",\"c\"]", + want: StringSlice{"a", "b", "c"}, + wantErr: false, + }, + { + name: "single string value", + data: "\"hello\"", + want: StringSlice{"hello"}, + wantErr: false, + }, + { + name: "invalid json", + data: "invalid", + want: nil, + wantErr: true, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + var s StringSlice + err := json.Unmarshal([]byte(tt.data), &s) + if (err != nil) != tt.wantErr { + t.Errorf("UnmarshalJSON() error = %v, wantErr %v", err, tt.wantErr) + return + } + if !tt.wantErr && len(s) != len(tt.want) { + t.Errorf("UnmarshalJSON() = %v, want %v", s, tt.want) + } + }) + } +} + +func TestStringSlice_Contains(t *testing.T) { + s := StringSlice{"a", "b", "c"} + + if !s.Contains("a") { + t.Error("Contains('a') should return true") + } + if !s.Contains("b") { + t.Error("Contains('b') should return true") + } + if s.Contains("d") { + t.Error("Contains('d') should return false") + } + if StringSlice(nil).Contains("a") { + t.Error("nil Contains() should return false") + } +} + +func TestStringSlice_ToMap(t *testing.T) { + s := StringSlice{"a", "b", "a", "c"} + + m := s.ToMap() + + if len(m) != 3 { + t.Errorf("ToMap() length = %d, want 3", len(m)) + } + if _, ok := m["a"]; !ok { + t.Error("ToMap() should contain 'a'") + } + if _, ok := m["b"]; !ok { + t.Error("ToMap() should contain 'b'") + } + if _, ok := m["c"]; !ok { + t.Error("ToMap() should contain 'c'") + } + if _, ok := m["d"]; ok { + t.Error("ToMap() should not contain 'd'") + } +} + +func TestStringSlice_Join(t *testing.T) { + s := StringSlice{"a", "b", "c"} + + if s.Join(",") != "a,b,c" { + t.Errorf("Join(',') = %s, want 'a,b,c'", s.Join(",")) + } + if s.Join("") != "abc" { + t.Errorf("Join('') = %s, want 'abc'", s.Join("")) + } + if s.Join("-") != "a-b-c" { + t.Errorf("Join('-') = %s, want 'a-b-c'", s.Join("-")) + } +} + +func TestStringSlice_ToSlice(t *testing.T) { + tests := []struct { + name string + s StringSlice + want []string + }{ + {"nil", nil, nil}, + {"empty", StringSlice{}, []string{}}, + {"single", StringSlice{"a"}, []string{"a"}}, + {"multiple", StringSlice{"a", "b", "c"}, []string{"a", "b", "c"}}, + } + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + got := tt.s.ToSlice() + if len(got) != len(tt.want) { + t.Fatalf("ToSlice() len = %d, want %d", len(got), len(tt.want)) + } + for i := range got { + if got[i] != tt.want[i] { + t.Errorf("ToSlice()[%d] = %s, want %s", i, got[i], tt.want[i]) + } + } + }) + } +} diff --git a/gormtypes/uint64_slice.go b/gormtypes/uint64_slice.go new file mode 100644 index 0000000..91db94a --- /dev/null +++ b/gormtypes/uint64_slice.go @@ -0,0 +1,99 @@ +// zogo/gormtypes/uint64_slice.go +package gormtypes + +import ( + "database/sql/driver" + "encoding/json" + "strconv" + "strings" +) + +// Uint64Slice []uint64 类型别名 +// 用于存储和传输 uint64 数组 +// - GORM 层:序列化为 JSON 字符串存储到数据库 +// - API 层:前端传 string[],返回 string[](避免 JavaScript 大数字精度问题) +type Uint64Slice []uint64 + +// Value 实现 driver.Valuer 接口,用于写入数据库 +func (s Uint64Slice) Value() (driver.Value, error) { + if s == nil { + return []byte("[]"), nil + } + return json.Marshal(s) +} + +// Scan 实现 sql.Scanner 接口,用于从数据库读取 +func (s *Uint64Slice) Scan(value interface{}) error { + if value == nil { + *s = Uint64Slice{} + return nil + } + + var data []byte + switch v := value.(type) { + case []byte: + data = v + case string: + data = []byte(v) + default: + return nil + } + + return json.Unmarshal(data, s) +} + +// MarshalJSON 序列化输出:string[] 格式 +func (s Uint64Slice) MarshalJSON() ([]byte, error) { + if s == nil { + return []byte("null"), nil + } + + var b strings.Builder + b.WriteByte('[') + for i, v := range s { + if i > 0 { + b.WriteByte(',') + } + b.WriteByte('"') + b.WriteString(strconv.FormatUint(v, 10)) + b.WriteByte('"') + } + b.WriteByte(']') + return []byte(b.String()), nil +} + +// UnmarshalJSON 反序列化支持 string[] 和 number[] 两种格式 +func (s *Uint64Slice) UnmarshalJSON(data []byte) error { + if string(data) == "null" { + *s = nil + return nil + } + + // 优先尝试解析为字符串数组 + var strArr []string + if err := json.Unmarshal(data, &strArr); err == nil { + res := make(Uint64Slice, 0, len(strArr)) + for _, v := range strArr { + i, err := strconv.ParseUint(v, 10, 64) + if err != nil { + return err + } + res = append(res, i) + } + *s = res + return nil + } + + // 兼容数字数组格式 + var numArr []uint64 + if err := json.Unmarshal(data, &numArr); err != nil { + return err + } + *s = numArr + return nil +} + +// ToSlice 转换为原生 []uint64 切片,用于 ... 展开 +func (s Uint64Slice) ToSlice() []uint64 { + return []uint64(s) +} diff --git a/gormtypes/uint64_slice_test.go b/gormtypes/uint64_slice_test.go new file mode 100644 index 0000000..a2d974f --- /dev/null +++ b/gormtypes/uint64_slice_test.go @@ -0,0 +1,186 @@ +// zogo/gormtypes/uint64_slice_test.go +package gormtypes + +import ( + "encoding/json" + "testing" +) + +func TestUint64Slice_Value(t *testing.T) { + tests := []struct { + name string + s Uint64Slice + want []byte + wantErr bool + }{ + {name: "nil slice", s: nil, want: []byte("[]"), wantErr: false}, + {name: "empty slice", s: Uint64Slice{}, want: []byte("[]"), wantErr: false}, + {name: "single element", s: Uint64Slice{1}, want: []byte("[\"1\"]"), wantErr: false}, + {name: "multiple elements", s: Uint64Slice{1, 2, 3}, want: []byte("[\"1\",\"2\",\"3\"]"), wantErr: false}, + {name: "large number", s: Uint64Slice{18446744073709551615}, want: []byte("[\"18446744073709551615\"]"), wantErr: false}, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + got, err := tt.s.Value() + if (err != nil) != tt.wantErr { + t.Errorf("Value() error = %v, wantErr %v", err, tt.wantErr) + return + } + gotBytes, ok := got.([]byte) + if !ok { + t.Errorf("Value() returned non-[]byte type: %T", got) + return + } + if string(gotBytes) != string(tt.want) { + t.Errorf("Value() = %v, want %v", string(gotBytes), string(tt.want)) + } + }) + } +} + +func TestUint64Slice_Scan(t *testing.T) { + tests := []struct { + name string + data interface{} + want Uint64Slice + wantErr bool + }{ + {name: "nil value", data: nil, want: Uint64Slice{}, wantErr: false}, + {name: "empty json", data: []byte("[]"), want: Uint64Slice{}, wantErr: false}, + {name: "single element", data: []byte("[1]"), want: Uint64Slice{1}, wantErr: false}, + {name: "multiple elements", data: []byte("[1,2,3]"), want: Uint64Slice{1, 2, 3}, wantErr: false}, + {name: "invalid json", data: []byte("invalid"), want: nil, wantErr: true}, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + var s Uint64Slice + err := s.Scan(tt.data) + if (err != nil) != tt.wantErr { + t.Errorf("Scan() error = %v, wantErr %v", err, tt.wantErr) + return + } + if !tt.wantErr && len(s) != len(tt.want) { + t.Errorf("Scan() = %v, want %v", s, tt.want) + } + }) + } +} + +func TestUint64Slice_MarshalJSON(t *testing.T) { + tests := []struct { + name string + s Uint64Slice + want string + }{ + {name: "nil slice", s: nil, want: "null"}, + {name: "empty slice", s: Uint64Slice{}, want: "[]"}, + {name: "single element", s: Uint64Slice{1}, want: "[\"1\"]"}, + {name: "multiple elements", s: Uint64Slice{1, 2, 3}, want: "[\"1\",\"2\",\"3\"]"}, + {name: "large number", s: Uint64Slice{18446744073709551615}, want: "[\"18446744073709551615\"]"}, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + got, err := json.Marshal(tt.s) + if err != nil { + t.Errorf("MarshalJSON() error = %v", err) + return + } + if string(got) != tt.want { + t.Errorf("MarshalJSON() = %v, want %v", string(got), tt.want) + } + }) + } +} + +func TestUint64Slice_UnmarshalJSON(t *testing.T) { + tests := []struct { + name string + data string + want Uint64Slice + wantErr bool + }{ + {name: "null", data: "null", want: nil, wantErr: false}, + {name: "empty array", data: "[]", want: Uint64Slice{}, wantErr: false}, + {name: "string array", data: "[\"1\",\"2\",\"3\"]", want: Uint64Slice{1, 2, 3}, wantErr: false}, + {name: "number array", data: "[1,2,3]", want: Uint64Slice{1, 2, 3}, wantErr: false}, + {name: "invalid string", data: "[\"invalid\"]", want: nil, wantErr: true}, + {name: "invalid json", data: "invalid", want: nil, wantErr: true}, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + var s Uint64Slice + err := json.Unmarshal([]byte(tt.data), &s) + if (err != nil) != tt.wantErr { + t.Errorf("UnmarshalJSON() error = %v, wantErr %v", err, tt.wantErr) + return + } + if !tt.wantErr && len(s) != len(tt.want) { + t.Errorf("UnmarshalJSON() = %v, want %v", s, tt.want) + } + }) + } +} + +func TestUint64Slice_RoundTrip(t *testing.T) { + original := Uint64Slice{1, 2, 3, 18446744073709551615} + + // GORM Value -> Scan + val, err := original.Value() + if err != nil { + t.Fatalf("Value() error = %v", err) + } + + var scanned Uint64Slice + if err := scanned.Scan(val); err != nil { + t.Fatalf("Scan() error = %v", err) + } + + if len(scanned) != len(original) { + t.Errorf("GORM round-trip failed: got %v, want %v", scanned, original) + } + + // JSON Marshal -> Unmarshal + jsonData, err := json.Marshal(original) + if err != nil { + t.Fatalf("MarshalJSON() error = %v", err) + } + + var unmarshaled Uint64Slice + if err := json.Unmarshal(jsonData, &unmarshaled); err != nil { + t.Fatalf("UnmarshalJSON() error = %v", err) + } + + if len(unmarshaled) != len(original) { + t.Errorf("JSON round-trip failed: got %v, want %v", unmarshaled, original) + } +} + +func TestUint64Slice_ToSlice(t *testing.T) { + tests := []struct { + name string + s Uint64Slice + want []uint64 + }{ + {"nil", nil, nil}, + {"empty", Uint64Slice{}, []uint64{}}, + {"single", Uint64Slice{1}, []uint64{1}}, + {"multiple", Uint64Slice{1, 2, 3}, []uint64{1, 2, 3}}, + } + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + got := tt.s.ToSlice() + if len(got) != len(tt.want) { + t.Fatalf("ToSlice() len = %d, want %d", len(got), len(tt.want)) + } + for i := range got { + if got[i] != tt.want[i] { + t.Errorf("ToSlice()[%d] = %d, want %d", i, got[i], tt.want[i]) + } + } + }) + } +}