This is an automated email from the ASF dual-hosted git repository. warren pushed a commit to branch main in repository https://gitbox.apache.org/repos/asf/incubator-devlake.git
commit 7b485ed07d871150946de3390fd7e3e6fb9e757a Author: Yingchu Chen <[email protected]> AuthorDate: Wed Jun 1 19:07:30 2022 +0800 unit test for connection Signed-off-by: Yingchu Chen <[email protected]> --- plugins/helper/connection.go | 49 +++++------ plugins/helper/connection_test.go | 178 ++++++++++++++++++++++++++++++++++++++ plugins/jira/tasks/api_client.go | 9 +- 3 files changed, 205 insertions(+), 31 deletions(-) diff --git a/plugins/helper/connection.go b/plugins/helper/connection.go index 85163068..17933a5c 100644 --- a/plugins/helper/connection.go +++ b/plugins/helper/connection.go @@ -21,7 +21,7 @@ type BaseConnection struct { type BasicAuth struct { Username string `mapstructure:"username" validate:"required" json:"username"` - Password string `mapstructure:"password" validate:"required" json:"password" encrypt:"yes"` + Password string `mapstructure:"password" validate:"required" json:"password" encryptField:"yes"` } func (ba BasicAuth) GetEncodedToken() string { @@ -29,7 +29,7 @@ func (ba BasicAuth) GetEncodedToken() string { } type AccessToken struct { - Token string `mapstructure:"token" validate:"required" json:"token" encrypt:"yes"` + Token string `mapstructure:"token" validate:"required" json:"token" encryptField:"yes"` } type RestConnection struct { @@ -65,11 +65,14 @@ func saveToDb(connection interface{}, db *gorm.DB) error { if dataVal.Kind() != reflect.Ptr { panic("entityPtr is not a pointer") } - + encKey, err := getEncKey() + if err != nil { + return err + } dataType := reflect.Indirect(dataVal).Type() - fieldName := getEncryptField(dataType, "encrypt") + fieldName := firstFieldNameWithTag(dataType, "encryptField") plainPwd := "" - err := doEncrypt(dataVal, fieldName) + err = encryptField(dataVal, fieldName, encKey) if err != nil { return err } @@ -78,7 +81,7 @@ func saveToDb(connection interface{}, db *gorm.DB) error { return err } - err = doDecrypt(dataVal, fieldName) + err = decryptField(dataVal, fieldName, encKey) if err != nil { return err } @@ -107,7 +110,7 @@ func mergeFieldsToConnection(specificConnection interface{}, connections ...map[ } func getEncKey() (string, error) { - // encrypt + // encryptField v := config.GetConfig() encKey := v.GetString(core.EncodeKeyEnvStr) if encKey == "" { @@ -142,8 +145,8 @@ func FindConnectionByInput(input *core.ApiResourceInput, connection interface{}, dataType := reflect.Indirect(dataVal).Type() - fieldName := getEncryptField(dataType, "encrypt") - return doDecrypt(dataVal, fieldName) + fieldName := firstFieldNameWithTag(dataType, "encryptField") + return decryptField(dataVal, fieldName, "") } @@ -156,12 +159,12 @@ func GetConnectionIdByInputParam(input *core.ApiResourceInput) (uint64, error) { return strconv.ParseUint(connectionId, 10, 64) } -func getEncryptField(t reflect.Type, tag string) string { +func firstFieldNameWithTag(t reflect.Type, tag string) string { fieldName := "" for i := 0; i < t.NumField(); i++ { field := t.Field(i) if field.Type.Kind() == reflect.Struct { - fieldName = getEncryptField(field.Type, tag) + fieldName = firstFieldNameWithTag(field.Type, tag) } else { if field.Tag.Get(tag) == "yes" { fieldName = field.Name @@ -177,34 +180,30 @@ func DecryptConnection(connection interface{}, fieldName string) error { if dataVal.Kind() != reflect.Ptr { panic("connection is not a pointer") } + encKey, err := getEncKey() + if err != nil { + return nil + } if len(fieldName) == 0 { dataType := reflect.Indirect(dataVal).Type() - fieldName = getEncryptField(dataType, "encrypt") + fieldName = firstFieldNameWithTag(dataType, "encryptField") } - return doDecrypt(dataVal, fieldName) + return decryptField(dataVal, fieldName, encKey) } -func doDecrypt(dataVal reflect.Value, fieldName string) error { - encryptCode, err := getEncKey() - if err != nil { - return err - } +func decryptField(dataVal reflect.Value, fieldName string, encKey string) error { if len(fieldName) > 0 { - decryptStr, _ := core.Decrypt(encryptCode, dataVal.Elem().FieldByName(fieldName).String()) + decryptStr, _ := core.Decrypt(encKey, dataVal.Elem().FieldByName(fieldName).String()) dataVal.Elem().FieldByName(fieldName).Set(reflect.ValueOf(decryptStr)) } return nil } -func doEncrypt(dataVal reflect.Value, fieldName string) error { - encryptCode, err := getEncKey() - if err != nil { - return err - } +func encryptField(dataVal reflect.Value, fieldName string, encKey string) error { if len(fieldName) > 0 { plainPwd := dataVal.Elem().FieldByName(fieldName).String() - encyptedStr, err := core.Encrypt(encryptCode, plainPwd) + encyptedStr, err := core.Encrypt(encKey, plainPwd) if err != nil { return err diff --git a/plugins/helper/connection_test.go b/plugins/helper/connection_test.go new file mode 100644 index 00000000..606730f3 --- /dev/null +++ b/plugins/helper/connection_test.go @@ -0,0 +1,178 @@ +/* +Licensed to the Apache Software Foundation (ASF) under one or more +contributor license agreements. See the NOTICE file distributed with +this work for additional information regarding copyright ownership. +The ASF licenses this file to You under the Apache License, Version 2.0 +(the "License"); you may not use this file except in compliance with +the License. You may obtain a copy of the License at + + http://www.apache.org/licenses/LICENSE-2.0 + +Unless required by applicable law or agreed to in writing, software +distributed under the License is distributed on an "AS IS" BASIS, +WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +See the License for the specific language governing permissions and +limitations under the License. +*/ + +package helper + +import ( + "github.com/apache/incubator-devlake/config" + "github.com/apache/incubator-devlake/plugins/core" + "reflect" + "testing" + + "github.com/stretchr/testify/assert" +) + +type TestConnection struct { + RestConnection `mapstructure:",squash"` + BasicAuth `mapstructure:",squash"` + EpicKeyField string `gorm:"type:varchar(50);" json:"epicKeyField"` + StoryPointField string `gorm:"type:varchar(50);" json:"storyPointField"` + RemotelinkCommitShaPattern string `gorm:"type:varchar(255);comment='golang regexp, the first group will be recognized as commit sha, ref https://github.com/google/re2/wiki/Syntax'" json:"remotelinkCommitShaPattern"` +} + +func TestMergeFieldsToConnection(t *testing.T) { + v := &TestConnection{ + RestConnection: RestConnection{ + BaseConnection: BaseConnection{ + Name: "1", + }, + Endpoint: "2", + Proxy: "3", + RateLimit: 0, + }, + BasicAuth: BasicAuth{ + Username: "4", + Password: "5", + }, + EpicKeyField: "6", + StoryPointField: "7", + RemotelinkCommitShaPattern: "8", + } + data := make(map[string]interface{}) + data["Endpoint"] = "2-2" + data["Username"] = "4-4" + data["Password"] = "5-5" + + err := mergeFieldsToConnection(v, data) + if err != nil { + return + } + + assert.Equal(t, "4-4", v.Username) + assert.Equal(t, "2-2", v.Endpoint) + assert.Equal(t, "5-5", v.Password) +} + +func TestDecryptAndEncrypt(t *testing.T) { + v := &TestConnection{ + RestConnection: RestConnection{ + BaseConnection: BaseConnection{ + Name: "1", + }, + Endpoint: "2", + Proxy: "3", + RateLimit: 0, + }, + BasicAuth: BasicAuth{ + Username: "4", + Password: "5", + }, + EpicKeyField: "6", + StoryPointField: "7", + RemotelinkCommitShaPattern: "8", + } + dataVal := reflect.ValueOf(v) + encKey := "test" + err := encryptField(dataVal, "Password", encKey) + if err != nil { + return + } + assert.NotEqual(t, "5", v.Password) + err = decryptField(dataVal, "Password", encKey) + if err != nil { + return + } + + assert.Equal(t, "5", v.Password) + +} + +func TestDecryptConnection(t *testing.T) { + v := &TestConnection{ + RestConnection: RestConnection{ + BaseConnection: BaseConnection{ + Name: "1", + }, + Endpoint: "2", + Proxy: "3", + RateLimit: 0, + }, + BasicAuth: BasicAuth{ + Username: "4", + Password: "5", + }, + EpicKeyField: "6", + StoryPointField: "7", + RemotelinkCommitShaPattern: "8", + } + encKey, err := getEncKey() + if err != nil { + return + } + dataVal := reflect.ValueOf(v) + err = encryptField(dataVal, "Password", encKey) + if err != nil { + return + } + encryptedPwd := v.Password + err = DecryptConnection(v, "Password") + if err != nil { + return + } + assert.NotEqual(t, encryptedPwd, v.Password) + assert.Equal(t, "5", v.Password) +} + +func TestGetEncKey(t *testing.T) { + // encryptField + v := config.GetConfig() + encKey := v.GetString(core.EncodeKeyEnvStr) + str, err := getEncKey() + if err != nil { + return + } + if len(encKey) > 0 { + assert.Equal(t, encKey, str) + } else { + assert.NotEqual(t, 0, len(str)) + } + +} + +func TestFirstFieldNameWithTag(t *testing.T) { + v := &TestConnection{ + RestConnection: RestConnection{ + BaseConnection: BaseConnection{ + Name: "1", + }, + Endpoint: "2", + Proxy: "3", + RateLimit: 0, + }, + BasicAuth: BasicAuth{ + Username: "4", + Password: "5", + }, + EpicKeyField: "6", + StoryPointField: "7", + RemotelinkCommitShaPattern: "8", + } + dataVal := reflect.ValueOf(v) + dataType := reflect.Indirect(dataVal).Type() + fieldName := firstFieldNameWithTag(dataType, "encryptField") + assert.Equal(t, "Password", fieldName) +} diff --git a/plugins/jira/tasks/api_client.go b/plugins/jira/tasks/api_client.go index 3673bee5..811c0b6a 100644 --- a/plugins/jira/tasks/api_client.go +++ b/plugins/jira/tasks/api_client.go @@ -27,18 +27,15 @@ import ( ) func NewJiraApiClient(taskCtx core.TaskContext, connection *models.JiraConnection) (*helper.ApiAsyncClient, error) { - // load configuration - encKey := taskCtx.GetConfig(core.EncodeKeyEnvStr) - decodedPassword, err := core.Decrypt(encKey, connection.Password) + // decrypt connection first + err := helper.DecryptConnection(connection, "Password") if err != nil { return nil, fmt.Errorf("Failed to decrypt Auth AccessToken: %w", err) } - connection.Password = decodedPassword - auth := connection.GetEncodedToken() // create synchronize api client so we can calculate api rate limit dynamically headers := map[string]string{ - "Authorization": fmt.Sprintf("Basic %v", auth), + "Authorization": fmt.Sprintf("Basic %v", connection.GetEncodedToken()), } apiClient, err := helper.NewApiClient(connection.Endpoint, headers, 0, connection.Proxy, taskCtx.GetContext()) if err != nil {
