package simpleorm import ( "database/sql" "errors" "reflect" log "gitlab.com/gdulai/simpleloglvl" ) type ORMExec[T any] struct { orm *ORM conn *DBConnection } func CreateExecution[T any](orm *ORM, conn *DBConnection) ORMExec[T] { return ORMExec[T]{orm: orm, conn: conn} } func InitDB(conn *DBConnection, orm *ORM) { // Force a real DB interaction if _, err := conn.db.Exec(orm.CreateDDL()); err != nil { log.LogFatalError("%", err) } } func (ormExec ORMExec[T]) Select(desc T, whereStmt string, args ...any) ([]T, error) { lookup := reflect.TypeOf(desc).Name() table, ok := ormExec.orm.cache.Get(lookup) if !ok { return []T{}, errors.New("Failed to get descriptor from cache! " + lookup) } sqlStr := table.ToSelectDML() if whereStmt != "" { sqlStr += " WHERE " + whereStmt } log.LogDebug("Preparing sql: %s, with args: %s", sqlStr, args) stmt, err := ormExec.conn.db.Prepare(sqlStr) if err != nil { return nil, err } defer stmt.Close() log.LogDebug("Executing statement: %s", stmt) var rows *sql.Rows if len(args) == 0 { rows, err = stmt.Query() } else { // Flattent args to make sure it can be parsed correctly var flatArgs []any for _, a := range args { if s, ok := a.([]any); ok { flatArgs = append(flatArgs, s...) } else { flatArgs = append(flatArgs, a) } } rows, err = stmt.Query(flatArgs...) } if err != nil { return nil, err } rowContainer := table.createSelectResultContainer() var results []T for rows.Next() { err = rows.Scan(rowContainer...) if err != nil { return nil, err } targetType := table.Type parsedResult := reflect.New(targetType) for i, fieldVal := range rowContainer { parsedResult.Elem().Field(i).Set(reflect.Indirect(reflect.ValueOf(fieldVal))) } parsedObj := reflect.Indirect(parsedResult).Interface().(T) results = append(results, parsedObj) } return results, nil } func (ormExec ORMExec[T]) Insert(src ...T) error { typeName := reflect.TypeOf(src[0]).Name() table, ok := ormExec.orm.cache.Get(typeName) if !ok { return errors.New("Could not find table for type " + typeName) } sql := table.ToInsertDML(len(src)) var params []any for i := range len(src) { actualParams := table.prepareParams(src[i]) if len(params) == 0 { params = make([]any, len(src)*len(actualParams)) } for j := range actualParams { params[(i*len(actualParams))+j] = actualParams[j] } } log.LogInfo("%s [%s]", sql, params) stmt, err := ormExec.conn.db.Prepare(sql) if err != nil { return err } defer stmt.Close() result, err := stmt.Exec(params...) if err != nil { return err } else { rowsAffected, _ := result.RowsAffected() log.LogInfo("Inserted %s row", rowsAffected) } return nil } func (ormExec ORMExec[T]) Update(src T) error { typeName := reflect.TypeOf(src).Name() table, ok := ormExec.orm.cache.Get(reflect.TypeOf(src).Name()) if !ok { return errors.New("Could not find table for type " + typeName) } sql := table.ToUpdateDML(src) params := table.prepareParams(src) pk, err := table.getPk(src) if err != nil { return err } // Put the pk back at the end params = append(params, pk) log.LogInfo("%s [%s]", sql, params) stmt, err := ormExec.conn.db.Prepare(sql) if err != nil { return err } defer stmt.Close() result, err := stmt.Exec(params...) if err != nil { return err } else { rowsAffected, _ := result.RowsAffected() log.LogInfo("Updated %s row", rowsAffected) } return nil } func (ormExec ORMExec[T]) Delete(src T) error { typeName := reflect.TypeOf(src).Name() table, ok := ormExec.orm.cache.Get(reflect.TypeOf(src).Name()) if !ok { return errors.New("Could not find table for type " + typeName) } sql := table.ToDeleteDML(src) pk, err := table.getPk(src) if err != nil { return err } log.LogInfo("%s [%s]", sql, pk) stmt, err := ormExec.conn.db.Prepare(sql) if err != nil { return err } defer stmt.Close() result, err := stmt.Exec(pk) if err != nil { return err } else { rowsAffected, _ := result.RowsAffected() log.LogInfo("Deleted %s row", rowsAffected) } return nil }