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