[MM-54456] Fix potential read after write issue when loading license (#24524)
* Fix potential read after write issue when loading license * Use upsert
Этот коммит содержится в:
коммит произвёл
GitHub
родитель
c53b5f7b2b
Коммит
b4a47803e6
@@ -4,6 +4,7 @@
|
||||
package storetest
|
||||
|
||||
import (
|
||||
"context"
|
||||
"testing"
|
||||
|
||||
"github.com/stretchr/testify/require"
|
||||
@@ -22,15 +23,15 @@ func testLicenseStoreSave(t *testing.T, ss store.Store) {
|
||||
l1.Id = model.NewId()
|
||||
l1.Bytes = "junk"
|
||||
|
||||
_, err := ss.License().Save(&l1)
|
||||
err := ss.License().Save(&l1)
|
||||
require.NoError(t, err, "couldn't save license record")
|
||||
|
||||
_, err = ss.License().Save(&l1)
|
||||
err = ss.License().Save(&l1)
|
||||
require.NoError(t, err, "shouldn't fail on trying to save existing license record")
|
||||
|
||||
l1.Id = ""
|
||||
|
||||
_, err = ss.License().Save(&l1)
|
||||
err = ss.License().Save(&l1)
|
||||
require.Error(t, err, "should fail on invalid license")
|
||||
}
|
||||
|
||||
@@ -39,14 +40,14 @@ func testLicenseStoreGet(t *testing.T, ss store.Store) {
|
||||
l1.Id = model.NewId()
|
||||
l1.Bytes = "junk"
|
||||
|
||||
_, err := ss.License().Save(&l1)
|
||||
err := ss.License().Save(&l1)
|
||||
require.NoError(t, err)
|
||||
|
||||
record, err := ss.License().Get(l1.Id)
|
||||
record, err := ss.License().Get(context.Background(), l1.Id)
|
||||
require.NoError(t, err, "couldn't get license")
|
||||
|
||||
require.Equal(t, record.Bytes, l1.Bytes, "license bytes didn't match")
|
||||
|
||||
_, err = ss.License().Get("missing")
|
||||
_, err = ss.License().Get(context.Background(), "missing")
|
||||
require.Error(t, err, "should fail on get license")
|
||||
}
|
||||
|
||||
@@ -5,6 +5,8 @@
|
||||
package mocks
|
||||
|
||||
import (
|
||||
context "context"
|
||||
|
||||
model "github.com/mattermost/mattermost/server/public/model"
|
||||
mock "github.com/stretchr/testify/mock"
|
||||
)
|
||||
@@ -14,25 +16,25 @@ type LicenseStore struct {
|
||||
mock.Mock
|
||||
}
|
||||
|
||||
// Get provides a mock function with given fields: id
|
||||
func (_m *LicenseStore) Get(id string) (*model.LicenseRecord, error) {
|
||||
ret := _m.Called(id)
|
||||
// Get provides a mock function with given fields: ctx, id
|
||||
func (_m *LicenseStore) Get(ctx context.Context, id string) (*model.LicenseRecord, error) {
|
||||
ret := _m.Called(ctx, id)
|
||||
|
||||
var r0 *model.LicenseRecord
|
||||
var r1 error
|
||||
if rf, ok := ret.Get(0).(func(string) (*model.LicenseRecord, error)); ok {
|
||||
return rf(id)
|
||||
if rf, ok := ret.Get(0).(func(context.Context, string) (*model.LicenseRecord, error)); ok {
|
||||
return rf(ctx, id)
|
||||
}
|
||||
if rf, ok := ret.Get(0).(func(string) *model.LicenseRecord); ok {
|
||||
r0 = rf(id)
|
||||
if rf, ok := ret.Get(0).(func(context.Context, string) *model.LicenseRecord); ok {
|
||||
r0 = rf(ctx, id)
|
||||
} else {
|
||||
if ret.Get(0) != nil {
|
||||
r0 = ret.Get(0).(*model.LicenseRecord)
|
||||
}
|
||||
}
|
||||
|
||||
if rf, ok := ret.Get(1).(func(string) error); ok {
|
||||
r1 = rf(id)
|
||||
if rf, ok := ret.Get(1).(func(context.Context, string) error); ok {
|
||||
r1 = rf(ctx, id)
|
||||
} else {
|
||||
r1 = ret.Error(1)
|
||||
}
|
||||
@@ -67,29 +69,17 @@ func (_m *LicenseStore) GetAll() ([]*model.LicenseRecord, error) {
|
||||
}
|
||||
|
||||
// Save provides a mock function with given fields: license
|
||||
func (_m *LicenseStore) Save(license *model.LicenseRecord) (*model.LicenseRecord, error) {
|
||||
func (_m *LicenseStore) Save(license *model.LicenseRecord) error {
|
||||
ret := _m.Called(license)
|
||||
|
||||
var r0 *model.LicenseRecord
|
||||
var r1 error
|
||||
if rf, ok := ret.Get(0).(func(*model.LicenseRecord) (*model.LicenseRecord, error)); ok {
|
||||
return rf(license)
|
||||
}
|
||||
if rf, ok := ret.Get(0).(func(*model.LicenseRecord) *model.LicenseRecord); ok {
|
||||
var r0 error
|
||||
if rf, ok := ret.Get(0).(func(*model.LicenseRecord) error); ok {
|
||||
r0 = rf(license)
|
||||
} else {
|
||||
if ret.Get(0) != nil {
|
||||
r0 = ret.Get(0).(*model.LicenseRecord)
|
||||
}
|
||||
r0 = ret.Error(0)
|
||||
}
|
||||
|
||||
if rf, ok := ret.Get(1).(func(*model.LicenseRecord) error); ok {
|
||||
r1 = rf(license)
|
||||
} else {
|
||||
r1 = ret.Error(1)
|
||||
}
|
||||
|
||||
return r0, r1
|
||||
return r0
|
||||
}
|
||||
|
||||
type mockConstructorTestingTNewLicenseStore interface {
|
||||
|
||||
Ссылка в новой задаче
Block a user