diff --git a/pkg/sqltypes/sql_types_scalar_test.go b/pkg/sqltypes/sql_types_scalar_test.go new file mode 100644 index 0000000..292a9f7 --- /dev/null +++ b/pkg/sqltypes/sql_types_scalar_test.go @@ -0,0 +1,157 @@ +package sqltypes + +import ( + "database/sql/driver" + "encoding/json" + "math" + "testing" + "time" +) + +func TestSqlNull_ValueScalarCases(t *testing.T) { + tests := []struct { + name string + input SqlNull[any] + want driver.Value + }{ + {name: "invalid", input: SqlNull[any]{}, want: nil}, + {name: "integer", input: Null[any](int64(42), true), want: int64(42)}, + {name: "string", input: Null[any]("hello", true), want: "hello"}, + {name: "boolean", input: Null[any](true, true), want: true}, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + got, err := tt.input.Value() + if err != nil { + t.Fatalf("Value returned error: %v", err) + } + if got != tt.want { + t.Errorf("Value() = %v (%T), want %v (%T)", got, got, tt.want, tt.want) + } + }) + } +} + +func TestSqlNull_Int64Conversions(t *testing.T) { + tests := []struct { + name string + input SqlNull[any] + want int64 + }{ + {name: "invalid", input: SqlNull[any]{}, want: 0}, + {name: "signed integer", input: Null[any](int32(-12), true), want: -12}, + {name: "unsigned integer", input: Null[any](uint16(12), true), want: 12}, + {name: "float truncates", input: Null[any](float64(12.9), true), want: 12}, + {name: "numeric string", input: Null[any]("123", true), want: 123}, + {name: "invalid string", input: Null[any]("not a number", true), want: 0}, + {name: "true", input: Null[any](true, true), want: 1}, + {name: "false", input: Null[any](false, true), want: 0}, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + if got := tt.input.Int64(); got != tt.want { + t.Errorf("Int64() = %d, want %d", got, tt.want) + } + }) + } +} + +func TestSqlNull_Float64Conversions(t *testing.T) { + tests := []struct { + name string + input SqlNull[any] + want float64 + }{ + {name: "invalid", input: SqlNull[any]{}, want: 0}, + {name: "float", input: Null[any](float32(1.25), true), want: 1.25}, + {name: "signed integer", input: Null[any](int64(-12), true), want: -12}, + {name: "unsigned integer", input: Null[any](uint16(12), true), want: 12}, + {name: "numeric string", input: Null[any]("12.5", true), want: 12.5}, + {name: "invalid string", input: Null[any]("not a number", true), want: 0}, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + if got := tt.input.Float64(); got != tt.want { + t.Errorf("Float64() = %v, want %v", got, tt.want) + } + }) + } +} + +func TestSqlDate_JSONNullAndInvalid(t *testing.T) { + tests := []struct { + name string + json string + valid bool + }{ + {name: "null", json: "null", valid: false}, + {name: "invalid date", json: `"not-a-date"`, valid: false}, + {name: "valid date", json: `"2024-01-15"`, valid: true}, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + var got SqlDate + if err := json.Unmarshal([]byte(tt.json), &got); err != nil { + t.Fatalf("UnmarshalJSON returned error: %v", err) + } + if got.Valid != tt.valid { + t.Errorf("Valid = %v, want %v", got.Valid, tt.valid) + } + }) + } + + if data, err := json.Marshal(SqlDate{}); err != nil { + t.Fatalf("MarshalJSON returned error: %v", err) + } else if string(data) != "null" { + t.Errorf("MarshalJSON() = %s, want null", data) + } +} + +func TestSqlTypeNowConstructors(t *testing.T) { + before := time.Now() + timestamp := SqlTimeStampNow() + date := SqlDateNow() + tm := SqlTimeNow() + after := time.Now() + + for name, got := range map[string]time.Time{ + "timestamp": timestamp.Time(), + "date": date.Time(), + "time": tm.Time(), + } { + if !got.After(before) && !got.Equal(before) || got.After(after) { + t.Errorf("%s constructor returned %v outside [%v, %v]", name, got, before, after) + } + } + if !timestamp.Valid || !date.Valid || !tm.Valid { + t.Fatal("Now constructors must return valid values") + } +} + +func TestNewSqlAndToJSONDT(t *testing.T) { + if got := NewSql[int64]("42"); !got.Valid || got.Val != 42 { + t.Errorf("NewSql[int64](\"42\") = %#v, want valid 42", got) + } + if got := NewSql[int64](nil); got.Valid { + t.Errorf("NewSql[int64](nil) = %#v, want invalid", got) + } + if got := NewSqlFloat32(1.5); !got.Valid || got.Val != 1.5 { + t.Errorf("NewSqlFloat32(1.5) = %#v, want valid 1.5", got) + } + + when := time.Date(2024, 1, 15, 10, 30, 45, 0, time.UTC) + if got := ToJSONDT(when); got != "2024-01-15T10:30:45Z" { + t.Errorf("ToJSONDT() = %q, want RFC3339 timestamp", got) + } +} + +func TestSqlNull_Float64PreservesInfinity(t *testing.T) { + got := Null[float64](math.Inf(1), true).Float64() + if !math.IsInf(got, 1) { + t.Errorf("Float64() = %v, want +Inf", got) + } +}