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

322 lines
8.1 KiB
Go

package lasebuche
import (
"database/sql"
"fmt"
"reflect"
"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)
ScanResult(row interface{ Scan(dest ...any) error }, fields []any, fieldTypes []dbfield) 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, tablename string) (Table[T], error) {
var instance T
t := reflect.TypeOf(instance)
fields := make([]dbfield, 0)
for field := range t.Fields() {
f := dbfield{
name: field.Name,
stype: field.Type.String(),
dbname: field.Tag.Get("db"),
}
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
}
_, 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
// Create field destinations
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()
}
// Use dialect to scan with proper type handling
if err := t.dialect.ScanResult(row, fields, t.fields); err != nil {
if err == sql.ErrNoRows {
return nil, sql.ErrNoRows
}
return nil, err
}
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 := t.dialect.ScanResult(row, fields, t.fields); err != nil {
if err == sql.ErrNoRows {
return nil, sql.ErrNoRows
}
return nil, err
}
return &obj, nil
}
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()
}
// Use dialect to scan with proper type handling
// *sql.Rows implements Scanner interface
if err := t.dialect.ScanResult(rows, fields, t.fields); 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]) 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)
}