From 9bc361f1d107ea379318fc058a56617a0b74b98c Mon Sep 17 00:00:00 2001 From: Mitja Zivkovic Date: Wed, 13 Mar 2019 20:14:49 +0100 Subject: [PATCH] fix(internal): rules tests are run in transaction --- Gopkg.lock | 4 +- internal/rules/main_test.go | 10 ---- internal/rules/resources_test.go | 32 +++++----- .../github.com/titpetric/factory/database.go | 60 +++++++++++++------ 4 files changed, 58 insertions(+), 48 deletions(-) diff --git a/Gopkg.lock b/Gopkg.lock index 2649d843b..fb962211e 100644 --- a/Gopkg.lock +++ b/Gopkg.lock @@ -378,14 +378,14 @@ [[projects]] branch = "master" - digest = "1:5e775c142bea35d760b5b5f877d04252d771f1a047218316f46d2ab3157feebc" + digest = "1:8a90d635c58170aff3184be7da859c76ad941d7d11fa37edb28943a133cae186" name = "github.com/titpetric/factory" packages = [ ".", "resputil", ] pruneopts = "UT" - revision = "5b058d689fe04625f102b1ada6c11724baa3e07c" + revision = "c1dcb85af748151b6cabfac413d7e243d7214b13" [[projects]] digest = "1:539aa04b8528cceeaa4e7f22034d1f463d0753ec5716b4571b766b7770ff79b5" diff --git a/internal/rules/main_test.go b/internal/rules/main_test.go index 99c107320..371b1e704 100644 --- a/internal/rules/main_test.go +++ b/internal/rules/main_test.go @@ -37,15 +37,5 @@ func TestMain(m *testing.M) { return } - // clean up tables - { - for _, name := range []string{"sys_user", "sys_role", "sys_role_member", "sys_rules"} { - _, err := db.Exec("truncate " + name) - if err != nil { - panic("Error when clearing " + name + ": " + err.Error()) - } - } - } - os.Exit(m.Run()) } diff --git a/internal/rules/resources_test.go b/internal/rules/resources_test.go index 5614e592f..6e3ab3437 100644 --- a/internal/rules/resources_test.go +++ b/internal/rules/resources_test.go @@ -32,10 +32,9 @@ func TestRules(t *testing.T) { // Create resources interface. resources := rules.NewResources(ctx, db) - // Run test with savepoint. - err := func() error { - db.Exec("SAVEPOINT rules_test") - + // Run tests in transaction to maintain DB state. + Error(t, db.Transaction(func() error { + db.Delete("sys_rules", "1=1") db.Insert("sys_user", user) db.Insert("sys_role", role) db.Insert("sys_role_member", types.RoleMember{RoleID: role.ID, UserID: user.ID}) @@ -43,7 +42,7 @@ func TestRules(t *testing.T) { // delete all for test roleID = 123456 { err := resources.Delete(role.ID) - NoError(t, err, "expected no error") + NoError(t, err, "expected no error, got %+v", err) } // default (unset=deny), forbidden check ...:* @@ -59,7 +58,7 @@ func TestRules(t *testing.T) { rules.Rule{Resource: "messaging:channel:2", Operation: "delete", Value: rules.Allow}, } err := resources.Grant(role.ID, list) - NoError(t, err, "expect no error") + NoError(t, err, "expect no error, got %+v", err) Expect(rules.Deny, resources.Check("messaging:channel:1", "update"), "messaging:channel:1 update - Deny") Expect(rules.Allow, resources.Check("messaging:channel:2", "update"), "messaging:channel:2 update - Allow") @@ -69,8 +68,8 @@ func TestRules(t *testing.T) { // list grants for test role { grants, err := resources.Read(role.ID) - NoError(t, err, "expect no error") - Assert(t, len(grants) == 2, "expected 2 grants") + NoError(t, err, "expect no error, got %+v", err) + Assert(t, len(grants) == 2, "expected 2 grants, got %v", len(grants)) for _, grant := range grants { Assert(t, grant.RoleID == role.ID, "expected RoleID == 123456, got %v", grant.RoleID) @@ -85,7 +84,7 @@ func TestRules(t *testing.T) { rules.Rule{Resource: "messaging:channel:1", Operation: "update", Value: rules.Deny}, } err := resources.Grant(role.ID, list) - NoError(t, err, "expect no error") + NoError(t, err, "expect no error, got %+v", err) Expect(rules.Deny, resources.Check("messaging:channel:1", "update"), "messaging:channel:1 update - Deny") Expect(rules.Allow, resources.Check("messaging:channel:2", "update"), "messaging:channel:2 update - Allow") @@ -101,7 +100,7 @@ func TestRules(t *testing.T) { rules.Rule{Resource: "messaging:channel:2", Operation: "delete", Value: rules.Inherit}, } err := resources.Grant(role.ID, list) - NoError(t, err, "expect no error") + NoError(t, err, "expect no error, got %+v", err) Expect(rules.Deny, resources.Check("messaging:channel:1", "update"), "messaging:channel:1 update - Deny") Expect(rules.Deny, resources.Check("messaging:channel:2", "update"), "messaging:channel:2 update - Deny") @@ -116,7 +115,7 @@ func TestRules(t *testing.T) { rules.Rule{Resource: "system", Operation: "organisation.create", Value: rules.Allow}, } err := resources.Grant(role.ID, list) - NoError(t, err, "expected no error") + NoError(t, err, "expect no error, got %+v", err) Expect(rules.Deny, resources.Check("messaging:channel:1", "update"), "messaging:channel:1 update - Deny") Expect(rules.Allow, resources.Check("messaging:channel:2", "update"), "messaging:channel:2 update - Allow") @@ -125,25 +124,22 @@ func TestRules(t *testing.T) { // list all by roleID { grants, err := resources.Read(role.ID) - NoError(t, err, "expected no error") + NoError(t, err, "expect no error, got %+v", err) Assert(t, len(grants) == 3, "expected grants == 3, got %v", len(grants)) } // delete all by roleID { err := resources.Delete(role.ID) - NoError(t, err, "expected no error") + NoError(t, err, "expect no error, got %+v", err) } // list all by roleID { grants, err := resources.Read(role.ID) - NoError(t, err, "expected no error") + NoError(t, err, "expect no error, got %+v", err) Assert(t, len(grants) == 0, "expected grants == 0, got %v", len(grants)) } return errors.New("Rollback") - }() - if err != nil { - db.Exec("ROLLBACK TO SAVEPOINT rules_test") - } + }), "expected rollback error") } diff --git a/vendor/github.com/titpetric/factory/database.go b/vendor/github.com/titpetric/factory/database.go index 1fd562ce0..1b38d3262 100644 --- a/vendor/github.com/titpetric/factory/database.go +++ b/vendor/github.com/titpetric/factory/database.go @@ -128,6 +128,7 @@ func (r *DatabaseFactory) Get(dbName ...string) (*DB, error) { r.instances[name] = &DB{ handle, context.Background(), + 0, nil, &sql.TxOptions{ ReadOnly: false, @@ -155,6 +156,7 @@ type DB struct { ctx context.Context + inTx int32 Tx *sqlx.Tx TxOpts *sql.TxOptions @@ -166,6 +168,7 @@ func (r *DB) Quiet() *DB { return &DB{ r.DB, r.ctx, + r.inTx, r.Tx, r.TxOpts, nil, @@ -177,6 +180,7 @@ func (r *DB) With(ctx context.Context) *DB { return &DB{ r.DB, ctx, + r.inTx, r.Tx, r.TxOpts, r.Profiler, @@ -184,17 +188,23 @@ func (r *DB) With(ctx context.Context) *DB { } // Begin will create a transaction in the DB with a context -func (r *DB) Begin() error { - var err error - if r.Tx != nil { - return errors.New("Transaction already started") +func (r *DB) Begin() (err error) { + if r.inTx > 0 { + _, err = r.Exec(fmt.Sprintf("SAVEPOINT sp_%d", r.inTx)) } - if r.ctx == nil { - r.Tx, err = r.DB.Beginx() + if r.inTx == 0 { + if r.ctx == nil { + r.Tx, err = r.DB.Beginx() + } else { + r.Tx, err = r.DB.BeginTxx(r.ctx, r.TxOpts) + } + } + if err != nil { return errors.WithStack(err) } - r.Tx, err = r.DB.BeginTxx(r.ctx, r.TxOpts) - return errors.WithStack(err) + + r.inTx++ + return nil } // Transaction will create a transaction and invoke a callback @@ -202,8 +212,8 @@ func (r *DB) Transaction(callback func() error) error { if err := r.Begin(); err != nil { return err } - defer r.Rollback() if err := callback(); err != nil { + r.Rollback() return err } return r.Commit() @@ -211,22 +221,36 @@ func (r *DB) Transaction(callback func() error) error { func (r *DB) Commit() error { if r.Tx != nil { - if err := r.Tx.Commit(); err != nil { - return errors.WithStack(err) + if r.inTx > 0 { + r.inTx-- } - r.Tx = nil - return nil + if r.inTx == 0 { + if err := r.Tx.Commit(); err != nil { + return errors.WithStack(err) + } + r.Tx = nil + return nil + } + _, err := r.Exec(fmt.Sprintf("RELEASE SAVEPOINT sp_%d", r.inTx)) + return errors.WithStack(err) } - return errors.New("No transation active") + return errors.New("No transaction active") } func (r *DB) Rollback() error { if r.Tx != nil { - if err := r.Tx.Rollback(); err != nil { - return errors.WithStack(err) + if r.inTx > 0 { + r.inTx-- } - r.Tx = nil - return nil + if r.inTx == 0 { + if err := r.Tx.Rollback(); err != nil { + return errors.WithStack(err) + } + r.Tx = nil + return nil + } + _, err := r.Exec(fmt.Sprintf("ROLLBACK SAVEPOINT sp_%d", r.inTx)) + return errors.WithStack(err) } return errors.New("No transaction active") }