diff --git a/builder.go b/builder.go index d8a4095..04bb9ab 100644 --- a/builder.go +++ b/builder.go @@ -1,4 +1,4 @@ -package orm +package lasebuche import ( "fmt" diff --git a/dialect_sqlite.go b/dialect_sqlite.go index a241076..2769ae0 100644 --- a/dialect_sqlite.go +++ b/dialect_sqlite.go @@ -1,4 +1,4 @@ -package orm +package lasebuche import ( "database/sql" diff --git a/orm.go b/orm.go index c73c994..b16b746 100644 --- a/orm.go +++ b/orm.go @@ -1,4 +1,4 @@ -package orm +package lasebuche import ( "database/sql" @@ -93,13 +93,13 @@ func (t Table[T]) Get(id string) (*T, error) { // 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 @@ -125,38 +125,38 @@ func (t Table[T]) getWithTimeHandling(id string) (*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 + 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 int + fieldValue reflect.Value + isTime bool + isTimePtr bool }{ - index: i, - fieldValue: fieldValue, - isTime: isTime, - isTimePtr: isTimePtr, + 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 @@ -169,7 +169,7 @@ func (t Table[T]) getWithTimeHandling(id string) (*T, error) { 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 { @@ -206,11 +206,11 @@ func (t Table[T]) SelectOne(where string, args ...any) (*T, error) { 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() { @@ -235,7 +235,7 @@ func (t Table[T]) selectOneWithTimeHandling(where string, args ...any) (*T, erro if row == nil { return nil, sql.ErrNoRows } - + return t.scanWithTimeHandling(row) } @@ -244,7 +244,7 @@ func (t Table[T]) SelectWhere(where string, args ...any) ([]*T, error) { if where != "" { query += " where " + where } - + rows, err := t.db.Query(query, args...) if err != nil { return nil, err @@ -252,7 +252,7 @@ func (t Table[T]) SelectWhere(where string, args ...any) ([]*T, error) { defer rows.Close() var results []*T - + for rows.Next() { var obj T ptrVal := reflect.ValueOf(&obj).Elem() @@ -264,7 +264,7 @@ func (t Table[T]) SelectWhere(where string, args ...any) ([]*T, error) { } 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 @@ -275,7 +275,7 @@ func (t Table[T]) SelectWhere(where string, args ...any) ([]*T, error) { } results = append(results, &obj) } - + if err := rows.Err(); err != nil { return nil, err } @@ -288,7 +288,7 @@ func (t Table[T]) selectWhereWithTimeHandling(where string, args ...any) ([]*T, if where != "" { query += " where " + where } - + rows, err := t.db.Query(query, args...) if err != nil { return nil, err @@ -296,7 +296,7 @@ func (t Table[T]) selectWhereWithTimeHandling(where string, args ...any) ([]*T, defer rows.Close() var results []*T - + for rows.Next() { obj, err := t.scanWithTimeHandling(rows) if err != nil { @@ -304,7 +304,7 @@ func (t Table[T]) selectWhereWithTimeHandling(where string, args ...any) ([]*T, } results = append(results, obj) } - + if err := rows.Err(); err != nil { return nil, err } @@ -314,43 +314,43 @@ func (t Table[T]) selectWhereWithTimeHandling(where string, args ...any) ([]*T, 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 + 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 int + fieldValue reflect.Value + isTime bool + isTimePtr bool }{ - index: i, - fieldValue: fieldValue, - isTime: isTime, - isTimePtr: isTimePtr, + 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 @@ -369,11 +369,11 @@ func (t Table[T]) scanWithTimeHandling(scanner any) (*T, error) { 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 { @@ -480,7 +480,7 @@ func (t Table[T]) Update(value *T) (*T, error) { return nil, fmt.Errorf("VersionId field not found") } oldVersionId := versionField.String() - + // Generate new VersionId for optimistic locking newVersionId := GenID() versionField.SetString(newVersionId) @@ -501,7 +501,7 @@ func (t Table[T]) Update(value *T) (*T, error) { // 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" { @@ -516,24 +516,24 @@ func (t Table[T]) Update(value *T) (*T, error) { // 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") } diff --git a/orm_test.go b/orm_test.go index a32ca12..5bc15ac 100644 --- a/orm_test.go +++ b/orm_test.go @@ -1,4 +1,4 @@ -package orm_test +package lasebuche_test import ( "database/sql" @@ -6,7 +6,7 @@ import ( "testing" "time" - orm "gitea.trankilou.fr/fabien/lasebuche" + "gitea.trankilou.fr/fabien/lasebuche" _ "modernc.org/sqlite" ) @@ -33,8 +33,8 @@ func TestHelloName(t *testing.T) { } defer db.Close() - dialect := orm.NewSqliteDialect() - tableUser, err := orm.NewTable[User](db, dialect, User{}, "user") + dialect := lasebuche.NewSqliteDialect() + tableUser, err := lasebuche.NewTable[User](db, dialect, User{}, "user") if err != nil { t.Error(err.Error()) } @@ -149,8 +149,8 @@ func TestCRUDOperations(t *testing.T) { } defer db.Close() - dialect := orm.NewSqliteDialect() - tableUser, err := orm.NewTable[User](db, dialect, User{}, "user") + dialect := lasebuche.NewSqliteDialect() + tableUser, err := lasebuche.NewTable[User](db, dialect, User{}, "user") if err != nil { t.Fatal(err.Error()) } @@ -300,8 +300,8 @@ func TestEdgeCases(t *testing.T) { } defer db.Close() - dialect := orm.NewSqliteDialect() - tableUser, err := orm.NewTable[User](db, dialect, User{}, "user") + dialect := lasebuche.NewSqliteDialect() + tableUser, err := lasebuche.NewTable[User](db, dialect, User{}, "user") if err != nil { t.Fatal(err.Error()) } @@ -374,8 +374,8 @@ func TestVersionLocking(t *testing.T) { } defer db.Close() - dialect := orm.NewSqliteDialect() - tableUser, err := orm.NewTable[User](db, dialect, User{}, "user") + dialect := lasebuche.NewSqliteDialect() + tableUser, err := lasebuche.NewTable[User](db, dialect, User{}, "user") if err != nil { t.Fatal(err.Error()) } diff --git a/utility.go b/utility.go index a6edf26..ad4a097 100644 --- a/utility.go +++ b/utility.go @@ -1,8 +1,9 @@ -package orm +package lasebuche import ( "crypto/rand" "encoding/hex" + "github.com/sixafter/nanoid" )