Fixed regexp issues

This commit is contained in:
Paul Berg 2018-11-20 11:47:09 +02:00 committed by Martin Holst Swende
parent 1bcec00190
commit ddefde6058
No known key found for this signature in database
GPG key ID: 683B438C05A5DDF0
2 changed files with 35 additions and 59 deletions

View file

@ -70,25 +70,25 @@ type ValidatorData struct {
} }
type TypedData struct { type TypedData struct {
Types EIP712Types `json:"types"` Types Types `json:"types"`
PrimaryType string `json:"primaryType"` PrimaryType string `json:"primaryType"`
Domain EIP712Domain `json:"domain"` Domain TypedDataDomain `json:"domain"`
Message EIP712Data `json:"message"` Message TypedDataMessage `json:"message"`
Output bytes.Buffer 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 Type string
Value uint Value uint
} }
type EIP712Data = map[string]interface{} type TypedDataMessage = map[string]interface{}
type EIP712Domain struct { type TypedDataDomain struct {
Name string `json:"name"` Name string `json:"name"`
Version string `json:"version"` Version string `json:"version"`
ChainId *big.Int `json:"chainId"` ChainId *big.Int `json:"chainId"`
@ -96,14 +96,7 @@ type EIP712Domain struct {
Salt string `json:"salt"` Salt string `json:"salt"`
} }
const ( 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)))(\[])?$`)
TypeAddress = "address"
TypeBool = "bool"
TypeBytes = "bytes"
TypeInt = "int"
TypeString = "string"
TypeUint = "uint"
)
// Sign receives a request and produces a signature // 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 // 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) encodedData, err := typedData.EncodeData(primaryType, data, 1)
if err != nil { if err != nil {
return nil, err return nil, err
@ -365,7 +358,7 @@ func (typedData *TypedData) TypeHash(primaryType string) hexutil.Bytes {
// //
// each encoded member is 32-byte long // each encoded member is 32-byte long
func (typedData *TypedData) EncodeData(primaryType string, data map[string]interface{}, depth int) (hexutil.Bytes, error) { func (typedData *TypedData) EncodeData(primaryType string, data map[string]interface{}, depth int) (hexutil.Bytes, error) {
encValues := []interface{}{} buffer := bytes.Buffer{}
// Verify extra data // Verify extra data
if len(typedData.Types[primaryType]) < len(data) { if len(typedData.Types[primaryType]) < len(data) {
@ -373,7 +366,7 @@ func (typedData *TypedData) EncodeData(primaryType string, data map[string]inter
} }
// Add typehash // Add typehash
encValues = append(encValues, typedData.TypeHash(primaryType)) buffer.Write(typedData.TypeHash(primaryType))
// Add field contents. Structs and arrays have special handlers. // Add field contents. Structs and arrays have special handlers.
for _, field := range typedData.Types[primaryType] { 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.Truncate(typedData.Output.Len() - 2)
typedData.Output.WriteString(fmt.Sprintf("\n%s],\n", strings.Repeat("\u00a0", depth*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 { } else if typedData.Types[field["type"]] != nil {
mapValue, ok := encValue.(map[string]interface{}) mapValue, ok := encValue.(map[string]interface{})
if !ok { if !ok {
@ -431,25 +424,19 @@ func (typedData *TypedData) EncodeData(primaryType string, data map[string]inter
typedData.Output.Truncate(typedData.Output.Len() - 2) typedData.Output.Truncate(typedData.Output.Len() - 2)
typedData.Output.WriteString(fmt.Sprintf("\n%s},\n", strings.Repeat("\u00a0", depth*2))) typedData.Output.WriteString(fmt.Sprintf("\n%s},\n", strings.Repeat("\u00a0", depth*2)))
encValue = crypto.Keccak256(encodedData) buffer.Write(crypto.Keccak256(encodedData))
encValues = append(encValues, encValue)
} else { } else {
primitiveEncValue, err := typedData.HandlePrimitiveValue(encType, encValue, depth) primitiveEncValue, err := typedData.HandlePrimitiveValue(encType, encValue, depth)
if err != nil { if err != nil {
return nil, err return nil, err
} }
encValues = append(encValues, primitiveEncValue) bytesValue, err := bytesValueOf(primitiveEncValue)
}
}
buffer := bytes.Buffer{}
for _, encValue := range encValues {
bytesValue, err := bytesValueOf(encValue)
if err != nil { if err != nil {
return nil, err return nil, err
} }
buffer.Write(bytesValue) buffer.Write(bytesValue)
} }
}
return buffer.Bytes(), nil // 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
} }
@ -539,6 +526,8 @@ func bytesValueOf(_interface interface{}) (hexutil.Bytes, error) {
switch reflect.TypeOf(_interface) { switch reflect.TypeOf(_interface) {
case reflect.TypeOf(hexutil.Bytes{}): case reflect.TypeOf(hexutil.Bytes{}):
return _interface.(hexutil.Bytes), nil return _interface.(hexutil.Bytes), nil
case reflect.TypeOf([]byte{}):
return hexutil.Bytes(_interface.([]byte)), nil
case reflect.TypeOf([]uint8{}): case reflect.TypeOf([]uint8{}):
return _interface.([]uint8), nil return _interface.([]uint8), nil
case reflect.TypeOf(string("")): case reflect.TypeOf(string("")):
@ -584,6 +573,9 @@ func UnmarshalValidatorData(data interface{}) (ValidatorData, error) {
raw := data.(map[string]interface{}) raw := data.(map[string]interface{})
addr, ok := raw["address"].(string) addr, ok := raw["address"].(string)
if !ok {
return ValidatorData{}, errors.New("validator address is not sent as a string")
}
addrBytes, err := hexutil.Decode(addr) addrBytes, err := hexutil.Decode(addr)
if err != nil { if err != nil {
return ValidatorData{}, err return ValidatorData{}, err
@ -593,6 +585,9 @@ func UnmarshalValidatorData(data interface{}) (ValidatorData, error) {
} }
message, ok := raw["message"].(string) message, ok := raw["message"].(string)
if !ok {
return ValidatorData{}, errors.New("message is not sent as a string")
}
messageBytes, err := hexutil.Decode(message) messageBytes, err := hexutil.Decode(message)
if err != nil { if err != nil {
return ValidatorData{}, err 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 // 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 typeKey, typeArr := range *types {
for _, typeObj := range typeArr { for _, typeObj := range typeArr {
typeVal := typeObj["type"] typeVal := typeObj["type"]
@ -643,7 +638,7 @@ func (types *EIP712Types) Validate() error {
return fmt.Errorf("referenced type '%s' is undefined", typeVal) return fmt.Errorf("referenced type '%s' is undefined", typeVal)
} }
} else { } else {
if !isStandardTypeStr(typeVal) { if !typedDataRegexp.MatchString(typeVal) {
if (*types)[typeVal] != nil { if (*types)[typeVal] != nil {
return fmt.Errorf("referenced type '%s' must be capitalized", typeVal) return fmt.Errorf("referenced type '%s' must be capitalized", typeVal)
} else { } else {
@ -656,28 +651,9 @@ func (types *EIP712Types) Validate() error {
return nil 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 // Validate checks if the given domain is valid, i.e. contains at least
// the minimum viable keys and values // the minimum viable keys and values
func (domain *EIP712Domain) Validate() error { func (domain *TypedDataDomain) Validate() error {
if domain.ChainId == big.NewInt(0) { if domain.ChainId == big.NewInt(0) {
return errors.New("chainId must be specified according to EIP-155") 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 // 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{}{ dataMap := map[string]interface{}{
"chainId": domain.ChainId, "chainId": domain.ChainId,
} }

View file

@ -28,7 +28,7 @@ import (
"github.com/ethereum/go-ethereum/common/hexutil" "github.com/ethereum/go-ethereum/common/hexutil"
) )
var typesStandard = EIP712Types{ var typesStandard = Types{
"EIP712Domain": { "EIP712Domain": {
{ {
"name": "name", "name": "name",
@ -75,7 +75,7 @@ var typesStandard = EIP712Types{
const primaryType = "Mail" const primaryType = "Mail"
var domainStandard = EIP712Domain{ var domainStandard = TypedDataDomain{
"Ether Mail", "Ether Mail",
"1", "1",
big.NewInt(1), big.NewInt(1),
@ -169,7 +169,7 @@ func TestHashStruct(t *testing.T) {
} }
domainHash := fmt.Sprintf("0x%s", common.Bytes2Hex(hash)) domainHash := fmt.Sprintf("0x%s", common.Bytes2Hex(hash))
if domainHash != "0xf2cee375fa42b42143804025fc449deafd50cc031ca257e0b194a650a912090f" { if domainHash != "0xf2cee375fa42b42143804025fc449deafd50cc031ca257e0b194a650a912090f" {
t.Errorf("Expected different hashStruct result (got %s)", domainHash) t.Errorf("Expected different domain hashStruct result (got %s)", domainHash)
} }
} }