коммит произвёл
GitHub
родитель
56163e9e0e
Коммит
f361e7d75a
@@ -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 {
|
||||
|
||||
Ссылка в новой задаче
Block a user