change package name

This commit is contained in:
2026-08-08 21:23:02 +02:00
parent 6fb4db0180
commit 4caff16e01
5 changed files with 71 additions and 70 deletions
+1 -1
View File
@@ -1,4 +1,4 @@
package orm package lasebuche
import ( import (
"fmt" "fmt"
+1 -1
View File
@@ -1,4 +1,4 @@
package orm package lasebuche
import ( import (
"database/sql" "database/sql"
+57 -57
View File
@@ -1,4 +1,4 @@
package orm package lasebuche
import ( import (
"database/sql" "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 // We need to pass pointers to each field in the order of the SQL columns
ptrVal := reflect.ValueOf(&obj).Elem() ptrVal := reflect.ValueOf(&obj).Elem()
fields := make([]any, len(t.fields)) fields := make([]any, len(t.fields))
for i, f := range t.fields { for i, f := range t.fields {
fieldValue := ptrVal.FieldByName(f.name) fieldValue := ptrVal.FieldByName(f.name)
if !fieldValue.IsValid() { if !fieldValue.IsValid() {
return nil, fmt.Errorf("field %s not found in struct", f.name) 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 all fields, pass the address so Scan can set the value
// For pointer fields (e.g., *string), this gives **T // For pointer fields (e.g., *string), this gives **T
// For value 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 var obj T
ptrVal := reflect.ValueOf(&obj).Elem() ptrVal := reflect.ValueOf(&obj).Elem()
// Create a slice to hold all values, using string for time fields // Create a slice to hold all values, using string for time fields
tempFields := make([]any, len(t.fields)) tempFields := make([]any, len(t.fields))
fieldInfos := make([]struct { fieldInfos := make([]struct {
index int index int
fieldValue reflect.Value fieldValue reflect.Value
isTime bool isTime bool
isTimePtr bool isTimePtr bool
}, len(t.fields)) }, len(t.fields))
for i, f := range t.fields { for i, f := range t.fields {
fieldValue := ptrVal.FieldByName(f.name) fieldValue := ptrVal.FieldByName(f.name)
if !fieldValue.IsValid() { if !fieldValue.IsValid() {
return nil, fmt.Errorf("field %s not found in struct", f.name) return nil, fmt.Errorf("field %s not found in struct", f.name)
} }
// Check if this is a time.Time or *time.Time field // Check if this is a time.Time or *time.Time field
isTime := fieldValue.Type() == reflect.TypeOf(time.Time{}) isTime := fieldValue.Type() == reflect.TypeOf(time.Time{})
isTimePtr := fieldValue.Kind() == reflect.Ptr && fieldValue.Type().Elem() == reflect.TypeOf(time.Time{}) isTimePtr := fieldValue.Kind() == reflect.Ptr && fieldValue.Type().Elem() == reflect.TypeOf(time.Time{})
fieldInfos[i] = struct { fieldInfos[i] = struct {
index int index int
fieldValue reflect.Value fieldValue reflect.Value
isTime bool isTime bool
isTimePtr bool isTimePtr bool
}{ }{
index: i, index: i,
fieldValue: fieldValue, fieldValue: fieldValue,
isTime: isTime, isTime: isTime,
isTimePtr: isTimePtr, isTimePtr: isTimePtr,
} }
// For time fields, use a string to capture the raw value // For time fields, use a string to capture the raw value
if isTime || isTimePtr { if isTime || isTimePtr {
var s sql.NullString var s sql.NullString
@@ -169,7 +169,7 @@ func (t Table[T]) getWithTimeHandling(id string) (*T, error) {
if err := row.Scan(tempFields...); err != nil { if err := row.Scan(tempFields...); err != nil {
return nil, err return nil, err
} }
// Copy values back, parsing time fields // Copy values back, parsing time fields
for _, info := range fieldInfos { for _, info := range fieldInfos {
if info.isTime || info.isTimePtr { if info.isTime || info.isTimePtr {
@@ -206,11 +206,11 @@ func (t Table[T]) SelectOne(where string, args ...any) (*T, error) {
if row == nil { if row == nil {
return nil, sql.ErrNoRows return nil, sql.ErrNoRows
} }
var obj T var obj T
ptrVal := reflect.ValueOf(&obj).Elem() ptrVal := reflect.ValueOf(&obj).Elem()
fields := make([]any, len(t.fields)) fields := make([]any, len(t.fields))
for i, f := range t.fields { for i, f := range t.fields {
fieldValue := ptrVal.FieldByName(f.name) fieldValue := ptrVal.FieldByName(f.name)
if !fieldValue.IsValid() { if !fieldValue.IsValid() {
@@ -235,7 +235,7 @@ func (t Table[T]) selectOneWithTimeHandling(where string, args ...any) (*T, erro
if row == nil { if row == nil {
return nil, sql.ErrNoRows return nil, sql.ErrNoRows
} }
return t.scanWithTimeHandling(row) return t.scanWithTimeHandling(row)
} }
@@ -244,7 +244,7 @@ func (t Table[T]) SelectWhere(where string, args ...any) ([]*T, error) {
if where != "" { if where != "" {
query += " where " + where query += " where " + where
} }
rows, err := t.db.Query(query, args...) rows, err := t.db.Query(query, args...)
if err != nil { if err != nil {
return nil, err return nil, err
@@ -252,7 +252,7 @@ func (t Table[T]) SelectWhere(where string, args ...any) ([]*T, error) {
defer rows.Close() defer rows.Close()
var results []*T var results []*T
for rows.Next() { for rows.Next() {
var obj T var obj T
ptrVal := reflect.ValueOf(&obj).Elem() 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() fields[i] = fieldValue.Addr().Interface()
} }
if err := rows.Scan(fields...); err != nil { if err := rows.Scan(fields...); err != nil {
if strings.Contains(err.Error(), "unsupported Scan, storing driver.Value type string into type *time.Time") { if strings.Contains(err.Error(), "unsupported Scan, storing driver.Value type string into type *time.Time") {
// Need to re-execute with time handling // 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) results = append(results, &obj)
} }
if err := rows.Err(); err != nil { if err := rows.Err(); err != nil {
return nil, err return nil, err
} }
@@ -288,7 +288,7 @@ func (t Table[T]) selectWhereWithTimeHandling(where string, args ...any) ([]*T,
if where != "" { if where != "" {
query += " where " + where query += " where " + where
} }
rows, err := t.db.Query(query, args...) rows, err := t.db.Query(query, args...)
if err != nil { if err != nil {
return nil, err return nil, err
@@ -296,7 +296,7 @@ func (t Table[T]) selectWhereWithTimeHandling(where string, args ...any) ([]*T,
defer rows.Close() defer rows.Close()
var results []*T var results []*T
for rows.Next() { for rows.Next() {
obj, err := t.scanWithTimeHandling(rows) obj, err := t.scanWithTimeHandling(rows)
if err != nil { if err != nil {
@@ -304,7 +304,7 @@ func (t Table[T]) selectWhereWithTimeHandling(where string, args ...any) ([]*T,
} }
results = append(results, obj) results = append(results, obj)
} }
if err := rows.Err(); err != nil { if err := rows.Err(); err != nil {
return nil, err 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) { func (t Table[T]) scanWithTimeHandling(scanner any) (*T, error) {
var obj T var obj T
ptrVal := reflect.ValueOf(&obj).Elem() ptrVal := reflect.ValueOf(&obj).Elem()
// Create a slice to hold all values, using string for time fields // Create a slice to hold all values, using string for time fields
tempFields := make([]any, len(t.fields)) tempFields := make([]any, len(t.fields))
fieldInfos := make([]struct { fieldInfos := make([]struct {
index int index int
fieldValue reflect.Value fieldValue reflect.Value
isTime bool isTime bool
isTimePtr bool isTimePtr bool
}, len(t.fields)) }, len(t.fields))
for i, f := range t.fields { for i, f := range t.fields {
fieldValue := ptrVal.FieldByName(f.name) fieldValue := ptrVal.FieldByName(f.name)
if !fieldValue.IsValid() { if !fieldValue.IsValid() {
return nil, fmt.Errorf("field %s not found in struct", f.name) return nil, fmt.Errorf("field %s not found in struct", f.name)
} }
// Check if this is a time.Time or *time.Time field // Check if this is a time.Time or *time.Time field
isTime := fieldValue.Type() == reflect.TypeOf(time.Time{}) isTime := fieldValue.Type() == reflect.TypeOf(time.Time{})
var isTimePtr bool var isTimePtr bool
if fieldValue.Kind() == reflect.Ptr { if fieldValue.Kind() == reflect.Ptr {
isTimePtr = fieldValue.Type().Elem() == reflect.TypeOf(time.Time{}) isTimePtr = fieldValue.Type().Elem() == reflect.TypeOf(time.Time{})
} }
fieldInfos[i] = struct { fieldInfos[i] = struct {
index int index int
fieldValue reflect.Value fieldValue reflect.Value
isTime bool isTime bool
isTimePtr bool isTimePtr bool
}{ }{
index: i, index: i,
fieldValue: fieldValue, fieldValue: fieldValue,
isTime: isTime, isTime: isTime,
isTimePtr: isTimePtr, isTimePtr: isTimePtr,
} }
// For time fields, use a string to capture the raw value // For time fields, use a string to capture the raw value
if isTime || isTimePtr { if isTime || isTimePtr {
var s sql.NullString var s sql.NullString
@@ -369,11 +369,11 @@ func (t Table[T]) scanWithTimeHandling(scanner any) (*T, error) {
default: default:
return nil, fmt.Errorf("unsupported scanner type") return nil, fmt.Errorf("unsupported scanner type")
} }
if err != nil { if err != nil {
return nil, err return nil, err
} }
// Copy values back, parsing time fields // Copy values back, parsing time fields
for _, info := range fieldInfos { for _, info := range fieldInfos {
if info.isTime || info.isTimePtr { 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") return nil, fmt.Errorf("VersionId field not found")
} }
oldVersionId := versionField.String() oldVersionId := versionField.String()
// Generate new VersionId for optimistic locking // Generate new VersionId for optimistic locking
newVersionId := GenID() newVersionId := GenID()
versionField.SetString(newVersionId) versionField.SetString(newVersionId)
@@ -501,7 +501,7 @@ func (t Table[T]) Update(value *T) (*T, error) {
// Build the values slice for the prepared statement // Build the values slice for the prepared statement
// UPDATE SQL expects: SET fields (excluding ID), then WHERE: ID and VersionId (old version) // UPDATE SQL expects: SET fields (excluding ID), then WHERE: ID and VersionId (old version)
vals := make([]any, 0, len(t.fields)) vals := make([]any, 0, len(t.fields))
// Add all fields except ID for SET clause (includes new VersionId) // Add all fields except ID for SET clause (includes new VersionId)
for _, f := range t.fields { for _, f := range t.fields {
if f.dbname == "id" { 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 // Add ID and OLD VersionId at the end for WHERE clause
idField := val.FieldByName("ID") idField := val.FieldByName("ID")
if !idField.IsValid() { if !idField.IsValid() {
return nil, fmt.Errorf("ID field not found") return nil, fmt.Errorf("ID field not found")
} }
vals = append(vals, idField.Interface(), oldVersionId) vals = append(vals, idField.Interface(), oldVersionId)
result, err := t.db.Exec(t.updatesql, vals...) result, err := t.db.Exec(t.updatesql, vals...)
if err != nil { if err != nil {
return nil, err return nil, err
} }
// Check if any row was affected // Check if any row was affected
rowsAffected, err := result.RowsAffected() rowsAffected, err := result.RowsAffected()
if err != nil { if err != nil {
return nil, err return nil, err
} }
if rowsAffected == 0 { if rowsAffected == 0 {
return nil, fmt.Errorf("no rows affected: record not found or version mismatch") return nil, fmt.Errorf("no rows affected: record not found or version mismatch")
} }
+10 -10
View File
@@ -1,4 +1,4 @@
package orm_test package lasebuche_test
import ( import (
"database/sql" "database/sql"
@@ -6,7 +6,7 @@ import (
"testing" "testing"
"time" "time"
orm "gitea.trankilou.fr/fabien/lasebuche" "gitea.trankilou.fr/fabien/lasebuche"
_ "modernc.org/sqlite" _ "modernc.org/sqlite"
) )
@@ -33,8 +33,8 @@ func TestHelloName(t *testing.T) {
} }
defer db.Close() defer db.Close()
dialect := orm.NewSqliteDialect() dialect := lasebuche.NewSqliteDialect()
tableUser, err := orm.NewTable[User](db, dialect, User{}, "user") tableUser, err := lasebuche.NewTable[User](db, dialect, User{}, "user")
if err != nil { if err != nil {
t.Error(err.Error()) t.Error(err.Error())
} }
@@ -149,8 +149,8 @@ func TestCRUDOperations(t *testing.T) {
} }
defer db.Close() defer db.Close()
dialect := orm.NewSqliteDialect() dialect := lasebuche.NewSqliteDialect()
tableUser, err := orm.NewTable[User](db, dialect, User{}, "user") tableUser, err := lasebuche.NewTable[User](db, dialect, User{}, "user")
if err != nil { if err != nil {
t.Fatal(err.Error()) t.Fatal(err.Error())
} }
@@ -300,8 +300,8 @@ func TestEdgeCases(t *testing.T) {
} }
defer db.Close() defer db.Close()
dialect := orm.NewSqliteDialect() dialect := lasebuche.NewSqliteDialect()
tableUser, err := orm.NewTable[User](db, dialect, User{}, "user") tableUser, err := lasebuche.NewTable[User](db, dialect, User{}, "user")
if err != nil { if err != nil {
t.Fatal(err.Error()) t.Fatal(err.Error())
} }
@@ -374,8 +374,8 @@ func TestVersionLocking(t *testing.T) {
} }
defer db.Close() defer db.Close()
dialect := orm.NewSqliteDialect() dialect := lasebuche.NewSqliteDialect()
tableUser, err := orm.NewTable[User](db, dialect, User{}, "user") tableUser, err := lasebuche.NewTable[User](db, dialect, User{}, "user")
if err != nil { if err != nil {
t.Fatal(err.Error()) t.Fatal(err.Error())
} }
+2 -1
View File
@@ -1,8 +1,9 @@
package orm package lasebuche
import ( import (
"crypto/rand" "crypto/rand"
"encoding/hex" "encoding/hex"
"github.com/sixafter/nanoid" "github.com/sixafter/nanoid"
) )