fix Insert Comment && add hooks feature

This commit is contained in:
gzdaijie
2020-02-29 00:28:12 +08:00
parent e709d288f9
commit d4b622d3ca
14 changed files with 164 additions and 24 deletions
@@ -11,7 +11,7 @@ import (
var TestDB *sql.DB
func TestMain(m *testing.M) {
TestDB, _ = sql.Open("sqlite3", "gee.db")
TestDB, _ = sql.Open("sqlite3", "../gee.db")
code := m.Run()
_ = TestDB.Close()
os.Exit(code)
@@ -16,7 +16,7 @@ var (
)
func TestMain(m *testing.M) {
TestDB, _ = sql.Open("sqlite3", "gee.db")
TestDB, _ = sql.Open("sqlite3", "../gee.db")
code := m.Run()
_ = TestDB.Close()
os.Exit(code)
+1 -1
View File
@@ -16,7 +16,7 @@ var (
)
func TestMain(m *testing.M) {
TestDB, _ = sql.Open("sqlite3", "gee.db")
TestDB, _ = sql.Open("sqlite3", "../gee.db")
code := m.Run()
_ = TestDB.Close()
os.Exit(code)
+1 -1
View File
@@ -5,7 +5,7 @@ import (
"reflect"
)
// Create one or more records in database
// Insert one or more records in database
func (s *Session) Insert(values ...interface{}) (int64, error) {
recordValues := make([]interface{}, 0)
for _, value := range values {
@@ -16,7 +16,7 @@ var (
)
func TestMain(m *testing.M) {
TestDB, _ = sql.Open("sqlite3", "gee.db")
TestDB, _ = sql.Open("sqlite3", "../gee.db")
code := m.Run()
_ = TestDB.Close()
os.Exit(code)
@@ -6,7 +6,7 @@ import (
"reflect"
)
// Create one or more records in database
// Insert one or more records in database
func (s *Session) Insert(values ...interface{}) (int64, error) {
recordValues := make([]interface{}, 0)
for _, value := range values {
+32
View File
@@ -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
}
+37
View File
@@ -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)
}
}
+1 -1
View File
@@ -16,7 +16,7 @@ var (
)
func TestMain(m *testing.M) {
TestDB, _ = sql.Open("sqlite3", "gee.db")
TestDB, _ = sql.Open("sqlite3", "../gee.db")
code := m.Run()
_ = TestDB.Close()
os.Exit(code)
+9 -14
View File
@@ -6,22 +6,11 @@ import (
"reflect"
)
// CallMethod calls the registered hooks
func (s *Session) CallMethod(method string) error {
if fm := reflect.ValueOf(s.refTable.Model).MethodByName(method); fm.IsValid() {
if v := fm.Call([]reflect.Value{}); len(v) > 0 {
if err, ok := v[0].Interface().(error); ok {
return err
}
}
}
return nil
}
// Create one or more records in database
// 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))
@@ -33,12 +22,13 @@ func (s *Session) Insert(values ...interface{}) (int64, error) {
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()
@@ -59,6 +49,7 @@ func (s *Session) Find(values interface{}) error {
if err := rows.Scan(values...); err != nil {
return err
}
s.CallMethod(AfterQuery, dest.Addr().Interface())
destSlice.Set(reflect.Append(destSlice, dest))
}
return rows.Close()
@@ -101,6 +92,7 @@ func (s *Session) OrderBy(desc string) *Session {
// 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{})
@@ -114,17 +106,20 @@ func (s *Session) Update(kv ...interface{}) (int64, error) {
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()
}
+32
View File
@@ -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)
}
}
+1 -1
View File
@@ -16,7 +16,7 @@ var (
)
func TestMain(m *testing.M) {
TestDB, _ = sql.Open("sqlite3", "gee.db")
TestDB, _ = sql.Open("sqlite3", "../gee.db")
code := m.Run()
_ = TestDB.Close()
os.Exit(code)
+9 -2
View File
@@ -6,10 +6,11 @@ import (
"reflect"
)
// Create one or more records in database
// 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))
@@ -21,12 +22,13 @@ func (s *Session) Insert(values ...interface{}) (int64, error) {
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()
@@ -47,6 +49,7 @@ func (s *Session) Find(values interface{}) error {
if err := rows.Scan(values...); err != nil {
return err
}
s.CallMethod(AfterQuery, dest.Addr().Interface())
destSlice.Set(reflect.Append(destSlice, dest))
}
return rows.Close()
@@ -89,6 +92,7 @@ func (s *Session) OrderBy(desc string) *Session {
// 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{})
@@ -102,17 +106,20 @@ func (s *Session) Update(kv ...interface{}) (int64, error) {
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()
}