diff --git a/dbconnection.go b/dbconnection.go index 0db9357..c2066d7 100644 --- a/dbconnection.go +++ b/dbconnection.go @@ -36,6 +36,10 @@ func (c *DBConnection) Prepare(sql string) (*sql.Stmt, error) { return c.db.Prepare(sql) } +func (c *DBConnection) Begin() (*sql.Tx, error) { + return c.db.Begin() +} + func (c *DBConnection) Close() (bool, error) { err := c.db.Close() if err != nil { diff --git a/exec/exec.go b/exec/exec.go index 09889fd..3d46575 100644 --- a/exec/exec.go +++ b/exec/exec.go @@ -10,6 +10,7 @@ import ( log "gitlab.com/gdulai/simpleloglvl" ) +// Created and executes the DDL created by [simpleorm.ORM] func ExecuteDDL(conn *simpleorm.DBConnection, orm *simpleorm.ORM) { schema, err := orm.CreateSchema() if err != nil { @@ -24,14 +25,143 @@ func ExecuteDDL(conn *simpleorm.DBConnection, orm *simpleorm.ORM) { } } +// Wraps and represents a DB transaction. +// Allows for all at once or separate execution of [exec.Exec] implementations. +type Transaction struct { + conn *simpleorm.DBConnection + executions []*Exec[any] + tx *sql.Tx + finished bool +} + +// Creates a new [exec.Transaction]. +// Can be initialized with a set of [exec.Exec]s. +func NewTransaction(conn *simpleorm.DBConnection, executions ...*Exec[any]) *Transaction { + return &Transaction{conn: conn, executions: executions, finished: false} +} + +// The executions assigned to this transaction. +func (t *Transaction) Executions() []*Exec[any] { + return t.executions +} + +// Executes all the [exec.Exec] implementations assigned to this transactions. +// Begins, commits or rollbacks the transaction. +// This is a terminal operation, the [exec.Transaction] is considered finished after calling this. +func (t *Transaction) ExecuteAtOnce() error { + if t.finished { + return errors.New("Transaction already finished.") + } + + if t.tx == nil { + tx, err := t.conn.Begin() + t.tx = tx + if err != nil { + t.finished = true + return err + } + } + + for _, exec := range t.executions { + err := (*exec).execute(nil, t.tx) + if err != nil { + rollbackErr := t.Rollback() + if rollbackErr != nil { + return rollbackErr + } + return err + } + + } + + err := t.tx.Commit() + if err != nil { + return err + } + + t.finished = true + return nil +} + +// Executes and assigns the passed [exec.Exec]s to the Transaction. +// Calls rollback in case of an error. +// This is NOT a terminal operation, transactions still has to be committed. +func (t *Transaction) Execute(execs ...Exec[any]) error { + if t.finished { + return errors.New("Transaction already finished.") + } + + for _, exec := range execs { + execPtr := &exec + t.executions = append(t.executions, execPtr) + if t.tx == nil { + tx, err := t.conn.Begin() + t.tx = tx + if err != nil { + t.finished = true + return err + } + } + + err := (*execPtr).execute(nil, t.tx) + if err != nil { + rollbackErr := t.Rollback() + if rollbackErr != nil { + return rollbackErr + } + return err + } + + } + + return nil +} + +// Finishes the transaction, commits the changes. +// This is terminal operation. +func (t *Transaction) Finish() error { + if t.finished { + return errors.New("Transaction already finished.") + } + + if t.tx == nil { + return errors.New("Transaction is nil.") + } + + t.tx.Commit() + t.finished = true + return nil +} + +// Rollbacks and aborts the transaction. +// This is a terminal operation. +func (t *Transaction) Rollback() error { + if t.finished { + return errors.New("Transaction already finished.") + } + + if t.tx == nil { + return errors.New("Transaction is nil.") + } + + rollbackErr := t.tx.Rollback() + if rollbackErr != nil { + return rollbackErr + } + t.finished = true + return nil +} + type Exec[T any] interface { - Execute(conn *simpleorm.DBConnection) ([]T, error) + Execute(conn *simpleorm.DBConnection) error + execute(conn *simpleorm.DBConnection, tx *sql.Tx) error } type Select[T any] struct { target schema.Table whereStmt string args []any + results []T } func CreateSelect[T any](orm *simpleorm.ORM, whereStmt string, args ...any) (Select[T], error) { @@ -43,10 +173,20 @@ func CreateSelect[T any](orm *simpleorm.ORM, whereStmt string, args ...any) (Sel return Select[T]{target: *table, whereStmt: whereStmt, args: args}, nil } -func (s Select[T]) Execute(conn *simpleorm.DBConnection) ([]T, error) { +func (s *Select[T]) Results() []T { + return s.results +} + +func (s *Select[T]) Execute(conn *simpleorm.DBConnection) error { + return s.execute(conn, nil) +} + +func (s *Select[T]) execute(conn *simpleorm.DBConnection, tx *sql.Tx) error { + // Reinit the results, new execution + s.results = []T{} dml, err := s.target.GetSelectDML() if err != nil { - return []T{}, err + return err } if s.whereStmt != "" { @@ -55,9 +195,15 @@ func (s Select[T]) Execute(conn *simpleorm.DBConnection) ([]T, error) { log.LogDebug("Preparing sql: %s, with args: %s", dml, s.args) - stmt, err := conn.Prepare(dml) + var stmt *sql.Stmt + if tx != nil { + stmt, err = tx.Prepare(dml) + } else { + stmt, err = conn.Prepare(dml) + } + if err != nil { - return nil, err + return err } defer stmt.Close() @@ -81,40 +227,21 @@ func (s Select[T]) Execute(conn *simpleorm.DBConnection) ([]T, error) { } if err != nil { - return nil, err + return 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) + s.results, err = readRows[T](s.target, rows) + if err != nil { + return err } - return results, nil + return nil } type Insert[T any] struct { target schema.Table toInsert []T + results []T } func NewInsert[T any](orm *simpleorm.ORM, toInsert ...T) (Insert[T], error) { @@ -126,10 +253,18 @@ func NewInsert[T any](orm *simpleorm.ORM, toInsert ...T) (Insert[T], error) { return Insert[T]{target: *table, toInsert: toInsert}, nil } -func (ins Insert[T]) Execute(conn *simpleorm.DBConnection) ([]T, error) { +func (ins *Insert[T]) Results() []T { + return ins.results +} + +func (ins *Insert[T]) Execute(conn *simpleorm.DBConnection) error { + return ins.execute(conn, nil) +} + +func (ins *Insert[T]) execute(conn *simpleorm.DBConnection, tx *sql.Tx) error { dml, err := ins.target.GetInsertDML(len(ins.toInsert)) if err != nil { - return []T{}, err + return err } var params []any @@ -145,21 +280,26 @@ func (ins Insert[T]) Execute(conn *simpleorm.DBConnection) ([]T, error) { log.LogInfo("%s [%s]", dml, params) - stmt, err := conn.Prepare(dml) - if err != nil { - return []T{}, err + var stmt *sql.Stmt + if tx != nil { + stmt, err = tx.Prepare(dml) + } else { + stmt, err = conn.Prepare(dml) } defer stmt.Close() - result, err := stmt.Exec(params...) + rows, err := stmt.Query(params...) + if err != nil { - } else { - rowsAffected, _ := result.RowsAffected() - log.LogInfo("Inserted %s row", rowsAffected) - lastInsertId, _ := result.LastInsertId() - log.LogInfo("Last ID: %s", lastInsertId) + return err } - return []T{}, nil + + ins.results, err = readRows[T](ins.target, rows) + if err != nil { + return err + } + + return nil } type Update[T any] struct { @@ -176,17 +316,21 @@ func NewUpdate[T any](orm *simpleorm.ORM, toUpdate T) (Update[T], error) { return Update[T]{target: *table, toUpdate: toUpdate}, nil } -func (u Update[T]) Execute(conn *simpleorm.DBConnection) ([]T, error) { +func (u Update[T]) Execute(conn *simpleorm.DBConnection) error { + return u.execute(conn, nil) +} + +func (u Update[T]) execute(conn *simpleorm.DBConnection, tx *sql.Tx) error { dml, err := u.target.GetUpdateDML() if err != nil { - return []T{}, err + return err } params := prepareParams(u.toUpdate, u.target) pkCols, err := getPk(u.toUpdate, u.target) if err != nil { - return []T{}, err + return err } // Put the pks back at the end for _, pk := range pkCols { @@ -195,20 +339,22 @@ func (u Update[T]) Execute(conn *simpleorm.DBConnection) ([]T, error) { log.LogInfo("%s [%s]", dml, params) - stmt, err := conn.Prepare(dml) - if err != nil { - return []T{}, err + var stmt *sql.Stmt + if tx != nil { + stmt, err = tx.Prepare(dml) + } else { + stmt, err = conn.Prepare(dml) } defer stmt.Close() result, err := stmt.Exec(params...) if err != nil { - return []T{}, err + return err } else { rowsAffected, _ := result.RowsAffected() log.LogInfo("Updated %s row", rowsAffected) } - return []T{u.toUpdate}, nil + return nil } type Delete[T any] struct { @@ -225,11 +371,15 @@ func NewDelete[T any](orm *simpleorm.ORM, toDelete []T) (Delete[T], error) { return Delete[T]{target: *table, toDelete: toDelete}, nil } -func (d Delete[T]) Execute(conn *simpleorm.DBConnection) ([]T, error) { +func (d *Delete[T]) Execute(conn *simpleorm.DBConnection) error { + return d.execute(conn, nil) +} + +func (d *Delete[T]) execute(conn *simpleorm.DBConnection, tx *sql.Tx) error { count := len(d.toDelete) dml, err := d.target.GetDeleteDML(count) if err != nil { - return []T{}, err + return err } var params []any @@ -237,7 +387,7 @@ func (d Delete[T]) Execute(conn *simpleorm.DBConnection) ([]T, error) { for _, del := range d.toDelete { pkCols, err := getPk(del, d.target) if err != nil { - return []T{}, err + return err } for _, pk := range pkCols { @@ -247,21 +397,23 @@ func (d Delete[T]) Execute(conn *simpleorm.DBConnection) ([]T, error) { log.LogInfo("%s [%s]", dml, params) - stmt, err := conn.Prepare(dml) - if err != nil { - return []T{}, err + var stmt *sql.Stmt + if tx != nil { + stmt, err = tx.Prepare(dml) + } else { + stmt, err = conn.Prepare(dml) } defer stmt.Close() result, err := stmt.Exec(params...) if err != nil { - return []T{}, err + return err } else { rowsAffected, _ := result.RowsAffected() log.LogInfo("Deleted %s row", rowsAffected) } - return d.toDelete, nil + return nil } @@ -322,3 +474,31 @@ func getPk(src any, t schema.Table) ([]any, error) { return values, nil } + +func readRows[T any](table schema.Table, rows *sql.Rows) ([]T, error) { + rowContainer := createSelectResultContainer(table) + + 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 { + col := table.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 +} diff --git a/repository/repository.go b/repository/repository.go index 2799381..f14a08b 100644 --- a/repository/repository.go +++ b/repository/repository.go @@ -29,15 +29,16 @@ func (r *Repository[T]) SelectAll() []*T { return []*T{} } - selecRes, err := selectExec.Execute(r.conn) + err = selectExec.Execute(r.conn) if err != nil { log.LogError("Failed to load projects: %s", err) return []*T{} } - res := make([]*T, 0, len(selecRes)) - for i := range selecRes { - res = append(res, &selecRes[i]) + selectRes := selectExec.Results() + res := make([]*T, 0, len(selectRes)) + for i := range selectRes { + res = append(res, &selectRes[i]) } return res @@ -73,55 +74,57 @@ func (r *Repository[T]) SelectByPk(pks ...any) *T { return nil } - res, err := selectExec.Execute(r.conn) + err = selectExec.Execute(r.conn) if err != nil { log.LogError("Failed to load entity: %s", err) return nil } - - if len(res) != 1 { - log.LogError("Invalid result number of results: %s", len(res)) + selectRes := selectExec.Results() + if len(selectRes) != 1 { + log.LogError("Invalid result number of results: %s", len(selectRes)) return nil } - return &res[0] + return &selectRes[0] } -func (r *Repository[T]) Save(entity T) { +func (r *Repository[T]) Save(entity *T) error { // Insert - if entity.IsInsertable() { - insertExec, err := exec.NewInsert[T](r.orm, entity) + if (*entity).IsInsertable() { + insertExec, err := exec.NewInsert[T](r.orm, *entity) if err != nil { - log.LogError("Failed to create insert execution: %s", err) - return + return err } - if _, err := insertExec.Execute(r.conn); err != nil { - log.LogError("Insert failed for entity: %s", entity) + err = insertExec.Execute(r.conn) + if err != nil { + return err } - - return + result := insertExec.Results() + // Point to the new inserted result + *entity = result[0] + return nil } - updateExec, err := exec.NewUpdate[T](r.orm, entity) + updateExec, err := exec.NewUpdate[T](r.orm, *entity) if err != nil { - log.LogError("Failed to create update execution: %s", err) - return + return err } - _, err = updateExec.Execute(r.conn) + err = updateExec.Execute(r.conn) if err != nil { - log.LogError("Update failed: %s", err) - return + return err } + + return nil } -func (r *Repository[T]) Delete(entity T) { - deleteExec, err := exec.NewDelete[T](r.orm, []T{entity}) +func (r *Repository[T]) Delete(entity *T) { + deleteExec, err := exec.NewDelete[T](r.orm, []T{*entity}) if err != nil { log.LogError("Failed to create delete execution: %s", err) return } - _, err = deleteExec.Execute(r.conn) + err = deleteExec.Execute(r.conn) if err != nil { log.LogError("Delete failed: %s", err) return diff --git a/schema/table.go b/schema/table.go index 9ad1ff8..76a0416 100644 --- a/schema/table.go +++ b/schema/table.go @@ -100,6 +100,8 @@ func (t Table) GetInsertDML(count int) (string, error) { } } + dml.WriteString("RETURNING *") + return dml.String(), nil } diff --git a/test/exec_test.go b/test/exec_test.go index 99a9b04..76d1b31 100644 --- a/test/exec_test.go +++ b/test/exec_test.go @@ -1,48 +1,36 @@ package simpleorm_test import ( - "os" "testing" "time" - simpleorm "git.gdulai.com/gdulai/simpleorm" "git.gdulai.com/gdulai/simpleorm/exec" log "gitlab.com/gdulai/simpleloglvl" ) -func cleanUp(dbFile string, conn *simpleorm.DBConnection) { - conn.Close() - os.Remove(dbFile) -} - func TestSelectEmpty(t *testing.T) { // GIVEN - conn := simpleorm.OpenConnection("sqlite3", "file:test.db") + orm, conn := testSetup() defer cleanUp("test.db", conn) - orm := simpleorm.NewORM(Test{}, TestWithFk{}, TestWithFkAndFkId{}, - TestWithCompositePk{}, TestWithCompositeFk{}) - - exec.ExecuteDDL(conn, orm) - // WHEN selectExec, err := exec.CreateSelect[Test](orm, "") // THEN - if err != nil { log.LogError("Failed to create select. %s", err) t.Fail() return } - res, err := selectExec.Execute(conn) - + err = selectExec.Execute(conn) if err != nil { log.LogError("Select failure. %s", err) t.Fail() return } + + res := selectExec.Results() if len(res) > 0 { log.LogError("Expected empty result.") t.Fail() @@ -50,17 +38,108 @@ func TestSelectEmpty(t *testing.T) { } } -func TestInsertAndSelectSingle(t *testing.T) { +func TestInsert(t *testing.T) { // GIVEN - conn := simpleorm.OpenConnection("sqlite3", "file:test.db") + orm, conn := testSetup() defer cleanUp("test.db", conn) - orm := simpleorm.NewORM(Test{}, TestWithFk{}, TestWithFkAndFkId{}, - TestWithCompositePk{}, TestWithCompositeFk{}) - - exec.ExecuteDDL(conn, orm) testObj := Test{Int64Field: 54, IntField: 12, StringField: "asdf"} + // WHEN + insertExec, err := exec.NewInsert(orm, testObj) + if err != nil { + log.LogError("TestInsertAndSelectSingle setup failed: %s", err) + t.Fail() + return + } + + insertErr := insertExec.Execute(conn) + if err != nil { + log.LogError("TestInsertAndSelectSingle insert failed: %s", err) + t.Fail() + return + } + insertRes := insertExec.Results() + + // THEN + if insertErr != nil { + log.LogError("TestInsertAndSelectSingle failure. %s", err) + t.Fail() + return + } + + if len(insertRes) != 1 { + log.LogError("TestInsertAndSelectSingle test failed. Expected: 1, Actual: %s", len(insertRes)) + t.Fail() + return + } + + singleRes := insertRes[0] + if singleRes.ID != 1 || singleRes.Int64Field != 54 || singleRes.IntField != 12 || singleRes.StringField != "asdf" { + log.LogError("TestInsertAndSelectSingle invalid result. Expected: %s, Actual: %s", testObj, singleRes) + t.Fail() + } +} + +func TestInsertMultiple(t *testing.T) { + // GIVEN + orm, conn := testSetup() + defer cleanUp("test.db", conn) + + testObj := Test{Int64Field: 54, IntField: 12, StringField: "asdf"} + testObjSecond := Test{Int64Field: 12, IntField: 54, StringField: "fdsa"} + + // WHEN + insertExec, err := exec.NewInsert[Test](orm, testObj, testObjSecond) + if err != nil { + log.LogError("TestInsertAndSelectSingle setup failed: %s", err) + t.Fail() + return + } + + insertErr := insertExec.Execute(conn) + if err != nil { + log.LogError("TestInsertAndSelectSingle insert failed: %s", err) + t.Fail() + return + } + insertRes := insertExec.Results() + + // THEN + if insertErr != nil { + log.LogError("TestInsertAndSelectSingle failure. %s", err) + t.Fail() + return + } + + if len(insertRes) != 2 { + log.LogError("TestInsertAndSelectSingle test failed. Expected: 1, Actual: %s", len(insertRes)) + t.Fail() + return + } + + entity := insertRes[0] + if entity.ID != 1 || entity.Int64Field != 54 || entity.IntField != 12 || entity.StringField != "asdf" { + log.LogError("TestInsertAndSelectSingle invalid result. Expected: %s, Actual: %s", testObj, entity) + t.Fail() + } + + entity = insertRes[1] + if entity.ID != 2 || entity.Int64Field != 12 || entity.IntField != 54 || entity.StringField != "fdsa" { + log.LogError("TestInsertAndSelectSingle invalid result. Expected: %s, Actual: %s", testObjSecond, entity) + t.Fail() + } + +} + +func TestInsertAndSelectSingle(t *testing.T) { + // GIVEN + orm, conn := testSetup() + defer cleanUp("test.db", conn) + + testObj := Test{Int64Field: 54, IntField: 12, StringField: "asdf"} + + // WHEN insertExec, err := exec.NewInsert(orm, testObj) if err != nil { log.LogError("TestInsertAndSelectSingle setup failed: %s", err) @@ -75,29 +154,43 @@ func TestInsertAndSelectSingle(t *testing.T) { return } - // WHEN - _, err = insertExec.Execute(conn) - if err != nil { + insertErr := insertExec.Execute(conn) + selectErr := selectExec.Execute(conn) + + // THEN + if insertErr != nil { log.LogError("TestInsertAndSelectSingle insert failed: %s", err) t.Fail() return } - res, err := selectExec.Execute(conn) - - // THEN - if err != nil { + if selectErr != nil { log.LogError("TestInsertAndSelectSingle failure. %s", err) t.Fail() return } - if len(res) != 1 { - log.LogError("TestInsertAndSelectSingle test failed. Expected: 1, Actual: %s", len(res)) + + insertRes := insertExec.Results() + if len(insertRes) != 1 { + log.LogError("TestInsertAndSelectSingle test failed. Expected: 1, Actual: %s", len(insertRes)) t.Fail() return } - singleRes := res[0] + selectRes := selectExec.Results() + if len(selectRes) != 1 { + log.LogError("TestInsertAndSelectSingle test failed. Expected: 1, Actual: %s", len(selectRes)) + t.Fail() + return + } + + singleRes := insertRes[0] + if singleRes.ID != 1 || singleRes.Int64Field != 54 || singleRes.IntField != 12 || singleRes.StringField != "asdf" { + log.LogError("TestInsertAndSelectSingle invalid result. Expected: %s, Actual: %s", testObj, singleRes) + t.Fail() + } + + singleRes = selectRes[0] if singleRes.ID != 1 || singleRes.Int64Field != 54 || singleRes.IntField != 12 || singleRes.StringField != "asdf" { log.LogError("TestInsertAndSelectSingle invalid result. Expected: %s, Actual: %s", testObj, singleRes) t.Fail() @@ -107,15 +200,12 @@ func TestInsertAndSelectSingle(t *testing.T) { func TestInsertSelectWithCompositePk(t *testing.T) { // GIVEN - conn := simpleorm.OpenConnection("sqlite3", "file:test.db") + orm, conn := testSetup() defer cleanUp("test.db", conn) - orm := simpleorm.NewORM(Test{}, TestWithFk{}, TestWithFkAndFkId{}, - TestWithCompositePk{}, TestWithCompositeFk{}) - - exec.ExecuteDDL(conn, orm) testObj := TestWithCompositePk{ID: 12, Name: "Test"} + // WHEN insertExec, err := exec.NewInsert(orm, testObj) if err != nil { log.LogError("TestInsertSelectWithCompositePk setup failed: %s", err) @@ -130,15 +220,14 @@ func TestInsertSelectWithCompositePk(t *testing.T) { return } - // WHEN - _, err = insertExec.Execute(conn) + err = insertExec.Execute(conn) if err != nil { log.LogError("TestInsertSelectWithCompositePk insert failed: %s", err) t.Fail() return } - res, err := selectExec.Execute(conn) + err = selectExec.Execute(conn) // THEN if err != nil { @@ -146,6 +235,8 @@ func TestInsertSelectWithCompositePk(t *testing.T) { t.Fail() return } + + res := selectExec.Results() if len(res) != 1 { log.LogError("TestInsertSelectWithCompositePk test failed. Expected: 1, Actual: %s", len(res)) t.Fail() @@ -162,14 +253,10 @@ func TestInsertSelectWithCompositePk(t *testing.T) { func TestInsertMultipleAndSelectWithParam(t *testing.T) { // GIVEN - conn := simpleorm.OpenConnection("sqlite3", "file:test.db") + orm, conn := testSetup() defer cleanUp("test.db", conn) - orm := simpleorm.NewORM(Test{}, TestWithFk{}, TestWithFkAndFkId{}, - TestWithCompositePk{}, TestWithCompositeFk{}) - - exec.ExecuteDDL(conn, orm) - + // WHEN testObj := Test{Int64Field: 54, IntField: 12, StringField: "asdf"} testObjSecond := Test{Int64Field: 12, IntField: 54, StringField: "fdsa"} @@ -194,23 +281,21 @@ func TestInsertMultipleAndSelectWithParam(t *testing.T) { return } - // WHEN - - _, err = insertExec.Execute(conn) + err = insertExec.Execute(conn) if err != nil { log.LogError("TestInsertMultipleAndSelectWithParam first insert failed: %s", err) t.Fail() return } - _, err = insertSecond.Execute(conn) + err = insertSecond.Execute(conn) if err != nil { log.LogError("TestInsertMultipleAndSelectWithParam second insert failed: %s", err) t.Fail() return } - res, err := selectExec.Execute(conn) + err = selectExec.Execute(conn) // THEN if err != nil { @@ -218,6 +303,8 @@ func TestInsertMultipleAndSelectWithParam(t *testing.T) { t.Fail() return } + + res := selectExec.Results() if len(res) != 1 { log.LogError(" TestInsertMultipleAndSelectWithParam test failed. Expected: 1, Actual: %s", len(res)) t.Fail() @@ -234,15 +321,12 @@ func TestInsertMultipleAndSelectWithParam(t *testing.T) { func TestInsertAndUpdate(t *testing.T) { // GIVEN - conn := simpleorm.OpenConnection("sqlite3", "file:test.db") + orm, conn := testSetup() defer cleanUp("test.db", conn) - orm := simpleorm.NewORM(Test{}, TestWithFk{}, TestWithFkAndFkId{}, - TestWithCompositePk{}, TestWithCompositeFk{}) - - exec.ExecuteDDL(conn, orm) testObj := Test{Int64Field: 54, IntField: 12, StringField: "asdf"} + // WHEN insertExec, err := exec.NewInsert(orm, testObj) if err != nil { log.LogError("TestInsertAndUpdate setup failed: %s", err) @@ -257,22 +341,21 @@ func TestInsertAndUpdate(t *testing.T) { return } - // WHEN - - _, err = insertExec.Execute(conn) + err = insertExec.Execute(conn) if err != nil { log.LogError("TestInsertAndUpdate insert failed: %s", err) t.Fail() return } - res, err := selectExec.Execute(conn) + err = selectExec.Execute(conn) if err != nil { log.LogError("TestInsertAndUpdate re-select failed: %s", err) t.Fail() return } + res := selectExec.Results() testObj = res[0] testObj.IntField = 42069 @@ -285,8 +368,8 @@ func TestInsertAndUpdate(t *testing.T) { return } - _, err = updateExec.Execute(conn) - res, selectErr := selectExec.Execute(conn) + err = updateExec.Execute(conn) + selectErr := selectExec.Execute(conn) // THEN if err != nil { @@ -301,6 +384,7 @@ func TestInsertAndUpdate(t *testing.T) { return } + res = selectExec.Results() if len(res) != 1 { log.LogError("TestInsertAndUpdate test failed. Expected: 1, Actual: %s", len(res)) t.Fail() @@ -317,15 +401,12 @@ func TestInsertAndUpdate(t *testing.T) { func TestInsertAndDelete(t *testing.T) { // GIVEN - conn := simpleorm.OpenConnection("sqlite3", "file:test.db") + orm, conn := testSetup() defer cleanUp("test.db", conn) - orm := simpleorm.NewORM(Test{}, TestWithFk{}, TestWithFkAndFkId{}, - TestWithCompositePk{}, TestWithCompositeFk{}) - - exec.ExecuteDDL(conn, orm) testObj := Test{Int64Field: 54, IntField: 12, StringField: "asdf"} + // WHEN insertExec, err := exec.NewInsert(orm, testObj) if err != nil { log.LogError("TestInsertAndDelete setup failed: %s", err) @@ -333,7 +414,7 @@ func TestInsertAndDelete(t *testing.T) { return } - _, err = insertExec.Execute(conn) + err = insertExec.Execute(conn) if err != nil { log.LogError("TestInsertAndDelete setup failed: %s", err) t.Fail() @@ -347,18 +428,15 @@ func TestInsertAndDelete(t *testing.T) { return } - selectRes, err := selectExec.Execute(conn) - + err = selectExec.Execute(conn) if err != nil { log.LogError("TestInsertAndDelete setup failed. %s", err) t.Fail() return } - deleteExec, err := exec.NewDelete(orm, selectRes) - - // WHEN - _, err = deleteExec.Execute(conn) + deleteExec, err := exec.NewDelete(orm, selectExec.Results()) + err = deleteExec.Execute(conn) // THEN if err != nil { @@ -367,7 +445,15 @@ func TestInsertAndDelete(t *testing.T) { return } - selectRes, err = selectExec.Execute(conn) + err = selectExec.Execute(conn) + // THEN + if err != nil { + log.LogError("TestInsertAndDelete delete failure. %s", err) + t.Fail() + return + } + + selectRes := selectExec.Results() if len(selectRes) != 0 { log.LogError("TestInsertAndDelete test failed. Expected: 0, Actual: %s", len(selectRes)) t.Fail() @@ -378,14 +464,12 @@ func TestInsertAndDelete(t *testing.T) { func TestBoolInsertAndSelectSingle(t *testing.T) { // GIVEN - conn := simpleorm.OpenConnection("sqlite3", "file:test.db") + orm, conn := testSetup() defer cleanUp("test.db", conn) - orm := simpleorm.NewORM(TestWithBool{}) - - exec.ExecuteDDL(conn, orm) testObj := TestWithBool{BoolField: true} + // WHEN insertExec, err := exec.NewInsert(orm, testObj) if err != nil { log.LogError("TestBoolInsertAndSelectSingle setup failed: %s", err) @@ -400,15 +484,14 @@ func TestBoolInsertAndSelectSingle(t *testing.T) { return } - // WHEN - _, err = insertExec.Execute(conn) + err = insertExec.Execute(conn) if err != nil { log.LogError("TestBoolInsertAndSelectSingle insert failed: %s", err) t.Fail() return } - res, err := selectExec.Execute(conn) + err = selectExec.Execute(conn) // THEN if err != nil { @@ -416,6 +499,8 @@ func TestBoolInsertAndSelectSingle(t *testing.T) { t.Fail() return } + + res := selectExec.Results() if len(res) != 1 { log.LogError("TestBoolInsertAndSelectSingle test failed. Expected: 1, Actual: %s", len(res)) t.Fail() @@ -432,15 +517,13 @@ func TestBoolInsertAndSelectSingle(t *testing.T) { func TestTimeInsertAndSelectSingle(t *testing.T) { // GIVEN - conn := simpleorm.OpenConnection("sqlite3", "file:test.db") + orm, conn := testSetup() defer cleanUp("test.db", conn) - orm := simpleorm.NewORM(TestWithTime{}) - - exec.ExecuteDDL(conn, orm) now := time.Now().UnixMilli() testObj := TestWithTime{TimeField: now} + // WHEN insertExec, err := exec.NewInsert(orm, testObj) if err != nil { log.LogError("TestTimeInsertAndSelectSingle setup failed: %s", err) @@ -455,15 +538,14 @@ func TestTimeInsertAndSelectSingle(t *testing.T) { return } - // WHEN - _, err = insertExec.Execute(conn) + err = insertExec.Execute(conn) if err != nil { log.LogError("TestTimeInsertAndSelectSingle insert failed: %s", err) t.Fail() return } - res, err := selectExec.Execute(conn) + err = selectExec.Execute(conn) // THEN if err != nil { @@ -471,6 +553,8 @@ func TestTimeInsertAndSelectSingle(t *testing.T) { t.Fail() return } + + res := selectExec.Results() if len(res) != 1 { log.LogError("TestTimeInsertAndSelectSingle test failed. Expected: 1, Actual: %s", len(res)) t.Fail() @@ -484,3 +568,116 @@ func TestTimeInsertAndSelectSingle(t *testing.T) { } } + +func TestTransactionExecuteSimple(t *testing.T) { + // GIVEN + orm, conn := testSetup() + defer cleanUp("test.db", conn) + + transaction := exec.NewTransaction(conn) + + testObj := Test{Int64Field: 54, IntField: 12, StringField: "asdf"} + + // WHEN + insertExec, err := exec.NewInsert[Test](orm, testObj) + if err != nil { + log.LogError("TestTransactionExecuteSimple failed: %s", err) + t.Fail() + return + } + + transErr := transaction.Execute(&insertExec) + if transErr != nil { + transErr = transaction.Finish() + } + + // THEN + if transErr != nil { + log.LogError("TestTransactionExecuteSimple failed: %s", err) + t.Fail() + return + } + + insertRes := insertExec.Results() + + if len(insertRes) != 1 { + log.LogError("TestInsertAndSelectSingle test failed. Expected: 1, Actual: %s", len(insertRes)) + t.Fail() + return + } + + singleRes := insertRes[0] + if singleRes.ID != 1 || singleRes.Int64Field != 54 || singleRes.IntField != 12 || singleRes.StringField != "asdf" { + log.LogError("TestInsertAndSelectSingle invalid result. Expected: %s, Actual: %s", testObj, singleRes) + t.Fail() + } +} + +func TestTransactionRollback(t *testing.T) { + // GIVEN + orm, conn := testSetup() + defer cleanUp("test.db", conn) + + transaction := exec.NewTransaction(conn) + + testObj := Test{Int64Field: 54, IntField: 12, StringField: "asdf"} + + // WHEN + insertExec, err := exec.NewInsert[Test](orm, testObj) + if err != nil { + log.LogError("TestTransactionRollback WHEN failed: %s", err) + t.Fail() + return + } + + transErr := transaction.Execute(&insertExec) + if transErr != nil { + log.LogError("TestTransactionRollback WHEN failed: %s", err) + t.Fail() + return + } + + selectExec, err := exec.CreateSelect[Test](orm, "id = ?", insertExec.Results()[0].ID) + if err != nil { + log.LogError("TestTransactionRollback WHEN failed: %s", err) + t.Fail() + return + } + + transErr = transaction.Execute(&selectExec) + if transErr != nil { + log.LogError("TestTransactionRollback WHEN failed: %s", err) + t.Fail() + return + } + + beforeRollbackSelectRes := selectExec.Results() + rollbackErr := transaction.Rollback() + afterRollbackErr := selectExec.Execute(conn) + afterRollbackSelectRes := selectExec.Results() + + // THEN + if len(beforeRollbackSelectRes) != 1 { + log.LogError("TestTransactionExecuteSimple-THEN: Expected 1, Actual: %s", len(beforeRollbackSelectRes)) + t.Fail() + return + } + + if rollbackErr != nil { + log.LogError("TestTransactionExecuteSimple-THEN: rollback error: %s", rollbackErr) + t.Fail() + return + } + + if afterRollbackErr != nil { + log.LogError("TestTransactionExecuteSimple-THEN: after rollback select error: %s", rollbackErr) + t.Fail() + return + } + + if len(afterRollbackSelectRes) != 0 { + log.LogError("TestTransactionExecuteSimple-THEN: Expected 0 after rollback select result") + t.Fail() + return + } +} diff --git a/test/main_test.go b/test/main_test.go index eb0680c..6d36dba 100644 --- a/test/main_test.go +++ b/test/main_test.go @@ -4,6 +4,8 @@ import ( "os" "testing" + simpleorm "git.gdulai.com/gdulai/simpleorm" + "git.gdulai.com/gdulai/simpleorm/exec" _ "github.com/mattn/go-sqlite3" log "gitlab.com/gdulai/simpleloglvl" ) @@ -61,3 +63,18 @@ func TestMain(m *testing.M) { os.Exit(code) } + +func testSetup() (*simpleorm.ORM, *simpleorm.DBConnection) { + conn := simpleorm.OpenConnection("sqlite3", "file:test.db") + + orm := simpleorm.NewORM(Test{}, TestWithFk{}, TestWithFkAndFkId{}, + TestWithCompositePk{}, TestWithCompositeFk{}, TestWithBool{}, TestWithTime{}) + exec.ExecuteDDL(conn, orm) + + return orm, conn +} + +func cleanUp(dbFile string, conn *simpleorm.DBConnection) { + conn.Close() + os.Remove(dbFile) +} diff --git a/test/repository_test.go b/test/repository_test.go index b5f8d90..c80b386 100644 --- a/test/repository_test.go +++ b/test/repository_test.go @@ -3,26 +3,14 @@ package simpleorm_test import ( "testing" - simpleorm "git.gdulai.com/gdulai/simpleorm" "git.gdulai.com/gdulai/simpleorm/exec" "git.gdulai.com/gdulai/simpleorm/repository" log "gitlab.com/gdulai/simpleloglvl" ) -func repoTestSetup() (*simpleorm.ORM, *simpleorm.DBConnection) { - conn := simpleorm.OpenConnection("sqlite3", "file:test.db") - - orm := simpleorm.NewORM(Test{}, TestWithFk{}, TestWithFkAndFkId{}, - TestWithCompositePk{}, TestWithCompositeFk{}) - - exec.ExecuteDDL(conn, orm) - - return orm, conn -} - func TestRepoSelectAll(t *testing.T) { // GIVEN - orm, conn := repoTestSetup() + orm, conn := testSetup() defer cleanUp("test.db", conn) testObj := Test{Int64Field: 54, IntField: 12, StringField: "asdf"} @@ -32,7 +20,7 @@ func TestRepoSelectAll(t *testing.T) { t.Fail() return } - _, err = insertExec.Execute(conn) + err = insertExec.Execute(conn) repo := repository.NewRepository[Test](conn, orm) // WHEN @@ -48,7 +36,7 @@ func TestRepoSelectAll(t *testing.T) { func TestRepoSelectByPk(t *testing.T) { // GIVEN - orm, conn := repoTestSetup() + orm, conn := testSetup() defer cleanUp("test.db", conn) testObj := Test{Int64Field: 54, IntField: 12, StringField: "asdf"} @@ -58,7 +46,7 @@ func TestRepoSelectByPk(t *testing.T) { t.Fail() return } - _, err = insertExec.Execute(conn) + err = insertExec.Execute(conn) repo := repository.NewRepository[Test](conn, orm) // WHEN @@ -74,7 +62,7 @@ func TestRepoSelectByPk(t *testing.T) { func TestRepoSelectByCompundPk(t *testing.T) { // GIVEN - orm, conn := repoTestSetup() + orm, conn := testSetup() defer cleanUp("test.db", conn) testObj := TestWithCompositePk{ID: 42069, Name: "Test"} @@ -84,7 +72,7 @@ func TestRepoSelectByCompundPk(t *testing.T) { t.Fail() return } - _, err = insertExec.Execute(conn) + err = insertExec.Execute(conn) repo := repository.NewRepository[TestWithCompositePk](conn, orm) // WHEN @@ -100,46 +88,45 @@ func TestRepoSelectByCompundPk(t *testing.T) { func TestRepoSave(t *testing.T) { // GIVEN - orm, conn := repoTestSetup() + orm, conn := testSetup() defer cleanUp("test.db", conn) - testObj := Test{Int64Field: 54, IntField: 12, StringField: "asdf"} + testObj := &Test{Int64Field: 54, IntField: 12, StringField: "asdf"} repo := repository.NewRepository[Test](conn, orm) // WHEN - repo.Save(testObj) + err := repo.Save(testObj) // THEN - res := repo.SelectByPk(1) - if res == nil { - log.LogError("TestRepoSave test failed. Result is nil!") + if err != nil { + log.LogError("TestRepoSave failed: %s", err) + } + if testObj.ID == 0 { + log.LogError("TestRepoSave failed! testObj.ID is 0!") t.Fail() - return } } func TestRepoDelete(t *testing.T) { // GIVEN - orm, conn := repoTestSetup() + orm, conn := testSetup() defer cleanUp("test.db", conn) - testObj := Test{Int64Field: 54, IntField: 12, StringField: "asdf"} + testObj := &Test{Int64Field: 54, IntField: 12, StringField: "asdf"} repo := repository.NewRepository[Test](conn, orm) - repo.Save(testObj) - - res := repo.SelectByPk(1) - if res == nil { + err := repo.Save(testObj) + if err != nil { log.LogError("TestRepoDelete setup failed! Test data not saved!") t.Fail() return } // WHEN - repo.Delete(*res) + repo.Delete(testObj) // THEN - res = repo.SelectByPk(1) + res := repo.SelectByPk(testObj.ID) if res != nil { log.LogError("TestRepoDelete test failed. Result is not nil!") t.Fail()