diff --git a/exec/exec.go b/exec/exec.go index 09889fd..c8447ff 100644 --- a/exec/exec.go +++ b/exec/exec.go @@ -113,8 +113,9 @@ func (s Select[T]) Execute(conn *simpleorm.DBConnection) ([]T, error) { } 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) { @@ -156,12 +157,16 @@ 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 ID: %s", ins.lastInsertId) } return []T{}, nil } +func (ins Insert[T]) LastInsertId() int64 { + return ins.lastInsertId +} + type Update[T any] struct { target schema.Table toUpdate T @@ -207,6 +212,7 @@ func (u Update[T]) Execute(conn *simpleorm.DBConnection) ([]T, error) { } else { rowsAffected, _ := result.RowsAffected() log.LogInfo("Updated %s row", rowsAffected) + result.LastInsertId() } return []T{u.toUpdate}, nil } diff --git a/repository/repository.go b/repository/repository.go index 2799381..b567cd5 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 { @@ -87,32 +89,36 @@ 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 { 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) if err != nil { log.LogError("Update failed: %s", err) - return + return nil, errors.New("") + } + + return &entity, nil } func (r *Repository[T]) Delete(entity T) { 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..9c3803d 100644 --- a/test/repository_test.go +++ b/test/repository_test.go @@ -33,7 +33,7 @@ func TestRepoSelectAll(t *testing.T) { return } _, err = insertExec.Execute(conn) - repo := repository.NewRepository[Test](conn, orm) + repo := repository.NewRepository[*Test](conn, orm) // WHEN res := repo.SelectAll() @@ -59,7 +59,7 @@ func TestRepoSelectByPk(t *testing.T) { return } _, err = insertExec.Execute(conn) - repo := repository.NewRepository[Test](conn, orm) + repo := repository.NewRepository[*Test](conn, orm) // WHEN res := repo.SelectByPk(1) @@ -85,7 +85,7 @@ func TestRepoSelectByCompundPk(t *testing.T) { return } _, err = insertExec.Execute(conn) - repo := repository.NewRepository[TestWithCompositePk](conn, orm) + 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 {