diff --git a/exec/exec.go b/exec/exec.go index 56ac02d..ec9dcf2 100644 --- a/exec/exec.go +++ b/exec/exec.go @@ -4,7 +4,6 @@ import ( "database/sql" "errors" "reflect" - "time" "git.gdulai.com/gdulai/simpleorm" "git.gdulai.com/gdulai/simpleorm/schema" @@ -98,7 +97,13 @@ func (s Select[T]) Execute(conn *simpleorm.DBConnection) ([]T, error) { targetType := s.target.Type parsedResult := reflect.New(targetType) for i, fieldVal := range rowContainer { - parsedResult.Elem().Field(i).Set(reflect.Indirect(reflect.ValueOf(fieldVal))) + 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) @@ -267,7 +272,7 @@ func createSelectResultContainer(t schema.Table) []any { case "string": var fieldContainer string vals[i] = &fieldContainer - case "int": + case "int", "bool": var fieldContainer int vals[i] = &fieldContainer case "int64", "time.Time": @@ -290,15 +295,8 @@ func prepareParams(src any, t schema.Table) []any { continue } fieldValue := reflect.ValueOf(src).FieldByIndex(field.Index) - typStr := field.Type.String() - switch typStr { - case "time.Time": - time := fieldValue.Interface().(time.Time) - params = append(params, time.UnixMilli()) - default: - params = append(params, fieldValue.Interface()) - } + params = append(params, col.Encode(fieldValue)) } return params } diff --git a/exec_test.go b/exec_test.go index de5a969..0ae7fa0 100644 --- a/exec_test.go +++ b/exec_test.go @@ -3,6 +3,7 @@ package simpleorm_test import ( "os" "testing" + "time" simpleorm "git.gdulai.com/gdulai/simpleorm" "git.gdulai.com/gdulai/simpleorm/exec" @@ -299,7 +300,7 @@ func TestInsertAndDelete(t *testing.T) { return } - deleteExec, err := exec.NewDelete[Test](orm, selectRes) + deleteExec, err := exec.NewDelete(orm, selectRes) // WHEN _, err = deleteExec.Execute(conn) @@ -319,3 +320,112 @@ func TestInsertAndDelete(t *testing.T) { } } + +func TestBoolInsertAndSelectSingle(t *testing.T) { + // GIVEN + conn := simpleorm.OpenConnection("sqlite3", "file:test.db") + defer cleanUp("test.db", conn) + orm := simpleorm.NewORM(TestWithBool{}) + + exec.ExecuteDDL(conn, orm) + + testObj := TestWithBool{BoolField: true} + + insertExec, err := exec.NewInsert(orm, testObj) + if err != nil { + log.LogError("TestBoolInsertAndSelectSingle setup failed: %s", err) + t.Fail() + return + } + + selectExec, err := exec.CreateSelect[TestWithBool](orm, "") + if err != nil { + log.LogError("TestBoolInsertAndSelectSingle setup failed: %s", err) + t.Fail() + return + } + + // WHEN + _, err = insertExec.Execute(conn) + if err != nil { + log.LogError("TestBoolInsertAndSelectSingle insert failed: %s", err) + t.Fail() + return + } + + res, err := selectExec.Execute(conn) + + // THEN + if err != nil { + log.LogError("TestBoolInsertAndSelectSingle failure. %s", err) + t.Fail() + return + } + if len(res) != 1 { + log.LogError("TestBoolInsertAndSelectSingle test failed. Expected: 1, Actual: %s", len(res)) + t.Fail() + return + } + + singleRes := res[0] + if singleRes.ID != 1 || !singleRes.BoolField { + log.LogError("TestBoolInsertAndSelectSingle invalid result. Expected: %s, Actual: %s", testObj, singleRes) + t.Fail() + } + +} + +func TestTimeInsertAndSelectSingle(t *testing.T) { + // GIVEN + conn := simpleorm.OpenConnection("sqlite3", "file:test.db") + defer cleanUp("test.db", conn) + orm := simpleorm.NewORM(TestWithTime{}) + + exec.ExecuteDDL(conn, orm) + + now := time.Now().UnixMilli() + testObj := TestWithTime{TimeField: now} + + insertExec, err := exec.NewInsert(orm, testObj) + if err != nil { + log.LogError("TestTimeInsertAndSelectSingle setup failed: %s", err) + t.Fail() + return + } + + selectExec, err := exec.CreateSelect[TestWithTime](orm, "") + if err != nil { + log.LogError("TestTimeInsertAndSelectSingle setup failed: %s", err) + t.Fail() + return + } + + // WHEN + _, err = insertExec.Execute(conn) + if err != nil { + log.LogError("TestTimeInsertAndSelectSingle insert failed: %s", err) + t.Fail() + return + } + + res, err := selectExec.Execute(conn) + + // THEN + if err != nil { + log.LogError("TestTimeInsertAndSelectSingle failure. %s", err) + t.Fail() + return + } + if len(res) != 1 { + log.LogError("TestTimeInsertAndSelectSingle test failed. Expected: 1, Actual: %s", len(res)) + t.Fail() + return + } + + singleRes := res[0] + if singleRes.ID != 1 || singleRes.TimeField != now { + log.LogError("TestTimeInsertAndSelectSingle invalid result. Expected: %s, Actual: %s", testObj, singleRes) + t.Fail() + } + +} diff --git a/main_test.go b/main_test.go index 835de30..371c8fd 100644 --- a/main_test.go +++ b/main_test.go @@ -36,6 +36,16 @@ type TestWithCompositeFk struct { CompositeFkName string `sql:"nn;fk=TestWithCompositePk.Name;"` } +type TestWithBool struct { + ID int `sql:"pk"` + BoolField bool `sql:"nn"` +} + +type TestWithTime struct { + ID int `sql:"pk"` + TimeField int64 `sql:"nn"` +} + func TestMain(m *testing.M) { log.SetupLogs("Info") diff --git a/parser/parser.go b/parser/parser.go index 5d335b4..7d8b15e 100644 --- a/parser/parser.go +++ b/parser/parser.go @@ -91,7 +91,7 @@ func determineType(typ reflect.Type) string { switch typStr { case "string": return "TEXT" - case "int": + case "int", "bool": return "INTEGER" case "time.Time", "int64": return "BIGINT" diff --git a/schema/column.go b/schema/column.go index 61ae9ac..cf08e58 100644 --- a/schema/column.go +++ b/schema/column.go @@ -4,6 +4,7 @@ import ( "errors" "reflect" "strings" + "time" ) type Column struct { @@ -51,7 +52,7 @@ func (c *Column) GetDDL() (string, error) { func (c *Column) inlineModifiers() []string { var mods []string - for k, _ := range c.Modifiers { + for k := range c.Modifiers { switch k { case "nn": mods = append(mods, "NOT NULL") @@ -80,3 +81,46 @@ func (c *Column) IsPK() bool { _, ok := c.Modifiers["pk"] return ok } + +// Encodes the value to be DB compatible +func (c Column) Encode(value reflect.Value) any { + typ := value.Type().String() + + var encoded any + switch typ { + case "time.Time": + timeVal := value.Interface().(time.Time) + encoded = timeVal.UnixMilli() + case "bool": + boolVal := value.Interface().(bool) + if boolVal { + encoded = 1 + } else { + encoded = false + } + default: + encoded = value.Interface() + } + + return encoded +} + +// Decodes the raw value from the DBs +func (c Column) Decode(expected string, rawValue reflect.Value) any { + var decoded any + switch expected { + case "time.Time": + decoded = time.UnixMilli(rawValue.Interface().(int64)) + case "bool": + rawValue := rawValue.Interface().(int) + if rawValue > 0 { + decoded = true + } else { + decoded = false + } + default: + decoded = rawValue.Interface() + } + + return decoded +}