package exec import ( "database/sql" "errors" "reflect" "git.gdulai.com/gdulai/simpleorm" "git.gdulai.com/gdulai/simpleorm/schema" log "gitlab.com/gdulai/simpleloglvl" ) func ExecuteDDL(conn *simpleorm.DBConnection, orm *simpleorm.ORM) { schema, err := orm.CreateSchema() if err != nil { log.LogError("Failed to run DDL: %s", schema) return } log.LogInfo("Executing DDL:\n%s", schema) if _, err := conn.Exec(schema); err != nil { log.LogFatalError("%", err) } } type Exec[T any] interface { Execute(conn *simpleorm.DBConnection) ([]T, error) } type Select[T any] struct { target schema.Table whereStmt string args []any } func CreateSelect[T any](orm *simpleorm.ORM, whereStmt string, args ...any) (Select[T], error) { table, ok := orm.Cache().Get(reflect.TypeFor[T]().Name()) if !ok { return Select[T]{}, errors.New("Failed to get table from schema cache") } return Select[T]{target: *table, whereStmt: whereStmt, args: args}, nil } func (s Select[T]) Execute(conn *simpleorm.DBConnection) ([]T, error) { dml, err := s.target.GetSelectDML() if err != nil { return []T{}, err } if s.whereStmt != "" { dml += " WHERE " + s.whereStmt } log.LogDebug("Preparing sql: %s, with args: %s", dml, s.args) stmt, err := conn.Prepare(dml) if err != nil { return nil, err } defer stmt.Close() log.LogDebug("Executing statement: %s", stmt) var rows *sql.Rows if len(s.args) == 0 { rows, err = stmt.Query() } else { // Flattent args to make sure it can be parsed correctly var flatArgs []any for _, a := range s.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 := createSelectResultContainer(s.target) var results []T for rows.Next() { err = rows.Scan(rowContainer...) if err != nil { return nil, err } targetType := s.target.Type parsedResult := reflect.New(targetType) for i, fieldVal := range rowContainer { col := s.target.Columns[i] targetField := parsedResult.Elem().Field(i) rawValue := reflect.Indirect(reflect.ValueOf(fieldVal)) decoded := reflect.ValueOf(col.Decode(targetField.Type().Name(), rawValue)) targetField.Set(decoded) } parsedObj := reflect.Indirect(parsedResult).Interface().(T) results = append(results, parsedObj) } return results, nil } type Insert[T any] struct { target schema.Table toInsert []T } func NewInsert[T any](orm *simpleorm.ORM, toInsert ...T) (Insert[T], error) { table, ok := orm.Cache().Get(reflect.TypeFor[T]().Name()) if !ok { return Insert[T]{}, errors.New("Failed to get table from schema cache") } return Insert[T]{target: *table, toInsert: toInsert}, nil } func (ins Insert[T]) Execute(conn *simpleorm.DBConnection) ([]T, error) { dml, err := ins.target.GetInsertDML(len(ins.toInsert)) if err != nil { return []T{}, err } var params []any for i := range len(ins.toInsert) { actualParams := prepareParams(ins.toInsert[i], ins.target) if len(params) == 0 { params = make([]any, len(ins.toInsert)*len(actualParams)) } for j := range actualParams { params[(i*len(actualParams))+j] = actualParams[j] } } log.LogInfo("%s [%s]", dml, params) stmt, err := conn.Prepare(dml) if err != nil { return []T{}, err } defer stmt.Close() result, err := stmt.Exec(params...) if err != nil { } else { rowsAffected, _ := result.RowsAffected() log.LogInfo("Inserted %s row", rowsAffected) lastInsertId, _ := result.LastInsertId() log.LogInfo("Last ID: %s", lastInsertId) } return []T{}, nil } type Update[T any] struct { target schema.Table toUpdate T } func NewUpdate[T any](orm *simpleorm.ORM, toUpdate T) (Update[T], error) { table, ok := orm.Cache().Get(reflect.TypeFor[T]().Name()) if !ok { return Update[T]{}, errors.New("Failed to get table from schema cache") } return Update[T]{target: *table, toUpdate: toUpdate}, nil } func (u Update[T]) Execute(conn *simpleorm.DBConnection) ([]T, error) { dml, err := u.target.GetUpdateDML() if err != nil { return []T{}, err } params := prepareParams(u.toUpdate, u.target) pkCols, err := getPk(u.toUpdate, u.target) if err != nil { return []T{}, err } // Put the pks back at the end for _, pk := range pkCols { params = append(params, pk) } log.LogInfo("%s [%s]", dml, params) stmt, err := conn.Prepare(dml) if err != nil { return []T{}, err } defer stmt.Close() result, err := stmt.Exec(params...) if err != nil { return []T{}, err } else { rowsAffected, _ := result.RowsAffected() log.LogInfo("Updated %s row", rowsAffected) } return []T{u.toUpdate}, nil } type Delete[T any] struct { target schema.Table toDelete []T } func NewDelete[T any](orm *simpleorm.ORM, toDelete []T) (Delete[T], error) { table, ok := orm.Cache().Get(reflect.TypeFor[T]().Name()) if !ok { return Delete[T]{}, errors.New("Failed to get table from schema cache") } return Delete[T]{target: *table, toDelete: toDelete}, nil } func (d Delete[T]) Execute(conn *simpleorm.DBConnection) ([]T, error) { count := len(d.toDelete) dml, err := d.target.GetDeleteDML(count) if err != nil { return []T{}, err } var params []any for _, del := range d.toDelete { pkCols, err := getPk(del, d.target) if err != nil { return []T{}, err } for _, pk := range pkCols { params = append(params, pk) } } log.LogInfo("%s [%s]", dml, params) stmt, err := conn.Prepare(dml) if err != nil { return []T{}, err } defer stmt.Close() result, err := stmt.Exec(params...) if err != nil { return []T{}, err } else { rowsAffected, _ := result.RowsAffected() log.LogInfo("Deleted %s row", rowsAffected) } return d.toDelete, nil } func createSelectResultContainer(t schema.Table) []any { vals := make([]any, t.Type.NumField()) for i := range vals { switch t.Type.Field(i).Type.Kind().String() { case "string": var fieldContainer string vals[i] = &fieldContainer case "int", "bool": var fieldContainer int vals[i] = &fieldContainer case "int64", "time.Time": var fieldContainer int64 vals[i] = &fieldContainer } } return vals } func prepareParams(src any, t schema.Table) []any { var params []any for _, col := range t.Columns { _, ok := col.Modifiers["pk"] if ok { continue } field, ok := t.Type.FieldByName(col.FieldName) if !ok { continue } fieldValue := reflect.ValueOf(src).FieldByIndex(field.Index) params = append(params, col.Encode(fieldValue)) } return params } func getPk(src any, t schema.Table) ([]any, error) { var values []any for _, constraint := range t.Constraints { if constraint.Type == "pk" { for _, col := range constraint.Columns { field, ok := t.Type.FieldByName(col.FieldName) if !ok { continue } fieldValue := reflect.ValueOf(src).FieldByIndex(field.Index) values = append(values, fieldValue.Interface()) } } } if len(values) == 0 { return nil, errors.New("Could not determine pk column!") } return values, nil }