Files
lasebuche/orm.go
T
2026-08-08 21:23:02 +02:00

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)
}