correction des pb de space : le spaceId est maintenant dans l'url uniquement
This commit is contained in:
1 parent
c52460bac0
commit
0bb8d1cb46
41 files changed
+1463
-121
No files matched your search
@@ -0,0 +1,373 @@
|
||||
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)
|
||||
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
|
||||
}
|
||||
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
|
||||
|
||||
// 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
|
||||
}
|
||||
|
||||
type query struct {
|
||||
where string
|
||||
args []any
|
||||
orderby string
|
||||
pagenum int
|
||||
pagesize int
|
||||
}
|
||||
|
||||
type Option func(*query)
|
||||
|
||||
func WithWhere(where string, args ...any) Option {
|
||||
return func(q *query) {
|
||||
q.where = where
|
||||
q.args = args
|
||||
}
|
||||
}
|
||||
|
||||
func WithOrder(orderby string) Option {
|
||||
return func(q *query) {
|
||||
q.orderby = orderby
|
||||
}
|
||||
}
|
||||
|
||||
func WithPagination(pagenum, pagesize int) Option {
|
||||
return func(q *query) {
|
||||
q.pagenum = pagenum
|
||||
q.pagesize = pagesize
|
||||
}
|
||||
}
|
||||
|
||||
func (t Table[T]) Select(options ...Option) ([]*T, error) {
|
||||
q := &query{}
|
||||
for _, o := range options {
|
||||
o(q)
|
||||
}
|
||||
query := t.selectwhere
|
||||
if q.where != "" {
|
||||
query += " where " + q.where
|
||||
}
|
||||
if q.orderby != "" {
|
||||
query += " order by " + q.orderby
|
||||
}
|
||||
if q.pagesize != 0 {
|
||||
query += fmt.Sprintf(" limit %d", q.pagesize)
|
||||
if q.pagenum != 0 {
|
||||
query += fmt.Sprintf(" offset %d", q.pagesize*q.pagenum)
|
||||
}
|
||||
}
|
||||
|
||||
var err error
|
||||
var rows *sql.Rows
|
||||
if q.args == nil {
|
||||
rows, err = t.db.Query(query)
|
||||
} else {
|
||||
rows, err = t.db.Query(query, q.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)
|
||||
}
|
||||
Reference in new issue
Block a user