552 lines
14 KiB
Go
552 lines
14 KiB
Go
package lasebuche
|
|
|
|
import (
|
|
"database/sql"
|
|
"fmt"
|
|
"reflect"
|
|
"strings"
|
|
"time"
|
|
)
|
|
|
|
type dbfield struct {
|
|
name string
|
|
stype string
|
|
dbname string
|
|
}
|
|
|
|
type Dialect interface {
|
|
TableExists(db *sql.DB, tableName string) (bool, error)
|
|
ColumnSpec(f dbfield) (string, error)
|
|
}
|
|
|
|
type Table[T any] struct {
|
|
db *sql.DB
|
|
dialect Dialect
|
|
tablename string
|
|
fields []dbfield
|
|
insertsql string
|
|
updatesql string
|
|
selectbyid string
|
|
deletebyid string
|
|
selectwhere string
|
|
deletewhere string
|
|
}
|
|
|
|
func NewTable[T any](db *sql.DB, dialect Dialect, sample T, tablename string) (Table[T], error) {
|
|
t := reflect.TypeOf(sample)
|
|
fields := make([]dbfield, 0)
|
|
for field := range t.Fields() {
|
|
f := dbfield{
|
|
name: field.Name,
|
|
stype: field.Type.String(),
|
|
dbname: field.Tag.Get("db"),
|
|
}
|
|
fmt.Println(f.name, f.stype, f.dbname)
|
|
fields = append(fields, f)
|
|
}
|
|
|
|
return Table[T]{
|
|
db: db,
|
|
dialect: dialect,
|
|
tablename: tablename,
|
|
fields: fields,
|
|
insertsql: buildInsertSql(tablename, fields, dialect),
|
|
updatesql: buildUpdateSql(tablename, fields, dialect),
|
|
selectbyid: buildSelectbyidSql(tablename, fields, dialect),
|
|
deletebyid: buildDeletebyidSql(tablename, fields, dialect),
|
|
selectwhere: buildSelectWhere(tablename, fields, dialect),
|
|
deletewhere: buildDeleteWhere(tablename, fields, dialect),
|
|
}, nil
|
|
}
|
|
|
|
func (t Table[T]) Sync() error {
|
|
if exists, err := t.dialect.TableExists(t.db, t.tablename); exists {
|
|
return nil
|
|
} else if err != nil {
|
|
return err
|
|
}
|
|
|
|
columns := make([]string, len(t.fields))
|
|
for i, f := range t.fields {
|
|
c, err := t.dialect.ColumnSpec(f)
|
|
if err != nil {
|
|
return err
|
|
}
|
|
columns[i] = c
|
|
}
|
|
sql := fmt.Sprintf("create table %s (%s)", t.tablename, strings.Join(columns, ","))
|
|
_, err := t.db.Exec(sql)
|
|
if err != nil {
|
|
return err
|
|
}
|
|
return nil
|
|
}
|
|
|
|
func (t Table[T]) Get(id string) (*T, error) {
|
|
row := t.db.QueryRow(t.selectbyid, id)
|
|
if row == nil {
|
|
return nil, sql.ErrNoRows
|
|
}
|
|
var obj T
|
|
|
|
// Use a slice of pointers to struct fields for scanning
|
|
// We need to pass pointers to each field in the order of the SQL columns
|
|
ptrVal := reflect.ValueOf(&obj).Elem()
|
|
fields := make([]any, len(t.fields))
|
|
|
|
for i, f := range t.fields {
|
|
fieldValue := ptrVal.FieldByName(f.name)
|
|
if !fieldValue.IsValid() {
|
|
return nil, fmt.Errorf("field %s not found in struct", f.name)
|
|
}
|
|
|
|
// For all fields, pass the address so Scan can set the value
|
|
// For pointer fields (e.g., *string), this gives **T
|
|
// For value fields (e.g., string), this gives *T
|
|
fields[i] = fieldValue.Addr().Interface()
|
|
}
|
|
|
|
if err := row.Scan(fields...); err != nil {
|
|
// Handle time scanning manually for SQLite which returns dates as strings
|
|
if strings.Contains(err.Error(), "unsupported Scan, storing driver.Value type string into type *time.Time") {
|
|
return t.getWithTimeHandling(id)
|
|
}
|
|
return nil, err
|
|
}
|
|
|
|
return &obj, nil
|
|
}
|
|
|
|
func (t Table[T]) getWithTimeHandling(id string) (*T, error) {
|
|
row := t.db.QueryRow(t.selectbyid, id)
|
|
if row == nil {
|
|
return nil, sql.ErrNoRows
|
|
}
|
|
var obj T
|
|
|
|
ptrVal := reflect.ValueOf(&obj).Elem()
|
|
|
|
// Create a slice to hold all values, using string for time fields
|
|
tempFields := make([]any, len(t.fields))
|
|
fieldInfos := make([]struct {
|
|
index int
|
|
fieldValue reflect.Value
|
|
isTime bool
|
|
isTimePtr bool
|
|
}, len(t.fields))
|
|
|
|
for i, f := range t.fields {
|
|
fieldValue := ptrVal.FieldByName(f.name)
|
|
if !fieldValue.IsValid() {
|
|
return nil, fmt.Errorf("field %s not found in struct", f.name)
|
|
}
|
|
|
|
// Check if this is a time.Time or *time.Time field
|
|
isTime := fieldValue.Type() == reflect.TypeOf(time.Time{})
|
|
isTimePtr := fieldValue.Kind() == reflect.Ptr && fieldValue.Type().Elem() == reflect.TypeOf(time.Time{})
|
|
|
|
fieldInfos[i] = struct {
|
|
index int
|
|
fieldValue reflect.Value
|
|
isTime bool
|
|
isTimePtr bool
|
|
}{
|
|
index: i,
|
|
fieldValue: fieldValue,
|
|
isTime: isTime,
|
|
isTimePtr: isTimePtr,
|
|
}
|
|
|
|
// For time fields, use a string to capture the raw value
|
|
if isTime || isTimePtr {
|
|
var s sql.NullString
|
|
tempFields[i] = &s
|
|
} else {
|
|
tempFields[i] = fieldValue.Addr().Interface()
|
|
}
|
|
}
|
|
|
|
if err := row.Scan(tempFields...); err != nil {
|
|
return nil, err
|
|
}
|
|
|
|
// Copy values back, parsing time fields
|
|
for _, info := range fieldInfos {
|
|
if info.isTime || info.isTimePtr {
|
|
s := tempFields[info.index].(*sql.NullString)
|
|
if s.Valid {
|
|
t, err := time.Parse(time.RFC3339, s.String)
|
|
if err != nil {
|
|
// Try other common formats
|
|
t, err = time.Parse("2006-01-02 15:04:05", s.String)
|
|
if err != nil {
|
|
t, err = time.Parse("2006-01-02T15:04:05Z", s.String)
|
|
if err != nil {
|
|
continue // Skip if we can't parse
|
|
}
|
|
}
|
|
}
|
|
if info.isTime {
|
|
info.fieldValue.Set(reflect.ValueOf(t))
|
|
} else if info.isTimePtr {
|
|
ptr := reflect.New(info.fieldValue.Type().Elem())
|
|
ptr.Elem().Set(reflect.ValueOf(t))
|
|
info.fieldValue.Set(ptr)
|
|
}
|
|
}
|
|
}
|
|
}
|
|
|
|
return &obj, nil
|
|
}
|
|
|
|
func (t Table[T]) SelectOne(where string, args ...any) (*T, error) {
|
|
query := t.selectwhere + " where " + where + " limit 1"
|
|
row := t.db.QueryRow(query, args...)
|
|
if row == nil {
|
|
return nil, sql.ErrNoRows
|
|
}
|
|
|
|
var obj T
|
|
ptrVal := reflect.ValueOf(&obj).Elem()
|
|
fields := make([]any, len(t.fields))
|
|
|
|
for i, f := range t.fields {
|
|
fieldValue := ptrVal.FieldByName(f.name)
|
|
if !fieldValue.IsValid() {
|
|
return nil, fmt.Errorf("field %s not found in struct", f.name)
|
|
}
|
|
fields[i] = fieldValue.Addr().Interface()
|
|
}
|
|
|
|
if err := row.Scan(fields...); err != nil {
|
|
if strings.Contains(err.Error(), "unsupported Scan, storing driver.Value type string into type *time.Time") {
|
|
return t.selectOneWithTimeHandling(where, args...)
|
|
}
|
|
return nil, err
|
|
}
|
|
|
|
return &obj, nil
|
|
}
|
|
|
|
func (t Table[T]) selectOneWithTimeHandling(where string, args ...any) (*T, error) {
|
|
query := t.selectwhere + " where " + where + " limit 1"
|
|
row := t.db.QueryRow(query, args...)
|
|
if row == nil {
|
|
return nil, sql.ErrNoRows
|
|
}
|
|
|
|
return t.scanWithTimeHandling(row)
|
|
}
|
|
|
|
func (t Table[T]) SelectWhere(where string, args ...any) ([]*T, error) {
|
|
query := t.selectwhere
|
|
if where != "" {
|
|
query += " where " + where
|
|
}
|
|
|
|
rows, err := t.db.Query(query, args...)
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
defer rows.Close()
|
|
|
|
var results []*T
|
|
|
|
for rows.Next() {
|
|
var obj T
|
|
ptrVal := reflect.ValueOf(&obj).Elem()
|
|
fields := make([]any, len(t.fields))
|
|
for i, f := range t.fields {
|
|
fieldValue := ptrVal.FieldByName(f.name)
|
|
if !fieldValue.IsValid() {
|
|
return nil, fmt.Errorf("field %s not found in struct", f.name)
|
|
}
|
|
fields[i] = fieldValue.Addr().Interface()
|
|
}
|
|
|
|
if err := rows.Scan(fields...); err != nil {
|
|
if strings.Contains(err.Error(), "unsupported Scan, storing driver.Value type string into type *time.Time") {
|
|
// Need to re-execute with time handling
|
|
rows.Close()
|
|
return t.selectWhereWithTimeHandling(where, args...)
|
|
}
|
|
return nil, err
|
|
}
|
|
results = append(results, &obj)
|
|
}
|
|
|
|
if err := rows.Err(); err != nil {
|
|
return nil, err
|
|
}
|
|
|
|
return results, nil
|
|
}
|
|
|
|
func (t Table[T]) selectWhereWithTimeHandling(where string, args ...any) ([]*T, error) {
|
|
query := t.selectwhere
|
|
if where != "" {
|
|
query += " where " + where
|
|
}
|
|
|
|
rows, err := t.db.Query(query, args...)
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
defer rows.Close()
|
|
|
|
var results []*T
|
|
|
|
for rows.Next() {
|
|
obj, err := t.scanWithTimeHandling(rows)
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
results = append(results, obj)
|
|
}
|
|
|
|
if err := rows.Err(); err != nil {
|
|
return nil, err
|
|
}
|
|
|
|
return results, nil
|
|
}
|
|
|
|
func (t Table[T]) scanWithTimeHandling(scanner any) (*T, error) {
|
|
var obj T
|
|
|
|
ptrVal := reflect.ValueOf(&obj).Elem()
|
|
|
|
// Create a slice to hold all values, using string for time fields
|
|
tempFields := make([]any, len(t.fields))
|
|
fieldInfos := make([]struct {
|
|
index int
|
|
fieldValue reflect.Value
|
|
isTime bool
|
|
isTimePtr bool
|
|
}, len(t.fields))
|
|
|
|
for i, f := range t.fields {
|
|
fieldValue := ptrVal.FieldByName(f.name)
|
|
if !fieldValue.IsValid() {
|
|
return nil, fmt.Errorf("field %s not found in struct", f.name)
|
|
}
|
|
|
|
// Check if this is a time.Time or *time.Time field
|
|
isTime := fieldValue.Type() == reflect.TypeOf(time.Time{})
|
|
var isTimePtr bool
|
|
if fieldValue.Kind() == reflect.Ptr {
|
|
isTimePtr = fieldValue.Type().Elem() == reflect.TypeOf(time.Time{})
|
|
}
|
|
|
|
fieldInfos[i] = struct {
|
|
index int
|
|
fieldValue reflect.Value
|
|
isTime bool
|
|
isTimePtr bool
|
|
}{
|
|
index: i,
|
|
fieldValue: fieldValue,
|
|
isTime: isTime,
|
|
isTimePtr: isTimePtr,
|
|
}
|
|
|
|
// For time fields, use a string to capture the raw value
|
|
if isTime || isTimePtr {
|
|
var s sql.NullString
|
|
tempFields[i] = &s
|
|
} else {
|
|
tempFields[i] = fieldValue.Addr().Interface()
|
|
}
|
|
}
|
|
|
|
var err error
|
|
switch s := scanner.(type) {
|
|
case *sql.Row:
|
|
err = s.Scan(tempFields...)
|
|
case *sql.Rows:
|
|
err = s.Scan(tempFields...)
|
|
default:
|
|
return nil, fmt.Errorf("unsupported scanner type")
|
|
}
|
|
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
|
|
// Copy values back, parsing time fields
|
|
for _, info := range fieldInfos {
|
|
if info.isTime || info.isTimePtr {
|
|
s := tempFields[info.index].(*sql.NullString)
|
|
if s.Valid {
|
|
t, err := time.Parse(time.RFC3339, s.String)
|
|
if err != nil {
|
|
// Try other common formats
|
|
t, err = time.Parse("2006-01-02 15:04:05", s.String)
|
|
if err != nil {
|
|
t, err = time.Parse("2006-01-02T15:04:05Z", s.String)
|
|
if err != nil {
|
|
// Try SQLite format: YYYY-MM-DD HH:MM:SS
|
|
t, err = time.Parse("2006-01-02 15:04:05.999999999-07:00", s.String)
|
|
if err != nil {
|
|
// Try simpler format
|
|
t, err = time.Parse("2006-01-02 15:04:05", s.String)
|
|
if err != nil {
|
|
continue
|
|
}
|
|
}
|
|
}
|
|
}
|
|
}
|
|
if info.isTime {
|
|
info.fieldValue.Set(reflect.ValueOf(t))
|
|
} else if info.isTimePtr {
|
|
ptr := reflect.New(info.fieldValue.Type().Elem())
|
|
ptr.Elem().Set(reflect.ValueOf(t))
|
|
info.fieldValue.Set(ptr)
|
|
}
|
|
}
|
|
}
|
|
}
|
|
|
|
return &obj, nil
|
|
}
|
|
|
|
func (t Table[T]) Delete(id string) error {
|
|
_, err := t.db.Exec(t.deletebyid, id)
|
|
return err
|
|
}
|
|
|
|
func (t Table[T]) DeleteWhere(where string, args ...any) error {
|
|
query := t.deletewhere
|
|
if where != "" {
|
|
query += " where " + where
|
|
}
|
|
_, err := t.db.Exec(query, args...)
|
|
return err
|
|
}
|
|
|
|
func (t Table[T]) Insert(value *T) (*T, error) {
|
|
// Get a settable reflect value of the struct
|
|
// value is *T, &value is **T, .Elem() gives *T (settables), .Elem() gives T (settables)
|
|
val := reflect.ValueOf(&value).Elem().Elem()
|
|
|
|
// Auto-fill ID field if it exists and is empty
|
|
if fieldID := val.FieldByName("ID"); fieldID.IsValid() && fieldID.CanSet() && fieldID.Kind() == reflect.String {
|
|
if fieldID.String() == "" {
|
|
fieldID.SetString(GenID())
|
|
}
|
|
}
|
|
|
|
// Auto-fill VersionId field if it exists and is empty
|
|
if fieldVersion := val.FieldByName("VersionId"); fieldVersion.IsValid() && fieldVersion.CanSet() && fieldVersion.Kind() == reflect.String {
|
|
if fieldVersion.String() == "" {
|
|
fieldVersion.SetString(GenID())
|
|
}
|
|
}
|
|
|
|
// Auto-fill DateCreated or DateCreate field if it exists
|
|
now := time.Now()
|
|
if fieldCreate := val.FieldByName("DateCreated"); fieldCreate.IsValid() && fieldCreate.CanSet() && fieldCreate.Type() == reflect.TypeOf(now) {
|
|
fieldCreate.Set(reflect.ValueOf(now))
|
|
} else if fieldCreate := val.FieldByName("DateCreate"); fieldCreate.IsValid() && fieldCreate.CanSet() && fieldCreate.Type() == reflect.TypeOf(now) {
|
|
fieldCreate.Set(reflect.ValueOf(now))
|
|
}
|
|
|
|
vals := make([]any, len(t.fields))
|
|
for i, f := range t.fields {
|
|
fieldValue := val.FieldByName(f.name)
|
|
if !fieldValue.IsValid() {
|
|
return nil, fmt.Errorf("field %s not found in struct", f.name)
|
|
}
|
|
vals[i] = fieldValue.Interface()
|
|
}
|
|
|
|
_, err := t.db.Exec(t.insertsql, vals...)
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
|
|
return value, nil
|
|
}
|
|
|
|
func (t Table[T]) Update(value *T) (*T, error) {
|
|
// Get a settable reflect value of the struct
|
|
val := reflect.ValueOf(&value).Elem().Elem()
|
|
|
|
// Save old VersionId for WHERE clause
|
|
versionField := val.FieldByName("VersionId")
|
|
if !versionField.IsValid() {
|
|
return nil, fmt.Errorf("VersionId field not found")
|
|
}
|
|
oldVersionId := versionField.String()
|
|
|
|
// Generate new VersionId for optimistic locking
|
|
newVersionId := GenID()
|
|
versionField.SetString(newVersionId)
|
|
|
|
// Update DateUpdated field if it exists
|
|
if fieldUpdated := val.FieldByName("DateUpdated"); fieldUpdated.IsValid() && fieldUpdated.CanSet() {
|
|
if fieldUpdated.Kind() == reflect.Ptr {
|
|
now := time.Now()
|
|
if fieldUpdated.IsNil() {
|
|
fieldUpdated.Set(reflect.New(fieldUpdated.Type().Elem()))
|
|
}
|
|
fieldUpdated.Elem().Set(reflect.ValueOf(now))
|
|
} else if fieldUpdated.Type() == reflect.TypeOf(time.Time{}) {
|
|
fieldUpdated.Set(reflect.ValueOf(time.Now()))
|
|
}
|
|
}
|
|
|
|
// Build the values slice for the prepared statement
|
|
// UPDATE SQL expects: SET fields (excluding ID), then WHERE: ID and VersionId (old version)
|
|
vals := make([]any, 0, len(t.fields))
|
|
|
|
// Add all fields except ID for SET clause (includes new VersionId)
|
|
for _, f := range t.fields {
|
|
if f.dbname == "id" {
|
|
continue // Skip ID field for SET clause
|
|
}
|
|
fieldValue := val.FieldByName(f.name)
|
|
if !fieldValue.IsValid() {
|
|
return nil, fmt.Errorf("field %s not found in struct", f.name)
|
|
}
|
|
vals = append(vals, fieldValue.Interface())
|
|
}
|
|
|
|
// Add ID and OLD VersionId at the end for WHERE clause
|
|
idField := val.FieldByName("ID")
|
|
|
|
if !idField.IsValid() {
|
|
return nil, fmt.Errorf("ID field not found")
|
|
}
|
|
|
|
vals = append(vals, idField.Interface(), oldVersionId)
|
|
|
|
result, err := t.db.Exec(t.updatesql, vals...)
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
|
|
// Check if any row was affected
|
|
rowsAffected, err := result.RowsAffected()
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
|
|
if rowsAffected == 0 {
|
|
return nil, fmt.Errorf("no rows affected: record not found or version mismatch")
|
|
}
|
|
|
|
return value, nil
|
|
}
|
|
|
|
func (t Table[T]) Debug() {
|
|
fmt.Println("insert: ", t.insertsql)
|
|
fmt.Println("update: ", t.updatesql)
|
|
fmt.Println("selectbyid: ", t.selectbyid)
|
|
fmt.Println("deletebyid: ", t.deletebyid)
|
|
fmt.Println("selectwhere: ", t.selectwhere)
|
|
fmt.Println("deletewhere: ", t.deletewhere)
|
|
}
|