[MM-54456] Fix potential read after write issue when loading license (#24524)

* Fix potential read after write issue when loading license

* Use upsert
Этот коммит содержится в:
Claudio Costa
2023-09-13 14:47:12 -06:00
коммит произвёл GitHub
родитель c53b5f7b2b
Коммит b4a47803e6
8 изменённых файлов: 71 добавлений и 81 удалений

Просмотреть файл

@@ -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 {