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