Automatic Merge
Этот коммит содержится в:
Mattermost Build
2025-10-27 12:59:15 +02:00
коммит произвёл GitHub
родитель 56163e9e0e
Коммит f361e7d75a
14 изменённых файлов: 352 добавлений и 32 удалений

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

@@ -14138,11 +14138,11 @@ func (s *RetryLayerTokenStore) Cleanup(expiryTime int64) {
}
func (s *RetryLayerTokenStore) ConsumeOnce(tokenStr string) (*model.Token, error) {
func (s *RetryLayerTokenStore) ConsumeOnce(tokenType string, tokenStr string) (*model.Token, error) {
tries := 0
for {
result, err := s.TokenStore.ConsumeOnce(tokenStr)
result, err := s.TokenStore.ConsumeOnce(tokenType, tokenStr)
if err == nil {
return result, nil
}

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

@@ -78,16 +78,41 @@ func (s SqlTokenStore) GetByToken(tokenString string) (*model.Token, error) {
return &token, nil
}
func (s SqlTokenStore) ConsumeOnce(tokenStr string) (*model.Token, error) {
func (s SqlTokenStore) ConsumeOnce(tokenType, tokenStr string) (*model.Token, error) {
var token model.Token
query := `DELETE FROM Tokens WHERE Token = ? RETURNING *`
if s.DriverName() == model.DatabaseDriverPostgres {
query := `DELETE FROM Tokens WHERE Type = ? AND Token = ? RETURNING *`
if err := s.GetMaster().Get(&token, query, tokenType, tokenStr); err != nil {
if err == sql.ErrNoRows {
return nil, store.NewErrNotFound("Token", tokenStr)
}
return nil, errors.Wrapf(err, "failed to consume token with type %s", tokenType)
}
return &token, nil
}
if err := s.GetMaster().Get(&token, query, tokenStr); err != nil {
transaction, err := s.GetMaster().Beginx()
if err != nil {
return nil, errors.Wrap(err, "failed to begin transaction")
}
defer finalizeTransactionX(transaction, &err)
query := `SELECT * FROM Tokens WHERE Type = ? AND Token = ? FOR UPDATE`
if err = transaction.Get(&token, query, tokenType, tokenStr); err != nil {
if err == sql.ErrNoRows {
return nil, store.NewErrNotFound("Token", tokenStr)
}
return nil, errors.Wrapf(err, "failed to consume token")
return nil, errors.Wrapf(err, "failed to select token with type %s", tokenType)
}
deleteQuery := `DELETE FROM Tokens WHERE Type = ? AND Token = ?`
if _, err = transaction.Exec(deleteQuery, tokenType, tokenStr); err != nil {
return nil, errors.Wrapf(err, "failed to delete token with type %s", tokenType)
}
if err = transaction.Commit(); err != nil {
return nil, errors.Wrap(err, "failed to commit transaction")
}
return &token, nil

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

@@ -693,7 +693,7 @@ type TokenStore interface {
Save(recovery *model.Token) error
Delete(token string) error
GetByToken(token string) (*model.Token, error)
ConsumeOnce(tokenStr string) (*model.Token, error)
ConsumeOnce(tokenType, tokenStr string) (*model.Token, error)
Cleanup(expiryTime int64)
GetAllTokensByType(tokenType string) ([]*model.Token, error)
RemoveAllTokensByType(tokenType string) error

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

@@ -19,9 +19,9 @@ func (_m *TokenStore) Cleanup(expiryTime int64) {
_m.Called(expiryTime)
}
// ConsumeOnce provides a mock function with given fields: tokenStr
func (_m *TokenStore) ConsumeOnce(tokenStr string) (*model.Token, error) {
ret := _m.Called(tokenStr)
// ConsumeOnce provides a mock function with given fields: tokenType, tokenStr
func (_m *TokenStore) ConsumeOnce(tokenType string, tokenStr string) (*model.Token, error) {
ret := _m.Called(tokenType, tokenStr)
if len(ret) == 0 {
panic("no return value specified for ConsumeOnce")
@@ -29,19 +29,19 @@ func (_m *TokenStore) ConsumeOnce(tokenStr string) (*model.Token, error) {
var r0 *model.Token
var r1 error
if rf, ok := ret.Get(0).(func(string) (*model.Token, error)); ok {
return rf(tokenStr)
if rf, ok := ret.Get(0).(func(string, string) (*model.Token, error)); ok {
return rf(tokenType, tokenStr)
}
if rf, ok := ret.Get(0).(func(string) *model.Token); ok {
r0 = rf(tokenStr)
if rf, ok := ret.Get(0).(func(string, string) *model.Token); ok {
r0 = rf(tokenType, tokenStr)
} else {
if ret.Get(0) != nil {
r0 = ret.Get(0).(*model.Token)
}
}
if rf, ok := ret.Get(1).(func(string) error); ok {
r1 = rf(tokenStr)
if rf, ok := ret.Get(1).(func(string, string) error); ok {
r1 = rf(tokenType, tokenStr)
} else {
r1 = ret.Error(1)
}

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

@@ -16,6 +16,7 @@ import (
func TestTokensStore(t *testing.T, rctx request.CTX, ss store.Store) {
t.Run("TokensCleanup", func(t *testing.T) { testTokensCleanup(t, rctx, ss) })
t.Run("ConsumeOnce", func(t *testing.T) { testConsumeOnce(t, rctx, ss) })
}
func testTokensCleanup(t *testing.T, rctx request.CTX, ss store.Store) {
@@ -41,3 +42,130 @@ func testTokensCleanup(t *testing.T, rctx request.CTX, ss store.Store) {
require.NoError(t, err)
assert.Len(t, tokens, 0)
}
func testConsumeOnce(t *testing.T, rctx request.CTX, ss store.Store) {
t.Run("successfully consume token once", func(t *testing.T) {
token := &model.Token{
Token: model.NewRandomString(model.TokenSize),
CreateAt: model.GetMillis(),
Type: model.TokenTypeOAuth,
Extra: "test-extra",
}
err := ss.Token().Save(token)
require.NoError(t, err)
consumedToken, err := ss.Token().ConsumeOnce(model.TokenTypeOAuth, token.Token)
require.NoError(t, err)
assert.Equal(t, token.Token, consumedToken.Token)
assert.Equal(t, token.Type, consumedToken.Type)
assert.Equal(t, token.Extra, consumedToken.Extra)
tokens, err := ss.Token().GetAllTokensByType(model.TokenTypeOAuth)
require.NoError(t, err)
assert.Len(t, tokens, 0)
})
t.Run("second consumption of same token fails", func(t *testing.T) {
token := &model.Token{
Token: model.NewRandomString(model.TokenSize),
CreateAt: model.GetMillis(),
Type: model.TokenTypeOAuth,
Extra: "test-extra",
}
err := ss.Token().Save(token)
require.NoError(t, err)
_, err = ss.Token().ConsumeOnce(model.TokenTypeOAuth, token.Token)
require.NoError(t, err)
_, err = ss.Token().ConsumeOnce(model.TokenTypeOAuth, token.Token)
require.Error(t, err)
var nfErr *store.ErrNotFound
assert.ErrorAs(t, err, &nfErr)
})
t.Run("consume with wrong type fails", func(t *testing.T) {
token := &model.Token{
Token: model.NewRandomString(model.TokenSize),
CreateAt: model.GetMillis(),
Type: model.TokenTypeOAuth,
Extra: "test-extra",
}
err := ss.Token().Save(token)
require.NoError(t, err)
_, err = ss.Token().ConsumeOnce(model.TokenTypeSSOCodeExchange, token.Token)
require.Error(t, err)
var nfErr *store.ErrNotFound
assert.ErrorAs(t, err, &nfErr)
tokens, err := ss.Token().GetAllTokensByType(model.TokenTypeOAuth)
require.NoError(t, err)
assert.Len(t, tokens, 1)
err = ss.Token().Delete(token.Token)
require.NoError(t, err)
})
t.Run("consume non-existent token fails", func(t *testing.T) {
nonExistentToken := model.NewRandomString(model.TokenSize)
_, err := ss.Token().ConsumeOnce(model.TokenTypeOAuth, nonExistentToken)
require.Error(t, err)
var nfErr *store.ErrNotFound
assert.ErrorAs(t, err, &nfErr)
})
t.Run("multiple tokens with same type can each be consumed once", func(t *testing.T) {
tokens := make([]*model.Token, 3)
for i := range tokens {
tokens[i] = &model.Token{
Token: model.NewRandomString(model.TokenSize),
CreateAt: model.GetMillis(),
Type: model.TokenTypeOAuth,
Extra: "test-extra",
}
err := ss.Token().Save(tokens[i])
require.NoError(t, err)
}
for _, token := range tokens {
consumedToken, err := ss.Token().ConsumeOnce(model.TokenTypeOAuth, token.Token)
require.NoError(t, err)
assert.Equal(t, token.Token, consumedToken.Token)
}
allTokens, err := ss.Token().GetAllTokensByType(model.TokenTypeOAuth)
require.NoError(t, err)
assert.Len(t, allTokens, 0)
})
t.Run("consuming token of different type leaves others intact", func(t *testing.T) {
oauthToken := &model.Token{
Token: model.NewRandomString(model.TokenSize),
CreateAt: model.GetMillis(),
Type: model.TokenTypeOAuth,
Extra: "oauth-extra",
}
codeExchangeToken := &model.Token{
Token: model.NewRandomString(model.TokenSize),
CreateAt: model.GetMillis(),
Type: model.TokenTypeSSOCodeExchange,
Extra: "password-extra",
}
err := ss.Token().Save(oauthToken)
require.NoError(t, err)
err = ss.Token().Save(codeExchangeToken)
require.NoError(t, err)
consumedToken, err := ss.Token().ConsumeOnce(model.TokenTypeOAuth, oauthToken.Token)
require.NoError(t, err)
assert.Equal(t, oauthToken.Token, consumedToken.Token)
codeExchangeTokens, err := ss.Token().GetAllTokensByType(model.TokenTypeSSOCodeExchange)
require.NoError(t, err)
assert.Len(t, codeExchangeTokens, 1)
err = ss.Token().Delete(codeExchangeToken.Token)
require.NoError(t, err)
})
}

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

@@ -11106,10 +11106,10 @@ func (s *TimerLayerTokenStore) Cleanup(expiryTime int64) {
}
}
func (s *TimerLayerTokenStore) ConsumeOnce(tokenStr string) (*model.Token, error) {
func (s *TimerLayerTokenStore) ConsumeOnce(tokenType string, tokenStr string) (*model.Token, error) {
start := time.Now()
result, err := s.TokenStore.ConsumeOnce(tokenStr)
result, err := s.TokenStore.ConsumeOnce(tokenType, tokenStr)
elapsed := float64(time.Since(start)) / float64(time.Second)
if s.Root.Metrics != nil {