From 7a535a8a9ef1141d68c84f8aa7b69ce77b3af5c1 Mon Sep 17 00:00:00 2001 From: Paul Berg Date: Tue, 20 Nov 2018 11:47:09 +0200 Subject: [PATCH] Fixed regexp issues --- signer/core/signed_data.go | 88 ++++++++++++--------------------- signer/core/signed_data_test.go | 6 +-- 2 files changed, 35 insertions(+), 59 deletions(-) diff --git a/signer/core/signed_data.go b/signer/core/signed_data.go index f2c8bc0057..2d26ec8049 100644 --- a/signer/core/signed_data.go +++ b/signer/core/signed_data.go @@ -70,25 +70,25 @@ type ValidatorData struct { } type TypedData struct { - Types EIP712Types `json:"types"` - PrimaryType string `json:"primaryType"` - Domain EIP712Domain `json:"domain"` - Message EIP712Data `json:"message"` + Types Types `json:"types"` + PrimaryType string `json:"primaryType"` + Domain TypedDataDomain `json:"domain"` + Message TypedDataMessage `json:"message"` Output bytes.Buffer } -type EIP712Type []map[string]string +type Type []map[string]string -type EIP712Types map[string]EIP712Type +type Types map[string]Type -type EIP712TypePriority struct { +type TypePriority struct { Type string Value uint } -type EIP712Data = map[string]interface{} +type TypedDataMessage = map[string]interface{} -type EIP712Domain struct { +type TypedDataDomain struct { Name string `json:"name"` Version string `json:"version"` ChainId *big.Int `json:"chainId"` @@ -96,14 +96,7 @@ type EIP712Domain struct { Salt string `json:"salt"` } -const ( - TypeAddress = "address" - TypeBool = "bool" - TypeBytes = "bytes" - TypeInt = "int" - TypeString = "string" - TypeUint = "uint" -) +var typedDataRegexp = regexp.MustCompile(`^((address|bool|bytes|string)|((bytes)([1-9]|[1-2][0-9]|3[0-2]))|((int|uint)(8|16|32|64|128|256)))(\[])?$`) // Sign receives a request and produces a signature @@ -292,7 +285,7 @@ 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, error) { +func (typedData *TypedData) HashStruct(primaryType string, data TypedDataMessage) (hexutil.Bytes, error) { encodedData, err := typedData.EncodeData(primaryType, data, 1) if err != nil { return nil, err @@ -365,7 +358,7 @@ func (typedData *TypedData) TypeHash(primaryType string) hexutil.Bytes { // // each encoded member is 32-byte long func (typedData *TypedData) EncodeData(primaryType string, data map[string]interface{}, depth int) (hexutil.Bytes, error) { - encValues := []interface{}{} + buffer := bytes.Buffer{} // Verify extra data if len(typedData.Types[primaryType]) < len(data) { @@ -373,7 +366,7 @@ func (typedData *TypedData) EncodeData(primaryType string, data map[string]inter } // Add typehash - encValues = append(encValues, typedData.TypeHash(primaryType)) + buffer.Write(typedData.TypeHash(primaryType)) // Add field contents. Structs and arrays have special handlers. for _, field := range typedData.Types[primaryType] { @@ -415,7 +408,7 @@ func (typedData *TypedData) EncodeData(primaryType string, data map[string]inter typedData.Output.Truncate(typedData.Output.Len() - 2) typedData.Output.WriteString(fmt.Sprintf("\n%s],\n", strings.Repeat("\u00a0", depth*2))) - encValues = append(encValues, crypto.Keccak256(arrayBuffer.Bytes())) + buffer.Write(crypto.Keccak256(arrayBuffer.Bytes())) } else if typedData.Types[field["type"]] != nil { mapValue, ok := encValue.(map[string]interface{}) if !ok { @@ -431,26 +424,20 @@ func (typedData *TypedData) EncodeData(primaryType string, data map[string]inter typedData.Output.Truncate(typedData.Output.Len() - 2) typedData.Output.WriteString(fmt.Sprintf("\n%s},\n", strings.Repeat("\u00a0", depth*2))) - encValue = crypto.Keccak256(encodedData) - encValues = append(encValues, encValue) + buffer.Write(crypto.Keccak256(encodedData)) } else { primitiveEncValue, err := typedData.HandlePrimitiveValue(encType, encValue, depth) if err != nil { return nil, err } - encValues = append(encValues, primitiveEncValue) + bytesValue, err := bytesValueOf(primitiveEncValue) + if err != nil { + return nil, err + } + buffer.Write(bytesValue) } } - buffer := bytes.Buffer{} - for _, encValue := range encValues { - bytesValue, err := bytesValueOf(encValue) - if err != nil { - return nil, err - } - buffer.Write(bytesValue) - } - return buffer.Bytes(), nil // https://github.com/ethereumjs/ethereumjs-abi/blob/master/lib/index.js#L336 } @@ -539,6 +526,8 @@ func bytesValueOf(_interface interface{}) (hexutil.Bytes, error) { switch reflect.TypeOf(_interface) { case reflect.TypeOf(hexutil.Bytes{}): return _interface.(hexutil.Bytes), nil + case reflect.TypeOf([]byte{}): + return hexutil.Bytes(_interface.([]byte)), nil case reflect.TypeOf([]uint8{}): return _interface.([]uint8), nil case reflect.TypeOf(string("")): @@ -584,6 +573,9 @@ func UnmarshalValidatorData(data interface{}) (ValidatorData, error) { raw := data.(map[string]interface{}) addr, ok := raw["address"].(string) + if !ok { + return ValidatorData{}, errors.New("validator address is not sent as a string") + } addrBytes, err := hexutil.Decode(addr) if err != nil { return ValidatorData{}, err @@ -593,6 +585,9 @@ func UnmarshalValidatorData(data interface{}) (ValidatorData, error) { } message, ok := raw["message"].(string) + if !ok { + return ValidatorData{}, errors.New("message is not sent as a string") + } messageBytes, err := hexutil.Decode(message) if err != nil { return ValidatorData{}, err @@ -630,7 +625,7 @@ func (typedData *TypedData) Map() map[string]interface{} { } // Validate checks if the types object is conformant to the specs -func (types *EIP712Types) Validate() error { +func (types *Types) Validate() error { for typeKey, typeArr := range *types { for _, typeObj := range typeArr { typeVal := typeObj["type"] @@ -643,7 +638,7 @@ func (types *EIP712Types) Validate() error { return fmt.Errorf("referenced type '%s' is undefined", typeVal) } } else { - if !isStandardTypeStr(typeVal) { + if !typedDataRegexp.MatchString(typeVal) { if (*types)[typeVal] != nil { return fmt.Errorf("referenced type '%s' must be capitalized", typeVal) } else { @@ -656,28 +651,9 @@ func (types *EIP712Types) Validate() error { return nil } -// isStandardType checks if the given type is a EIP712 conformant type -func isStandardTypeStr(encType string) bool { - // Atomic types - exp, _ := regexp.Compile(`^(address|bool|bytes|string)$`) - if exp.MatchString(encType) { - return true - } - - // Dynamic types - exp, _ = regexp.Compile(`^(bytes|int|uint)(\d+)$`) - if exp.MatchString(encType) { - return true - } - - // Arrays - exp, _ = regexp.Compile(`^(address|bool|bytes|string|((bytes|int|uint)(\d+)))\[]$`) - return exp.MatchString(encType) -} - // Validate checks if the given domain is valid, i.e. contains at least // the minimum viable keys and values -func (domain *EIP712Domain) Validate() error { +func (domain *TypedDataDomain) Validate() error { if domain.ChainId == big.NewInt(0) { return errors.New("chainId must be specified according to EIP-155") } @@ -690,7 +666,7 @@ func (domain *EIP712Domain) Validate() error { } // Map is a helper function to generate a map version of the domain -func (domain *EIP712Domain) Map() map[string]interface{} { +func (domain *TypedDataDomain) Map() map[string]interface{} { dataMap := map[string]interface{}{ "chainId": domain.ChainId, } diff --git a/signer/core/signed_data_test.go b/signer/core/signed_data_test.go index 936f0fbb36..e6b6535eff 100644 --- a/signer/core/signed_data_test.go +++ b/signer/core/signed_data_test.go @@ -28,7 +28,7 @@ import ( "github.com/ethereum/go-ethereum/common/hexutil" ) -var typesStandard = EIP712Types{ +var typesStandard = Types{ "EIP712Domain": { { "name": "name", @@ -75,7 +75,7 @@ var typesStandard = EIP712Types{ const primaryType = "Mail" -var domainStandard = EIP712Domain{ +var domainStandard = TypedDataDomain{ "Ether Mail", "1", big.NewInt(1), @@ -169,7 +169,7 @@ func TestHashStruct(t *testing.T) { } domainHash := fmt.Sprintf("0x%s", common.Bytes2Hex(hash)) if domainHash != "0xf2cee375fa42b42143804025fc449deafd50cc031ca257e0b194a650a912090f" { - t.Errorf("Expected different hashStruct result (got %s)", domainHash) + t.Errorf("Expected different domain hashStruct result (got %s)", domainHash) } }