diff --git a/cache.go b/cache.go index 0b16f0f..205cd04 100644 --- a/cache.go +++ b/cache.go @@ -23,13 +23,13 @@ func NewOrmCache(tables []*Table) *OrmCache { func (o *OrmCache) add(table *Table) { o.mu.Lock() defer o.mu.Unlock() - o.data[table.TypeName] = table + o.data[table.Type.Name()] = table } -func (o *OrmCache) Get(tableName string) (*Table, bool) { +func (o *OrmCache) Get(typeName string) (*Table, bool) { o.mu.RLock() defer o.mu.RUnlock() - table, ok := o.data[tableName] + table, ok := o.data[typeName] return table, ok } diff --git a/dbconnection.go b/dbconnection.go new file mode 100644 index 0000000..c88c0e3 --- /dev/null +++ b/dbconnection.go @@ -0,0 +1,32 @@ +package simpleorm + +import ( + "database/sql" + + log "gitlab.com/gdulai/simpleloglvl" +) + +type DBConnection struct { + db *sql.DB + State int +} + +func OpenConnection(driver string, dsn string) *DBConnection { + // Open encrypted database + db, err := sql.Open(driver, dsn) + if err != nil { + log.LogError("Failed to open db connection: %s", err) + return &DBConnection{db: nil, State: -1} + } + + return &DBConnection{db: db, State: 1} +} + +func (c *DBConnection) Close() (bool, error) { + err := c.db.Close() + if err != nil { + return false, err + } + c = nil + return true, nil +} diff --git a/go.mod b/go.mod index 50ffbdd..8b3b789 100644 --- a/go.mod +++ b/go.mod @@ -2,4 +2,7 @@ module simpleorm go 1.26.2 -require gitlab.com/gdulai/simpleloglvl v0.0.0-20260418080844-d5cca4888d97 // indirect +require ( + github.com/mattn/go-sqlite3 v1.14.44 + gitlab.com/gdulai/simpleloglvl v0.0.0-20260418080844-d5cca4888d97 // indirect +) diff --git a/go.sum b/go.sum index 256fbf7..0168c0b 100644 --- a/go.sum +++ b/go.sum @@ -1,2 +1,4 @@ +github.com/mattn/go-sqlite3 v1.14.44 h1:3VSe+xafpbzsLbdr2AWlAZk9yRHiBhTBakioXaCKTF8= +github.com/mattn/go-sqlite3 v1.14.44/go.mod h1:pjEuOr8IwzLJP2MfGeTb0A35jauH+C2kbHKBr7yXKVQ= gitlab.com/gdulai/simpleloglvl v0.0.0-20260418080844-d5cca4888d97 h1:f6UvrrilTIgQ2jWMJEEObZwAQQPykoWY2ju3Bq2qaGA= gitlab.com/gdulai/simpleloglvl v0.0.0-20260418080844-d5cca4888d97/go.mod h1:H7XPunUrSyAvPa9nx8UbKnThEQJDmj3mdt+9ZtqDth4= diff --git a/main_test.go b/main_test.go index dd922e7..835de30 100644 --- a/main_test.go +++ b/main_test.go @@ -4,11 +4,40 @@ import ( "os" "testing" + _ "github.com/mattn/go-sqlite3" log "gitlab.com/gdulai/simpleloglvl" ) +type Test struct { + ID int `sql:"pk"` + Int64Field int64 `sql:"nn"` + IntField int `sql:"nn"` + StringField string `sql:"nn"` +} + +type TestWithFk struct { + ID int `sql:"pk"` + TestID int `sql:"nn;fk=Test.ID"` +} + +type TestWithFkAndFkId struct { + ID int `sql:"pk"` + TestID int `sql:"nn;fk=Test.ID;fk_id=custom_fk"` +} + +type TestWithCompositePk struct { + ID int `sql:"nn;pk"` + Name string `sql:"nn;pk"` +} + +type TestWithCompositeFk struct { + ID int `sql:"nn;pk"` + CompositeFkId int `sql:"nn;fk=TestWithCompositePk.ID;"` + CompositeFkName string `sql:"nn;fk=TestWithCompositePk.Name;"` +} + func TestMain(m *testing.M) { - log.SetupLogs("Debug") + log.SetupLogs("Info") code := m.Run() diff --git a/orm.go b/orm.go index abea45d..1d4ae59 100644 --- a/orm.go +++ b/orm.go @@ -43,7 +43,8 @@ func NewORM(objs ...any) *ORM { return &ORM{cache: cache} } -func (orm *ORM) CreateDDL() string { +// Builds the DDL and returns it as a string +func (orm ORM) CreateDDL() string { var ddl strings.Builder tables := orm.cache.GetAll() diff --git a/orm_test.go b/orm_test.go index 33a0526..8ad5d95 100644 --- a/orm_test.go +++ b/orm_test.go @@ -8,13 +8,6 @@ import ( log "gitlab.com/gdulai/simpleloglvl" ) -type Test struct { - ID int `sql:"pk"` - Int64Field int64 `sql:"nn"` - IntField int `sql:"nn"` - StringField string `sql:"nn"` -} - func TestParseSimple(t *testing.T) { log.LogInfo("Test start...") // GIVEN @@ -31,11 +24,6 @@ func TestParseSimple(t *testing.T) { log.LogInfo("Test finished!") } -type TestWithFk struct { - ID int `sql:"pk"` - TestID int `sql:"nn;fk=Test.ID"` -} - func TestParseFk(t *testing.T) { log.LogInfo("Test start...") // GIVEN @@ -54,11 +42,6 @@ func TestParseFk(t *testing.T) { log.LogInfo("Test finished!") } -type TestWithFkAndFkId struct { - ID int `sql:"pk"` - TestID int `sql:"nn;fk=Test.ID;fk_id=custom_fk"` -} - func TestParseFkAndFkId(t *testing.T) { // GIVEN orm := simpleorm.NewORM(Test{}, TestWithFkAndFkId{}) @@ -75,11 +58,6 @@ func TestParseFkAndFkId(t *testing.T) { } } -type TestWithCompositePk struct { - ID int `sql:"nn;pk"` - Name string `sql:"nn;pk"` -} - func TestParseCompositePk(t *testing.T) { // GIVEN orm := simpleorm.NewORM(TestWithCompositePk{}) @@ -94,12 +72,6 @@ func TestParseCompositePk(t *testing.T) { } } -type TestWithCompositeFk struct { - ID int `sql:"nn;pk"` - CompositeFkId int `sql:"nn;fk=TestWithCompositePk.ID;"` - CompositeFkName string `sql:"nn;fk=TestWithCompositePk.Name;"` -} - func TestParseCompositeFk(t *testing.T) { // GIVEN orm := simpleorm.NewORM(TestWithCompositePk{}, TestWithCompositeFk{}) diff --git a/ormexec.go b/ormexec.go new file mode 100644 index 0000000..ca725f7 --- /dev/null +++ b/ormexec.go @@ -0,0 +1,191 @@ +package simpleorm + +import ( + "database/sql" + "errors" + "reflect" + + log "gitlab.com/gdulai/simpleloglvl" +) + +type ORMExec[T any] struct { + orm *ORM + conn *DBConnection +} + +func CreateExecution[T any](orm *ORM, conn *DBConnection) ORMExec[T] { + return ORMExec[T]{orm: orm, conn: conn} +} + +func InitDB(conn *DBConnection, orm *ORM) { + // Force a real DB interaction + if _, err := conn.db.Exec(orm.CreateDDL()); err != nil { + log.LogFatalError("%", err) + } +} + +func (ormExec ORMExec[T]) Select(desc T, whereStmt string, args ...any) ([]T, error) { + lookup := reflect.TypeOf(desc).Name() + table, ok := ormExec.orm.cache.Get(lookup) + if !ok { + return []T{}, errors.New("Failed to get descriptor from cache! " + lookup) + } + sqlStr := table.ToSelectDML() + if whereStmt != "" { + sqlStr += " WHERE " + whereStmt + } + log.LogDebug("Preparing sql: %s, with args: %s", sqlStr, args) + + stmt, err := ormExec.conn.db.Prepare(sqlStr) + if err != nil { + return nil, err + } + defer stmt.Close() + + log.LogDebug("Executing statement: %s", stmt) + + var rows *sql.Rows + if len(args) == 0 { + rows, err = stmt.Query() + } else { + // Flattent args to make sure it can be parsed correctly + var flatArgs []any + for _, a := range args { + if s, ok := a.([]any); ok { + flatArgs = append(flatArgs, s...) + } else { + flatArgs = append(flatArgs, a) + } + } + rows, err = stmt.Query(flatArgs...) + + } + + if err != nil { + return nil, err + } + + rowContainer := table.createSelectResultContainer() + + 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 { + parsedResult.Elem().Field(i).Set(reflect.Indirect(reflect.ValueOf(fieldVal))) + } + parsedObj := reflect.Indirect(parsedResult).Interface().(T) + results = append(results, parsedObj) + } + + return results, nil +} + +func (ormExec ORMExec[T]) Insert(src ...T) error { + typeName := reflect.TypeOf(src[0]).Name() + table, ok := ormExec.orm.cache.Get(typeName) + if !ok { + return errors.New("Could not find table for type " + typeName) + } + + sql := table.ToInsertDML(len(src)) + var params []any + for i := range len(src) { + actualParams := table.prepareParams(src[i]) + if len(params) == 0 { + params = make([]any, len(src)*len(actualParams)) + } + for j := range actualParams { + params[(i*len(actualParams))+j] = actualParams[j] + } + } + + log.LogInfo("%s [%s]", sql, params) + + stmt, err := ormExec.conn.db.Prepare(sql) + if err != nil { + return err + } + defer stmt.Close() + + result, err := stmt.Exec(params...) + if err != nil { + return err + } else { + rowsAffected, _ := result.RowsAffected() + log.LogInfo("Inserted %s row", rowsAffected) + } + return nil +} + +func (ormExec ORMExec[T]) Update(src T) error { + typeName := reflect.TypeOf(src).Name() + table, ok := ormExec.orm.cache.Get(reflect.TypeOf(src).Name()) + if !ok { + return errors.New("Could not find table for type " + typeName) + } + + sql := table.ToUpdateDML(src) + params := table.prepareParams(src) + + pk, err := table.getPk(src) + if err != nil { + return err + } + // Put the pk back at the end + params = append(params, pk) + + log.LogInfo("%s [%s]", sql, params) + + stmt, err := ormExec.conn.db.Prepare(sql) + if err != nil { + return err + } + defer stmt.Close() + + result, err := stmt.Exec(params...) + if err != nil { + return err + } else { + rowsAffected, _ := result.RowsAffected() + log.LogInfo("Updated %s row", rowsAffected) + } + return nil +} + +func (ormExec ORMExec[T]) Delete(src T) error { + typeName := reflect.TypeOf(src).Name() + table, ok := ormExec.orm.cache.Get(reflect.TypeOf(src).Name()) + if !ok { + return errors.New("Could not find table for type " + typeName) + } + + sql := table.ToDeleteDML(src) + pk, err := table.getPk(src) + if err != nil { + return err + } + + log.LogInfo("%s [%s]", sql, pk) + + stmt, err := ormExec.conn.db.Prepare(sql) + if err != nil { + return err + } + defer stmt.Close() + + result, err := stmt.Exec(pk) + if err != nil { + return err + } else { + rowsAffected, _ := result.RowsAffected() + log.LogInfo("Deleted %s row", rowsAffected) + } + return nil +} diff --git a/ormexec_test.go b/ormexec_test.go new file mode 100644 index 0000000..4413bd1 --- /dev/null +++ b/ormexec_test.go @@ -0,0 +1,114 @@ +package simpleorm_test + +import ( + "os" + "simpleorm" + "testing" + + log "gitlab.com/gdulai/simpleloglvl" +) + +func TestSelectEmpty(t *testing.T) { + // GIVEN + conn := simpleorm.OpenConnection("sqlite3", "file:test.db") + defer cleanUp("test.db", conn) + + orm := simpleorm.NewORM(Test{}, TestWithFk{}, TestWithFkAndFkId{}, + TestWithCompositePk{}, TestWithCompositeFk{}) + + simpleorm.InitDB(conn, orm) + + exec := simpleorm.CreateExecution[Test](orm, conn) + // WHEN + res, err := exec.Select(Test{}, "") + // THEN + if err != nil { + log.LogError("Select failure. %s", err) + t.Fail() + } + if len(res) > 0 { + log.LogError("Expected empty result.") + t.Fail() + } +} + +func TestInsertAndSelectSingle(t *testing.T) { + // GIVEN + conn := simpleorm.OpenConnection("sqlite3", "file:test.db") + defer cleanUp("test.db", conn) + + orm := simpleorm.NewORM(Test{}, TestWithFk{}, TestWithFkAndFkId{}, + TestWithCompositePk{}, TestWithCompositeFk{}) + + simpleorm.InitDB(conn, orm) + + exec := simpleorm.CreateExecution[Test](orm, conn) + + testObj := Test{Int64Field: 54, IntField: 12, StringField: "asdf"} + err := exec.Insert(testObj) + if err != nil { + log.LogError("Insert failure. %s", err) + t.Fail() + } + // WHEN + res, err := exec.Select(Test{}, "") + // THEN + if err != nil { + log.LogError("Select failure. %s", err) + t.Fail() + } + if len(res) != 1 { + log.LogError("Selec test failed. Expected: 1, Actual: %s", len(res)) + t.Fail() + } + + singleRes := res[0] + if singleRes.ID != 1 || singleRes.Int64Field != 54 || singleRes.IntField != 12 || singleRes.StringField != "asdf" { + log.LogError("Invalid result. Expected: %s, Actual: %s", testObj, singleRes) + t.Fail() + } + +} + +func TestInsertMultipleAndSelectWithParam(t *testing.T) { + // GIVEN + conn := simpleorm.OpenConnection("sqlite3", "file:test.db") + defer cleanUp("test.db", conn) + + orm := simpleorm.NewORM(Test{}, TestWithFk{}, TestWithFkAndFkId{}, + TestWithCompositePk{}, TestWithCompositeFk{}) + + simpleorm.InitDB(conn, orm) + + exec := simpleorm.CreateExecution[Test](orm, conn) + + testObj := Test{Int64Field: 54, IntField: 12, StringField: "asdf"} + testObjSecond := Test{Int64Field: 12, IntField: 54, StringField: "fdsa"} + err := exec.Insert(testObj, testObjSecond) + if err != nil { + log.LogError("Insert failure. %s", err) + t.Fail() + } + // WHEN + res, err := exec.Select(Test{}, "string_field = ?", "fdsa") + // THEN + if err != nil { + log.LogError("Select failure. %s", err) + t.Fail() + } + if len(res) != 1 { + log.LogError("Select test failed. Expected: 1, Actual: %s", len(res)) + t.Fail() + } + + singleRes := res[0] + if singleRes.ID != 2 || singleRes.Int64Field != 12 || singleRes.IntField != 54 || singleRes.StringField != "fdsa" { + log.LogError("Invalid result. Expected: %s, Actual: %s", testObjSecond, singleRes) + t.Fail() + } +} + +func cleanUp(dbFile string, conn *simpleorm.DBConnection) { + conn.Close() + os.Remove(dbFile) +} diff --git a/parser.go b/parser.go index c74cbc8..72cf0c8 100644 --- a/parser.go +++ b/parser.go @@ -13,14 +13,14 @@ type Parser struct { Table *Table } -func NewParser(obj any) *Parser { +func NewParser[T any](obj T) *Parser { objType := reflect.TypeOf(obj) return &Parser{typ: objType} } // This is step 1 of the parsing, it creates the table instance and func (p *Parser) ParseColumns() *Table { - table := Table{Name: camelToSnake(p.typ.Name()), TypeName: p.typ.Name()} + table := Table{Name: camelToSnake(p.typ.Name()), Type: p.typ} var columns []Column for field := range p.typ.Fields() { diff --git a/table.go b/table.go index 617786d..3e36d65 100644 --- a/table.go +++ b/table.go @@ -1,17 +1,20 @@ package simpleorm import ( + "errors" + "reflect" "strings" + "time" ) type Table struct { Name string - TypeName string + Type reflect.Type Columns []Column Constraints []Constraint } -func (t *Table) ToDDL() string { +func (t Table) ToDDL() string { var ddl strings.Builder ddl.WriteString("CREATE TABLE IF NOT EXISTS " + t.Name + " (") for i, col := range t.Columns { @@ -27,3 +30,156 @@ func (t *Table) ToDDL() string { ddl.WriteString(");") return ddl.String() } + +func (t Table) ToSelectDML() string { + var dml strings.Builder + dml.WriteString("SELECT ") + + for i, col := range t.Columns { + if i != 0 { + dml.WriteString(", ") + } + dml.WriteString(col.Name) + } + + dml.WriteString(" FROM " + t.Name) + return dml.String() +} + +func (t Table) ToInsertDML(count int) string { + var dml strings.Builder + dml.WriteString("INSERT INTO " + camelToSnake(t.Type.Name()) + " (") + + columnsLen := len(t.Columns) + effectiveColumnsLen := 0 + for i, col := range t.Columns { + _, ok := col.Modifiers["pk"] + if ok { + continue + } + + dml.WriteString(col.Name) + if i != columnsLen-1 { + dml.WriteString(", ") + } + effectiveColumnsLen++ + } + + dml.WriteString(") VALUES") + + for i := range count { + if i != 0 { + dml.WriteString(", ") + } + for j := range effectiveColumnsLen { + if j == 0 { + dml.WriteString("(") + } + if j != effectiveColumnsLen-1 { + dml.WriteString("?, ") + } else { + dml.WriteString("?)") + } + } + } + + return dml.String() +} + +func (t Table) ToUpdateDML(src any) string { + var dml strings.Builder + dml.WriteString("UPDATE " + camelToSnake(t.Type.Name()) + " SET ") + + var pkCol string + // Colum names + for i := 0; i < t.Type.NumField(); i++ { + field := t.Type.Field(i) + modifiers := determineModifiers(field.Tag) + + _, ok := modifiers["pk"] + if ok { + pkCol = camelToSnake(field.Name) + continue + } + + if i != t.Type.NumField()-1 { + dml.WriteString(camelToSnake(field.Name) + " = ?, ") + } else { + dml.WriteString(camelToSnake(field.Name) + " = ? ") + } + + } + + dml.WriteString("WHERE " + pkCol + " = ?") + + return dml.String() +} + +func (t Table) ToDeleteDML(src any) string { + var dml strings.Builder + dml.WriteString("DELETE FROM " + camelToSnake(t.Type.Name()) + " WHERE ") + + for i := 0; i < t.Type.NumField(); i++ { + field := t.Type.Field(i) + modifiers := determineModifiers(field.Tag) + _, ok := modifiers["pk"] + if ok { + dml.WriteString(camelToSnake(field.Name) + " = ?") + return dml.String() + } + } + + return "" +} + +func (t Table) createSelectResultContainer() []any { + vals := make([]any, t.Type.NumField()) + for i := range vals { + switch t.Type.Field(i).Type.Kind().String() { + case "string": + var fieldContainer string + vals[i] = &fieldContainer + case "int": + var fieldContainer int + vals[i] = &fieldContainer + case "int64", "time.Time": + var fieldContainer int64 + vals[i] = &fieldContainer + } + } + return vals +} + +func (t Table) prepareParams(src any) []any { + var params []any + for i := 0; i < t.Type.NumField(); i++ { + field := t.Type.Field(i) + modifiers := determineModifiers(field.Tag) + _, ok := modifiers["pk"] + if ok { + continue + } + typStr := field.Type.String() + switch typStr { + case "time.Time": + time := reflect.ValueOf(src).Field(i).Interface().(time.Time) + params = append(params, time.UnixMilli()) + default: + params = append(params, reflect.ValueOf(src).Field(i).Interface()) + } + } + + return params +} + +func (t Table) getPk(src any) (any, error) { + for i := 0; i < t.Type.NumField(); i++ { + field := t.Type.Field(i) + modifiers := determineModifiers(field.Tag) + _, ok := modifiers["pk"] + if ok { + return reflect.ValueOf(src).Field(i).Interface(), nil + } + } + return nil, errors.New("Could not determine pk column!") +}