mirror of
https://github.com/geektutu/7days-golang.git
synced 2024-04-21 12:32:11 +00:00
add day7 migrate
This commit is contained in:
@@ -40,7 +40,15 @@
|
||||
|
||||
[GeeORM] 是一个模仿 [gorm](https://github.com/jinzhu/gorm) 和 [xorm](https://github.com/go-xorm/xorm) 的 ORM 框架
|
||||
|
||||
gorm 准备推出完全重写的 v2 版本(目前还在开发中),相对 gorm-v1 来说,xorm 的设计更容易理解,所以 geeorm 接口设计上主要参考了 xorm,具体实现参考了 gorm。
|
||||
gorm 准备推出完全重写的 v2 版本(目前还在开发中),相对 gorm-v1 来说,xorm 的设计更容易理解,所以 geeorm 接口设计上主要参考了 xorm,一些细节实现上参考了 gorm。
|
||||
|
||||
- 第一天:database/sql 基础 | [Code](gee-cache/day1-database-sql)
|
||||
- 第二天:对象表结构映射 | [Code](gee-cache/day2-reflect-schema)
|
||||
- 第三天:插入/查询记录 | [Code](gee-cache/day3-save-query)
|
||||
- 第四天:链式操作与更新删除 | [Code](gee-cache/day4-chain-operation)
|
||||
- 第五天:实现钩子(Hooks) | [Code](gee-cache/day5-hooks)
|
||||
- 第六天:支持事务(Transaction) | [Code](gee-cache/day6-transaction)
|
||||
- 第七天:数据库迁移(Migrate) | [Code](gee-cache/day7-migrate)
|
||||
|
||||
### WebAssembly 使用示例
|
||||
|
||||
@@ -80,12 +88,20 @@ What can I write in 7 days? A gin-like web framework? A distributed cache like g
|
||||
- Day 6 - Cache Breakdown & Single Flight | [Code](gee-cache/day6-single-flight)
|
||||
- Day 7 - Use Protobuf as RPC Data Exchange Type | [Code](gee-cache/day7-proto-buf)
|
||||
|
||||
## Object Relational Mapping - GeeOrm
|
||||
## Object Relational Mapping - GeeORM
|
||||
|
||||
[GeeOrm] is a [gorm](https://github.com/jinzhu/gorm)-like and [xorm](https://github.com/go-xorm/xorm)-like object relational mapping library
|
||||
[GeeORM] is a [gorm](https://github.com/jinzhu/gorm)-like and [xorm](https://github.com/go-xorm/xorm)-like object relational mapping library
|
||||
|
||||
Xorm's desgin is easier to understand than gorm-v1, so the main designs references xorm and some detailed implementions references gorm-v1.
|
||||
|
||||
- Day 1 - database/sql Basic | [Code](gee-cache/day1-database-sql)
|
||||
- Day 2 - Object Schame Mapping | [Code](gee-cache/day2-reflect-schema)
|
||||
- Day 3 - Insert and Query | [Code](gee-cache/day3-save-query)
|
||||
- Day 4 - Chain, Delete and Update | [Code](gee-cache/day4-chain-operation)
|
||||
- Day 5 - Support Hooks | [Code](gee-cache/day5-hooks)
|
||||
- Day 6 - Support Transaction | [Code](gee-cache/day6-transaction)
|
||||
- Day 7 - Migrate Database | [Code](gee-cache/day7-migrate)
|
||||
|
||||
## Golang WebAssembly Demo
|
||||
|
||||
- Demo 1 - Hello World [Code](demo-wasm/hello-world)
|
||||
|
||||
@@ -24,7 +24,7 @@ func NewSession() *Session {
|
||||
func TestSession_Exec(t *testing.T) {
|
||||
s := NewSession()
|
||||
_, _ = s.Raw("DROP TABLE IF EXISTS User;").Exec()
|
||||
_, _ = s.Raw("CREATE TABLE User(name text);").Exec()
|
||||
_, _ = s.Raw("CREATE TABLE User(Name text);").Exec()
|
||||
result, _ := s.Raw("INSERT INTO User(`Name`) values (?), (?)", "Tom", "Sam").Exec()
|
||||
if count, err := result.RowsAffected(); err != nil || count != 2 {
|
||||
t.Fatal("expect 2, but got", count)
|
||||
|
||||
@@ -29,7 +29,7 @@ func NewSession() *Session {
|
||||
func TestSession_Exec(t *testing.T) {
|
||||
s := NewSession()
|
||||
_, _ = s.Raw("DROP TABLE IF EXISTS User;").Exec()
|
||||
_, _ = s.Raw("CREATE TABLE User(name text);").Exec()
|
||||
_, _ = s.Raw("CREATE TABLE User(Name text);").Exec()
|
||||
result, _ := s.Raw("INSERT INTO User(`Name`) values (?), (?)", "Tom", "Sam").Exec()
|
||||
if count, err := result.RowsAffected(); err != nil || count != 2 {
|
||||
t.Fatal("expect 2, but got", count)
|
||||
|
||||
@@ -29,7 +29,7 @@ func NewSession() *Session {
|
||||
func TestSession_Exec(t *testing.T) {
|
||||
s := NewSession()
|
||||
_, _ = s.Raw("DROP TABLE IF EXISTS User;").Exec()
|
||||
_, _ = s.Raw("CREATE TABLE User(name text);").Exec()
|
||||
_, _ = s.Raw("CREATE TABLE User(Name text);").Exec()
|
||||
result, _ := s.Raw("INSERT INTO User(`Name`) values (?), (?)", "Tom", "Sam").Exec()
|
||||
if count, err := result.RowsAffected(); err != nil || count != 2 {
|
||||
t.Fatal("expect 2, but got", count)
|
||||
|
||||
@@ -34,14 +34,14 @@ func testSelect(t *testing.T) {
|
||||
|
||||
func testUpdate(t *testing.T) {
|
||||
var clause Clause
|
||||
clause.Set(UPDATE, "User", map[string]interface{}{"Age": 30, "Name": "Tommy"})
|
||||
clause.Set(UPDATE, "User", map[string]interface{}{"Age": 30})
|
||||
clause.Set(WHERE, "Name = ?", "Tom")
|
||||
sql, vars := clause.Build(UPDATE, WHERE)
|
||||
t.Log(sql, vars)
|
||||
if sql != "UPDATE User SET Age = ?, Name = ? WHERE Name = ?" {
|
||||
if sql != "UPDATE User SET Age = ? WHERE Name = ?" {
|
||||
t.Fatal("failed to build SQL")
|
||||
}
|
||||
if !reflect.DeepEqual(vars, []interface{}{30, "Tommy", "Tom"}) {
|
||||
if !reflect.DeepEqual(vars, []interface{}{30, "Tom"}) {
|
||||
t.Fatal("failed to build SQLVars")
|
||||
}
|
||||
}
|
||||
|
||||
@@ -29,7 +29,7 @@ func NewSession() *Session {
|
||||
func TestSession_Exec(t *testing.T) {
|
||||
s := NewSession()
|
||||
_, _ = s.Raw("DROP TABLE IF EXISTS User;").Exec()
|
||||
_, _ = s.Raw("CREATE TABLE User(name text);").Exec()
|
||||
_, _ = s.Raw("CREATE TABLE User(Name text);").Exec()
|
||||
result, _ := s.Raw("INSERT INTO User(`Name`) values (?), (?)", "Tom", "Sam").Exec()
|
||||
if count, err := result.RowsAffected(); err != nil || count != 2 {
|
||||
t.Fatal("expect 2, but got", count)
|
||||
|
||||
@@ -34,14 +34,14 @@ func testSelect(t *testing.T) {
|
||||
|
||||
func testUpdate(t *testing.T) {
|
||||
var clause Clause
|
||||
clause.Set(UPDATE, "User", map[string]interface{}{"Age": 30, "Name": "Tommy"})
|
||||
clause.Set(UPDATE, "User", map[string]interface{}{"Age": 30})
|
||||
clause.Set(WHERE, "Name = ?", "Tom")
|
||||
sql, vars := clause.Build(UPDATE, WHERE)
|
||||
t.Log(sql, vars)
|
||||
if sql != "UPDATE User SET Age = ?, Name = ? WHERE Name = ?" {
|
||||
if sql != "UPDATE User SET Age = ? WHERE Name = ?" {
|
||||
t.Fatal("failed to build SQL")
|
||||
}
|
||||
if !reflect.DeepEqual(vars, []interface{}{30, "Tommy", "Tom"}) {
|
||||
if !reflect.DeepEqual(vars, []interface{}{30, "Tom"}) {
|
||||
t.Fatal("failed to build SQLVars")
|
||||
}
|
||||
}
|
||||
|
||||
@@ -29,7 +29,7 @@ func NewSession() *Session {
|
||||
func TestSession_Exec(t *testing.T) {
|
||||
s := NewSession()
|
||||
_, _ = s.Raw("DROP TABLE IF EXISTS User;").Exec()
|
||||
_, _ = s.Raw("CREATE TABLE User(name text);").Exec()
|
||||
_, _ = s.Raw("CREATE TABLE User(Name text);").Exec()
|
||||
result, _ := s.Raw("INSERT INTO User(`Name`) values (?), (?)", "Tom", "Sam").Exec()
|
||||
if count, err := result.RowsAffected(); err != nil || count != 2 {
|
||||
t.Fatal("expect 2, but got", count)
|
||||
|
||||
@@ -34,14 +34,14 @@ func testSelect(t *testing.T) {
|
||||
|
||||
func testUpdate(t *testing.T) {
|
||||
var clause Clause
|
||||
clause.Set(UPDATE, "User", map[string]interface{}{"Age": 30, "Name": "Tommy"})
|
||||
clause.Set(UPDATE, "User", map[string]interface{}{"Age": 30})
|
||||
clause.Set(WHERE, "Name = ?", "Tom")
|
||||
sql, vars := clause.Build(UPDATE, WHERE)
|
||||
t.Log(sql, vars)
|
||||
if sql != "UPDATE User SET Age = ?, Name = ? WHERE Name = ?" {
|
||||
if sql != "UPDATE User SET Age = ? WHERE Name = ?" {
|
||||
t.Fatal("failed to build SQL")
|
||||
}
|
||||
if !reflect.DeepEqual(vars, []interface{}{30, "Tommy", "Tom"}) {
|
||||
if !reflect.DeepEqual(vars, []interface{}{30, "Tom"}) {
|
||||
t.Fatal("failed to build SQLVars")
|
||||
}
|
||||
}
|
||||
|
||||
@@ -29,7 +29,7 @@ func NewSession() *Session {
|
||||
func TestSession_Exec(t *testing.T) {
|
||||
s := NewSession()
|
||||
_, _ = s.Raw("DROP TABLE IF EXISTS User;").Exec()
|
||||
_, _ = s.Raw("CREATE TABLE User(name text);").Exec()
|
||||
_, _ = s.Raw("CREATE TABLE User(Name text);").Exec()
|
||||
result, _ := s.Raw("INSERT INTO User(`Name`) values (?), (?)", "Tom", "Sam").Exec()
|
||||
if count, err := result.RowsAffected(); err != nil || count != 2 {
|
||||
t.Fatal("expect 2, but got", count)
|
||||
|
||||
@@ -0,0 +1,51 @@
|
||||
package clause
|
||||
|
||||
import (
|
||||
"strings"
|
||||
)
|
||||
|
||||
// Clause contains SQL conditions
|
||||
type Clause struct {
|
||||
sql map[Type]string
|
||||
sqlVars map[Type][]interface{}
|
||||
}
|
||||
|
||||
// Type is the type of Clause
|
||||
type Type int
|
||||
|
||||
// Support types for Clause
|
||||
const (
|
||||
INSERT Type = iota
|
||||
VALUES
|
||||
SELECT
|
||||
LIMIT
|
||||
WHERE
|
||||
ORDERBY
|
||||
UPDATE
|
||||
DELETE
|
||||
COUNT
|
||||
)
|
||||
|
||||
// Set adds a sub clause of specific type
|
||||
func (c *Clause) Set(name Type, vars ...interface{}) {
|
||||
if c.sql == nil {
|
||||
c.sql = make(map[Type]string)
|
||||
c.sqlVars = make(map[Type][]interface{})
|
||||
}
|
||||
sql, vars := generators[name](vars...)
|
||||
c.sql[name] = sql
|
||||
c.sqlVars[name] = vars
|
||||
}
|
||||
|
||||
// Build generate the final SQL and SQLVars
|
||||
func (c *Clause) Build(orders ...Type) (string, []interface{}) {
|
||||
var sqls []string
|
||||
var vars []interface{}
|
||||
for _, order := range orders {
|
||||
if sql, ok := c.sql[order]; ok {
|
||||
sqls = append(sqls, sql)
|
||||
vars = append(vars, c.sqlVars[order]...)
|
||||
}
|
||||
}
|
||||
return strings.Join(sqls, " "), vars
|
||||
}
|
||||
@@ -0,0 +1,74 @@
|
||||
package clause
|
||||
|
||||
import (
|
||||
"reflect"
|
||||
"testing"
|
||||
)
|
||||
|
||||
func TestClause_Set(t *testing.T) {
|
||||
var clause Clause
|
||||
clause.Set(INSERT, "User", []string{"Name", "Age"})
|
||||
sql := clause.sql[INSERT]
|
||||
vars := clause.sqlVars[INSERT]
|
||||
t.Log(sql, vars)
|
||||
if sql != "INSERT INTO User (Name,Age)" || len(vars) != 0 {
|
||||
t.Fatal("failed to get clause")
|
||||
}
|
||||
}
|
||||
|
||||
func testSelect(t *testing.T) {
|
||||
var clause Clause
|
||||
clause.Set(LIMIT, 3)
|
||||
clause.Set(SELECT, "User", []string{"*"})
|
||||
clause.Set(WHERE, "Name = ?", "Tom")
|
||||
clause.Set(ORDERBY, "Age ASC")
|
||||
sql, vars := clause.Build(SELECT, WHERE, ORDERBY, LIMIT)
|
||||
t.Log(sql, vars)
|
||||
if sql != "SELECT * FROM User WHERE Name = ? ORDER BY Age ASC LIMIT ?" {
|
||||
t.Fatal("failed to build SQL")
|
||||
}
|
||||
if !reflect.DeepEqual(vars, []interface{}{"Tom", 3}) {
|
||||
t.Fatal("failed to build SQLVars")
|
||||
}
|
||||
}
|
||||
|
||||
func testUpdate(t *testing.T) {
|
||||
var clause Clause
|
||||
clause.Set(UPDATE, "User", map[string]interface{}{"Age": 30})
|
||||
clause.Set(WHERE, "Name = ?", "Tom")
|
||||
sql, vars := clause.Build(UPDATE, WHERE)
|
||||
t.Log(sql, vars)
|
||||
if sql != "UPDATE User SET Age = ? WHERE Name = ?" {
|
||||
t.Fatal("failed to build SQL")
|
||||
}
|
||||
if !reflect.DeepEqual(vars, []interface{}{30, "Tom"}) {
|
||||
t.Fatal("failed to build SQLVars")
|
||||
}
|
||||
}
|
||||
|
||||
func testDelete(t *testing.T) {
|
||||
var clause Clause
|
||||
clause.Set(DELETE, "User")
|
||||
clause.Set(WHERE, "Name = ?", "Tom")
|
||||
|
||||
sql, vars := clause.Build(DELETE, WHERE)
|
||||
t.Log(sql, vars)
|
||||
if sql != "DELETE FROM User WHERE Name = ?" {
|
||||
t.Fatal("failed to build SQL")
|
||||
}
|
||||
if !reflect.DeepEqual(vars, []interface{}{"Tom"}) {
|
||||
t.Fatal("failed to build SQLVars")
|
||||
}
|
||||
}
|
||||
|
||||
func TestClause_Build(t *testing.T) {
|
||||
t.Run("select", func(t *testing.T) {
|
||||
testSelect(t)
|
||||
})
|
||||
t.Run("update", func(t *testing.T) {
|
||||
testUpdate(t)
|
||||
})
|
||||
t.Run("delete", func(t *testing.T) {
|
||||
testDelete(t)
|
||||
})
|
||||
}
|
||||
@@ -0,0 +1,101 @@
|
||||
package clause
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"strings"
|
||||
)
|
||||
|
||||
type generator func(values ...interface{}) (string, []interface{})
|
||||
|
||||
var generators map[Type]generator
|
||||
|
||||
func init() {
|
||||
generators = make(map[Type]generator)
|
||||
generators[INSERT] = _insert
|
||||
generators[VALUES] = _values
|
||||
generators[SELECT] = _select
|
||||
generators[LIMIT] = _limit
|
||||
generators[WHERE] = _where
|
||||
generators[ORDERBY] = _orderBy
|
||||
generators[UPDATE] = _update
|
||||
generators[DELETE] = _delete
|
||||
generators[COUNT] = _count
|
||||
}
|
||||
|
||||
func genBindVars(num int) string {
|
||||
var vars []string
|
||||
for i := 0; i < num; i++ {
|
||||
vars = append(vars, "?")
|
||||
}
|
||||
return strings.Join(vars, ", ")
|
||||
}
|
||||
|
||||
func _insert(values ...interface{}) (string, []interface{}) {
|
||||
// INSERT INTO $tableName ($fields)
|
||||
tableName := values[0]
|
||||
fields := strings.Join(values[1].([]string), ",")
|
||||
return fmt.Sprintf("INSERT INTO %s (%v)", tableName, fields), []interface{}{}
|
||||
}
|
||||
|
||||
func _values(values ...interface{}) (string, []interface{}) {
|
||||
// VALUES ($v1), (&v2), ...
|
||||
var bindStr string
|
||||
var sql strings.Builder
|
||||
var vars []interface{}
|
||||
sql.WriteString("VALUES ")
|
||||
for i, value := range values {
|
||||
v := value.([]interface{})
|
||||
if bindStr == "" {
|
||||
bindStr = genBindVars(len(v))
|
||||
}
|
||||
sql.WriteString(fmt.Sprintf("(%v)", bindStr))
|
||||
if i+1 != len(values) {
|
||||
sql.WriteString(", ")
|
||||
}
|
||||
vars = append(vars, v...)
|
||||
}
|
||||
return sql.String(), vars
|
||||
|
||||
}
|
||||
|
||||
func _select(values ...interface{}) (string, []interface{}) {
|
||||
// SELECT $fields FROM $tableName
|
||||
tableName := values[0]
|
||||
fields := strings.Join(values[1].([]string), ",")
|
||||
return fmt.Sprintf("SELECT %v FROM %s", fields, tableName), []interface{}{}
|
||||
}
|
||||
|
||||
func _limit(values ...interface{}) (string, []interface{}) {
|
||||
// LIMIT $num
|
||||
return "LIMIT ?", values
|
||||
}
|
||||
|
||||
func _where(values ...interface{}) (string, []interface{}) {
|
||||
// WHERE $desc
|
||||
desc, vars := values[0], values[1:]
|
||||
return fmt.Sprintf("WHERE %s", desc), vars
|
||||
}
|
||||
|
||||
func _orderBy(values ...interface{}) (string, []interface{}) {
|
||||
return fmt.Sprintf("ORDER BY %s", values[0]), []interface{}{}
|
||||
}
|
||||
|
||||
func _update(values ...interface{}) (string, []interface{}) {
|
||||
tableName := values[0]
|
||||
m := values[1].(map[string]interface{})
|
||||
var keys []string
|
||||
var vars []interface{}
|
||||
for k, v := range m {
|
||||
keys = append(keys, k+" = ?")
|
||||
vars = append(vars, v)
|
||||
}
|
||||
return fmt.Sprintf("UPDATE %s SET %s", tableName, strings.Join(keys, ", ")), vars
|
||||
}
|
||||
|
||||
func _delete(values ...interface{}) (string, []interface{}) {
|
||||
return fmt.Sprintf("DELETE FROM %s", values[0]), []interface{}{}
|
||||
}
|
||||
|
||||
func _count(values ...interface{}) (string, []interface{}) {
|
||||
return _select(values[0], []string{"count(*)"})
|
||||
}
|
||||
@@ -0,0 +1,22 @@
|
||||
package dialect
|
||||
|
||||
import "reflect"
|
||||
|
||||
var dialectsMap = map[string]Dialect{}
|
||||
|
||||
// Dialect is an interface contains methods that a dialect has to implement
|
||||
type Dialect interface {
|
||||
DataTypeOf(typ reflect.Value) string
|
||||
TableExistSQL(tableName string) (string, []interface{})
|
||||
}
|
||||
|
||||
// RegisterDialect register a dialect to the global variable
|
||||
func RegisterDialect(name string, dialect Dialect) {
|
||||
dialectsMap[name] = dialect
|
||||
}
|
||||
|
||||
// Get the dialect from global variable if it exists
|
||||
func GetDialect(name string) (dialect Dialect, ok bool) {
|
||||
dialect, ok = dialectsMap[name]
|
||||
return
|
||||
}
|
||||
@@ -0,0 +1,45 @@
|
||||
package dialect
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"reflect"
|
||||
"time"
|
||||
)
|
||||
|
||||
type sqlite3 struct{}
|
||||
|
||||
var _ Dialect = (*sqlite3)(nil)
|
||||
|
||||
func init() {
|
||||
RegisterDialect("sqlite3", &sqlite3{})
|
||||
}
|
||||
|
||||
// Get Data Type for sqlite3 Dialect
|
||||
func (s *sqlite3) DataTypeOf(typ reflect.Value) string {
|
||||
switch typ.Kind() {
|
||||
case reflect.Bool:
|
||||
return "bool"
|
||||
case reflect.Int, reflect.Int8, reflect.Int16, reflect.Int32,
|
||||
reflect.Uint, reflect.Uint8, reflect.Uint16, reflect.Uint32, reflect.Uintptr:
|
||||
return "integer"
|
||||
case reflect.Int64, reflect.Uint64:
|
||||
return "bigint"
|
||||
case reflect.Float32, reflect.Float64:
|
||||
return "real"
|
||||
case reflect.String:
|
||||
return "text"
|
||||
case reflect.Array, reflect.Slice:
|
||||
return "blob"
|
||||
case reflect.Struct:
|
||||
if _, ok := typ.Interface().(time.Time); ok {
|
||||
return "datetime"
|
||||
}
|
||||
}
|
||||
panic(fmt.Sprintf("invalid sql type %s (%s)", typ.Type().Name(), typ.Kind()))
|
||||
}
|
||||
|
||||
// TableExistSQL returns SQL that judge whether the table exists in database
|
||||
func (s *sqlite3) TableExistSQL(tableName string) (string, []interface{}) {
|
||||
args := []interface{}{tableName}
|
||||
return "SELECT name FROM sqlite_master WHERE type='table' and name = ?", args
|
||||
}
|
||||
@@ -0,0 +1,25 @@
|
||||
package dialect
|
||||
|
||||
import (
|
||||
"reflect"
|
||||
"testing"
|
||||
)
|
||||
|
||||
func TestDataTypeOf(t *testing.T) {
|
||||
dial := &sqlite3{}
|
||||
cases := []struct {
|
||||
Value interface{}
|
||||
Type string
|
||||
}{
|
||||
{"Tom", "text"},
|
||||
{123, "integer"},
|
||||
{1.2, "real"},
|
||||
{[]int{1, 2, 3}, "blob"},
|
||||
}
|
||||
|
||||
for _, c := range cases {
|
||||
if typ := dial.DataTypeOf(reflect.ValueOf(c.Value)); typ != c.Type {
|
||||
t.Fatalf("expect %s, but got %s", c.Type, typ)
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,127 @@
|
||||
package geeorm
|
||||
|
||||
import (
|
||||
"database/sql"
|
||||
"fmt"
|
||||
"geeorm/dialect"
|
||||
"geeorm/log"
|
||||
"geeorm/session"
|
||||
"strings"
|
||||
)
|
||||
|
||||
// Engine is the main struct of geeorm, manages all db sessions and transactions.
|
||||
type Engine struct {
|
||||
db *sql.DB
|
||||
dialect dialect.Dialect
|
||||
}
|
||||
|
||||
// NewEngine create a instance of Engine
|
||||
// connect database and ping it to test whether it's alive
|
||||
func NewEngine(driver, source string) (e *Engine, err error) {
|
||||
db, err := sql.Open(driver, source)
|
||||
if err != nil {
|
||||
log.Error(err)
|
||||
return
|
||||
}
|
||||
// Send a ping to make sure the database connection is alive.
|
||||
if err = db.Ping(); err != nil {
|
||||
log.Error(err)
|
||||
return
|
||||
}
|
||||
// make sure the specific dialect exists
|
||||
dial, ok := dialect.GetDialect(driver)
|
||||
if !ok {
|
||||
log.Errorf("dialect %s Not Found", driver)
|
||||
return
|
||||
}
|
||||
e = &Engine{db: db, dialect: dial}
|
||||
log.Info("Connect database success")
|
||||
return
|
||||
}
|
||||
|
||||
// Close database connection
|
||||
func (engine *Engine) Close() {
|
||||
if err := engine.db.Close(); err != nil {
|
||||
log.Error("Failed to close database")
|
||||
}
|
||||
log.Info("Close database success")
|
||||
}
|
||||
|
||||
// NewSession creates a new session for next operations
|
||||
func (engine *Engine) NewSession() *session.Session {
|
||||
return session.New(engine.db, engine.dialect)
|
||||
}
|
||||
|
||||
// TxFunc will be called between tx.Begin() and tx.Commit()
|
||||
// https://stackoverflow.com/questions/16184238/database-sql-tx-detecting-commit-or-rollback
|
||||
type TxFunc func(*session.Session) (interface{}, error)
|
||||
|
||||
// Transaction executes sql wrapped in a transaction, then automatically commit if no error occurs
|
||||
func (engine *Engine) Transaction(f TxFunc) (result interface{}, err error) {
|
||||
s := engine.NewSession()
|
||||
if err := s.Begin(); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
defer func() {
|
||||
if p := recover(); p != nil {
|
||||
_ = s.Rollback()
|
||||
panic(p) // re-throw panic after Rollback
|
||||
} else if err != nil {
|
||||
_ = s.Rollback() // err is non-nil; don't change it
|
||||
} else {
|
||||
err = s.Commit() // err is nil; if Commit returns error update err
|
||||
}
|
||||
}()
|
||||
|
||||
return f(s)
|
||||
}
|
||||
|
||||
// difference returns a - b
|
||||
func difference(a []string, b []string) (diff []string) {
|
||||
mapB := make(map[string]bool)
|
||||
for _, v := range b {
|
||||
mapB[v] = true
|
||||
}
|
||||
for _, v := range a {
|
||||
if _, ok := mapB[v]; !ok {
|
||||
diff = append(diff, v)
|
||||
}
|
||||
}
|
||||
return
|
||||
}
|
||||
|
||||
// Migrate table
|
||||
func (engine *Engine) Migrate(value interface{}) error {
|
||||
_, err := engine.Transaction(func(s *session.Session) (result interface{}, err error) {
|
||||
if !s.Model(value).HasTable() {
|
||||
log.Infof("table %s doesn't exist", s.RefTable().Name)
|
||||
return nil, s.CreateTable()
|
||||
}
|
||||
table := s.RefTable()
|
||||
rows, _ := s.Raw(fmt.Sprintf("SELECT * FROM %s LIMIT 1", table.Name)).QueryRows()
|
||||
columns, _ := rows.Columns()
|
||||
addCols := difference(table.FieldNames, columns)
|
||||
delCols := difference(columns, table.FieldNames)
|
||||
log.Infof("added cols %v, deleted cols %v", addCols, delCols)
|
||||
|
||||
for _, col := range addCols {
|
||||
f := table.GetField(col)
|
||||
sqlStr := fmt.Sprintf("ALTER TABLE %s ADD COLUMN %s %s;", table.Name, f.Name, f.Tag)
|
||||
if _, err = s.Raw(sqlStr).Exec(); err != nil {
|
||||
return
|
||||
}
|
||||
}
|
||||
|
||||
if len(delCols) == 0 {
|
||||
return
|
||||
}
|
||||
tmp := "tmp_" + table.Name
|
||||
fieldStr := strings.Join(table.FieldNames, ", ")
|
||||
s.Raw(fmt.Sprintf("CREATE TABLE %s AS SELECT %s from %s;", tmp, fieldStr, table.Name))
|
||||
s.Raw(fmt.Sprintf("DROP TABLE %s;", table.Name))
|
||||
s.Raw(fmt.Sprintf("ALTER TABLE %s RENAME TO %s;", tmp, table.Name))
|
||||
_, err = s.Exec()
|
||||
return
|
||||
})
|
||||
return err
|
||||
}
|
||||
@@ -0,0 +1,86 @@
|
||||
package geeorm
|
||||
|
||||
import (
|
||||
"errors"
|
||||
"geeorm/session"
|
||||
"reflect"
|
||||
"testing"
|
||||
|
||||
_ "github.com/mattn/go-sqlite3"
|
||||
)
|
||||
|
||||
func OpenDB(t *testing.T) *Engine {
|
||||
t.Helper()
|
||||
engine, err := NewEngine("sqlite3", "gee.db")
|
||||
if err != nil {
|
||||
t.Fatal("failed to connect", err)
|
||||
}
|
||||
return engine
|
||||
}
|
||||
|
||||
func TestNewEngine(t *testing.T) {
|
||||
engine := OpenDB(t)
|
||||
defer engine.Close()
|
||||
}
|
||||
|
||||
type User struct {
|
||||
Name string `geeorm:"PRIMARY KEY"`
|
||||
Age int
|
||||
}
|
||||
|
||||
func transactionRollback(t *testing.T) {
|
||||
engine := OpenDB(t)
|
||||
defer engine.Close()
|
||||
s := engine.NewSession()
|
||||
_ = s.Model(&User{}).DropTable()
|
||||
_, err := engine.Transaction(func(s *session.Session) (result interface{}, err error) {
|
||||
_ = s.Model(&User{}).CreateTable()
|
||||
_, err = s.Insert(&User{"Tom", 18})
|
||||
return nil, errors.New("Error")
|
||||
})
|
||||
if err == nil || s.HasTable() {
|
||||
t.Fatal("failed to rollback")
|
||||
}
|
||||
}
|
||||
|
||||
func transactionCommit(t *testing.T) {
|
||||
engine := OpenDB(t)
|
||||
defer engine.Close()
|
||||
s := engine.NewSession()
|
||||
_ = s.Model(&User{}).DropTable()
|
||||
_, err := engine.Transaction(func(s *session.Session) (result interface{}, err error) {
|
||||
_ = s.Model(&User{}).CreateTable()
|
||||
_, err = s.Insert(&User{"Tom", 18})
|
||||
return
|
||||
})
|
||||
u := &User{}
|
||||
_ = s.First(u)
|
||||
if err != nil || u.Name != "Tom" {
|
||||
t.Fatal("failed to commit")
|
||||
}
|
||||
}
|
||||
|
||||
func TestEngine_Transaction(t *testing.T) {
|
||||
t.Run("rollback", func(t *testing.T) {
|
||||
transactionRollback(t)
|
||||
})
|
||||
t.Run("commit", func(t *testing.T) {
|
||||
transactionCommit(t)
|
||||
})
|
||||
}
|
||||
|
||||
func TestEngine_Migrate(t *testing.T) {
|
||||
engine := OpenDB(t)
|
||||
defer engine.Close()
|
||||
s := engine.NewSession()
|
||||
_, _ = s.Raw("DROP TABLE IF EXISTS User;").Exec()
|
||||
_, _ = s.Raw("CREATE TABLE User(Name text, XXX integer);").Exec()
|
||||
_, _ = s.Raw("INSERT INTO User(`Name`) values (?), (?)", "Tom", "Sam").Exec()
|
||||
engine.Migrate(&User{})
|
||||
|
||||
rows, _ := s.Raw("SELECT * FROM User").QueryRows()
|
||||
columns, _ := rows.Columns()
|
||||
if !reflect.DeepEqual(columns, []string{"Name", "Age"}) {
|
||||
t.Fatal("Failed to migrate table User, got columns", columns)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,5 @@
|
||||
module geeorm
|
||||
|
||||
go 1.13
|
||||
|
||||
require github.com/mattn/go-sqlite3 v2.0.3+incompatible
|
||||
@@ -0,0 +1,47 @@
|
||||
package log
|
||||
|
||||
import (
|
||||
"io/ioutil"
|
||||
"log"
|
||||
"os"
|
||||
"sync"
|
||||
)
|
||||
|
||||
var (
|
||||
errorLog = log.New(os.Stdout, "\033[31m[error]\033[0m ", log.LstdFlags|log.Lshortfile)
|
||||
infoLog = log.New(os.Stdout, "\033[34m[info ]\033[0m ", log.LstdFlags|log.Lshortfile)
|
||||
loggers = []*log.Logger{errorLog, infoLog}
|
||||
mu sync.Mutex
|
||||
)
|
||||
|
||||
// log methods
|
||||
var (
|
||||
Error = errorLog.Println
|
||||
Errorf = errorLog.Printf
|
||||
Info = infoLog.Println
|
||||
Infof = infoLog.Printf
|
||||
)
|
||||
|
||||
// log levels
|
||||
const (
|
||||
InfoLevel = iota
|
||||
ErrorLevel
|
||||
Disabled
|
||||
)
|
||||
|
||||
// SetLevel controls log level
|
||||
func SetLevel(level int) {
|
||||
mu.Lock()
|
||||
defer mu.Unlock()
|
||||
|
||||
for _, logger := range loggers {
|
||||
logger.SetOutput(os.Stdout)
|
||||
}
|
||||
|
||||
if ErrorLevel < level {
|
||||
errorLog.SetOutput(ioutil.Discard)
|
||||
}
|
||||
if InfoLevel < level {
|
||||
infoLog.SetOutput(ioutil.Discard)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,17 @@
|
||||
package log
|
||||
|
||||
import (
|
||||
"os"
|
||||
"testing"
|
||||
)
|
||||
|
||||
func TestSetLevel(t *testing.T) {
|
||||
SetLevel(ErrorLevel)
|
||||
if infoLog.Writer() == os.Stdout || errorLog.Writer() != os.Stdout {
|
||||
t.Fatal("failed to set log level")
|
||||
}
|
||||
SetLevel(Disabled)
|
||||
if infoLog.Writer() == os.Stdout || errorLog.Writer() == os.Stdout {
|
||||
t.Fatal("failed to set log level")
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,65 @@
|
||||
package schema
|
||||
|
||||
import (
|
||||
"geeorm/dialect"
|
||||
"go/ast"
|
||||
"reflect"
|
||||
)
|
||||
|
||||
// Field represents a column of database
|
||||
type Field struct {
|
||||
Name string
|
||||
Type string
|
||||
Tag string
|
||||
}
|
||||
|
||||
// Schema represents a table of database
|
||||
type Schema struct {
|
||||
Model interface{}
|
||||
Name string
|
||||
Fields []*Field
|
||||
FieldNames []string
|
||||
fieldMap map[string]*Field
|
||||
}
|
||||
|
||||
// GetField returns field by name
|
||||
func (schema *Schema) GetField(name string) *Field {
|
||||
return schema.fieldMap[name]
|
||||
}
|
||||
|
||||
// Values return the values of dest's member variables
|
||||
func (schema *Schema) RecordValues(dest interface{}) []interface{} {
|
||||
destValue := reflect.Indirect(reflect.ValueOf(dest))
|
||||
var fieldValues []interface{}
|
||||
for _, field := range schema.Fields {
|
||||
fieldValues = append(fieldValues, destValue.FieldByName(field.Name).Interface())
|
||||
}
|
||||
return fieldValues
|
||||
}
|
||||
|
||||
// Parse a struct to a Schema instance
|
||||
func Parse(dest interface{}, d dialect.Dialect) *Schema {
|
||||
modelType := reflect.Indirect(reflect.ValueOf(dest)).Type()
|
||||
schema := &Schema{
|
||||
Model: dest,
|
||||
Name: modelType.Name(),
|
||||
fieldMap: make(map[string]*Field),
|
||||
}
|
||||
|
||||
for i := 0; i < modelType.NumField(); i++ {
|
||||
p := modelType.Field(i)
|
||||
if !p.Anonymous && ast.IsExported(p.Name) {
|
||||
field := &Field{
|
||||
Name: p.Name,
|
||||
Type: d.DataTypeOf(reflect.Indirect(reflect.New(p.Type))),
|
||||
}
|
||||
if v, ok := p.Tag.Lookup("geeorm"); ok {
|
||||
field.Tag = v
|
||||
}
|
||||
schema.Fields = append(schema.Fields, field)
|
||||
schema.FieldNames = append(schema.FieldNames, p.Name)
|
||||
schema.fieldMap[p.Name] = field
|
||||
}
|
||||
}
|
||||
return schema
|
||||
}
|
||||
@@ -0,0 +1,35 @@
|
||||
package schema
|
||||
|
||||
import (
|
||||
"geeorm/dialect"
|
||||
"testing"
|
||||
)
|
||||
|
||||
type User struct {
|
||||
Name string `geeorm:"PRIMARY KEY"`
|
||||
Age int
|
||||
}
|
||||
|
||||
var TestDial, _ = dialect.GetDialect("sqlite3")
|
||||
|
||||
func TestParse(t *testing.T) {
|
||||
schema := Parse(&User{}, TestDial)
|
||||
if schema.Name != "User" || len(schema.Fields) != 2 {
|
||||
t.Fatal("failed to parse User struct")
|
||||
}
|
||||
if schema.GetField("Name").Tag != "PRIMARY KEY" {
|
||||
t.Fatal("failed to parse primary key")
|
||||
}
|
||||
}
|
||||
|
||||
func TestSchema_RecordValues(t *testing.T) {
|
||||
schema := Parse(&User{}, TestDial)
|
||||
values := schema.RecordValues(&User{"Tom", 18})
|
||||
|
||||
name := values[0].(string)
|
||||
age := values[1].(int)
|
||||
|
||||
if name != "Tom" || age != 18 {
|
||||
t.Fatal("failed to get values")
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,32 @@
|
||||
package session
|
||||
|
||||
import "reflect"
|
||||
|
||||
// Hooks constants
|
||||
const (
|
||||
BeforeQuery = "BeforeQuery"
|
||||
AfterQuery = "AfterQuery"
|
||||
BeforeUpdate = "BeforeUpdate"
|
||||
AfterUpdate = "AfterUpate"
|
||||
BeforeDelete = "BeforeDelete"
|
||||
AfterDelete = "AfterDelete"
|
||||
BeforeInsert = "BeforeInsert"
|
||||
AfterInsert = "AfterInsert"
|
||||
)
|
||||
|
||||
// CallMethod calls the registered hooks
|
||||
func (s *Session) CallMethod(method string, value interface{}) error {
|
||||
fm := reflect.ValueOf(s.RefTable().Model).MethodByName(method)
|
||||
if value != nil {
|
||||
fm = reflect.ValueOf(value).MethodByName(method)
|
||||
}
|
||||
param := []reflect.Value{reflect.ValueOf(s)}
|
||||
if fm.IsValid() {
|
||||
if v := fm.Call(param); len(v) > 0 {
|
||||
if err, ok := v[0].Interface().(error); ok {
|
||||
return err
|
||||
}
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
@@ -0,0 +1,37 @@
|
||||
package session
|
||||
|
||||
import (
|
||||
"geeorm/log"
|
||||
"testing"
|
||||
)
|
||||
|
||||
type Account struct {
|
||||
ID int `geeorm:"PRIMARY KEY"`
|
||||
Password string
|
||||
}
|
||||
|
||||
func (account *Account) BeforeInsert(s *Session) error {
|
||||
log.Info("before inert", account)
|
||||
account.ID += 1000
|
||||
return nil
|
||||
}
|
||||
|
||||
func (account *Account) AfterQuery(s *Session) error {
|
||||
log.Info("after query", account)
|
||||
account.Password = "******"
|
||||
return nil
|
||||
}
|
||||
|
||||
func TestSession_CallMethod(t *testing.T) {
|
||||
s := NewSession().Model(&Account{})
|
||||
_ = s.DropTable()
|
||||
_ = s.CreateTable()
|
||||
_, _ = s.Insert(&Account{1, "123456"}, &Account{2, "qwerty"})
|
||||
|
||||
u := &Account{}
|
||||
|
||||
err := s.First(u)
|
||||
if err != nil || u.ID != 1001 || u.Password != "******" {
|
||||
t.Fatal("Failed to call hooks after query, got", u)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,90 @@
|
||||
package session
|
||||
|
||||
import (
|
||||
"database/sql"
|
||||
"geeorm/clause"
|
||||
"geeorm/dialect"
|
||||
"geeorm/log"
|
||||
"geeorm/schema"
|
||||
"strings"
|
||||
)
|
||||
|
||||
// Session keep a pointer to sql.DB and provides all execution of all
|
||||
// kind of database operations.
|
||||
type Session struct {
|
||||
db *sql.DB
|
||||
dialect dialect.Dialect
|
||||
tx *sql.Tx
|
||||
refTable *schema.Schema
|
||||
clause clause.Clause
|
||||
sql strings.Builder
|
||||
sqlVars []interface{}
|
||||
}
|
||||
|
||||
// New creates a instance of Session
|
||||
func New(db *sql.DB, dialect dialect.Dialect) *Session {
|
||||
return &Session{
|
||||
db: db,
|
||||
dialect: dialect,
|
||||
}
|
||||
}
|
||||
|
||||
// Clear initialize the state of a session
|
||||
func (s *Session) Clear() {
|
||||
s.sql.Reset()
|
||||
s.sqlVars = nil
|
||||
s.clause = clause.Clause{}
|
||||
}
|
||||
|
||||
// CommonDB is a minimal function set of db
|
||||
type CommonDB interface {
|
||||
Query(query string, args ...interface{}) (*sql.Rows, error)
|
||||
QueryRow(query string, args ...interface{}) *sql.Row
|
||||
Exec(query string, args ...interface{}) (sql.Result, error)
|
||||
}
|
||||
|
||||
var _ CommonDB = (*sql.DB)(nil)
|
||||
var _ CommonDB = (*sql.Tx)(nil)
|
||||
|
||||
// DB returns tx if a tx begins. otherwise return *sql.DB
|
||||
func (s *Session) DB() CommonDB {
|
||||
if s.tx != nil {
|
||||
return s.tx
|
||||
}
|
||||
return s.db
|
||||
}
|
||||
|
||||
// Exec raw sql with sqlVars
|
||||
func (s *Session) Exec() (result sql.Result, err error) {
|
||||
defer s.Clear()
|
||||
log.Info(s.sql.String(), s.sqlVars)
|
||||
if result, err = s.DB().Exec(s.sql.String(), s.sqlVars...); err != nil {
|
||||
log.Error(err)
|
||||
}
|
||||
return
|
||||
}
|
||||
|
||||
// QueryRow gets a record from db
|
||||
func (s *Session) QueryRow() *sql.Row {
|
||||
defer s.Clear()
|
||||
log.Info(s.sql.String(), s.sqlVars)
|
||||
return s.DB().QueryRow(s.sql.String(), s.sqlVars...)
|
||||
}
|
||||
|
||||
// QueryRows gets a list of records from db
|
||||
func (s *Session) QueryRows() (rows *sql.Rows, err error) {
|
||||
defer s.Clear()
|
||||
log.Info(s.sql.String(), s.sqlVars)
|
||||
if rows, err = s.DB().Query(s.sql.String(), s.sqlVars...); err != nil {
|
||||
log.Error(err)
|
||||
}
|
||||
return
|
||||
}
|
||||
|
||||
// Raw appends sql and sqlVars
|
||||
func (s *Session) Raw(sql string, values ...interface{}) *Session {
|
||||
s.sql.WriteString(sql)
|
||||
s.sql.WriteString(" ")
|
||||
s.sqlVars = append(s.sqlVars, values...)
|
||||
return s
|
||||
}
|
||||
@@ -0,0 +1,48 @@
|
||||
package session
|
||||
|
||||
import (
|
||||
"database/sql"
|
||||
"os"
|
||||
"testing"
|
||||
|
||||
"geeorm/dialect"
|
||||
|
||||
_ "github.com/mattn/go-sqlite3"
|
||||
)
|
||||
|
||||
var (
|
||||
TestDB *sql.DB
|
||||
TestDial, _ = dialect.GetDialect("sqlite3")
|
||||
)
|
||||
|
||||
func TestMain(m *testing.M) {
|
||||
TestDB, _ = sql.Open("sqlite3", "../gee.db")
|
||||
code := m.Run()
|
||||
_ = TestDB.Close()
|
||||
os.Exit(code)
|
||||
}
|
||||
|
||||
func NewSession() *Session {
|
||||
return &Session{db: TestDB, dialect: TestDial}
|
||||
}
|
||||
|
||||
func TestSession_Exec(t *testing.T) {
|
||||
s := NewSession()
|
||||
_, _ = s.Raw("DROP TABLE IF EXISTS User;").Exec()
|
||||
_, _ = s.Raw("CREATE TABLE User(Name text);").Exec()
|
||||
result, _ := s.Raw("INSERT INTO User(`Name`) values (?), (?)", "Tom", "Sam").Exec()
|
||||
if count, err := result.RowsAffected(); err != nil || count != 2 {
|
||||
t.Fatal("expect 2, but got", count)
|
||||
}
|
||||
}
|
||||
|
||||
func TestSession_QueryRows(t *testing.T) {
|
||||
s := NewSession()
|
||||
_, _ = s.Raw("DROP TABLE IF EXISTS User;").Exec()
|
||||
_, _ = s.Raw("CREATE TABLE User(Name text);").Exec()
|
||||
row := s.Raw("SELECT count(*) FROM User").QueryRow()
|
||||
var count int
|
||||
if err := row.Scan(&count); err != nil || count != 0 {
|
||||
t.Fatal("failed to query db", err)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,136 @@
|
||||
package session
|
||||
|
||||
import (
|
||||
"errors"
|
||||
"geeorm/clause"
|
||||
"reflect"
|
||||
)
|
||||
|
||||
// Insert one or more records in database
|
||||
func (s *Session) Insert(values ...interface{}) (int64, error) {
|
||||
recordValues := make([]interface{}, 0)
|
||||
for _, value := range values {
|
||||
s.CallMethod(BeforeInsert, value)
|
||||
table := s.Model(value).RefTable()
|
||||
s.clause.Set(clause.INSERT, table.Name, table.FieldNames)
|
||||
recordValues = append(recordValues, table.RecordValues(value))
|
||||
}
|
||||
|
||||
s.clause.Set(clause.VALUES, recordValues...)
|
||||
sql, vars := s.clause.Build(clause.INSERT, clause.VALUES)
|
||||
result, err := s.Raw(sql, vars...).Exec()
|
||||
if err != nil {
|
||||
return 0, err
|
||||
}
|
||||
s.CallMethod(AfterInsert, nil)
|
||||
return result.RowsAffected()
|
||||
}
|
||||
|
||||
// Find gets all eligible records
|
||||
func (s *Session) Find(values interface{}) error {
|
||||
s.CallMethod(BeforeQuery, nil)
|
||||
destSlice := reflect.Indirect(reflect.ValueOf(values))
|
||||
destType := destSlice.Type().Elem()
|
||||
table := s.Model(reflect.New(destType).Elem().Interface()).RefTable()
|
||||
|
||||
s.clause.Set(clause.SELECT, table.Name, table.FieldNames)
|
||||
sql, vars := s.clause.Build(clause.SELECT, clause.WHERE, clause.ORDERBY, clause.LIMIT)
|
||||
rows, err := s.Raw(sql, vars...).QueryRows()
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
for rows.Next() {
|
||||
dest := reflect.New(destType).Elem()
|
||||
var values []interface{}
|
||||
for _, name := range table.FieldNames {
|
||||
values = append(values, dest.FieldByName(name).Addr().Interface())
|
||||
}
|
||||
if err := rows.Scan(values...); err != nil {
|
||||
return err
|
||||
}
|
||||
s.CallMethod(AfterQuery, dest.Addr().Interface())
|
||||
destSlice.Set(reflect.Append(destSlice, dest))
|
||||
}
|
||||
return rows.Close()
|
||||
}
|
||||
|
||||
// First gets the 1st row
|
||||
func (s *Session) First(value interface{}) error {
|
||||
dest := reflect.Indirect(reflect.ValueOf(value))
|
||||
destSlice := reflect.New(reflect.SliceOf(dest.Type())).Elem()
|
||||
if err := s.Limit(1).Find(destSlice.Addr().Interface()); err != nil {
|
||||
return err
|
||||
}
|
||||
if destSlice.Len() == 0 {
|
||||
return errors.New("NOT FOUND")
|
||||
}
|
||||
dest.Set(destSlice.Index(0))
|
||||
return nil
|
||||
}
|
||||
|
||||
// Limit adds limit condition to clause
|
||||
func (s *Session) Limit(num int) *Session {
|
||||
s.clause.Set(clause.LIMIT, num)
|
||||
return s
|
||||
}
|
||||
|
||||
// Where adds limit condition to clause
|
||||
func (s *Session) Where(desc string, args ...interface{}) *Session {
|
||||
var vars []interface{}
|
||||
s.clause.Set(clause.WHERE, append(append(vars, desc), args...)...)
|
||||
return s
|
||||
}
|
||||
|
||||
// OrderBy adds order by condition to clause
|
||||
func (s *Session) OrderBy(desc string) *Session {
|
||||
s.clause.Set(clause.ORDERBY, desc)
|
||||
return s
|
||||
}
|
||||
|
||||
// Update records with where clause
|
||||
// support map[string]interface{}
|
||||
// also support kv list: "Name", "Tom", "Age", 18, ....
|
||||
func (s *Session) Update(kv ...interface{}) (int64, error) {
|
||||
s.CallMethod(BeforeUpdate, nil)
|
||||
m, ok := kv[0].(map[string]interface{})
|
||||
if !ok {
|
||||
m = make(map[string]interface{})
|
||||
for i := 0; i < len(kv); i += 2 {
|
||||
m[kv[i].(string)] = kv[i+1]
|
||||
}
|
||||
}
|
||||
s.clause.Set(clause.UPDATE, s.RefTable().Name, m)
|
||||
sql, vars := s.clause.Build(clause.UPDATE, clause.WHERE)
|
||||
result, err := s.Raw(sql, vars...).Exec()
|
||||
if err != nil {
|
||||
return 0, err
|
||||
}
|
||||
s.CallMethod(AfterUpdate, nil)
|
||||
return result.RowsAffected()
|
||||
}
|
||||
|
||||
// Delete records with where clause
|
||||
func (s *Session) Delete() (int64, error) {
|
||||
s.CallMethod(BeforeDelete, nil)
|
||||
s.clause.Set(clause.DELETE, s.RefTable().Name)
|
||||
sql, vars := s.clause.Build(clause.DELETE, clause.WHERE)
|
||||
result, err := s.Raw(sql, vars...).Exec()
|
||||
if err != nil {
|
||||
return 0, err
|
||||
}
|
||||
s.CallMethod(AfterDelete, nil)
|
||||
return result.RowsAffected()
|
||||
}
|
||||
|
||||
// Count records with where clause
|
||||
func (s *Session) Count() (int64, error) {
|
||||
s.clause.Set(clause.COUNT, s.RefTable().Name)
|
||||
sql, vars := s.clause.Build(clause.COUNT, clause.WHERE)
|
||||
row := s.Raw(sql, vars...).QueryRow()
|
||||
var tmp int64
|
||||
if err := row.Scan(&tmp); err != nil {
|
||||
return 0, err
|
||||
}
|
||||
return tmp, nil
|
||||
}
|
||||
@@ -0,0 +1,97 @@
|
||||
package session
|
||||
|
||||
import "testing"
|
||||
|
||||
var (
|
||||
user1 = &User{"Tom", 18}
|
||||
user2 = &User{"Sam", 25}
|
||||
user3 = &User{"Jack", 25}
|
||||
)
|
||||
|
||||
func testRecordInit(t *testing.T) *Session {
|
||||
t.Helper()
|
||||
s := NewSession().Model(&User{})
|
||||
err1 := s.DropTable()
|
||||
err2 := s.CreateTable()
|
||||
_, err3 := s.Insert(user1, user2)
|
||||
if err1 != nil || err2 != nil || err3 != nil {
|
||||
t.Fatal("failed init test records")
|
||||
}
|
||||
return s
|
||||
}
|
||||
|
||||
func TestSession_Insert(t *testing.T) {
|
||||
s := testRecordInit(t)
|
||||
affected, err := s.Insert(user3)
|
||||
if err != nil || affected != 1 {
|
||||
t.Fatal("failed to create record")
|
||||
}
|
||||
}
|
||||
|
||||
func TestSession_Find(t *testing.T) {
|
||||
s := testRecordInit(t)
|
||||
var users []User
|
||||
if err := s.Find(&users); err != nil || len(users) != 2 {
|
||||
t.Fatal("failed to query all")
|
||||
}
|
||||
}
|
||||
|
||||
func TestSession_First(t *testing.T) {
|
||||
s := testRecordInit(t)
|
||||
u := &User{}
|
||||
err := s.First(u)
|
||||
if err != nil || u.Name != "Tom" || u.Age != 18 {
|
||||
t.Fatal("failed to query first")
|
||||
}
|
||||
}
|
||||
|
||||
func TestSession_Limit(t *testing.T) {
|
||||
s := testRecordInit(t)
|
||||
var users []User
|
||||
err := s.Limit(1).Find(&users)
|
||||
if err != nil || len(users) != 1 {
|
||||
t.Fatal("failed to query with limit condition")
|
||||
}
|
||||
}
|
||||
|
||||
func TestSession_Where(t *testing.T) {
|
||||
s := testRecordInit(t)
|
||||
var users []User
|
||||
_, err1 := s.Insert(user3)
|
||||
err2 := s.Where("Age = ?", 25).Find(&users)
|
||||
|
||||
if err1 != nil || err2 != nil || len(users) != 2 {
|
||||
t.Fatal("failed to query with where condition")
|
||||
}
|
||||
}
|
||||
|
||||
func TestSession_OrderBy(t *testing.T) {
|
||||
s := testRecordInit(t)
|
||||
u := &User{}
|
||||
err := s.OrderBy("Age DESC").First(u)
|
||||
|
||||
if err != nil || u.Age != 25 {
|
||||
t.Fatal("failed to query with order by condition")
|
||||
}
|
||||
}
|
||||
|
||||
func TestSession_Update(t *testing.T) {
|
||||
s := testRecordInit(t)
|
||||
affected, _ := s.Where("Name = ?", "Tom").Update("Age", 30)
|
||||
u := &User{}
|
||||
_ = s.OrderBy("Age DESC").First(u)
|
||||
|
||||
if affected != 1 || u.Age != 30 {
|
||||
t.Fatal("failed to update")
|
||||
}
|
||||
}
|
||||
|
||||
func TestSession_DeleteAndCount(t *testing.T) {
|
||||
s := testRecordInit(t)
|
||||
affected, _ := s.Where("Name = ?", "Tom").Delete()
|
||||
count, _ := s.Count()
|
||||
|
||||
if affected != 1 || count != 1 {
|
||||
t.Fatal("failed to delete or count")
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,54 @@
|
||||
package session
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"geeorm/log"
|
||||
"reflect"
|
||||
"strings"
|
||||
|
||||
"geeorm/schema"
|
||||
)
|
||||
|
||||
// Model assigns refTable
|
||||
func (s *Session) Model(value interface{}) *Session {
|
||||
// nil or different model, update refTable
|
||||
if s.refTable == nil || reflect.TypeOf(value) != reflect.TypeOf(s.refTable.Model) {
|
||||
s.refTable = schema.Parse(value, s.dialect)
|
||||
}
|
||||
return s
|
||||
}
|
||||
|
||||
// RefTable returns a Schema instance that contains all parsed fields
|
||||
func (s *Session) RefTable() *schema.Schema {
|
||||
if s.refTable == nil {
|
||||
log.Error("Model is not set")
|
||||
}
|
||||
return s.refTable
|
||||
}
|
||||
|
||||
// CreateTable create a table in database with a model
|
||||
func (s *Session) CreateTable() error {
|
||||
table := s.RefTable()
|
||||
var columns []string
|
||||
for _, field := range table.Fields {
|
||||
columns = append(columns, fmt.Sprintf("%s %s %s", field.Name, field.Type, field.Tag))
|
||||
}
|
||||
desc := strings.Join(columns, ",")
|
||||
_, err := s.Raw(fmt.Sprintf("CREATE TABLE %s (%s);", table.Name, desc)).Exec()
|
||||
return err
|
||||
}
|
||||
|
||||
// DropTable drops a table with the name of model
|
||||
func (s *Session) DropTable() error {
|
||||
_, err := s.Raw(fmt.Sprintf("DROP TABLE IF EXISTS %s", s.RefTable().Name)).Exec()
|
||||
return err
|
||||
}
|
||||
|
||||
// HasTable returns true of the table exists
|
||||
func (s *Session) HasTable() bool {
|
||||
sql, values := s.dialect.TableExistSQL(s.RefTable().Name)
|
||||
row := s.Raw(sql, values...).QueryRow()
|
||||
var tmp string
|
||||
_ = row.Scan(&tmp)
|
||||
return tmp == s.RefTable().Name
|
||||
}
|
||||
@@ -0,0 +1,28 @@
|
||||
package session
|
||||
|
||||
import (
|
||||
"testing"
|
||||
)
|
||||
|
||||
type User struct {
|
||||
Name string `geeorm:"PRIMARY KEY"`
|
||||
Age int
|
||||
}
|
||||
|
||||
func TestSession_CreateTable(t *testing.T) {
|
||||
s := NewSession().Model(&User{})
|
||||
_ = s.DropTable()
|
||||
_ = s.CreateTable()
|
||||
if !s.HasTable() {
|
||||
t.Fatal("Failed to create table User")
|
||||
}
|
||||
}
|
||||
|
||||
func TestSession_Model(t *testing.T) {
|
||||
s := NewSession().Model(&User{})
|
||||
table := s.RefTable()
|
||||
s.Model(&Session{})
|
||||
if table.Name != "User" || s.RefTable().Name != "Session" {
|
||||
t.Fatal("Failed to change model")
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,31 @@
|
||||
package session
|
||||
|
||||
import "geeorm/log"
|
||||
|
||||
// Begin a transaction
|
||||
func (s *Session) Begin() (err error) {
|
||||
log.Info("transaction begin")
|
||||
if s.tx, err = s.db.Begin(); err != nil {
|
||||
log.Error(err)
|
||||
return
|
||||
}
|
||||
return
|
||||
}
|
||||
|
||||
// Commit a transaction
|
||||
func (s *Session) Commit() (err error) {
|
||||
log.Info("transaction commit")
|
||||
if err = s.tx.Commit(); err != nil {
|
||||
log.Error(err)
|
||||
}
|
||||
return
|
||||
}
|
||||
|
||||
// Rollback a transaction
|
||||
func (s *Session) Rollback() (err error) {
|
||||
log.Info("transaction rollback")
|
||||
if err = s.tx.Rollback(); err != nil {
|
||||
log.Error(err)
|
||||
}
|
||||
return
|
||||
}
|
||||
Reference in New Issue
Block a user