feat(sqltypes): add nullable SQL types with automatic casting
CI / build-and-test (push) Successful in 1m54s
Release / release (push) Successful in 3m12s

* Introduced SqlNull type for nullable values with auto-casting.
* Implemented JSON, YAML, and XML marshaling/unmarshaling.
* Added specific types for common SQL types (e.g., SqlInt64, SqlString).
* Included utility functions for creating nullable types.
* Added SqlTimeStamp, SqlDate, and SqlTime types with custom formatting.
This commit is contained in:
Hein
2026-07-15 12:55:15 +02:00
parent f94fddddb1
commit 3d4e6d0939
40 changed files with 1240 additions and 330 deletions
@@ -1,755 +0,0 @@
package spectypes
import (
"database/sql/driver"
"encoding/json"
"fmt"
"strconv"
"strings"
"github.com/google/uuid"
)
// parsePostgresArrayElements parses a PostgreSQL array literal (e.g. `{a,"b,c",d}`)
// into a slice of raw string elements. Each element retains its unquoted/unescaped value.
func parsePostgresArrayElements(s string) ([]string, error) {
s = strings.TrimSpace(s)
if s == "" || strings.EqualFold(s, "null") || strings.EqualFold(s, "NULL") {
return nil, nil
}
if !strings.HasPrefix(s, "{") || !strings.HasSuffix(s, "}") {
return nil, fmt.Errorf("not a valid PostgreSQL array literal: %q", s)
}
inner := s[1 : len(s)-1]
if inner == "" {
return []string{}, nil
}
var result []string
var cur strings.Builder
inQuotes := false
i := 0
for i < len(inner) {
c := inner[i]
switch {
case c == '"' && !inQuotes:
inQuotes = true
case c == '"' && inQuotes:
if i+1 < len(inner) && inner[i+1] == '"' {
cur.WriteByte('"')
i++
} else {
inQuotes = false
}
case c == '\\' && inQuotes:
if i+1 < len(inner) {
cur.WriteByte(inner[i+1])
i++
}
case c == ',' && !inQuotes:
result = append(result, cur.String())
cur.Reset()
default:
cur.WriteByte(c)
}
i++
}
result = append(result, cur.String())
return result, nil
}
// formatPostgresStringArray formats a []string back into a PostgreSQL array literal.
func formatPostgresStringArray(vals []string) string {
if vals == nil {
return "NULL"
}
parts := make([]string, len(vals))
for i, v := range vals {
// Quote if value contains comma, double-quote, backslash, braces, whitespace, or is empty.
needsQuote := v == "" || strings.ContainsAny(v, `,"\\{}`+"\t\n\r ")
if needsQuote {
v = strings.ReplaceAll(v, `\`, `\\`)
v = strings.ReplaceAll(v, `"`, `""`)
parts[i] = `"` + v + `"`
} else {
parts[i] = v
}
}
return "{" + strings.Join(parts, ",") + "}"
}
// ── SqlStringArray ───────────────────────────────────────────────────────────
// SqlStringArray is a nullable PostgreSQL text[] / varchar[] array.
type SqlStringArray struct {
Val []string
Valid bool
}
func (a *SqlStringArray) Scan(value any) error {
if value == nil {
a.Valid = false
a.Val = nil
return nil
}
var s string
switch v := value.(type) {
case string:
s = v
case []byte:
s = string(v)
default:
return fmt.Errorf("SqlStringArray: cannot scan type %T", value)
}
elems, err := parsePostgresArrayElements(s)
if err != nil {
return err
}
a.Val = elems
a.Valid = true
return nil
}
func (a SqlStringArray) Value() (driver.Value, error) {
if !a.Valid {
return nil, nil
}
return formatPostgresStringArray(a.Val), nil
}
func (a SqlStringArray) MarshalJSON() ([]byte, error) {
if !a.Valid {
return []byte("null"), nil
}
return json.Marshal(a.Val)
}
func (a *SqlStringArray) UnmarshalJSON(b []byte) error {
s := strings.TrimSpace(string(b))
if s == "null" {
a.Valid = false
a.Val = nil
return nil
}
var vals []string
if err := json.Unmarshal(b, &vals); err != nil {
return err
}
a.Val = vals
a.Valid = true
return nil
}
func NewSqlStringArray(v []string) SqlStringArray {
return SqlStringArray{Val: v, Valid: true}
}
// ── SqlInt16Array ────────────────────────────────────────────────────────────
type SqlInt16Array struct {
Val []int16
Valid bool
}
func (a *SqlInt16Array) Scan(value any) error {
if value == nil {
a.Valid = false
a.Val = nil
return nil
}
var s string
switch v := value.(type) {
case string:
s = v
case []byte:
s = string(v)
default:
return fmt.Errorf("SqlInt16Array: cannot scan type %T", value)
}
elems, err := parsePostgresArrayElements(s)
if err != nil {
return err
}
a.Val = make([]int16, len(elems))
for i, e := range elems {
n, err := strconv.ParseInt(strings.TrimSpace(e), 10, 16)
if err != nil {
return fmt.Errorf("SqlInt16Array: element %d %q: %w", i, e, err)
}
a.Val[i] = int16(n)
}
a.Valid = true
return nil
}
func (a SqlInt16Array) Value() (driver.Value, error) {
if !a.Valid {
return nil, nil
}
parts := make([]string, len(a.Val))
for i, v := range a.Val {
parts[i] = strconv.FormatInt(int64(v), 10)
}
return "{" + strings.Join(parts, ",") + "}", nil
}
func (a SqlInt16Array) MarshalJSON() ([]byte, error) {
if !a.Valid {
return []byte("null"), nil
}
return json.Marshal(a.Val)
}
func (a *SqlInt16Array) UnmarshalJSON(b []byte) error {
if strings.TrimSpace(string(b)) == "null" {
a.Valid = false
a.Val = nil
return nil
}
var vals []int16
if err := json.Unmarshal(b, &vals); err != nil {
return err
}
a.Val = vals
a.Valid = true
return nil
}
func NewSqlInt16Array(v []int16) SqlInt16Array {
return SqlInt16Array{Val: v, Valid: true}
}
// ── SqlInt32Array ────────────────────────────────────────────────────────────
type SqlInt32Array struct {
Val []int32
Valid bool
}
func (a *SqlInt32Array) Scan(value any) error {
if value == nil {
a.Valid = false
a.Val = nil
return nil
}
var s string
switch v := value.(type) {
case string:
s = v
case []byte:
s = string(v)
default:
return fmt.Errorf("SqlInt32Array: cannot scan type %T", value)
}
elems, err := parsePostgresArrayElements(s)
if err != nil {
return err
}
a.Val = make([]int32, len(elems))
for i, e := range elems {
n, err := strconv.ParseInt(strings.TrimSpace(e), 10, 32)
if err != nil {
return fmt.Errorf("SqlInt32Array: element %d %q: %w", i, e, err)
}
a.Val[i] = int32(n)
}
a.Valid = true
return nil
}
func (a SqlInt32Array) Value() (driver.Value, error) {
if !a.Valid {
return nil, nil
}
parts := make([]string, len(a.Val))
for i, v := range a.Val {
parts[i] = strconv.FormatInt(int64(v), 10)
}
return "{" + strings.Join(parts, ",") + "}", nil
}
func (a SqlInt32Array) MarshalJSON() ([]byte, error) {
if !a.Valid {
return []byte("null"), nil
}
return json.Marshal(a.Val)
}
func (a *SqlInt32Array) UnmarshalJSON(b []byte) error {
if strings.TrimSpace(string(b)) == "null" {
a.Valid = false
a.Val = nil
return nil
}
var vals []int32
if err := json.Unmarshal(b, &vals); err != nil {
return err
}
a.Val = vals
a.Valid = true
return nil
}
func NewSqlInt32Array(v []int32) SqlInt32Array {
return SqlInt32Array{Val: v, Valid: true}
}
// ── SqlInt64Array ────────────────────────────────────────────────────────────
type SqlInt64Array struct {
Val []int64
Valid bool
}
func (a *SqlInt64Array) Scan(value any) error {
if value == nil {
a.Valid = false
a.Val = nil
return nil
}
var s string
switch v := value.(type) {
case string:
s = v
case []byte:
s = string(v)
default:
return fmt.Errorf("SqlInt64Array: cannot scan type %T", value)
}
elems, err := parsePostgresArrayElements(s)
if err != nil {
return err
}
a.Val = make([]int64, len(elems))
for i, e := range elems {
n, err := strconv.ParseInt(strings.TrimSpace(e), 10, 64)
if err != nil {
return fmt.Errorf("SqlInt64Array: element %d %q: %w", i, e, err)
}
a.Val[i] = n
}
a.Valid = true
return nil
}
func (a SqlInt64Array) Value() (driver.Value, error) {
if !a.Valid {
return nil, nil
}
parts := make([]string, len(a.Val))
for i, v := range a.Val {
parts[i] = strconv.FormatInt(v, 10)
}
return "{" + strings.Join(parts, ",") + "}", nil
}
func (a SqlInt64Array) MarshalJSON() ([]byte, error) {
if !a.Valid {
return []byte("null"), nil
}
return json.Marshal(a.Val)
}
func (a *SqlInt64Array) UnmarshalJSON(b []byte) error {
if strings.TrimSpace(string(b)) == "null" {
a.Valid = false
a.Val = nil
return nil
}
var vals []int64
if err := json.Unmarshal(b, &vals); err != nil {
return err
}
a.Val = vals
a.Valid = true
return nil
}
func NewSqlInt64Array(v []int64) SqlInt64Array {
return SqlInt64Array{Val: v, Valid: true}
}
// ── SqlFloat32Array ──────────────────────────────────────────────────────────
type SqlFloat32Array struct {
Val []float32
Valid bool
}
func (a *SqlFloat32Array) Scan(value any) error {
if value == nil {
a.Valid = false
a.Val = nil
return nil
}
var s string
switch v := value.(type) {
case string:
s = v
case []byte:
s = string(v)
default:
return fmt.Errorf("SqlFloat32Array: cannot scan type %T", value)
}
elems, err := parsePostgresArrayElements(s)
if err != nil {
return err
}
a.Val = make([]float32, len(elems))
for i, e := range elems {
f, err := strconv.ParseFloat(strings.TrimSpace(e), 32)
if err != nil {
return fmt.Errorf("SqlFloat32Array: element %d %q: %w", i, e, err)
}
a.Val[i] = float32(f)
}
a.Valid = true
return nil
}
func (a SqlFloat32Array) Value() (driver.Value, error) {
if !a.Valid {
return nil, nil
}
parts := make([]string, len(a.Val))
for i, v := range a.Val {
parts[i] = strconv.FormatFloat(float64(v), 'f', -1, 32)
}
return "{" + strings.Join(parts, ",") + "}", nil
}
func (a SqlFloat32Array) MarshalJSON() ([]byte, error) {
if !a.Valid {
return []byte("null"), nil
}
return json.Marshal(a.Val)
}
func (a *SqlFloat32Array) UnmarshalJSON(b []byte) error {
if strings.TrimSpace(string(b)) == "null" {
a.Valid = false
a.Val = nil
return nil
}
var vals []float32
if err := json.Unmarshal(b, &vals); err != nil {
return err
}
a.Val = vals
a.Valid = true
return nil
}
func NewSqlFloat32Array(v []float32) SqlFloat32Array {
return SqlFloat32Array{Val: v, Valid: true}
}
// ── SqlFloat64Array ──────────────────────────────────────────────────────────
type SqlFloat64Array struct {
Val []float64
Valid bool
}
func (a *SqlFloat64Array) Scan(value any) error {
if value == nil {
a.Valid = false
a.Val = nil
return nil
}
var s string
switch v := value.(type) {
case string:
s = v
case []byte:
s = string(v)
default:
return fmt.Errorf("SqlFloat64Array: cannot scan type %T", value)
}
elems, err := parsePostgresArrayElements(s)
if err != nil {
return err
}
a.Val = make([]float64, len(elems))
for i, e := range elems {
f, err := strconv.ParseFloat(strings.TrimSpace(e), 64)
if err != nil {
return fmt.Errorf("SqlFloat64Array: element %d %q: %w", i, e, err)
}
a.Val[i] = f
}
a.Valid = true
return nil
}
func (a SqlFloat64Array) Value() (driver.Value, error) {
if !a.Valid {
return nil, nil
}
parts := make([]string, len(a.Val))
for i, v := range a.Val {
parts[i] = strconv.FormatFloat(v, 'f', -1, 64)
}
return "{" + strings.Join(parts, ",") + "}", nil
}
func (a SqlFloat64Array) MarshalJSON() ([]byte, error) {
if !a.Valid {
return []byte("null"), nil
}
return json.Marshal(a.Val)
}
func (a *SqlFloat64Array) UnmarshalJSON(b []byte) error {
if strings.TrimSpace(string(b)) == "null" {
a.Valid = false
a.Val = nil
return nil
}
var vals []float64
if err := json.Unmarshal(b, &vals); err != nil {
return err
}
a.Val = vals
a.Valid = true
return nil
}
func NewSqlFloat64Array(v []float64) SqlFloat64Array {
return SqlFloat64Array{Val: v, Valid: true}
}
// ── SqlBoolArray ─────────────────────────────────────────────────────────────
type SqlBoolArray struct {
Val []bool
Valid bool
}
func (a *SqlBoolArray) Scan(value any) error {
if value == nil {
a.Valid = false
a.Val = nil
return nil
}
var s string
switch v := value.(type) {
case string:
s = v
case []byte:
s = string(v)
default:
return fmt.Errorf("SqlBoolArray: cannot scan type %T", value)
}
elems, err := parsePostgresArrayElements(s)
if err != nil {
return err
}
a.Val = make([]bool, len(elems))
for i, e := range elems {
e = strings.ToLower(strings.TrimSpace(e))
a.Val[i] = e == "t" || e == "true" || e == "1" || e == "yes"
}
a.Valid = true
return nil
}
func (a SqlBoolArray) Value() (driver.Value, error) {
if !a.Valid {
return nil, nil
}
parts := make([]string, len(a.Val))
for i, v := range a.Val {
if v {
parts[i] = "t"
} else {
parts[i] = "f"
}
}
return "{" + strings.Join(parts, ",") + "}", nil
}
func (a SqlBoolArray) MarshalJSON() ([]byte, error) {
if !a.Valid {
return []byte("null"), nil
}
return json.Marshal(a.Val)
}
func (a *SqlBoolArray) UnmarshalJSON(b []byte) error {
if strings.TrimSpace(string(b)) == "null" {
a.Valid = false
a.Val = nil
return nil
}
var vals []bool
if err := json.Unmarshal(b, &vals); err != nil {
return err
}
a.Val = vals
a.Valid = true
return nil
}
func NewSqlBoolArray(v []bool) SqlBoolArray {
return SqlBoolArray{Val: v, Valid: true}
}
// ── SqlUUIDArray ─────────────────────────────────────────────────────────────
type SqlUUIDArray struct {
Val []uuid.UUID
Valid bool
}
func (a *SqlUUIDArray) Scan(value any) error {
if value == nil {
a.Valid = false
a.Val = nil
return nil
}
var s string
switch v := value.(type) {
case string:
s = v
case []byte:
s = string(v)
default:
return fmt.Errorf("SqlUUIDArray: cannot scan type %T", value)
}
elems, err := parsePostgresArrayElements(s)
if err != nil {
return err
}
a.Val = make([]uuid.UUID, len(elems))
for i, e := range elems {
u, err := uuid.Parse(strings.TrimSpace(e))
if err != nil {
return fmt.Errorf("SqlUUIDArray: element %d %q: %w", i, e, err)
}
a.Val[i] = u
}
a.Valid = true
return nil
}
func (a SqlUUIDArray) Value() (driver.Value, error) {
if !a.Valid {
return nil, nil
}
parts := make([]string, len(a.Val))
for i, v := range a.Val {
parts[i] = v.String()
}
return "{" + strings.Join(parts, ",") + "}", nil
}
func (a SqlUUIDArray) MarshalJSON() ([]byte, error) {
if !a.Valid {
return []byte("null"), nil
}
return json.Marshal(a.Val)
}
func (a *SqlUUIDArray) UnmarshalJSON(b []byte) error {
if strings.TrimSpace(string(b)) == "null" {
a.Valid = false
a.Val = nil
return nil
}
var vals []uuid.UUID
if err := json.Unmarshal(b, &vals); err != nil {
return err
}
a.Val = vals
a.Valid = true
return nil
}
func NewSqlUUIDArray(v []uuid.UUID) SqlUUIDArray {
return SqlUUIDArray{Val: v, Valid: true}
}
// ── SqlVector ────────────────────────────────────────────────────────────────
// SqlVector is a nullable pgvector `vector` type backed by []float32.
// Wire format: `[1.0,2.0,3.0]` (square brackets, comma-separated floats).
type SqlVector struct {
Val []float32
Valid bool
}
func (v *SqlVector) Scan(value any) error {
if value == nil {
v.Valid = false
v.Val = nil
return nil
}
var s string
switch val := value.(type) {
case string:
s = val
case []byte:
s = string(val)
default:
return fmt.Errorf("SqlVector: cannot scan type %T", value)
}
s = strings.TrimSpace(s)
if !strings.HasPrefix(s, "[") || !strings.HasSuffix(s, "]") {
return fmt.Errorf("SqlVector: not a valid vector literal: %q", s)
}
inner := s[1 : len(s)-1]
if inner == "" {
v.Val = []float32{}
v.Valid = true
return nil
}
parts := strings.Split(inner, ",")
v.Val = make([]float32, len(parts))
for i, p := range parts {
f, err := strconv.ParseFloat(strings.TrimSpace(p), 32)
if err != nil {
return fmt.Errorf("SqlVector: element %d %q: %w", i, p, err)
}
v.Val[i] = float32(f)
}
v.Valid = true
return nil
}
func (v SqlVector) Value() (driver.Value, error) {
if !v.Valid {
return nil, nil
}
parts := make([]string, len(v.Val))
for i, f := range v.Val {
parts[i] = strconv.FormatFloat(float64(f), 'f', -1, 32)
}
return "[" + strings.Join(parts, ",") + "]", nil
}
func (v SqlVector) MarshalJSON() ([]byte, error) {
if !v.Valid {
return []byte("null"), nil
}
return json.Marshal(v.Val)
}
func (v *SqlVector) UnmarshalJSON(b []byte) error {
if strings.TrimSpace(string(b)) == "null" {
v.Valid = false
v.Val = nil
return nil
}
var vals []float32
if err := json.Unmarshal(b, &vals); err != nil {
return err
}
v.Val = vals
v.Valid = true
return nil
}
func NewSqlVector(val []float32) SqlVector {
return SqlVector{Val: val, Valid: true}
}
-658
View File
@@ -1,658 +0,0 @@
// Package spectypes provides nullable SQL types with automatic casting and conversion methods.
package spectypes
import (
"database/sql"
"database/sql/driver"
"encoding/base64"
"encoding/json"
"fmt"
"reflect"
"strconv"
"strings"
"time"
"github.com/google/uuid"
)
// tryParseDT attempts to parse a string into a time.Time using various formats.
func tryParseDT(str string) (time.Time, error) {
var lasterror error
tryFormats := []string{
time.RFC3339,
"2006-01-02T15:04:05.000-0700",
"2006-01-02T15:04:05.000",
"06-01-02T15:04:05.000",
"2006-01-02T15:04:05",
"2006-01-02 15:04:05",
"02/01/2006",
"02-01-2006",
"2006-01-02",
"15:04:05.000",
"15:04:05",
"15:04",
}
for _, f := range tryFormats {
tx, err := time.Parse(f, str)
if err == nil {
return tx, nil
}
lasterror = err
}
return time.Time{}, lasterror // Return zero time on failure
}
// ToJSONDT formats a time.Time to RFC3339 string.
func ToJSONDT(dt time.Time) string {
return dt.Format(time.RFC3339)
}
// SqlNull is a generic nullable type that behaves like sql.NullXXX with auto-casting.
type SqlNull[T any] struct {
Val T
Valid bool
}
// Scan implements sql.Scanner.
func (n *SqlNull[T]) Scan(value any) error {
if value == nil {
n.Valid = false
n.Val = *new(T)
return nil
}
// Check if T is []byte, and decode base64 if applicable
// Do this BEFORE trying sql.Null to ensure base64 is handled
var zero T
if _, ok := any(zero).([]byte); ok {
// For []byte types, try to decode from base64
var strVal string
switch v := value.(type) {
case string:
strVal = v
case []byte:
strVal = string(v)
default:
strVal = fmt.Sprintf("%v", value)
}
// Try base64 decode
if decoded, err := base64.StdEncoding.DecodeString(strVal); err == nil {
n.Val = any(decoded).(T)
n.Valid = true
return nil
}
// Fallback to raw bytes
n.Val = any([]byte(strVal)).(T)
n.Valid = true
return nil
}
// Try standard sql.Null[T] for other types.
var sqlNull sql.Null[T]
if err := sqlNull.Scan(value); err == nil {
n.Val = sqlNull.V
n.Valid = sqlNull.Valid
return nil
}
// Fallback: parse from string/bytes.
switch v := value.(type) {
case string:
return n.FromString(v)
case []byte:
return n.FromString(string(v))
case float32, float64:
return n.FromString(fmt.Sprintf("%f", value))
case int, int8, int16, int32, int64, uint, uint8, uint16, uint32, uint64:
return n.FromString(fmt.Sprintf("%d", value))
default:
return n.FromString(fmt.Sprintf("%v", value))
}
}
func (n *SqlNull[T]) FromString(s string) error {
s = strings.TrimSpace(s)
n.Valid = false
n.Val = *new(T)
if s == "" || strings.EqualFold(s, "null") {
return nil
}
var zero T
switch any(zero).(type) {
case int, int8, int16, int32, int64, uint, uint8, uint16, uint32, uint64:
if i, err := strconv.ParseInt(s, 10, 64); err == nil {
reflect.ValueOf(&n.Val).Elem().SetInt(i)
n.Valid = true
}
if f, err := strconv.ParseFloat(s, 64); err == nil {
reflect.ValueOf(&n.Val).Elem().SetInt(int64(f))
n.Valid = true
}
case float32, float64:
if f, err := strconv.ParseFloat(s, 64); err == nil {
reflect.ValueOf(&n.Val).Elem().SetFloat(f)
n.Valid = true
}
case bool:
if b, err := strconv.ParseBool(s); err == nil {
n.Val = any(b).(T)
n.Valid = true
}
case time.Time:
if t, err := tryParseDT(s); err == nil && !t.IsZero() {
n.Val = any(t).(T)
n.Valid = true
}
case uuid.UUID:
if u, err := uuid.Parse(s); err == nil {
n.Val = any(u).(T)
n.Valid = true
}
case []byte:
n.Val = any([]byte(s)).(T)
n.Valid = true
case string:
n.Val = any(s).(T)
n.Valid = true
}
return nil
}
// Value implements driver.Valuer.
func (n SqlNull[T]) Value() (driver.Value, error) {
if !n.Valid {
return nil, nil
}
// Check if the type implements fmt.Stringer (e.g., uuid.UUID, custom types)
// Convert to string for driver compatibility
if stringer, ok := any(n.Val).(fmt.Stringer); ok {
return stringer.String(), nil
}
return any(n.Val), nil
}
// MarshalJSON implements json.Marshaler.
func (n SqlNull[T]) MarshalJSON() ([]byte, error) {
if !n.Valid {
return []byte("null"), nil
}
// Check if T is []byte, and encode to base64
if _, ok := any(n.Val).([]byte); ok {
// Encode []byte as base64
encoded := base64.StdEncoding.EncodeToString(any(n.Val).([]byte))
return json.Marshal(encoded)
}
return json.Marshal(n.Val)
}
// UnmarshalJSON implements json.Unmarshaler.
func (n *SqlNull[T]) UnmarshalJSON(b []byte) error {
if len(b) == 0 || string(b) == "null" || strings.TrimSpace(string(b)) == "" {
n.Valid = false
n.Val = *new(T)
return nil
}
// Check if T is []byte, and decode from base64
var val T
if _, ok := any(val).([]byte); ok {
// Unmarshal as string first (JSON representation)
var s string
if err := json.Unmarshal(b, &s); err == nil {
// Decode from base64
if decoded, err := base64.StdEncoding.DecodeString(s); err == nil {
n.Val = any(decoded).(T)
n.Valid = true
return nil
}
// Fallback to raw string as bytes
n.Val = any([]byte(s)).(T)
n.Valid = true
return nil
}
}
if err := json.Unmarshal(b, &val); err == nil {
n.Val = val
n.Valid = true
return nil
}
// Fallback: unmarshal as string and parse.
var s string
if err := json.Unmarshal(b, &s); err == nil {
return n.FromString(s)
}
return fmt.Errorf("cannot unmarshal %s into SqlNull[%T]", b, n.Val)
}
// String implements fmt.Stringer.
func (n SqlNull[T]) String() string {
if !n.Valid {
return ""
}
// Check if the type implements fmt.Stringer for better string representation
if stringer, ok := any(n.Val).(fmt.Stringer); ok {
return stringer.String()
}
return fmt.Sprintf("%v", n.Val)
}
// Int64 converts to int64 or 0 if invalid.
func (n SqlNull[T]) Int64() int64 {
if !n.Valid {
return 0
}
v := reflect.ValueOf(any(n.Val))
switch v.Kind() {
case reflect.Int, reflect.Int8, reflect.Int16, reflect.Int32, reflect.Int64:
return v.Int()
case reflect.Uint, reflect.Uint8, reflect.Uint16, reflect.Uint32, reflect.Uint64:
return int64(v.Uint())
case reflect.Float32, reflect.Float64:
return int64(v.Float())
case reflect.String:
i, _ := strconv.ParseInt(v.String(), 10, 64)
return i
case reflect.Bool:
if v.Bool() {
return 1
}
return 0
}
return 0
}
// Float64 converts to float64 or 0.0 if invalid.
func (n SqlNull[T]) Float64() float64 {
if !n.Valid {
return 0.0
}
v := reflect.ValueOf(any(n.Val))
switch v.Kind() {
case reflect.Float32, reflect.Float64:
return v.Float()
case reflect.Int, reflect.Int8, reflect.Int16, reflect.Int32, reflect.Int64:
return float64(v.Int())
case reflect.Uint, reflect.Uint8, reflect.Uint16, reflect.Uint32, reflect.Uint64:
return float64(v.Uint())
case reflect.String:
f, _ := strconv.ParseFloat(v.String(), 64)
return f
}
return 0.0
}
// Bool converts to bool or false if invalid.
func (n SqlNull[T]) Bool() bool {
if !n.Valid {
return false
}
v := reflect.ValueOf(any(n.Val))
if v.Kind() == reflect.Bool {
return v.Bool()
}
s := strings.ToLower(strings.TrimSpace(fmt.Sprint(n.Val)))
return s == "true" || s == "t" || s == "1" || s == "yes" || s == "on"
}
// Time converts to time.Time or zero if invalid.
func (n SqlNull[T]) Time() time.Time {
if !n.Valid {
return time.Time{}
}
if t, ok := any(n.Val).(time.Time); ok {
return t
}
return time.Time{}
}
// UUID converts to uuid.UUID or Nil if invalid.
func (n SqlNull[T]) UUID() uuid.UUID {
if !n.Valid {
return uuid.Nil
}
if u, ok := any(n.Val).(uuid.UUID); ok {
return u
}
return uuid.Nil
}
// Type aliases for common types.
type (
SqlInt16 = SqlNull[int16]
SqlInt32 = SqlNull[int32]
SqlInt64 = SqlNull[int64]
SqlFloat64 = SqlNull[float64]
SqlBool = SqlNull[bool]
SqlString = SqlNull[string]
SqlByteArray = SqlNull[[]byte]
SqlUUID = SqlNull[uuid.UUID]
)
// SqlTimeStamp - Timestamp with custom formatting (YYYY-MM-DDTHH:MM:SS).
type SqlTimeStamp struct{ SqlNull[time.Time] }
func (t SqlTimeStamp) MarshalJSON() ([]byte, error) {
if !t.Valid || t.Val.IsZero() || t.Val.Before(time.Date(0002, 1, 1, 0, 0, 0, 0, time.UTC)) {
return []byte("null"), nil
}
return []byte(fmt.Sprintf(`"%s"`, t.Val.Format("2006-01-02T15:04:05"))), nil
}
func (t *SqlTimeStamp) UnmarshalJSON(b []byte) error {
if err := t.SqlNull.UnmarshalJSON(b); err != nil {
return err
}
if t.Valid && (t.Val.IsZero() || t.Val.Format("2006-01-02T15:04:05") == "0001-01-01T00:00:00") {
t.Valid = false
}
return nil
}
func (t SqlTimeStamp) Value() (driver.Value, error) {
if !t.Valid || t.Val.IsZero() || t.Val.Before(time.Date(0002, 1, 1, 0, 0, 0, 0, time.UTC)) {
return nil, nil
}
return t.Val.Format("2006-01-02T15:04:05"), nil
}
func SqlTimeStampNow() SqlTimeStamp {
return SqlTimeStamp{SqlNull: SqlNull[time.Time]{Val: time.Now(), Valid: true}}
}
// SqlDate - Date only (YYYY-MM-DD).
type SqlDate struct{ SqlNull[time.Time] }
func (d SqlDate) MarshalJSON() ([]byte, error) {
if !d.Valid || d.Val.IsZero() {
return []byte("null"), nil
}
s := d.Val.Format("2006-01-02")
if strings.HasPrefix(s, "0001-01-01") {
return []byte("null"), nil
}
return []byte(fmt.Sprintf(`"%s"`, s)), nil
}
func (d *SqlDate) UnmarshalJSON(b []byte) error {
if err := d.SqlNull.UnmarshalJSON(b); err != nil {
return err
}
if d.Valid && d.Val.Format("2006-01-02") <= "0001-01-01" {
d.Valid = false
}
return nil
}
func (d SqlDate) Value() (driver.Value, error) {
if !d.Valid || d.Val.IsZero() {
return nil, nil
}
s := d.Val.Format("2006-01-02")
if s <= "0001-01-01" {
return nil, nil
}
return s, nil
}
func (d SqlDate) String() string {
if !d.Valid {
return ""
}
s := d.Val.Format("2006-01-02")
if strings.HasPrefix(s, "0001-01-01") || strings.HasPrefix(s, "1800-12-31") {
return ""
}
return s
}
func SqlDateNow() SqlDate {
return SqlDate{SqlNull: SqlNull[time.Time]{Val: time.Now(), Valid: true}}
}
// SqlTime - Time only (HH:MM:SS).
type SqlTime struct{ SqlNull[time.Time] }
func (t SqlTime) MarshalJSON() ([]byte, error) {
if !t.Valid || t.Val.IsZero() {
return []byte("null"), nil
}
s := t.Val.Format("15:04:05")
if s == "00:00:00" {
return []byte("null"), nil
}
return []byte(fmt.Sprintf(`"%s"`, s)), nil
}
func (t *SqlTime) UnmarshalJSON(b []byte) error {
if err := t.SqlNull.UnmarshalJSON(b); err != nil {
return err
}
if t.Valid && t.Val.Format("15:04:05") == "00:00:00" {
t.Valid = false
}
return nil
}
func (t SqlTime) Value() (driver.Value, error) {
if !t.Valid || t.Val.IsZero() {
return nil, nil
}
return t.Val.Format("15:04:05"), nil
}
func (t SqlTime) String() string {
if !t.Valid {
return ""
}
return t.Val.Format("15:04:05")
}
func SqlTimeNow() SqlTime {
return SqlTime{SqlNull: SqlNull[time.Time]{Val: time.Now(), Valid: true}}
}
// SqlJSONB - Nullable JSONB as []byte.
type SqlJSONB []byte
// Scan implements sql.Scanner.
func (n *SqlJSONB) Scan(value any) error {
if value == nil {
*n = nil
return nil
}
switch v := value.(type) {
case string:
*n = []byte(v)
case []byte:
*n = v
default:
dat, err := json.Marshal(value)
if err != nil {
return fmt.Errorf("failed to marshal value to JSON: %v", err)
}
*n = dat
}
return nil
}
// Value implements driver.Valuer.
func (n SqlJSONB) Value() (driver.Value, error) {
if len(n) == 0 {
return nil, nil
}
var js any
if err := json.Unmarshal(n, &js); err != nil {
return nil, fmt.Errorf("invalid JSON: %v", err)
}
return string(n), nil
}
// MarshalJSON implements json.Marshaler.
func (n SqlJSONB) MarshalJSON() ([]byte, error) {
if len(n) == 0 {
return []byte("null"), nil
}
var obj any
if err := json.Unmarshal(n, &obj); err != nil {
return []byte("null"), nil
}
return n, nil
}
// UnmarshalJSON implements json.Unmarshaler.
func (n *SqlJSONB) UnmarshalJSON(b []byte) error {
s := strings.TrimSpace(string(b))
if s == "null" || s == "" || (!strings.HasPrefix(s, "{") && !strings.HasPrefix(s, "[")) {
*n = nil
return nil
}
*n = b
return nil
}
func (n SqlJSONB) AsMap() (map[string]any, error) {
if len(n) == 0 {
return nil, nil
}
js := make(map[string]any)
if err := json.Unmarshal(n, &js); err != nil {
return nil, fmt.Errorf("invalid JSON: %v", err)
}
return js, nil
}
func (n SqlJSONB) AsSlice() ([]any, error) {
if len(n) == 0 {
return nil, nil
}
js := make([]any, 0)
if err := json.Unmarshal(n, &js); err != nil {
return nil, fmt.Errorf("invalid JSON: %v", err)
}
return js, nil
}
// TryIfInt64 tries to parse any value to int64 with default.
func TryIfInt64(v any, def int64) int64 {
switch val := v.(type) {
case string:
i, err := strconv.ParseInt(val, 10, 64)
if err != nil {
return def
}
return i
case int:
return int64(val)
case int8:
return int64(val)
case int16:
return int64(val)
case int32:
return int64(val)
case int64:
return val
case uint:
return int64(val)
case uint8:
return int64(val)
case uint16:
return int64(val)
case uint32:
return int64(val)
case uint64:
return int64(val)
case float32:
return int64(val)
case float64:
return int64(val)
case []byte:
i, err := strconv.ParseInt(string(val), 10, 64)
if err != nil {
return def
}
return i
default:
return def
}
}
// Constructor helpers - clean and fast value creation
func Null[T any](v T, valid bool) SqlNull[T] {
return SqlNull[T]{Val: v, Valid: valid}
}
func NewSql[T any](value any) SqlNull[T] {
n := SqlNull[T]{}
if value == nil {
return n
}
// Fast path: exact match
if v, ok := value.(T); ok {
n.Val = v
n.Valid = true
return n
}
// Try from another SqlNull
if sn, ok := value.(SqlNull[T]); ok {
return sn
}
// Convert via string
_ = n.FromString(fmt.Sprintf("%v", value))
return n
}
func NewSqlInt16(v int16) SqlInt16 {
return SqlInt16{Val: v, Valid: true}
}
func NewSqlInt32(v int32) SqlInt32 {
return SqlInt32{Val: v, Valid: true}
}
func NewSqlInt64(v int64) SqlInt64 {
return SqlInt64{Val: v, Valid: true}
}
func NewSqlFloat64(v float64) SqlFloat64 {
return SqlFloat64{Val: v, Valid: true}
}
func NewSqlBool(v bool) SqlBool {
return SqlBool{Val: v, Valid: true}
}
func NewSqlString(v string) SqlString {
return SqlString{Val: v, Valid: true}
}
func NewSqlByteArray(v []byte) SqlByteArray {
return SqlByteArray{Val: v, Valid: true}
}
func NewSqlUUID(v uuid.UUID) SqlUUID {
return SqlUUID{Val: v, Valid: true}
}
func NewSqlTimeStamp(v time.Time) SqlTimeStamp {
return SqlTimeStamp{SqlNull: SqlNull[time.Time]{Val: v, Valid: true}}
}
func NewSqlDate(v time.Time) SqlDate {
return SqlDate{SqlNull: SqlNull[time.Time]{Val: v, Valid: true}}
}
func NewSqlTime(v time.Time) SqlTime {
return SqlTime{SqlNull: SqlNull[time.Time]{Val: v, Valid: true}}
}