fix(internal): rules tests are run in transaction
This commit is contained in:
committed by
Tit Petric
parent
d631e28894
commit
9bc361f1d1
Generated
+2
-2
@@ -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"
|
||||
|
||||
@@ -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())
|
||||
}
|
||||
|
||||
@@ -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")
|
||||
}
|
||||
|
||||
+42
-18
@@ -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")
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user