diff --git a/exec/exec.go b/exec/exec.go index ea77d56..20e4833 100644 --- a/exec/exec.go +++ b/exec/exec.go @@ -36,9 +36,10 @@ type Select[T any] struct { } func CreateSelect[T any](orm *simpleorm.ORM, whereStmt string, args ...any) (Select[T], error) { - table, ok := orm.Cache().Get(reflect.TypeFor[T]().Name()) + typeName := reflect.TypeFor[T]().Name() + table, ok := orm.Cache().Get(typeName) if !ok { - return Select[T]{}, errors.New("Failed to get table from schema cache") + return Select[T]{}, errors.New("Failed to get table from schema cache: " + typeName) } return Select[T]{target: *table, whereStmt: whereStmt, args: args}, nil @@ -48,7 +49,10 @@ func (s Select[T]) GetResult() []T { return s.result } -func (s Select[T]) Execute(conn *simpleorm.DBConnection) error { +func (s *Select[T]) Execute(conn *simpleorm.DBConnection) error { + // Clear before executing + s.result = []T{} + dml, err := s.target.GetSelectDML() if err != nil { return err @@ -122,7 +126,9 @@ type Insert[T any] struct { } func NewInsert[T any](orm *simpleorm.ORM, toInsert ...T) (Insert[T], error) { - table, ok := orm.Cache().Get(reflect.TypeFor[T]().Name()) + typeName := reflect.TypeFor[T]().Name() + log.LogDebug("Searching Table cache by type: %s", typeName) + table, ok := orm.Cache().Get(typeName) if !ok { return Insert[T]{}, errors.New("Failed to get table from schema cache") } @@ -130,7 +136,7 @@ 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) error { +func (ins *Insert[T]) Execute(conn *simpleorm.DBConnection) error { dml, err := ins.target.GetInsertDML(len(ins.toInsert)) if err != nil { return err @@ -170,6 +176,10 @@ func (ins Insert[T]) LastInsertId() int64 { return ins.lastInsertId } +func (ins Insert[T]) Target() schema.Table { + return ins.target +} + type Update[T any] struct { target schema.Table toUpdate T diff --git a/repository/repository.go b/repository/repository.go index 65c8d4d..e2b08b8 100644 --- a/repository/repository.go +++ b/repository/repository.go @@ -12,7 +12,6 @@ import ( type HasPK interface { IsInsertable() bool - SetPk(pks ...any) } type Repository[T HasPK] struct { @@ -103,7 +102,7 @@ func (r *Repository[T]) Save(entity T) (*T, error) { log.LogError("Insert failed for entity: %s", entity) } - return nil, errors.New("") + return &entity, nil } updateExec, err := exec.NewUpdate[T](r.orm, entity) diff --git a/schema/table.go b/schema/table.go index 9ad1ff8..ce40e7f 100644 --- a/schema/table.go +++ b/schema/table.go @@ -1,6 +1,7 @@ package schema import ( + "errors" "reflect" "strings" @@ -186,3 +187,23 @@ func (t Table) IsPkAuto() bool { } return true } + +func (t Table) GetPkColumns() ([]Column, error) { + var values []Column + for _, constraint := range t.Constraints { + if constraint.Type == "pk" { + for _, col := range constraint.Columns { + _, ok := t.Type.FieldByName(col.FieldName) + if !ok { + continue + } + values = append(values, col) + } + } + } + if len(values) == 0 { + return nil, errors.New("Could not determine pk column!") + } + + return values, nil +} diff --git a/test/exec_test.go b/test/exec_test.go index 14db5c0..ea9a0c3 100644 --- a/test/exec_test.go +++ b/test/exec_test.go @@ -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[Test](orm, testObj) if err != nil { log.LogError("TestInsertAndSelectSingle setup failed: %s", err) t.Fail() @@ -118,7 +118,7 @@ func TestInsertSelectWithCompositePk(t *testing.T) { testObj := TestWithCompositePk{ID: 12, Name: "Test"} - insertExec, err := exec.NewInsert(orm, &testObj) + insertExec, err := exec.NewInsert[TestWithCompositePk](orm, testObj) if err != nil { log.LogError("TestInsertSelectWithCompositePk setup failed: %s", err) t.Fail() @@ -176,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[Test](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[Test](orm, testObjSecond) if err != nil { log.LogError("TestInsertMultipleAndSelectWithParam setup failed: %s", err) t.Fail() @@ -247,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[Test](orm, testObj) if err != nil { log.LogError("TestInsertAndUpdate setup failed: %s", err) t.Fail() @@ -333,7 +333,7 @@ func TestInsertAndDelete(t *testing.T) { testObj := Test{Int64Field: 54, IntField: 12, StringField: "asdf"} - insertExec, err := exec.NewInsert(orm, &testObj) + insertExec, err := exec.NewInsert[Test](orm, testObj) if err != nil { log.LogError("TestInsertAndDelete setup failed: %s", err) t.Fail() @@ -363,7 +363,7 @@ func TestInsertAndDelete(t *testing.T) { return } - deleteExec, err := exec.NewDelete(orm, selectRes) + deleteExec, err := exec.NewDelete[Test](orm, selectRes...) // WHEN err = deleteExec.Execute(conn) @@ -396,7 +396,7 @@ func TestBoolInsertAndSelectSingle(t *testing.T) { testObj := TestWithBool{BoolField: true} - insertExec, err := exec.NewInsert(orm, &testObj) + insertExec, err := exec.NewInsert[TestWithBool](orm, testObj) if err != nil { log.LogError("TestBoolInsertAndSelectSingle setup failed: %s", err) t.Fail() @@ -452,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[TestWithTime](orm, testObj) if err != nil { log.LogError("TestTimeInsertAndSelectSingle setup failed: %s", err) t.Fail() diff --git a/test/main_test.go b/test/main_test.go index debbb47..ae22bb1 100644 --- a/test/main_test.go +++ b/test/main_test.go @@ -15,14 +15,10 @@ 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"` @@ -32,10 +28,6 @@ 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"` @@ -54,15 +46,10 @@ type TestWithCompositePk struct { 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;"` @@ -73,10 +60,6 @@ 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"` @@ -86,10 +69,6 @@ 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"` @@ -99,10 +78,6 @@ 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 50d181e..c2e8ed1 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[Test](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) + repo := repository.NewRepository[Test](conn, orm) // WHEN res := repo.SelectAll() @@ -52,17 +52,17 @@ 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[Test](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) + repo := repository.NewRepository[Test](conn, orm) // WHEN - res := repo.SelectByPk(1) + res := repo.SelectByPk(insertExec.LastInsertId()) // THEN if res == nil { @@ -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[TestWithCompositePk](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) + 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 {