From 98eb62e625e9761f5d2f4b8633371e510a4630f4 Mon Sep 17 00:00:00 2001 From: Paul Berg Date: Fri, 9 Nov 2018 20:35:05 +0200 Subject: [PATCH] Solved malformed data panics and also wrote tests --- signer/core/cliui.go | 2 +- signer/core/signed_data.go | 141 ++++++++++++----- signer/core/signed_data_test.go | 266 +++++++++++++++++++++++++++++++- 3 files changed, 369 insertions(+), 40 deletions(-) diff --git a/signer/core/cliui.go b/signer/core/cliui.go index 7fefaabd75..1e5927b242 100644 --- a/signer/core/cliui.go +++ b/signer/core/cliui.go @@ -165,7 +165,7 @@ func (ui *CommandlineUI) ApproveSignData(request *SignDataRequest) (SignDataResp fmt.Printf("-------- Sign data request--------------\n") fmt.Printf("Account: %s\n", request.Address.String()) fmt.Printf("message: \n%q\n", request.Message) - fmt.Printf("raw data: \n%v\n", request.Rawdata) + fmt.Printf("raw data: \n%v\n", request.Rawdata) fmt.Printf("message hash: %v\n", request.Hash) fmt.Printf("-------------------------------------------\n") showMetadata(request.Meta) diff --git a/signer/core/signed_data.go b/signer/core/signed_data.go index 5a736481cf..95454acfdc 100644 --- a/signer/core/signed_data.go +++ b/signer/core/signed_data.go @@ -19,7 +19,6 @@ package core import ( "bytes" "context" - "encoding/json" "errors" "fmt" "math/big" @@ -264,9 +263,14 @@ func SignTextPlain(data hexutil.Bytes) (hexutil.Bytes, string) { // SignTypedData signs EIP-712 conformant typed data // hash = keccak256("\x19${byteVersion}${domainSeparator}${hashStruct(message)}") func (api *SignerAPI) SignTypedData(ctx context.Context, addr common.MixedcaseAddress, typedData TypedData) (hexutil.Bytes, error) { - domainSeparator := typedData.HashStruct("EIP712Domain", typedData.Domain.Map()) - typedDataHash := typedData.HashStruct(typedData.PrimaryType, typedData.Message) - _, err := json.Marshal(typedData.Map()) + if err := typedData.IsValid(); err != nil { + return nil, err + } + domainSeparator, err := typedData.HashStruct("EIP712Domain", typedData.Domain.Map()) + if err != nil { + return nil, err + } + typedDataHash, err := typedData.HashStruct(typedData.PrimaryType, typedData.Message) if err != nil { return nil, err } @@ -282,8 +286,12 @@ func (api *SignerAPI) SignTypedData(ctx context.Context, addr common.MixedcaseAd } // HashStruct generates a keccak256 hash of the encoding of the provided data -func (typedData *TypedData) HashStruct(primaryType string, data EIP712Data) hexutil.Bytes { - return crypto.Keccak256(typedData.EncodeData(primaryType, data)) +func (typedData *TypedData) HashStruct(primaryType string, data EIP712Data) (hexutil.Bytes, error) { + encodedData, err := typedData.EncodeData(primaryType, data) + if err != nil { + return nil, err + } + return crypto.Keccak256(encodedData), nil } // Dependencies returns an array of custom types ordered by their hierarchical reference tree @@ -354,7 +362,7 @@ func (typedData *TypedData) TypeHash(primaryType string) hexutil.Bytes { // `enc(value₁) ‖ enc(value₂) ‖ … ‖ enc(valueₙ)` // // each encoded member is 32-byte long -func (typedData *TypedData) EncodeData(primaryType string, data map[string]interface{}) hexutil.Bytes { +func (typedData *TypedData) EncodeData(primaryType string, data map[string]interface{}) (hexutil.Bytes, error) { encTypes := []string{} encValues := []interface{}{} @@ -362,8 +370,13 @@ func (typedData *TypedData) EncodeData(primaryType string, data map[string]inter encTypes = append(encTypes, "bytes32") encValues = append(encValues, typedData.TypeHash(primaryType)) + // Generate error for a mismatch between the provided type and data + dataMismatchError := func(encType string, encValue interface{}) error { + return fmt.Errorf("provided data '%v' doesn't match type '%s'", encValue, encType) + } + // Handle primitive values - handlePrimitiveValue := func(encType string, encValue interface{}) (string, interface{}) { + handlePrimitiveValue := func(encType string, encValue interface{}) (string, interface{}, error) { var primitiveEncType string var primitiveEncValue interface{} @@ -374,20 +387,32 @@ func (typedData *TypedData) EncodeData(primaryType string, data map[string]inter for i := 0; i < 12; i++ { bytesValue = append(bytesValue, 0) } - for _, _byte := range common.HexToAddress(encValue.(string)) { + stringValue, ok := encValue.(string) + if !ok || !common.IsHexAddress(stringValue) { + return "", nil, dataMismatchError(encType, encValue) + } + for _, _byte := range common.HexToAddress(stringValue) { bytesValue = append(bytesValue, _byte) } primitiveEncValue = bytesValue case "bool": primitiveEncType = "uint256" var int64Val int64 - if encValue.(bool) { + boolValue, ok := encValue.(bool) + if !ok { + return "", nil, dataMismatchError(encType, encValue) + } + if boolValue { int64Val = 1 } primitiveEncValue = abi.U256(big.NewInt(int64Val)) case "bytes", "string": primitiveEncType = "bytes32" - primitiveEncValue = crypto.Keccak256(bytesValueOf(encValue)) + bytesValue, err := bytesValueOf(encValue) + if err != nil { + return "", nil, dataMismatchError(encType, encValue) + } + primitiveEncValue = crypto.Keccak256(bytesValue) default: if strings.HasPrefix(encType, "bytes") { encTypes = append(encTypes, "bytes32") @@ -397,14 +422,21 @@ func (typedData *TypedData) EncodeData(primaryType string, data map[string]inter for i := 0; i < 32-size; i++ { bytesValue = append(bytesValue, 0) } + if _, ok := encValue.(hexutil.Bytes); !ok { + return "", nil, dataMismatchError(encType, encValue) + } bytesValue = append(bytesValue, encValue.(hexutil.Bytes)...) primitiveEncValue = bytesValue } else if strings.HasPrefix(encType, "uint") || strings.HasPrefix(encType, "int") { primitiveEncType = "uint256" - primitiveEncValue = abi.U256(encValue.(*big.Int)) + bigIntValue, ok := encValue.(*big.Int) + if !ok { + return "", nil, dataMismatchError(encType, encValue) + } + primitiveEncValue = abi.U256(bigIntValue) } } - return primitiveEncType, primitiveEncValue + return primitiveEncType, primitiveEncValue, nil } // Add field contents. Structs and arrays have special handlings. @@ -414,24 +446,49 @@ func (typedData *TypedData) EncodeData(primaryType string, data map[string]inter if encType[len(encType)-1:] == "]" { encTypes = append(encTypes, "bytes32") parsedType := strings.Split(encType, "[")[0] + arrayBuffer := bytes.Buffer{} for _, item := range encValue.([]interface{}) { if typedData.Types[parsedType] != nil { - encoding := typedData.EncodeData(parsedType, item.(map[string]interface{})) - arrayBuffer.Write(encoding) + mapValue, ok := item.(map[string]interface{}) + if !ok { + return nil, dataMismatchError(parsedType, item) + } + encodedData, err := typedData.EncodeData(parsedType, mapValue) + if err != nil { + return nil, err + } + arrayBuffer.Write(encodedData) } else { - _, encValue := handlePrimitiveValue(encType, encValue) - arrayBuffer.Write(bytesValueOf(encValue)) + _, encValue, err := handlePrimitiveValue(encType, encValue) + if err != nil { + return nil, err + } + bytesValue, err := bytesValueOf(encValue) + if err != nil { + return nil, err + } + arrayBuffer.Write(bytesValue) } } encValues = append(encValues, crypto.Keccak256(arrayBuffer.Bytes())) } else if typedData.Types[field["type"]] != nil { encTypes = append(encTypes, "bytes32") - mapValue := encValue.(map[string]interface{}) - encValue = crypto.Keccak256(typedData.EncodeData(field["type"], mapValue)) + mapValue, ok := encValue.(map[string]interface{}) + if !ok { + return nil, dataMismatchError(encType, encValue) + } + encodedData, err := typedData.EncodeData(field["type"], mapValue) + if err != nil { + return nil, err + } + encValue = crypto.Keccak256(encodedData) encValues = append(encValues, encValue) } else { - primitiveEncType, primitiveEncValue := handlePrimitiveValue(encType, encValue) + primitiveEncType, primitiveEncValue, err := handlePrimitiveValue(encType, encValue) + if err != nil { + return nil, err + } encTypes = append(encTypes, primitiveEncType) encValues = append(encValues, primitiveEncValue) } @@ -439,31 +496,34 @@ func (typedData *TypedData) EncodeData(primaryType string, data map[string]inter buffer := bytes.Buffer{} for _, encValue := range encValues { - buffer.Write(bytesValueOf(encValue)) + bytesValue, err := bytesValueOf(encValue) + if err != nil { + return nil, err + } + buffer.Write(bytesValue) } - return buffer.Bytes() // https://github.com/ethereumjs/ethereumjs-abi/blob/master/lib/index.js#L336 + return buffer.Bytes(), nil // https://github.com/ethereumjs/ethereumjs-abi/blob/master/lib/index.js#L336 } -func bytesValueOf(_interface interface{}) hexutil.Bytes { +func bytesValueOf(_interface interface{}) (hexutil.Bytes, error) { bytesValue, ok := _interface.(hexutil.Bytes) if ok { - return bytesValue + return bytesValue, nil } switch reflect.TypeOf(_interface) { case reflect.TypeOf(hexutil.Bytes{}): - return _interface.(hexutil.Bytes) + return _interface.(hexutil.Bytes), nil case reflect.TypeOf([]uint8{}): - return _interface.([]uint8) + return _interface.([]uint8), nil case reflect.TypeOf(string("")): - return hexutil.Bytes(_interface.(string)) + return hexutil.Bytes(_interface.(string)), nil default: break } - panic(fmt.Errorf("unrecognized interface type %T", _interface)) - return hexutil.Bytes{} + return nil, fmt.Errorf("unrecognized interface type %T", _interface) } // EcRecover recovers the address associated with the given sig. @@ -523,6 +583,17 @@ func UnmarshalValidatorData(data interface{}) (ValidatorData, error) { }, nil } +// IsValid checks if the typed data is sound +func (typedData *TypedData) IsValid() error { + if err := typedData.Types.IsValid(); err != nil { + return err + } + if err := typedData.Domain.IsValid(); err != nil { + return err + } + return nil +} + // Map is a helper function to generate a map version of the typed data func (typedData *TypedData) Map() map[string]interface{} { dataMap := map[string]interface{}{ @@ -535,27 +606,25 @@ func (typedData *TypedData) Map() map[string]interface{} { return dataMap } -// IsValid checks if the given types object is conformant to the specs +// IsValid checks if the types object is conformant to the specs func (types *EIP712Types) IsValid() error { for typeKey, typeArr := range *types { for _, typeObj := range typeArr { typeVal := typeObj["type"] if typeKey == typeVal { - panic(fmt.Errorf("type %s cannot reference itself", typeVal)) + return fmt.Errorf("type '%s' cannot reference itself", typeVal) } - firstChar := []rune(typeVal)[0] if unicode.IsUpper(firstChar) { if (*types)[typeVal] == nil { - return fmt.Errorf("referenced type %s is undefined", typeVal) + return fmt.Errorf("referenced type '%s' is undefined", typeVal) } } else { - // TODO: better type checking if !isStandardTypeStr(typeVal) { if (*types)[typeVal] != nil { - return fmt.Errorf("custom type %s must be capitalized", typeVal) + return fmt.Errorf("referenced type '%s' must be capitalized", typeVal) } else { - return fmt.Errorf("unknown type %s", typeVal) + return fmt.Errorf("unknown atomic type '%s'", typeVal) } } } diff --git a/signer/core/signed_data_test.go b/signer/core/signed_data_test.go index ac0aab0cd7..745a3141c0 100644 --- a/signer/core/signed_data_test.go +++ b/signer/core/signed_data_test.go @@ -18,6 +18,7 @@ package core import ( "context" + "encoding/json" "fmt" "math/big" "testing" @@ -153,12 +154,20 @@ func TestSignData(t *testing.T) { } func TestHashStruct(t *testing.T) { - mainHash := fmt.Sprintf("0x%s", common.Bytes2Hex(typedData.HashStruct(typedData.PrimaryType, typedData.Message))) + hash, err := typedData.HashStruct(typedData.PrimaryType, typedData.Message) + if err != nil { + t.Fatal(err) + } + mainHash := fmt.Sprintf("0x%s", common.Bytes2Hex(hash)) if mainHash != "0xc52c0ee5d84264471806290a3f2c4cecfc5490626bf912d01f240d7a274b371e" { t.Errorf("Expected different hashStruct result (got %s)", mainHash) } - domainHash := fmt.Sprintf("0x%s", common.Bytes2Hex(typedData.HashStruct("EIP712Domain", typedData.Domain.Map()))) + hash, err = typedData.HashStruct("EIP712Domain", typedData.Domain.Map()) + if err != nil { + t.Error(err) + } + domainHash := fmt.Sprintf("0x%s", common.Bytes2Hex(hash)) if domainHash != "0xf2cee375fa42b42143804025fc449deafd50cc031ca257e0b194a650a912090f" { t.Errorf("Expected different hashStruct result (got %s)", domainHash) } @@ -184,8 +193,259 @@ func TestTypeHash(t *testing.T) { } func TestEncodeData(t *testing.T) { - dataEncoding := fmt.Sprintf("0x%s", common.Bytes2Hex(typedData.EncodeData(typedData.PrimaryType, typedData.Message))) + hash, err := typedData.EncodeData(typedData.PrimaryType, typedData.Message) + if err != nil { + t.Fatal(err) + } + dataEncoding := fmt.Sprintf("0x%s", common.Bytes2Hex(hash)) if dataEncoding != "0xa0cedeb2dc280ba39b857546d74f5549c3a1d7bdc2dd96bf881f76108e23dac2fc71e5fa27ff56c350aa531bc129ebdf613b772b6604664f5d8dbe21b85eb0c8cd54f074a4af31b4411ff6a60c9719dbd559c221c8ac3492d9d872b041d703d1b5aadf3154a261abdd9086fc627b61efca26ae5702701d05cd2305f7c52a2fc8" { t.Errorf("Expected different encodeData result (got %s)", dataEncoding) } } + +func TestMalformedData1(t *testing.T) { + var data = ` + { + "types": { + "EIP712Domain": [ + { + "name": "name", + "type": "string" + }, + { + "name": "version", + "type": "string" + }, + { + "name": "chainId", + "type": "uint256" + }, + { + "name": "verifyingContract", + "type": "address" + } + ], + "Person": [ + { + "name": "name", + "type": "string" + }, + { + "name": "wallet", + "type": "address" + } + ], + "Mail": [ + { + "name": "from", + "type": "Person" + }, + { + "name": "to", + "type": "Person" + }, + { + "name": "contents", + "type": "Person" + } + ] + }, + "primaryType": "Mail", + "domain": { + "name": "Ether Mail", + "version": "1", + "chainId": 1, + "verifyingContract": "0xCcCCccccCCCCcCCCCCCcCcCccCcCCCcCcccccccC" + }, + "message": { + "from": { + "name": "Cow", + "wallet": "0xCD2a3d9F938E13CD947Ec05AbC7FE734Df8DD826" + }, + "to": { + "name": "Bob", + "wallet": "0xbBbBBBBbbBBBbbbBbbBbbbbBBbBbbbbBbBbbBBbB" + }, + "contents": "Hello, Bob!" + } + } + +` + var typedData TypedData + err := json.Unmarshal([]byte(data), &typedData) + if err != nil { + t.Fatalf("unmarshalling failed %v", err) + } + err = typedData.IsValid() + if err != nil { + t.Fatalf("Expected no error, got %v", err) + } + _, err = typedData.HashStruct(typedData.PrimaryType, typedData.Message) + if err.Error() != "provided data 'Hello, Bob!' doesn't match type 'Person'" { + t.Errorf("Expected `provided data 'Hello, Bob!' doesn't match type 'Person'`, got %v", err) + } +} + +func TestMalformedDomainData(t *testing.T) { + var data = ` +{ + "types": { + "EIP712Domain": [ + { + "name": "name", + "type": "string" + }, + { + "name": "version", + "type": "string" + }, + { + "name": "chainId", + "type": "uint256" + }, + { + "name": "verifyingContract", + "type": "address" + } + ], + "Person": [ + { + "name": "name", + "type": "string" + }, + { + "name": "wallet", + "type": "address" + } + ], + "Mail": [ + { + "name": "from", + "type": "Person" + }, + { + "name": "to", + "type": "Person" + }, + { + "name": "contents", + "type": "Blahonga" + } + ] + }, + "primaryType": "Mail", + "domain": { + "name": "Ether Mail", + "version": "1", + "chainId": 1, + "verifyingContract": "0xCcCCccccCCCCcCCCCCCcCcCccCcCCCcCcccccccC" + }, + "message": { + "from": { + "name": "Cow", + "wallet": "0xCD2a3d9F938E13CD947Ec05AbC7FE734Df8DD826" + }, + "to": { + "name": "Bob", + "wallet": "0xbBbBBBBbbBBBbbbBbbBbbbbBBbBbbbbBbBbbBBbB" + }, + "contents": "Hello, Bob!" + } + }` + var typedData TypedData + err := json.Unmarshal([]byte(data), &typedData) + if err != nil { + t.Fatalf("unmarshalling failed %v", err) + } + err = typedData.IsValid() + if err == nil { + t.Fatalf("Expected `referenced type 'Blahonga' is undefined`, got %v", err) + } + _, err = typedData.HashStruct(typedData.PrimaryType, typedData.Message) + if err.Error() != "unrecognized interface type " { + t.Errorf("Expected `unrecognized interface type `, got %v", err) + } +} + +func TestMalformedData3(t *testing.T) { + var data = ` + { + "types": { + "EIP712Domain": [ + { + "name": "name", + "type": "string" + }, + { + "name": "version", + "type": "string" + }, + { + "name": "chainId", + "type": "uint256" + }, + { + "name": "verifyingContract", + "type": "address" + } + ], + "Person": [ + { + "name": "name", + "type": "string" + }, + { + "name": "wallet", + "type": "address" + } + ], + "Mail": [ + { + "name": "from", + "type": "Person" + }, + { + "name": "to", + "type": "Person" + }, + { + "name": "contents", + "type": "string" + } + ] + }, + "primaryType": "Mail", + "domain": { + "name": "Ether Mail", + "version": "1", + "chainId": 1, + "vxerifyingContract": "0xCcCCccccCCCCcCCCCCCcCcCccCcCCCcCcccccccC" + }, + "message": { + "from": { + "name": "Cow", + "wallet": "0xCD2a3d9F938E13CD947Ec05AbC7FE734Df8DD826" + }, + "to": { + "name": "Bob", + "wallet": "0xbBbBBBBbbBBBbbbBbbBbbbbBBbBbbbbBbBbbBBbB" + }, + "contents": "Hello, Bob!" + } + } + +` + var typedData TypedData + err := json.Unmarshal([]byte(data), &typedData) + if err != nil { + t.Fatalf("unmarshalling failed %v", err) + } + err = typedData.IsValid() + if err != nil { + t.Fatalf("Expected no error, got %v", err) + } + _, err = typedData.HashStruct("EIP712Domain", typedData.Domain.Map()) + if err.Error() != "provided data '' doesn't match type 'address'" { + t.Errorf("Expected `provided data '' doesn't match type 'address'`, got %v", err) + } +}