diff --git a/exec/exec.go b/exec/exec.go index 09889fd..ea77d56 100644 --- a/exec/exec.go +++ b/exec/exec.go @@ -25,13 +25,14 @@ func ExecuteDDL(conn *simpleorm.DBConnection, orm *simpleorm.ORM) { } type Exec[T any] interface { - Execute(conn *simpleorm.DBConnection) ([]T, error) + Execute(conn *simpleorm.DBConnection) error } type Select[T any] struct { target schema.Table whereStmt string args []any + result []T } func CreateSelect[T any](orm *simpleorm.ORM, whereStmt string, args ...any) (Select[T], error) { @@ -43,10 +44,14 @@ 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]) GetResult() []T { + return s.result +} + +func (s Select[T]) Execute(conn *simpleorm.DBConnection) error { dml, err := s.target.GetSelectDML() if err != nil { - return []T{}, err + return err } if s.whereStmt != "" { @@ -57,7 +62,7 @@ func (s Select[T]) Execute(conn *simpleorm.DBConnection) ([]T, error) { stmt, err := conn.Prepare(dml) if err != nil { - return nil, err + return err } defer stmt.Close() @@ -81,17 +86,15 @@ 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 + return err } targetType := s.target.Type @@ -106,15 +109,16 @@ func (s Select[T]) Execute(conn *simpleorm.DBConnection) ([]T, error) { targetField.Set(decoded) } parsedObj := reflect.Indirect(parsedResult).Interface().(T) - results = append(results, parsedObj) + s.result = append(s.result, parsedObj) } - return results, nil + return nil } type Insert[T any] struct { - target schema.Table - toInsert []T + target schema.Table + toInsert []T + lastInsertId int64 } func NewInsert[T any](orm *simpleorm.ORM, toInsert ...T) (Insert[T], error) { @@ -126,10 +130,10 @@ 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]) Execute(conn *simpleorm.DBConnection) error { dml, err := ins.target.GetInsertDML(len(ins.toInsert)) if err != nil { - return []T{}, err + return err } var params []any @@ -147,7 +151,7 @@ func (ins Insert[T]) Execute(conn *simpleorm.DBConnection) ([]T, error) { stmt, err := conn.Prepare(dml) if err != nil { - return []T{}, err + return err } defer stmt.Close() @@ -156,10 +160,14 @@ func (ins Insert[T]) Execute(conn *simpleorm.DBConnection) ([]T, error) { } else { rowsAffected, _ := result.RowsAffected() log.LogInfo("Inserted %s row", rowsAffected) - lastInsertId, _ := result.LastInsertId() - log.LogInfo("Last ID: %s", lastInsertId) + ins.lastInsertId, _ = result.LastInsertId() + log.LogInfo("Last inserted ID: %s", ins.lastInsertId) } - return []T{}, nil + return nil +} + +func (ins Insert[T]) LastInsertId() int64 { + return ins.lastInsertId } type Update[T any] struct { @@ -176,17 +184,17 @@ 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 { 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 { @@ -197,18 +205,19 @@ func (u Update[T]) Execute(conn *simpleorm.DBConnection) ([]T, error) { stmt, err := conn.Prepare(dml) if err != nil { - return []T{}, err + return err } 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) + result.LastInsertId() } - return []T{u.toUpdate}, nil + return nil } type Delete[T any] struct { @@ -216,7 +225,7 @@ type Delete[T any] struct { toDelete []T } -func NewDelete[T any](orm *simpleorm.ORM, toDelete []T) (Delete[T], error) { +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") @@ -225,11 +234,11 @@ 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 { count := len(d.toDelete) dml, err := d.target.GetDeleteDML(count) if err != nil { - return []T{}, err + return err } var params []any @@ -237,7 +246,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 { @@ -249,19 +258,19 @@ func (d Delete[T]) Execute(conn *simpleorm.DBConnection) ([]T, error) { stmt, err := conn.Prepare(dml) if err != nil { - return []T{}, err + return err } 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 } diff --git a/repository/repository.go b/repository/repository.go index 2799381..65c8d4d 100644 --- a/repository/repository.go +++ b/repository/repository.go @@ -1,6 +1,7 @@ package repository import ( + "errors" "reflect" "strings" @@ -11,6 +12,7 @@ import ( type HasPK interface { IsInsertable() bool + SetPk(pks ...any) } type Repository[T HasPK] struct { @@ -29,11 +31,12 @@ 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{} } + selecRes := selectExec.GetResult() res := make([]*T, 0, len(selecRes)) for i := range selecRes { @@ -73,7 +76,8 @@ func (r *Repository[T]) SelectByPk(pks ...any) *T { return nil } - res, err := selectExec.Execute(r.conn) + err = selectExec.Execute(r.conn) + res := selectExec.GetResult() if err != nil { log.LogError("Failed to load entity: %s", err) return nil @@ -87,41 +91,45 @@ func (r *Repository[T]) SelectByPk(pks ...any) *T { return &res[0] } -func (r *Repository[T]) Save(entity T) { +func (r *Repository[T]) Save(entity T) (*T, error) { // Insert 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 nil, errors.New("") } - if _, err := insertExec.Execute(r.conn); err != nil { + if err = insertExec.Execute(r.conn); err != nil { log.LogError("Insert failed for entity: %s", entity) } - return + return nil, errors.New("") } updateExec, err := exec.NewUpdate[T](r.orm, entity) if err != nil { log.LogError("Failed to create update execution: %s", err) - return + return nil, errors.New("") + } - _, err = updateExec.Execute(r.conn) + err = updateExec.Execute(r.conn) if err != nil { log.LogError("Update failed: %s", err) - return + return nil, errors.New("") + } + + return &entity, nil } func (r *Repository[T]) Delete(entity T) { - deleteExec, err := exec.NewDelete[T](r.orm, []T{entity}) + deleteExec, err := exec.NewDelete[T](r.orm, 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/test/exec_test.go b/test/exec_test.go index 99a9b04..14db5c0 100644 --- a/test/exec_test.go +++ b/test/exec_test.go @@ -36,14 +36,14 @@ func TestSelectEmpty(t *testing.T) { return } - res, err := selectExec.Execute(conn) + err = selectExec.Execute(conn) if err != nil { log.LogError("Select failure. %s", err) t.Fail() return } - if len(res) > 0 { + if len(selectExec.GetResult()) > 0 { log.LogError("Expected empty result.") t.Fail() return @@ -61,7 +61,7 @@ func TestInsertAndSelectSingle(t *testing.T) { testObj := Test{Int64Field: 54, IntField: 12, StringField: "asdf"} - insertExec, err := exec.NewInsert(orm, testObj) + insertExec, err := exec.NewInsert(orm, &testObj) if err != nil { log.LogError("TestInsertAndSelectSingle setup failed: %s", err) t.Fail() @@ -76,16 +76,18 @@ func TestInsertAndSelectSingle(t *testing.T) { } // WHEN - _, err = insertExec.Execute(conn) + err = insertExec.Execute(conn) if err != nil { log.LogError("TestInsertAndSelectSingle insert failed: %s", err) t.Fail() return } - res, err := selectExec.Execute(conn) + err = selectExec.Execute(conn) // THEN + res := selectExec.GetResult() + if err != nil { log.LogError("TestInsertAndSelectSingle failure. %s", err) t.Fail() @@ -116,7 +118,7 @@ func TestInsertSelectWithCompositePk(t *testing.T) { testObj := TestWithCompositePk{ID: 12, Name: "Test"} - insertExec, err := exec.NewInsert(orm, testObj) + insertExec, err := exec.NewInsert(orm, &testObj) if err != nil { log.LogError("TestInsertSelectWithCompositePk setup failed: %s", err) t.Fail() @@ -131,14 +133,15 @@ func TestInsertSelectWithCompositePk(t *testing.T) { } // 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) + res := selectExec.GetResult() // THEN if err != nil { @@ -173,14 +176,14 @@ func TestInsertMultipleAndSelectWithParam(t *testing.T) { testObj := Test{Int64Field: 54, IntField: 12, StringField: "asdf"} testObjSecond := Test{Int64Field: 12, IntField: 54, StringField: "fdsa"} - insertExec, err := exec.NewInsert(orm, testObj) + insertExec, err := exec.NewInsert(orm, &testObj) if err != nil { log.LogError("TestInsertMultipleAndSelectWithParam setup failed: %s", err) t.Fail() return } - insertSecond, err := exec.NewInsert(orm, testObjSecond) + insertSecond, err := exec.NewInsert(orm, &testObjSecond) if err != nil { log.LogError("TestInsertMultipleAndSelectWithParam setup failed: %s", err) t.Fail() @@ -196,21 +199,22 @@ func TestInsertMultipleAndSelectWithParam(t *testing.T) { // 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) + res := selectExec.GetResult() // THEN if err != nil { @@ -243,7 +247,7 @@ func TestInsertAndUpdate(t *testing.T) { testObj := Test{Int64Field: 54, IntField: 12, StringField: "asdf"} - insertExec, err := exec.NewInsert(orm, testObj) + insertExec, err := exec.NewInsert(orm, &testObj) if err != nil { log.LogError("TestInsertAndUpdate setup failed: %s", err) t.Fail() @@ -259,14 +263,16 @@ func TestInsertAndUpdate(t *testing.T) { // 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) + res := selectExec.GetResult() + if err != nil { log.LogError("TestInsertAndUpdate re-select failed: %s", err) t.Fail() @@ -285,8 +291,9 @@ func TestInsertAndUpdate(t *testing.T) { return } - _, err = updateExec.Execute(conn) - res, selectErr := selectExec.Execute(conn) + err = updateExec.Execute(conn) + selectErr := selectExec.Execute(conn) + res = selectExec.GetResult() // THEN if err != nil { @@ -326,14 +333,14 @@ func TestInsertAndDelete(t *testing.T) { testObj := Test{Int64Field: 54, IntField: 12, StringField: "asdf"} - insertExec, err := exec.NewInsert(orm, testObj) + insertExec, err := exec.NewInsert(orm, &testObj) if err != nil { log.LogError("TestInsertAndDelete setup failed: %s", err) t.Fail() return } - _, err = insertExec.Execute(conn) + err = insertExec.Execute(conn) if err != nil { log.LogError("TestInsertAndDelete setup failed: %s", err) t.Fail() @@ -347,7 +354,8 @@ func TestInsertAndDelete(t *testing.T) { return } - selectRes, err := selectExec.Execute(conn) + err = selectExec.Execute(conn) + selectRes := selectExec.GetResult() if err != nil { log.LogError("TestInsertAndDelete setup failed. %s", err) @@ -358,7 +366,7 @@ func TestInsertAndDelete(t *testing.T) { deleteExec, err := exec.NewDelete(orm, selectRes) // WHEN - _, err = deleteExec.Execute(conn) + err = deleteExec.Execute(conn) // THEN if err != nil { @@ -367,7 +375,9 @@ func TestInsertAndDelete(t *testing.T) { return } - selectRes, err = selectExec.Execute(conn) + err = selectExec.Execute(conn) + selectRes = selectExec.GetResult() + if len(selectRes) != 0 { log.LogError("TestInsertAndDelete test failed. Expected: 0, Actual: %s", len(selectRes)) t.Fail() @@ -386,7 +396,7 @@ func TestBoolInsertAndSelectSingle(t *testing.T) { testObj := TestWithBool{BoolField: true} - insertExec, err := exec.NewInsert(orm, testObj) + insertExec, err := exec.NewInsert(orm, &testObj) if err != nil { log.LogError("TestBoolInsertAndSelectSingle setup failed: %s", err) t.Fail() @@ -401,14 +411,15 @@ func TestBoolInsertAndSelectSingle(t *testing.T) { } // 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) + res := selectExec.GetResult() // THEN if err != nil { @@ -441,7 +452,7 @@ func TestTimeInsertAndSelectSingle(t *testing.T) { now := time.Now().UnixMilli() testObj := TestWithTime{TimeField: now} - insertExec, err := exec.NewInsert(orm, testObj) + insertExec, err := exec.NewInsert(orm, &testObj) if err != nil { log.LogError("TestTimeInsertAndSelectSingle setup failed: %s", err) t.Fail() @@ -456,14 +467,15 @@ func TestTimeInsertAndSelectSingle(t *testing.T) { } // 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) + res := selectExec.GetResult() // THEN if err != nil { diff --git a/test/main_test.go b/test/main_test.go index eb0680c..debbb47 100644 --- a/test/main_test.go +++ b/test/main_test.go @@ -15,45 +15,94 @@ type Test struct { StringField string `sql:"nn"` } -func (t Test) IsInsertable() bool { +func (t *Test) IsInsertable() bool { return t.ID == 0 } +func (t *Test) SetPk(pks ...any) { + t.ID = pks[0].(int) +} + type TestWithFk struct { ID int `sql:"pk"` TestID int `sql:"nn;fk=Test.ID"` } +func (t *TestWithFk) IsInsertable() bool { + return t.ID == 0 +} + +func (t *TestWithFk) SetPk(pks ...any) { + t.ID = pks[0].(int) +} + type TestWithFkAndFkId struct { ID int `sql:"pk"` TestID int `sql:"nn;fk=Test.ID;fk_id=custom_fk"` } +func (t *TestWithFkAndFkId) IsInsertable() bool { + return t.ID == 0 +} + +func (t *TestWithFkAndFkId) SetPk(pks ...any) { + t.ID = pks[0].(int) +} + type TestWithCompositePk struct { ID int `sql:"nn;pk;"` Name string `sql:"nn;pk"` } -func (t TestWithCompositePk) IsInsertable() bool { +func (t *TestWithCompositePk) IsInsertable() bool { return t.ID == 0 && t.Name == "" } +func (t *TestWithCompositePk) SetPk(pks ...any) { + t.ID = pks[0].(int) + t.Name = pks[1].(string) +} + type TestWithCompositeFk struct { ID int `sql:"nn;pk"` CompositeFkId int `sql:"nn;fk=TestWithCompositePk.ID;"` CompositeFkName string `sql:"nn;fk=TestWithCompositePk.Name;"` } +func (t *TestWithCompositeFk) IsInsertable() bool { + return t.ID == 0 +} + +func (t *TestWithCompositeFk) SetPk(pks ...any) { + t.ID = pks[0].(int) +} + type TestWithBool struct { ID int `sql:"pk"` BoolField bool `sql:"nn"` } +func (t *TestWithBool) IsInsertable() bool { + return t.ID == 0 +} + +func (t *TestWithBool) SetPk(pks ...any) { + t.ID = pks[0].(int) +} + type TestWithTime struct { ID int `sql:"pk"` TimeField int64 `sql:"nn"` } +func (t *TestWithTime) IsInsertable() bool { + return t.ID == 0 +} + +func (t *TestWithTime) SetPk(pks ...any) { + t.ID = pks[0].(int) +} + func TestMain(m *testing.M) { log.SetupLogs("Info") diff --git a/test/repository_test.go b/test/repository_test.go index b5f8d90..50d181e 100644 --- a/test/repository_test.go +++ b/test/repository_test.go @@ -26,14 +26,14 @@ func TestRepoSelectAll(t *testing.T) { defer cleanUp("test.db", conn) testObj := Test{Int64Field: 54, IntField: 12, StringField: "asdf"} - insertExec, err := exec.NewInsert(orm, testObj) + insertExec, err := exec.NewInsert(orm, &testObj) if err != nil { log.LogError("TestRepoSelectAll setup failed: %s", err) t.Fail() return } - _, err = insertExec.Execute(conn) - repo := repository.NewRepository[Test](conn, orm) + err = insertExec.Execute(conn) + repo := repository.NewRepository[*Test](conn, orm) // WHEN res := repo.SelectAll() @@ -52,14 +52,14 @@ func TestRepoSelectByPk(t *testing.T) { defer cleanUp("test.db", conn) testObj := Test{Int64Field: 54, IntField: 12, StringField: "asdf"} - insertExec, err := exec.NewInsert(orm, testObj) + insertExec, err := exec.NewInsert(orm, &testObj) if err != nil { log.LogError("TestRepoSelectByPk setup failed: %s", err) t.Fail() return } - _, err = insertExec.Execute(conn) - repo := repository.NewRepository[Test](conn, orm) + err = insertExec.Execute(conn) + repo := repository.NewRepository[*Test](conn, orm) // WHEN res := repo.SelectByPk(1) @@ -78,14 +78,14 @@ func TestRepoSelectByCompundPk(t *testing.T) { defer cleanUp("test.db", conn) testObj := TestWithCompositePk{ID: 42069, Name: "Test"} - insertExec, err := exec.NewInsert(orm, testObj) + insertExec, err := exec.NewInsert(orm, &testObj) if err != nil { log.LogError("TestRepoSelectByCompundPk setup failed: %s", err) t.Fail() return } - _, err = insertExec.Execute(conn) - repo := repository.NewRepository[TestWithCompositePk](conn, orm) + err = insertExec.Execute(conn) + repo := repository.NewRepository[*TestWithCompositePk](conn, orm) // WHEN res := repo.SelectByPk(42069, "Test") @@ -104,10 +104,10 @@ func TestRepoSave(t *testing.T) { defer cleanUp("test.db", conn) testObj := Test{Int64Field: 54, IntField: 12, StringField: "asdf"} - repo := repository.NewRepository[Test](conn, orm) + repo := repository.NewRepository[*Test](conn, orm) // WHEN - repo.Save(testObj) + repo.Save(&testObj) // THEN res := repo.SelectByPk(1) @@ -124,9 +124,9 @@ func TestRepoDelete(t *testing.T) { defer cleanUp("test.db", conn) testObj := Test{Int64Field: 54, IntField: 12, StringField: "asdf"} - repo := repository.NewRepository[Test](conn, orm) + repo := repository.NewRepository[*Test](conn, orm) - repo.Save(testObj) + repo.Save(&testObj) res := repo.SelectByPk(1) if res == nil {