mirror of
https://github.com/geektutu/7days-golang.git
synced 2024-04-21 12:32:11 +00:00
fix Insert Comment && add hooks feature
This commit is contained in:
@@ -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)
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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 {
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
@@ -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,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()
|
||||
}
|
||||
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
@@ -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,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()
|
||||
}
|
||||
|
||||
|
||||
Reference in New Issue
Block a user