change package name
This commit is contained in:
@@ -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")
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user