From 354581a15808f86d2cf64044f9ea59f633375c2f Mon Sep 17 00:00:00 2001 From: Michael Standen Date: Fri, 30 Jan 2026 08:08:21 +1300 Subject: [PATCH] Malleable Sapient lib --- lib/sapient/malleable/README.md | 66 +++ lib/sapient/malleable/builder.go | 269 +++++++++++ lib/sapient/malleable/builder_test.go | 65 +++ lib/sapient/malleable/imagehash.go | 227 +++++++++ lib/sapient/malleable/imagehash_test.go | 117 +++++ lib/sapient/malleable/locators_abi.go | 72 +++ lib/sapient/malleable/locators_abi_test.go | 72 +++ lib/sapient/malleable/locators_packedcalls.go | 108 +++++ .../malleable/locators_packedcalls_test.go | 119 +++++ lib/sapient/malleable/path.go | 450 ++++++++++++++++++ lib/sapient/malleable/path_test.go | 265 +++++++++++ lib/sapient/malleable/signature.go | 102 ++++ lib/sapient/malleable/spans.go | 57 +++ 13 files changed, 1989 insertions(+) create mode 100644 lib/sapient/malleable/README.md create mode 100644 lib/sapient/malleable/builder.go create mode 100644 lib/sapient/malleable/builder_test.go create mode 100644 lib/sapient/malleable/imagehash.go create mode 100644 lib/sapient/malleable/imagehash_test.go create mode 100644 lib/sapient/malleable/locators_abi.go create mode 100644 lib/sapient/malleable/locators_abi_test.go create mode 100644 lib/sapient/malleable/locators_packedcalls.go create mode 100644 lib/sapient/malleable/locators_packedcalls_test.go create mode 100644 lib/sapient/malleable/path.go create mode 100644 lib/sapient/malleable/path_test.go create mode 100644 lib/sapient/malleable/signature.go create mode 100644 lib/sapient/malleable/spans.go diff --git a/lib/sapient/malleable/README.md b/lib/sapient/malleable/README.md new file mode 100644 index 00000000..9b477268 --- /dev/null +++ b/lib/sapient/malleable/README.md @@ -0,0 +1,66 @@ +# Malleable Sapient (Go) + +Build MalleableSapient signatures by locating byte ranges in call data, and optionally compute the image hash. + +## Usage + +```go +payload := v3.NewCallsPayload(...) + +permitValue := malleable.NewPath(). + CallData(0). + ABI(trailsABI, "hydrateExecute"). + ArgBytesData("packedPayload"). + EncodedCallsPayload(). + EncodedCallData(0). + ABI(erc2612ABI, "permit"). + ArgSlot("value"). + AsSelector() + +transferValue := malleable.NewPath(). + CallData(0). + ABI(trailsABI, "hydrateExecute"). + ArgBytesData("packedPayload"). + EncodedCallsPayload(). + EncodedCallData(1). + ABI(erc20ABI, "transferFrom"). + ArgSlot("_value"). + AsSelector() + +b := malleable.NewBuilder(payload, &malleable.BuilderOptions{ + ValidateRepeats: true, + MergeAdjacentStatic: true, +}) + +b.Repeat(permitValue, transferValue) // repeat constraint + +// mark other malleable fields +b.Malleable(malleable.NewPath(). + CallData(0). + ABI(trailsABI, "hydrateExecute"). + ArgBytesData("packedPayload"). + EncodedCallsPayload(). + EncodedCallData(0). + ABI(erc2612ABI, "permit"). + ArgSlot("deadline"). + AsSelector(), +) + +sig, _, err := b.Build() +``` + +If ABI params are unnamed, use index-based selectors: + +```go +value := malleable.NewPath(). + CallData(0). + ABI(erc20ABI, "transferFrom"). + ArgSlotIndex(2). + AsSelector() +``` + +Compute the image hash: + +```go +hash, err := malleable.ComputeImageHash(payload, sig, chainID) +``` diff --git a/lib/sapient/malleable/builder.go b/lib/sapient/malleable/builder.go new file mode 100644 index 00000000..f8041c13 --- /dev/null +++ b/lib/sapient/malleable/builder.go @@ -0,0 +1,269 @@ +package malleable + +import ( + "fmt" + "sort" + + "github.com/0xsequence/ethkit/go-ethereum/crypto" + v3 "github.com/0xsequence/go-sequence/core/v3" +) + +type BuilderOptions struct { + ValidateRepeats bool + MergeAdjacentStatic bool + MaxOffset uint32 + MaxSize uint32 +} + +type Builder struct { + payload *v3.CallsPayload + options BuilderOptions + malleable []Selector + repeats []repeatSelector +} + +type repeatSelector struct { + a Selector + b Selector +} + +type Plan struct { + Static []StaticSection + Repeat []RepeatSection +} + +func (p *Plan) DebugString() string { + out := "static:" + for _, s := range p.Static { + out += fmt.Sprintf(" [t=%d c=%d s=%d]", s.TIndex, s.CIndex, s.Size) + } + out += " repeat:" + for _, r := range p.Repeat { + out += fmt.Sprintf(" [t=%d c=%d s=%d t2=%d c2=%d]", r.TIndex, r.CIndex, r.Size, r.TIndex2, r.CIndex2) + } + return out +} + +func NewBuilder(payload *v3.CallsPayload, opts *BuilderOptions) *Builder { + options := BuilderOptions{ + MaxOffset: 0xFFFF, + MaxSize: 0xFFFF, + ValidateRepeats: false, + MergeAdjacentStatic: true, + } + if opts != nil { + if opts.MaxOffset != 0 { + options.MaxOffset = opts.MaxOffset + } + if opts.MaxSize != 0 { + options.MaxSize = opts.MaxSize + } + options.ValidateRepeats = opts.ValidateRepeats + options.MergeAdjacentStatic = opts.MergeAdjacentStatic + } + + return &Builder{ + payload: payload, + options: options, + } +} + +func (b *Builder) Malleable(sel Selector) *Builder { + b.malleable = append(b.malleable, sel) + return b +} + +func (b *Builder) Repeat(a Selector, b2 Selector) *Builder { + b.repeats = append(b.repeats, repeatSelector{a: a, b: b2}) + return b +} + +func (b *Builder) Build() ([]byte, *Plan, error) { + if b.payload == nil { + return nil, nil, fmt.Errorf("payload is nil") + } + if len(b.payload.Calls) > 128 { + return nil, nil, fmt.Errorf("too many calls (%d): tindex is 7-bit", len(b.payload.Calls)) + } + + exByCall := make([][]Span, len(b.payload.Calls)) + + addExclude := func(r ByteRange) error { + if r.CallIndex < 0 || r.CallIndex >= len(b.payload.Calls) { + return fmt.Errorf("tindex out of range: %d", r.CallIndex) + } + dataLen := len(b.payload.Calls[r.CallIndex].Data) + if r.Offset < 0 || r.Size < 0 || r.Offset+r.Size > dataLen { + return fmt.Errorf("span out of bounds: [%d,%d) > %d", r.Offset, r.Offset+r.Size, dataLen) + } + exByCall[r.CallIndex] = append(exByCall[r.CallIndex], Span{Start: r.Offset, Len: r.Size}) + return nil + } + + for _, sel := range b.malleable { + ranges, err := sel.Resolve(b.payload) + if err != nil { + return nil, nil, fmt.Errorf("malleable %s: %w", sel.String(), err) + } + for _, r := range ranges { + if err := addExclude(r); err != nil { + return nil, nil, fmt.Errorf("malleable %s: %w", sel.String(), err) + } + } + } + + var repeats []RepeatSection + for _, rp := range b.repeats { + rangesA, err := rp.a.Resolve(b.payload) + if err != nil { + return nil, nil, fmt.Errorf("repeat a %s: %w", rp.a.String(), err) + } + rangesB, err := rp.b.Resolve(b.payload) + if err != nil { + return nil, nil, fmt.Errorf("repeat b %s: %w", rp.b.String(), err) + } + if len(rangesA) != 1 || len(rangesB) != 1 { + return nil, nil, fmt.Errorf("repeat expects single range per selector") + } + a := rangesA[0] + bb := rangesB[0] + if a.Size != bb.Size { + return nil, nil, fmt.Errorf("repeat size mismatch: %d vs %d", a.Size, bb.Size) + } + if err := addExclude(a); err != nil { + return nil, nil, fmt.Errorf("repeat a: %w", err) + } + if err := addExclude(bb); err != nil { + return nil, nil, fmt.Errorf("repeat b: %w", err) + } + if b.options.ValidateRepeats { + sectionA, err := a.Slice(b.payload) + if err != nil { + return nil, nil, err + } + sectionB, err := bb.Slice(b.payload) + if err != nil { + return nil, nil, err + } + if crypto.Keccak256Hash(sectionA) != crypto.Keccak256Hash(sectionB) { + return nil, nil, fmt.Errorf("repeat section mismatch") + } + } + repeats = append(repeats, RepeatSection{ + TIndex: uint8(a.CallIndex), + CIndex: uint16(a.Offset), + Size: uint16(a.Size), + TIndex2: uint8(bb.CallIndex), + CIndex2: uint16(bb.Offset), + }) + } + + var statics []SpanWithCall + for t := range b.payload.Calls { + length := len(b.payload.Calls[t].Data) + ex := mergeSpans(exByCall[t]) + cursor := 0 + for _, s := range ex { + if cursor < s.Start { + statics = append(statics, SpanWithCall{CallIndex: t, Span: Span{Start: cursor, Len: s.Start - cursor}}) + } + cursor = max(cursor, s.Start+s.Len) + } + if cursor < length { + statics = append(statics, SpanWithCall{CallIndex: t, Span: Span{Start: cursor, Len: length - cursor}}) + } + } + + sort.Slice(statics, func(i, j int) bool { + if statics[i].CallIndex != statics[j].CallIndex { + return statics[i].CallIndex < statics[j].CallIndex + } + return statics[i].Start < statics[j].Start + }) + + if b.options.MergeAdjacentStatic { + statics = mergeAdjacentStatics(statics) + } + + staticSections, err := b.encodeStaticSections(statics) + if err != nil { + return nil, nil, err + } + + signature, err := EncodeSignature(staticSections, repeats) + if err != nil { + return nil, nil, err + } + + return signature, &Plan{Static: staticSections, Repeat: repeats}, nil +} + +type SpanWithCall struct { + CallIndex int + Span +} + +func (b *Builder) encodeStaticSections(statics []SpanWithCall) ([]StaticSection, error) { + var sections []StaticSection + for _, s := range statics { + if s.Len == 0 { + continue + } + if s.CallIndex < 0 || s.CallIndex >= len(b.payload.Calls) { + return nil, fmt.Errorf("tindex out of range: %d", s.CallIndex) + } + offset := s.Start + length := s.Len + for length > 0 { + chunk := length + if uint32(chunk) > b.options.MaxSize { + chunk = int(b.options.MaxSize) + } + if uint32(offset) > b.options.MaxOffset { + return nil, fmt.Errorf("cindex too large: %d", offset) + } + sections = append(sections, StaticSection{ + TIndex: uint8(s.CallIndex), + CIndex: uint16(offset), + Size: uint16(chunk), + }) + offset += chunk + length -= chunk + } + } + return sections, nil +} + +func mergeSpans(spans []Span) []Span { + if len(spans) == 0 { + return nil + } + sort.Slice(spans, func(i, j int) bool { return spans[i].Start < spans[j].Start }) + out := []Span{spans[0]} + for _, s := range spans[1:] { + last := &out[len(out)-1] + if s.Start <= last.Start+last.Len { + end := max(last.Start+last.Len, s.Start+s.Len) + last.Len = end - last.Start + } else { + out = append(out, s) + } + } + return out +} + +func mergeAdjacentStatics(statics []SpanWithCall) []SpanWithCall { + if len(statics) == 0 { + return statics + } + out := []SpanWithCall{statics[0]} + for _, s := range statics[1:] { + last := &out[len(out)-1] + if s.CallIndex == last.CallIndex && s.Start == last.Start+last.Len { + last.Len += s.Len + } else { + out = append(out, s) + } + } + return out +} diff --git a/lib/sapient/malleable/builder_test.go b/lib/sapient/malleable/builder_test.go new file mode 100644 index 00000000..b1e5a855 --- /dev/null +++ b/lib/sapient/malleable/builder_test.go @@ -0,0 +1,65 @@ +package malleable + +import ( + "math/big" + "testing" + + "github.com/0xsequence/ethkit/go-ethereum/common" + v3 "github.com/0xsequence/go-sequence/core/v3" + "github.com/stretchr/testify/require" +) + +func TestBuilder_StaticComplement(t *testing.T) { + payload := v3.NewCallsPayload(common.Address{}, big.NewInt(1), []v3.Call{ + {Data: []byte{0, 1, 2, 3, 4, 5, 6, 7, 8, 9}}, + }, big.NewInt(0), big.NewInt(0)) + + builder := NewBuilder(&payload, &BuilderOptions{MergeAdjacentStatic: true}) + builder.Malleable(NewRangeSelector(0, 2, 2)) + builder.Malleable(NewRangeSelector(0, 7, 2)) + + sig, plan, err := builder.Build() + require.NoError(t, err) + require.NotEmpty(t, sig) + + require.Equal(t, []StaticSection{ + {TIndex: 0, CIndex: 0, Size: 2}, + {TIndex: 0, CIndex: 4, Size: 3}, + {TIndex: 0, CIndex: 9, Size: 1}, + }, plan.Static) + require.Empty(t, plan.Repeat) +} + +func TestBuilder_RepeatValidation(t *testing.T) { + payload := v3.NewCallsPayload(common.Address{}, big.NewInt(1), []v3.Call{ + {Data: []byte{0xaa, 0xbb, 0xcc, 0xdd}}, + {Data: []byte{0xaa, 0xbb, 0xee, 0xff}}, + }, big.NewInt(0), big.NewInt(0)) + + builder := NewBuilder(&payload, &BuilderOptions{ValidateRepeats: true}) + builder.Repeat(NewRangeSelector(0, 0, 2), NewRangeSelector(1, 0, 2)) + _, _, err := builder.Build() + require.NoError(t, err) + + builder = NewBuilder(&payload, &BuilderOptions{ValidateRepeats: true}) + builder.Repeat(NewRangeSelector(0, 0, 2), NewRangeSelector(1, 2, 2)) + _, _, err = builder.Build() + require.Error(t, err) +} + +func TestEncodeDecodeSignature_RoundTrip(t *testing.T) { + statics := []StaticSection{ + {TIndex: 0, CIndex: 1, Size: 2}, + {TIndex: 1, CIndex: 3, Size: 4}, + } + repeats := []RepeatSection{ + {TIndex: 0, CIndex: 5, Size: 2, TIndex2: 1, CIndex2: 6}, + } + + sig, err := EncodeSignature(statics, repeats) + require.NoError(t, err) + + sections, err := DecodeSignature(sig) + require.NoError(t, err) + require.Len(t, sections, 3) +} diff --git a/lib/sapient/malleable/imagehash.go b/lib/sapient/malleable/imagehash.go new file mode 100644 index 00000000..2050796e --- /dev/null +++ b/lib/sapient/malleable/imagehash.go @@ -0,0 +1,227 @@ +package malleable + +import ( + "encoding/binary" + "fmt" + "math/big" + + "github.com/0xsequence/ethkit/go-ethereum/accounts/abi" + "github.com/0xsequence/ethkit/go-ethereum/common" + "github.com/0xsequence/ethkit/go-ethereum/crypto" + v3 "github.com/0xsequence/go-sequence/core/v3" +) + +func ComputeImageHash(payload *v3.CallsPayload, signature []byte, chainID *big.Int) (common.Hash, error) { + if payload == nil { + return common.Hash{}, fmt.Errorf("payload is nil") + } + if len(payload.Calls) > 128 { + return common.Hash{}, fmt.Errorf("too many calls (%d)", len(payload.Calls)) + } + + space := nz(payload.Space) + nonce := nz(payload.Nonce) + root := fkeccak256(u256Bytes32(space), u256Bytes32(nonce)) + + payloadChainID := payload.ChainID() + noChainID := payloadChainID == nil || payloadChainID.Sign() == 0 + if noChainID { + root = fkeccak256(root, common.Hash{}) + } else { + if chainID == nil { + chainID = payloadChainID + } + root = fkeccak256(root, u256Bytes32(nz(chainID))) + } + + stringTy, _ := abi.NewType("string", "", nil) + u256Ty, _ := abi.NewType("uint256", "", nil) + addrTy, _ := abi.NewType("address", "", nil) + boolTy, _ := abi.NewType("bool", "", nil) + bytesTy, _ := abi.NewType("bytes", "", nil) + + callMetaArgs := abi.Arguments{ + {Type: stringTy}, + {Type: u256Ty}, + {Type: addrTy}, + {Type: u256Ty}, + {Type: u256Ty}, + {Type: boolTy}, + {Type: boolTy}, + {Type: u256Ty}, + } + + staticArgs := abi.Arguments{ + {Type: stringTy}, + {Type: u256Ty}, + {Type: u256Ty}, + {Type: bytesTy}, + } + + repeatArgs := abi.Arguments{ + {Type: stringTy}, + {Type: u256Ty}, + {Type: u256Ty}, + {Type: u256Ty}, + {Type: u256Ty}, + {Type: u256Ty}, + } + + for i := 0; i < len(payload.Calls); i++ { + c := payload.Calls[i] + b, err := callMetaArgs.Pack( + "call", + big.NewInt(int64(i)), + c.To, + nz(c.Value), + nz(c.GasLimit), + c.DelegateCall, + c.OnlyFallback, + big.NewInt(int64(c.BehaviorOnError)), + ) + if err != nil { + return common.Hash{}, err + } + root = fkeccak256(root, crypto.Keccak256Hash(b)) + } + + r := 0 + for r < len(signature) { + if r+5 > len(signature) { + return common.Hash{}, fmt.Errorf("signature truncated at %d", r) + } + tRaw := signature[r] + r++ + cindex := binary.BigEndian.Uint16(signature[r : r+2]) + r += 2 + size := binary.BigEndian.Uint16(signature[r : r+2]) + r += 2 + + repeat := (tRaw & 0x80) != 0 + t := int(tRaw & 0x7F) + if t >= len(payload.Calls) { + return common.Hash{}, fmt.Errorf("tindex out of range: %d", t) + } + if int(cindex)+int(size) > len(payload.Calls[t].Data) { + return common.Hash{}, fmt.Errorf("section out of bounds (t=%d)", t) + } + section := payload.Calls[t].Data[cindex : cindex+size] + + if repeat { + if r+3 > len(signature) { + return common.Hash{}, fmt.Errorf("repeat truncated at %d", r) + } + t2 := int(signature[r]) + r++ + c2 := binary.BigEndian.Uint16(signature[r : r+2]) + r += 2 + + if t2 >= len(payload.Calls) { + return common.Hash{}, fmt.Errorf("tindex2 out of range: %d", t2) + } + if int(c2)+int(size) > len(payload.Calls[t2].Data) { + return common.Hash{}, fmt.Errorf("repeat section2 out of bounds") + } + section2 := payload.Calls[t2].Data[c2 : c2+size] + if crypto.Keccak256Hash(section) != crypto.Keccak256Hash(section2) { + return common.Hash{}, fmt.Errorf("repeat section mismatch") + } + + b, err := repeatArgs.Pack( + "repeat-section", + big.NewInt(int64(t)), + new(big.Int).SetUint64(uint64(cindex)), + new(big.Int).SetUint64(uint64(size)), + big.NewInt(int64(t2)), + new(big.Int).SetUint64(uint64(c2)), + ) + if err != nil { + return common.Hash{}, err + } + root = fkeccak256(root, crypto.Keccak256Hash(b)) + } else { + b, err := staticArgs.Pack( + "static-section", + big.NewInt(int64(t)), + new(big.Int).SetUint64(uint64(cindex)), + section, + ) + if err != nil { + return common.Hash{}, err + } + root = fkeccak256(root, crypto.Keccak256Hash(b)) + } + } + + return root, nil +} + +func ValidateSignature(payload *v3.CallsPayload, signature []byte) error { + if payload == nil { + return fmt.Errorf("payload is nil") + } + r := 0 + for r < len(signature) { + if r+5 > len(signature) { + return fmt.Errorf("signature truncated at %d", r) + } + tRaw := signature[r] + r++ + cindex := binary.BigEndian.Uint16(signature[r : r+2]) + r += 2 + size := binary.BigEndian.Uint16(signature[r : r+2]) + r += 2 + + t := int(tRaw & 0x7F) + if t >= len(payload.Calls) { + return fmt.Errorf("tindex out of range: %d", t) + } + if int(cindex)+int(size) > len(payload.Calls[t].Data) { + return fmt.Errorf("section out of bounds (t=%d)", t) + } + + if tRaw&0x80 != 0 { + if r+3 > len(signature) { + return fmt.Errorf("repeat truncated at %d", r) + } + t2 := int(signature[r]) + r++ + c2 := binary.BigEndian.Uint16(signature[r : r+2]) + r += 2 + + if t2 >= len(payload.Calls) { + return fmt.Errorf("tindex2 out of range: %d", t2) + } + if int(c2)+int(size) > len(payload.Calls[t2].Data) { + return fmt.Errorf("repeat section2 out of bounds") + } + section := payload.Calls[t].Data[cindex : cindex+size] + section2 := payload.Calls[t2].Data[c2 : c2+size] + if crypto.Keccak256Hash(section) != crypto.Keccak256Hash(section2) { + return fmt.Errorf("repeat section mismatch") + } + } + } + return nil +} + +func fkeccak256(a, b common.Hash) common.Hash { + return crypto.Keccak256Hash(a[:], b[:]) +} + +func u256Bytes32(x *big.Int) common.Hash { + var out common.Hash + if x == nil { + return out + } + b := x.Bytes() + copy(out[32-len(b):], b) + return out +} + +func nz(x *big.Int) *big.Int { + if x == nil { + return big.NewInt(0) + } + return x +} diff --git a/lib/sapient/malleable/imagehash_test.go b/lib/sapient/malleable/imagehash_test.go new file mode 100644 index 00000000..7d90ba96 --- /dev/null +++ b/lib/sapient/malleable/imagehash_test.go @@ -0,0 +1,117 @@ +package malleable + +import ( + "context" + "math/big" + "testing" + + "github.com/0xsequence/ethkit/ethrpc" + "github.com/0xsequence/ethkit/go-ethereum/accounts/abi/bind" + "github.com/0xsequence/ethkit/go-ethereum/common" + "github.com/0xsequence/go-sequence/contracts/gen/trailsutils" + v3 "github.com/0xsequence/go-sequence/core/v3" + "github.com/stretchr/testify/require" +) + +const ( + trailsUtils = "0x0000000066c426Fe13962e276f894F11Aa6ebbF2" + baseRPC = "https://nodes.sequence.app/base" + baseChainID = 8453 +) + +func TestComputeImageHash_DifferentSignatures(t *testing.T) { + payload := v3.NewCallsPayload(common.Address{}, big.NewInt(1), []v3.Call{ + {Data: []byte{0x01, 0x02, 0x03, 0x04, 0x05}}, + }, big.NewInt(0), big.NewInt(0)) + + hash1, err := ComputeImageHash(&payload, []byte{}, big.NewInt(1)) + require.NoError(t, err) + + signature := []byte{0x00, 0x00, 0x00, 0x00, 0x03} + hash2, err := ComputeImageHash(&payload, signature, big.NewInt(1)) + require.NoError(t, err) + + require.NotEqual(t, hash1, hash2) +} + +func TestValidateSignature_RepeatMismatch(t *testing.T) { + payload := v3.NewCallsPayload(common.Address{}, big.NewInt(1), []v3.Call{ + {Data: []byte{0x01, 0x02, 0x03}}, + {Data: []byte{0x01, 0xff, 0x03}}, + }, big.NewInt(0), big.NewInt(0)) + + repeat := RepeatSection{TIndex: 0, CIndex: 0, Size: 2, TIndex2: 1, CIndex2: 0} + sig, err := EncodeSignature(nil, []RepeatSection{repeat}) + require.NoError(t, err) + + err = ValidateSignature(&payload, sig) + require.Error(t, err) +} + +func TestComputeImageHash_RecoverSapientSignature(t *testing.T) { + payload := v3.NewCallsPayload(common.Address{}, big.NewInt(baseChainID), []v3.Call{ + { + To: common.HexToAddress("0x1111111111111111111111111111111111111111"), + Value: big.NewInt(0), + Data: []byte{0x01, 0x02}, + GasLimit: big.NewInt(0), + DelegateCall: false, + OnlyFallback: false, + BehaviorOnError: v3.BehaviorOnErrorRevert, + }, + }, big.NewInt(0), big.NewInt(0)) + + signature := []byte{} + expected, err := ComputeImageHash(&payload, signature, big.NewInt(baseChainID)) + require.NoError(t, err) + + provider, err := ethrpc.NewProvider(baseRPC) + require.NoError(t, err) + + contract, err := trailsutils.NewTrailsUtilsCaller(common.HexToAddress(trailsUtils), provider) + require.NoError(t, err) + + callOpts := &bind.CallOpts{Context: context.Background()} + onChainHash, err := contract.RecoverSapientSignature(callOpts, toTrailsUtilsPayload(payload), signature) + require.NoError(t, err) + + require.Equal(t, expected, common.BytesToHash(onChainHash[:])) +} + +func toTrailsUtilsPayload(payload v3.CallsPayload) trailsutils.PayloadDecoded { + calls := make([]trailsutils.PayloadCall, len(payload.Calls)) + for i, call := range payload.Calls { + value := call.Value + if value == nil { + value = big.NewInt(0) + } + gasLimit := call.GasLimit + if gasLimit == nil { + gasLimit = big.NewInt(0) + } + + calls[i] = trailsutils.PayloadCall{ + To: call.To, + Value: value, + Data: call.Data, + GasLimit: gasLimit, + DelegateCall: call.DelegateCall, + OnlyFallback: call.OnlyFallback, + BehaviorOnError: big.NewInt(int64(call.BehaviorOnError)), + } + } + + noChainID := payload.ChainID().Sign() == 0 + + return trailsutils.PayloadDecoded{ + Kind: v3.KindTransactions, + NoChainId: noChainID, + Calls: calls, + Space: payload.Space, + Nonce: payload.Nonce, + Message: []byte{}, + ImageHash: [32]byte{}, + Digest: [32]byte{}, + ParentWallets: []common.Address{}, + } +} diff --git a/lib/sapient/malleable/locators_abi.go b/lib/sapient/malleable/locators_abi.go new file mode 100644 index 00000000..5fa1f8a0 --- /dev/null +++ b/lib/sapient/malleable/locators_abi.go @@ -0,0 +1,72 @@ +package malleable + +import ( + "fmt" + "math/big" + + "github.com/0xsequence/ethkit/go-ethereum/accounts/abi" +) + +func CalldataStaticWord(method abi.Method, argIndex int) (start, length int, err error) { + if argIndex < 0 || argIndex >= len(method.Inputs) { + return 0, 0, fmt.Errorf("argIndex out of range") + } + return 4 + 32*argIndex, 32, nil +} + +// CalldataBytesContent returns the raw bytes/string content (excludes length word and padding). +func CalldataBytesContent(calldata []byte, method abi.Method, argIndex int) (start, length int, err error) { + tailStart, dataLen, err := calldataBytesTail(calldata, method, argIndex) + if err != nil { + return 0, 0, err + } + contentStart := tailStart + 32 + if contentStart+dataLen > len(calldata) { + return 0, 0, fmt.Errorf("calldata too short for bytes content") + } + return contentStart, dataLen, nil +} + +// CalldataBytesEncoded returns the full ABI-encoded tail: length word + data + padding. +func CalldataBytesEncoded(calldata []byte, method abi.Method, argIndex int) (start, length int, err error) { + tailStart, dataLen, err := calldataBytesTail(calldata, method, argIndex) + if err != nil { + return 0, 0, err + } + padded := ((dataLen + 31) / 32) * 32 + total := 32 + padded + if tailStart+total > len(calldata) { + return 0, 0, fmt.Errorf("calldata too short for bytes encoded") + } + return tailStart, total, nil +} + +func calldataBytesTail(calldata []byte, method abi.Method, argIndex int) (tailStart int, dataLen int, err error) { + if argIndex < 0 || argIndex >= len(method.Inputs) { + return 0, 0, fmt.Errorf("argIndex out of range") + } + t := method.Inputs[argIndex].Type + if t.T != abi.BytesTy && t.T != abi.StringTy { + return 0, 0, fmt.Errorf("arg %d is not bytes/string (got %s)", argIndex, t.String()) + } + head := 4 + 32*argIndex + if head+32 > len(calldata) { + return 0, 0, fmt.Errorf("calldata too short for head word") + } + + off := new(big.Int).SetBytes(calldata[head : head+32]) + if !off.IsInt64() { + return 0, 0, fmt.Errorf("dynamic offset too large") + } + tailStart = 4 + int(off.Int64()) + if tailStart+32 > len(calldata) { + return 0, 0, fmt.Errorf("calldata too short for tail length word") + } + + l := new(big.Int).SetBytes(calldata[tailStart : tailStart+32]) + if !l.IsInt64() { + return 0, 0, fmt.Errorf("bytes length too large") + } + dataLen = int(l.Int64()) + return tailStart, dataLen, nil +} diff --git a/lib/sapient/malleable/locators_abi_test.go b/lib/sapient/malleable/locators_abi_test.go new file mode 100644 index 00000000..6b337c0d --- /dev/null +++ b/lib/sapient/malleable/locators_abi_test.go @@ -0,0 +1,72 @@ +package malleable + +import ( + "math/big" + "strings" + "testing" + + "github.com/0xsequence/ethkit/go-ethereum/accounts/abi" + "github.com/0xsequence/ethkit/go-ethereum/common" + "github.com/stretchr/testify/require" +) + +func TestCalldataBytesContent(t *testing.T) { + abiDef := `[{"name":"hydrateExecute","type":"function","inputs":[{"name":"payload","type":"bytes"},{"name":"hydrateData","type":"bytes"}]}]` + parsedABI, err := abi.JSON(strings.NewReader(abiDef)) + require.NoError(t, err) + + payloadBytes := []byte{0x11, 0x22, 0x33, 0x44} + hydrateData := []byte{0xaa} + calldata, err := parsedABI.Pack("hydrateExecute", payloadBytes, hydrateData) + require.NoError(t, err) + + method := parsedABI.Methods["hydrateExecute"] + start, length, err := CalldataBytesContent(calldata, method, 0) + require.NoError(t, err) + require.Equal(t, payloadBytes, calldata[start:start+length]) +} + +func TestCalldataBytesEncoded(t *testing.T) { + abiDef := `[{"name":"hydrateExecute","type":"function","inputs":[{"name":"payload","type":"bytes"},{"name":"hydrateData","type":"bytes"}]}]` + parsedABI, err := abi.JSON(strings.NewReader(abiDef)) + require.NoError(t, err) + + payloadBytes := []byte{0x11, 0x22, 0x33, 0x44} + hydrateData := []byte{0xaa} + calldata, err := parsedABI.Pack("hydrateExecute", payloadBytes, hydrateData) + require.NoError(t, err) + + lengthWord := common.BigToHash(big.NewInt(int64(len(payloadBytes)))).Bytes() + paddedLen := ((len(payloadBytes) + 31) / 32) * 32 + padding := make([]byte, paddedLen-len(payloadBytes)) + encodedPayloadBytes := append(append(lengthWord, payloadBytes...), padding...) + + method := parsedABI.Methods["hydrateExecute"] + start, length, err := CalldataBytesEncoded(calldata, method, 0) + require.NoError(t, err) + require.Equal(t, encodedPayloadBytes, calldata[start:start+length]) +} + +func TestCalldataStaticWord(t *testing.T) { + abiDef := `[{"name":"permit","type":"function","inputs":[{"name":"owner","type":"address"},{"name":"spender","type":"address"},{"name":"value","type":"uint256"},{"name":"deadline","type":"uint256"},{"name":"v","type":"uint8"},{"name":"r","type":"bytes32"},{"name":"s","type":"bytes32"}]}]` + parsedABI, err := abi.JSON(strings.NewReader(abiDef)) + require.NoError(t, err) + + calldata, err := parsedABI.Pack( + "permit", + common.HexToAddress("0x1111111111111111111111111111111111111111"), + common.HexToAddress("0x2222222222222222222222222222222222222222"), + big.NewInt(123), + big.NewInt(456), + uint8(27), + common.HexToHash("0x01"), + common.HexToHash("0x02"), + ) + require.NoError(t, err) + + method := parsedABI.Methods["permit"] + start, length, err := CalldataStaticWord(method, 2) + require.NoError(t, err) + require.Equal(t, 32, length) + require.Equal(t, calldata[4+32*2:4+32*3], calldata[start:start+length]) +} diff --git a/lib/sapient/malleable/locators_packedcalls.go b/lib/sapient/malleable/locators_packedcalls.go new file mode 100644 index 00000000..a8d55f19 --- /dev/null +++ b/lib/sapient/malleable/locators_packedcalls.go @@ -0,0 +1,108 @@ +package malleable + +import ( + "encoding/binary" + "fmt" +) + +type PackedCallsLayout struct { + GlobalFlag byte + NumCalls int + CallData []Span +} + +// ParsePackedCalls parses the packed calls layout from the packed payload.calls data. +func ParsePackedCalls(packed []byte) (*PackedCallsLayout, error) { + if len(packed) < 1 { + return nil, fmt.Errorf("packed calls too short") + } + p := 0 + globalFlag := packed[p] + p++ + + if globalFlag&0x01 == 0x00 { + if p+20 > len(packed) { + return nil, fmt.Errorf("packed calls truncated reading space") + } + p += 20 + } + + nonceSize := int((globalFlag >> 1) & 0x07) + if nonceSize > 0 { + if p+nonceSize > len(packed) { + return nil, fmt.Errorf("packed calls truncated reading nonce (%d bytes)", nonceSize) + } + p += nonceSize + } + + var numCalls int + if globalFlag&0x10 == 0x10 { + numCalls = 1 + } else { + if globalFlag&0x20 == 0x20 { + if p+2 > len(packed) { + return nil, fmt.Errorf("packed calls truncated reading numCalls (u16)") + } + numCalls = int(binary.BigEndian.Uint16(packed[p : p+2])) + p += 2 + } else { + if p+1 > len(packed) { + return nil, fmt.Errorf("packed calls truncated reading numCalls (u8)") + } + numCalls = int(packed[p]) + p++ + } + } + + callData := make([]Span, numCalls) + for i := 0; i < numCalls; i++ { + if p+1 > len(packed) { + return nil, fmt.Errorf("packed calls truncated reading call flags (i=%d)", i) + } + flags := packed[p] + p++ + + if flags&0x01 == 0x00 { + if p+20 > len(packed) { + return nil, fmt.Errorf("packed calls truncated reading to (i=%d)", i) + } + p += 20 + } + + if flags&0x02 == 0x02 { + if p+32 > len(packed) { + return nil, fmt.Errorf("packed calls truncated reading value (i=%d)", i) + } + p += 32 + } + + if flags&0x04 == 0x04 { + if p+3 > len(packed) { + return nil, fmt.Errorf("packed calls truncated reading calldataSize (i=%d)", i) + } + calldataSize := int(packed[p])<<16 | int(packed[p+1])<<8 | int(packed[p+2]) + p += 3 + + if p+calldataSize > len(packed) { + return nil, fmt.Errorf("packed calls truncated reading calldata bytes (i=%d size=%d)", i, calldataSize) + } + callData[i] = Span{Start: p, Len: calldataSize} + p += calldataSize + } else { + callData[i] = Span{Start: -1, Len: 0} + } + + if flags&0x08 == 0x08 { + if p+32 > len(packed) { + return nil, fmt.Errorf("packed calls truncated reading gasLimit (i=%d)", i) + } + p += 32 + } + } + + return &PackedCallsLayout{ + GlobalFlag: globalFlag, + NumCalls: numCalls, + CallData: callData, + }, nil +} diff --git a/lib/sapient/malleable/locators_packedcalls_test.go b/lib/sapient/malleable/locators_packedcalls_test.go new file mode 100644 index 00000000..db819396 --- /dev/null +++ b/lib/sapient/malleable/locators_packedcalls_test.go @@ -0,0 +1,119 @@ +package malleable + +import ( + "math/big" + "testing" + + "github.com/0xsequence/ethkit/go-ethereum/common" + v3 "github.com/0xsequence/go-sequence/core/v3" + "github.com/stretchr/testify/require" +) + +func TestParsePackedCalls_SingleCallWithData(t *testing.T) { + encodeAddr := common.HexToAddress("0x1111111111111111111111111111111111111111") + callData := []byte{0x01, 0x02, 0x03, 0x04} + payload := v3.NewCallsPayload(encodeAddr, big.NewInt(1), []v3.Call{ + { + To: common.HexToAddress("0x2222222222222222222222222222222222222222"), + Value: big.NewInt(5), + Data: callData, + GasLimit: big.NewInt(7), + }, + }, big.NewInt(0), big.NewInt(0)) + + packed := payload.Encode(encodeAddr) + layout, err := ParsePackedCalls(packed) + require.NoError(t, err) + require.Equal(t, 1, layout.NumCalls) + require.Equal(t, 1, len(layout.CallData)) + + span := layout.CallData[0] + require.True(t, span.Start >= 0) + require.Equal(t, callData, packed[span.Start:span.Start+span.Len]) +} + +func TestParsePackedCalls_MultipleCalls_MixedData(t *testing.T) { + encodeAddr := common.HexToAddress("0x3333333333333333333333333333333333333333") + payload := v3.NewCallsPayload(encodeAddr, big.NewInt(1), []v3.Call{ + { + To: common.HexToAddress("0x4444444444444444444444444444444444444444"), + Data: []byte{0xaa, 0xbb}, + }, + { + To: common.HexToAddress("0x5555555555555555555555555555555555555555"), + Data: nil, + }, + }, big.NewInt(0), big.NewInt(1)) + + packed := payload.Encode(encodeAddr) + layout, err := ParsePackedCalls(packed) + require.NoError(t, err) + require.Equal(t, 2, layout.NumCalls) + + span0 := layout.CallData[0] + require.True(t, span0.Start >= 0) + require.Equal(t, []byte{0xaa, 0xbb}, packed[span0.Start:span0.Start+span0.Len]) + + span1 := layout.CallData[1] + require.Equal(t, -1, span1.Start) + require.Equal(t, 0, span1.Len) +} + +func TestParsePackedCalls_TruncatedPacked(t *testing.T) { + encodeAddr := common.HexToAddress("0x6666666666666666666666666666666666666666") + payload := v3.NewCallsPayload(encodeAddr, big.NewInt(1), []v3.Call{ + { + To: common.HexToAddress("0x7777777777777777777777777777777777777777"), + Data: []byte{0xaa, 0xbb, 0xcc}, + }, + }, big.NewInt(0), big.NewInt(0)) + + packed := payload.Encode(encodeAddr) + truncated := packed[:len(packed)-1] + _, err := ParsePackedCalls(truncated) + if err != nil { + t.Logf("parse error: %v", err) + } + require.Error(t, err) +} + +func TestParsePackedCalls_LargeCallCount(t *testing.T) { + encodeAddr := common.HexToAddress("0x8888888888888888888888888888888888888888") + const callCount = 300 + calls := make([]v3.Call, callCount) + for i := 0; i < callCount; i++ { + calls[i] = v3.Call{ + To: common.HexToAddress("0x9999999999999999999999999999999999999999"), + } + } + payload := v3.NewCallsPayload(encodeAddr, big.NewInt(1), calls, big.NewInt(0), big.NewInt(0)) + + packed := payload.Encode(encodeAddr) + layout, err := ParsePackedCalls(packed) + require.NoError(t, err) + require.Equal(t, callCount, layout.NumCalls) + require.Len(t, layout.CallData, callCount) + for _, span := range layout.CallData { + require.Equal(t, -1, span.Start) + require.Equal(t, 0, span.Len) + } +} + +func TestParsePackedCalls_WithNonceAndSpace(t *testing.T) { + encodeAddr := common.HexToAddress("0xaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaa") + callData := []byte{0x01} + payload := v3.NewCallsPayload(encodeAddr, big.NewInt(1), []v3.Call{ + { + To: common.HexToAddress("0xbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbb"), + Data: callData, + }, + }, big.NewInt(0x1234), big.NewInt(0x42)) + + packed := payload.Encode(encodeAddr) + layout, err := ParsePackedCalls(packed) + require.NoError(t, err) + require.Equal(t, 1, layout.NumCalls) + span := layout.CallData[0] + require.True(t, span.Start >= 0) + require.Equal(t, callData, packed[span.Start:span.Start+span.Len]) +} diff --git a/lib/sapient/malleable/path.go b/lib/sapient/malleable/path.go new file mode 100644 index 00000000..c4afeb3a --- /dev/null +++ b/lib/sapient/malleable/path.go @@ -0,0 +1,450 @@ +package malleable + +import ( + "fmt" + "strings" + + "github.com/0xsequence/ethkit/go-ethereum/accounts/abi" + v3 "github.com/0xsequence/go-sequence/core/v3" +) + +type Selector interface { + Resolve(payload *v3.CallsPayload) ([]ByteRange, error) + String() string +} + +type Path struct { + steps []pathStep + desc []string +} + +func NewPath() *Path { + return &Path{} +} + +func (p *Path) CallData(i int) *Path { + p.steps = append(p.steps, callDataStep{index: i}) + p.desc = append(p.desc, fmt.Sprintf("callData(%d)", i)) + return p +} + +func (p *Path) Slice(offset uint32, size uint32) *Path { + p.steps = append(p.steps, sliceStep{offset: int(offset), size: int(size)}) + p.desc = append(p.desc, fmt.Sprintf("slice(%d,%d)", offset, size)) + return p +} + +func (p *Path) ABI(contractABI *abi.ABI, method string) *Path { + p.steps = append(p.steps, abiStep{contractABI: contractABI, method: method}) + p.desc = append(p.desc, fmt.Sprintf("abi(%s)", method)) + return p +} + +func (p *Path) ArgSlot(argName string) *Path { + p.steps = append(p.steps, argSlotStep{name: argName}) + p.desc = append(p.desc, fmt.Sprintf("argSlot(%s)", argName)) + return p +} + +func (p *Path) ArgSlotIndex(argIndex int) *Path { + p.steps = append(p.steps, argSlotIndexStep{index: argIndex}) + p.desc = append(p.desc, fmt.Sprintf("argSlotIndex(%d)", argIndex)) + return p +} + +func (p *Path) ArgBytesData(argName string) *Path { + p.steps = append(p.steps, argBytesDataStep{name: argName}) + p.desc = append(p.desc, fmt.Sprintf("argBytesData(%s)", argName)) + return p +} + +func (p *Path) ArgBytesDataIndex(argIndex int) *Path { + p.steps = append(p.steps, argBytesDataIndexStep{index: argIndex}) + p.desc = append(p.desc, fmt.Sprintf("argBytesDataIndex(%d)", argIndex)) + return p +} + +func (p *Path) ArgBytesEncoded(argName string) *Path { + p.steps = append(p.steps, argBytesEncodedStep{name: argName}) + p.desc = append(p.desc, fmt.Sprintf("argBytesEncoded(%s)", argName)) + return p +} + +func (p *Path) EncodedCallsPayload() *Path { + p.steps = append(p.steps, encodedCallsPayloadStep{}) + p.desc = append(p.desc, "encodedCallsPayload()") + return p +} + +func (p *Path) EncodedCallData(i int) *Path { + p.steps = append(p.steps, encodedCallDataStep{index: i}) + p.desc = append(p.desc, fmt.Sprintf("encodedCallData(%d)", i)) + return p +} + +func (p *Path) AsSelector() Selector { + return p +} + +func (p *Path) Resolve(payload *v3.CallsPayload) ([]ByteRange, error) { + state := pathState{} + for _, step := range p.steps { + if err := step.apply(payload, &state); err != nil { + return nil, err + } + } + if len(state.ranges) == 0 { + return nil, fmt.Errorf("path resolved to empty ranges") + } + return state.ranges, nil +} + +func (p *Path) String() string { + if len(p.desc) == 0 { + return "path()" + } + return "path(" + strings.Join(p.desc, " -> ") + ")" +} + +type pathState struct { + ranges []ByteRange + method *abi.Method +} + +type pathStep interface { + apply(payload *v3.CallsPayload, state *pathState) error +} + +type callDataStep struct { + index int +} + +func (s callDataStep) apply(payload *v3.CallsPayload, state *pathState) error { + if payload == nil { + return fmt.Errorf("payload is nil") + } + if s.index < 0 || s.index >= len(payload.Calls) { + return fmt.Errorf("call index out of range: %d", s.index) + } + state.ranges = []ByteRange{{ + CallIndex: s.index, + Offset: 0, + Size: len(payload.Calls[s.index].Data), + }} + state.method = nil + return nil +} + +type sliceStep struct { + offset int + size int +} + +func (s sliceStep) apply(payload *v3.CallsPayload, state *pathState) error { + if len(state.ranges) == 0 { + return fmt.Errorf("slice step has no active ranges") + } + var out []ByteRange + for _, r := range state.ranges { + if s.offset < 0 || s.size < 0 || s.offset+s.size > r.Size { + return fmt.Errorf("slice out of bounds: [%d,%d) within %d", s.offset, s.offset+s.size, r.Size) + } + out = append(out, ByteRange{ + CallIndex: r.CallIndex, + Offset: r.Offset + s.offset, + Size: s.size, + }) + } + state.ranges = out + return nil +} + +type abiStep struct { + contractABI *abi.ABI + method string +} + +func (s abiStep) apply(payload *v3.CallsPayload, state *pathState) error { + if s.contractABI == nil { + return fmt.Errorf("abi step missing contract ABI") + } + method, ok := s.contractABI.Methods[s.method] + if !ok { + return fmt.Errorf("method not found: %s", s.method) + } + if len(state.ranges) == 0 { + return fmt.Errorf("abi step has no active ranges") + } + for _, r := range state.ranges { + if r.Size < 4 { + return fmt.Errorf("calldata too short for selector") + } + data, err := r.Slice(payload) + if err != nil { + return err + } + if len(data) < 4 { + return fmt.Errorf("calldata too short for selector") + } + if !bytesEqual(data[:4], method.ID) { + return fmt.Errorf("calldata selector mismatch for %s", s.method) + } + } + state.method = &method + return nil +} + +type argSlotStep struct { + name string +} + +func (s argSlotStep) apply(payload *v3.CallsPayload, state *pathState) error { + if state.method == nil { + return fmt.Errorf("argSlot step requires ABI context") + } + argIndex, err := argIndexByName(*state.method, s.name) + if err != nil { + return err + } + argType := state.method.Inputs[argIndex].Type + if isDynamicType(argType) { + return fmt.Errorf("arg %s is dynamic (%s)", s.name, argType.String()) + } + var out []ByteRange + for _, r := range state.ranges { + start := r.Offset + 4 + 32*argIndex + if start+32 > r.Offset+r.Size { + return fmt.Errorf("arg slot out of bounds for %s", s.name) + } + out = append(out, ByteRange{ + CallIndex: r.CallIndex, + Offset: start, + Size: 32, + }) + } + state.ranges = out + return nil +} + +type argSlotIndexStep struct { + index int +} + +func (s argSlotIndexStep) apply(payload *v3.CallsPayload, state *pathState) error { + if state.method == nil { + return fmt.Errorf("argSlotIndex step requires ABI context") + } + argIndex := s.index + if argIndex < 0 || argIndex >= len(state.method.Inputs) { + return fmt.Errorf("arg index out of range: %d", argIndex) + } + argType := state.method.Inputs[argIndex].Type + if isDynamicType(argType) { + return fmt.Errorf("arg %d is dynamic (%s)", argIndex, argType.String()) + } + var out []ByteRange + for _, r := range state.ranges { + start := r.Offset + 4 + 32*argIndex + if start+32 > r.Offset+r.Size { + return fmt.Errorf("arg slot out of bounds for %d", argIndex) + } + out = append(out, ByteRange{ + CallIndex: r.CallIndex, + Offset: start, + Size: 32, + }) + } + state.ranges = out + return nil +} + +type argBytesDataStep struct { + name string +} + +func (s argBytesDataStep) apply(payload *v3.CallsPayload, state *pathState) error { + if state.method == nil { + return fmt.Errorf("argBytesData step requires ABI context") + } + argIndex, err := argIndexByName(*state.method, s.name) + if err != nil { + return err + } + argType := state.method.Inputs[argIndex].Type + if argType.T != abi.BytesTy && argType.T != abi.StringTy { + return fmt.Errorf("arg %s is not bytes/string (%s)", s.name, argType.String()) + } + var out []ByteRange + for _, r := range state.ranges { + data, err := r.Slice(payload) + if err != nil { + return err + } + start, length, err := CalldataBytesContent(data, *state.method, argIndex) + if err != nil { + return err + } + out = append(out, ByteRange{ + CallIndex: r.CallIndex, + Offset: r.Offset + start, + Size: length, + }) + } + state.ranges = out + return nil +} + +type argBytesDataIndexStep struct { + index int +} + +func (s argBytesDataIndexStep) apply(payload *v3.CallsPayload, state *pathState) error { + if state.method == nil { + return fmt.Errorf("argBytesDataIndex step requires ABI context") + } + argIndex := s.index + if argIndex < 0 || argIndex >= len(state.method.Inputs) { + return fmt.Errorf("arg index out of range: %d", argIndex) + } + argType := state.method.Inputs[argIndex].Type + if argType.T != abi.BytesTy && argType.T != abi.StringTy { + return fmt.Errorf("arg %d is not bytes/string (%s)", argIndex, argType.String()) + } + var out []ByteRange + for _, r := range state.ranges { + data, err := r.Slice(payload) + if err != nil { + return err + } + start, length, err := CalldataBytesContent(data, *state.method, argIndex) + if err != nil { + return err + } + out = append(out, ByteRange{ + CallIndex: r.CallIndex, + Offset: r.Offset + start, + Size: length, + }) + } + state.ranges = out + return nil +} + +type argBytesEncodedStep struct { + name string +} + +func (s argBytesEncodedStep) apply(payload *v3.CallsPayload, state *pathState) error { + if state.method == nil { + return fmt.Errorf("argBytesEncoded step requires ABI context") + } + argIndex, err := argIndexByName(*state.method, s.name) + if err != nil { + return err + } + argType := state.method.Inputs[argIndex].Type + if argType.T != abi.BytesTy && argType.T != abi.StringTy { + return fmt.Errorf("arg %s is not bytes/string (%s)", s.name, argType.String()) + } + var out []ByteRange + for _, r := range state.ranges { + data, err := r.Slice(payload) + if err != nil { + return err + } + start, length, err := CalldataBytesEncoded(data, *state.method, argIndex) + if err != nil { + return err + } + out = append(out, ByteRange{ + CallIndex: r.CallIndex, + Offset: r.Offset + start, + Size: length, + }) + } + state.ranges = out + return nil +} + +type encodedCallsPayloadStep struct{} + +func (s encodedCallsPayloadStep) apply(payload *v3.CallsPayload, state *pathState) error { + if len(state.ranges) == 0 { + return fmt.Errorf("encodedCallsPayload step has no active ranges") + } + return nil +} + +type encodedCallDataStep struct { + index int +} + +func (s encodedCallDataStep) apply(payload *v3.CallsPayload, state *pathState) error { + if len(state.ranges) == 0 { + return fmt.Errorf("encodedCallData step has no active ranges") + } + var out []ByteRange + for _, r := range state.ranges { + data, err := r.Slice(payload) + if err != nil { + return err + } + layout, err := ParsePackedCalls(data) + if err != nil { + return err + } + if s.index < 0 || s.index >= layout.NumCalls { + return fmt.Errorf("packed call index out of range: %d", s.index) + } + span := layout.CallData[s.index] + if span.Start < 0 || span.Len == 0 { + return fmt.Errorf("packed call %d has no calldata", s.index) + } + out = append(out, ByteRange{ + CallIndex: r.CallIndex, + Offset: r.Offset + span.Start, + Size: span.Len, + }) + } + state.ranges = out + return nil +} + +func argIndexByName(method abi.Method, name string) (int, error) { + for i, input := range method.Inputs { + if input.Name == name { + return i, nil + } + } + return -1, fmt.Errorf("arg not found: %s", name) +} + +func isDynamicType(t abi.Type) bool { + switch t.T { + case abi.BytesTy, abi.StringTy, abi.SliceTy: + return true + case abi.ArrayTy: + return t.Size == 0 || isDynamicType(*t.Elem) + case abi.TupleTy: + for _, elem := range t.TupleElems { + if isDynamicType(*elem) { + return true + } + } + return false + default: + return false + } +} + +func bytesEqual(a, b []byte) bool { + if len(a) != len(b) { + return false + } + for i := range a { + if a[i] != b[i] { + return false + } + } + return true +} diff --git a/lib/sapient/malleable/path_test.go b/lib/sapient/malleable/path_test.go new file mode 100644 index 00000000..19f849e3 --- /dev/null +++ b/lib/sapient/malleable/path_test.go @@ -0,0 +1,265 @@ +package malleable + +import ( + "math/big" + "strings" + "testing" + + "github.com/0xsequence/ethkit/go-ethereum/accounts/abi" + "github.com/0xsequence/ethkit/go-ethereum/common" + v3 "github.com/0xsequence/go-sequence/core/v3" + "github.com/stretchr/testify/require" +) + +func TestPath_ResolveNestedPackedCalls(t *testing.T) { + trailsABIJSON := `[{"name":"hydrateExecute","type":"function","inputs":[{"name":"payload","type":"bytes"},{"name":"hydrateData","type":"bytes"}]}]` + erc2612ABIJSON := `[{"name":"permit","type":"function","inputs":[{"name":"owner","type":"address"},{"name":"spender","type":"address"},{"name":"value","type":"uint256"},{"name":"deadline","type":"uint256"},{"name":"v","type":"uint8"},{"name":"r","type":"bytes32"},{"name":"s","type":"bytes32"}]}]` + erc20ABIJSON := `[{"name":"transferFrom","type":"function","inputs":[{"name":"from","type":"address"},{"name":"to","type":"address"},{"name":"value","type":"uint256"}]}]` + + trailsABI, err := abi.JSON(strings.NewReader(trailsABIJSON)) + require.NoError(t, err) + erc2612ABI, err := abi.JSON(strings.NewReader(erc2612ABIJSON)) + require.NoError(t, err) + erc20ABI, err := abi.JSON(strings.NewReader(erc20ABIJSON)) + require.NoError(t, err) + + permitCalldata, err := erc2612ABI.Pack( + "permit", + common.HexToAddress("0x1111111111111111111111111111111111111111"), + common.HexToAddress("0x2222222222222222222222222222222222222222"), + big.NewInt(7), + big.NewInt(8), + uint8(27), + common.HexToHash("0x01"), + common.HexToHash("0x02"), + ) + require.NoError(t, err) + + transferCalldata, err := erc20ABI.Pack( + "transferFrom", + common.HexToAddress("0x3333333333333333333333333333333333333333"), + common.HexToAddress("0x4444444444444444444444444444444444444444"), + big.NewInt(7), + ) + require.NoError(t, err) + + encodeAddr := common.HexToAddress("0xaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaa") + innerPayload := v3.NewCallsPayload(encodeAddr, big.NewInt(1), []v3.Call{ + {To: common.HexToAddress("0xbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbb"), Data: permitCalldata}, + {To: common.HexToAddress("0xcccccccccccccccccccccccccccccccccccccccc"), Data: transferCalldata}, + }, big.NewInt(0), big.NewInt(0)) + + packed := innerPayload.Encode(encodeAddr) + hydrateData := []byte{0x55} + outerCallData, err := trailsABI.Pack("hydrateExecute", packed, hydrateData) + require.NoError(t, err) + + payload := v3.NewCallsPayload(common.Address{}, big.NewInt(1), []v3.Call{ + {To: common.HexToAddress("0xdddddddddddddddddddddddddddddddddddddddd"), Data: outerCallData}, + }, big.NewInt(0), big.NewInt(0)) + + permitValueSel := NewPath(). + CallData(0). + ABI(&trailsABI, "hydrateExecute"). + ArgBytesData("payload"). + EncodedCallsPayload(). + EncodedCallData(0). + ABI(&erc2612ABI, "permit"). + ArgSlot("value"). + AsSelector() + t.Logf("permitValueSel: %s", permitValueSel.String()) + + transferValueSel := NewPath(). + CallData(0). + ABI(&trailsABI, "hydrateExecute"). + ArgBytesData("payload"). + EncodedCallsPayload(). + EncodedCallData(1). + ABI(&erc20ABI, "transferFrom"). + ArgSlot("value"). + AsSelector() + t.Logf("transferValueSel: %s", transferValueSel.String()) + + permitRanges, err := permitValueSel.Resolve(&payload) + require.NoError(t, err) + require.Len(t, permitRanges, 1) + + transferRanges, err := transferValueSel.Resolve(&payload) + require.NoError(t, err) + require.Len(t, transferRanges, 1) + + permitSlice, err := permitRanges[0].Slice(&payload) + require.NoError(t, err) + transferSlice, err := transferRanges[0].Slice(&payload) + require.NoError(t, err) + + require.Equal(t, permitCalldata[4+32*2:4+32*3], permitSlice) + require.Equal(t, transferCalldata[4+32*2:4+32*3], transferSlice) +} + +func TestPath_ResolveDirectTransferFromValue(t *testing.T) { + erc20ABIJSON := `[{"name":"transferFrom","type":"function","inputs":[{"name":"_from","type":"address"},{"name":"_to","type":"address"},{"name":"_value","type":"uint256"}]}]` + erc20ABI, err := abi.JSON(strings.NewReader(erc20ABIJSON)) + require.NoError(t, err) + + transferCalldata, err := erc20ABI.Pack( + "transferFrom", + common.HexToAddress("0x3333333333333333333333333333333333333333"), + common.HexToAddress("0x4444444444444444444444444444444444444444"), + big.NewInt(7), + ) + require.NoError(t, err) + + payload := v3.NewCallsPayload(common.Address{}, big.NewInt(1), []v3.Call{ + {To: common.HexToAddress("0x5555555555555555555555555555555555555555"), Data: transferCalldata}, + }, big.NewInt(0), big.NewInt(0)) + + valueSel := NewPath(). + CallData(0). + ABI(&erc20ABI, "transferFrom"). + ArgSlot("_value"). + AsSelector() + t.Logf("valueSel: %s", valueSel.String()) + + ranges, err := valueSel.Resolve(&payload) + require.NoError(t, err) + require.Len(t, ranges, 1) + + valueSlice, err := ranges[0].Slice(&payload) + require.NoError(t, err) + require.Equal(t, transferCalldata[4+32*2:4+32*3], valueSlice) +} + +func TestPath_ResolveDirectTransferFromValueByIndex(t *testing.T) { + erc20ABIJSON := `[{"name":"transferFrom","type":"function","inputs":[{"name":"","type":"address"},{"name":"","type":"address"},{"name":"","type":"uint256"}]}]` + erc20ABI, err := abi.JSON(strings.NewReader(erc20ABIJSON)) + require.NoError(t, err) + + transferCalldata, err := erc20ABI.Pack( + "transferFrom", + common.HexToAddress("0x3333333333333333333333333333333333333333"), + common.HexToAddress("0x4444444444444444444444444444444444444444"), + big.NewInt(7), + ) + require.NoError(t, err) + + payload := v3.NewCallsPayload(common.Address{}, big.NewInt(1), []v3.Call{ + {To: common.HexToAddress("0x5555555555555555555555555555555555555555"), Data: transferCalldata}, + }, big.NewInt(0), big.NewInt(0)) + + valueSel := NewPath(). + CallData(0). + ABI(&erc20ABI, "transferFrom"). + ArgSlotIndex(2). + AsSelector() + t.Logf("valueSelByIndex: %s", valueSel.String()) + + ranges, err := valueSel.Resolve(&payload) + require.NoError(t, err) + require.Len(t, ranges, 1) + + valueSlice, err := ranges[0].Slice(&payload) + require.NoError(t, err) + require.Equal(t, transferCalldata[4+32*2:4+32*3], valueSlice) +} + +func TestPath_ResolveFailsWithSelectorMismatch(t *testing.T) { + erc20ABIJSON := `[{"name":"transferFrom","type":"function","inputs":[{"name":"_from","type":"address"},{"name":"_to","type":"address"},{"name":"_value","type":"uint256"}]}]` + erc20ABI, err := abi.JSON(strings.NewReader(erc20ABIJSON)) + require.NoError(t, err) + + approveABIJSON := `[{"name":"approve","type":"function","inputs":[{"name":"_spender","type":"address"},{"name":"_value","type":"uint256"}]}]` + approveABI, err := abi.JSON(strings.NewReader(approveABIJSON)) + require.NoError(t, err) + + approveCalldata, err := approveABI.Pack( + "approve", + common.HexToAddress("0x3333333333333333333333333333333333333333"), + big.NewInt(7), + ) + require.NoError(t, err) + + payload := v3.NewCallsPayload(common.Address{}, big.NewInt(1), []v3.Call{ + {To: common.HexToAddress("0x5555555555555555555555555555555555555555"), Data: approveCalldata}, + }, big.NewInt(0), big.NewInt(0)) + + valueSel := NewPath(). + CallData(0). + ABI(&erc20ABI, "transferFrom"). + ArgSlot("_value"). + AsSelector() + t.Logf("mismatchedSelector: %s", valueSel.String()) + + _, err = valueSel.Resolve(&payload) + if err != nil { + t.Logf("resolve error: %v", err) + } + require.Error(t, err) +} + +func TestPath_ResolveFailsWithOutOfRangeCallIndex(t *testing.T) { + erc20ABIJSON := `[{"name":"transferFrom","type":"function","inputs":[{"name":"_from","type":"address"},{"name":"_to","type":"address"},{"name":"_value","type":"uint256"}]}]` + erc20ABI, err := abi.JSON(strings.NewReader(erc20ABIJSON)) + require.NoError(t, err) + + transferCalldata, err := erc20ABI.Pack( + "transferFrom", + common.HexToAddress("0x3333333333333333333333333333333333333333"), + common.HexToAddress("0x4444444444444444444444444444444444444444"), + big.NewInt(7), + ) + require.NoError(t, err) + + payload := v3.NewCallsPayload(common.Address{}, big.NewInt(1), []v3.Call{ + {To: common.HexToAddress("0x5555555555555555555555555555555555555555"), Data: transferCalldata}, + }, big.NewInt(0), big.NewInt(0)) + + valueSel := NewPath(). + CallData(1). + ABI(&erc20ABI, "transferFrom"). + ArgSlot("_value"). + AsSelector() + t.Logf("outOfRangeCallIndex: %s", valueSel.String()) + + _, err = valueSel.Resolve(&payload) + if err != nil { + t.Logf("resolve error: %v", err) + } + require.Error(t, err) +} + +func TestPath_ResolveFailsWithInvalidEncodedCallsPayload(t *testing.T) { + trailsABIJSON := `[{"name":"hydrateExecute","type":"function","inputs":[{"name":"payload","type":"bytes"},{"name":"hydrateData","type":"bytes"}]}]` + trailsABI, err := abi.JSON(strings.NewReader(trailsABIJSON)) + require.NoError(t, err) + + erc20ABIJSON := `[{"name":"transferFrom","type":"function","inputs":[{"name":"_from","type":"address"},{"name":"_to","type":"address"},{"name":"_value","type":"uint256"}]}]` + erc20ABI, err := abi.JSON(strings.NewReader(erc20ABIJSON)) + require.NoError(t, err) + + invalidPacked := []byte{} + hydrateData := []byte{0x55} + outerCallData, err := trailsABI.Pack("hydrateExecute", invalidPacked, hydrateData) + require.NoError(t, err) + + payload := v3.NewCallsPayload(common.Address{}, big.NewInt(1), []v3.Call{ + {To: common.HexToAddress("0xdddddddddddddddddddddddddddddddddddddddd"), Data: outerCallData}, + }, big.NewInt(0), big.NewInt(0)) + + valueSel := NewPath(). + CallData(0). + ABI(&trailsABI, "hydrateExecute"). + ArgBytesData("payload"). + EncodedCallsPayload(). + EncodedCallData(0). + ABI(&erc20ABI, "transferFrom"). + ArgSlot("_value"). + AsSelector() + t.Logf("invalidEncodedCallsPayload: %s", valueSel.String()) + + _, err = valueSel.Resolve(&payload) + if err != nil { + t.Logf("resolve error: %v", err) + } + require.Error(t, err) +} diff --git a/lib/sapient/malleable/signature.go b/lib/sapient/malleable/signature.go new file mode 100644 index 00000000..7a6f6da9 --- /dev/null +++ b/lib/sapient/malleable/signature.go @@ -0,0 +1,102 @@ +package malleable + +import ( + "bytes" + "encoding/binary" + "fmt" +) + +type SectionKind uint8 + +const ( + SectionStatic SectionKind = iota + SectionRepeat +) + +type Section interface { + Kind() SectionKind +} + +type StaticSection struct { + TIndex uint8 + CIndex uint16 + Size uint16 +} + +func (s StaticSection) Kind() SectionKind { return SectionStatic } + +type RepeatSection struct { + TIndex uint8 + CIndex uint16 + Size uint16 + TIndex2 uint8 + CIndex2 uint16 +} + +func (s RepeatSection) Kind() SectionKind { return SectionRepeat } + +func EncodeSignature(statics []StaticSection, repeats []RepeatSection) ([]byte, error) { + var out bytes.Buffer + + for _, s := range statics { + if s.TIndex > 0x7F { + return nil, fmt.Errorf("tindex out of range: %d", s.TIndex) + } + out.WriteByte(byte(s.TIndex & 0x7F)) + _ = binary.Write(&out, binary.BigEndian, s.CIndex) + _ = binary.Write(&out, binary.BigEndian, s.Size) + } + + for _, r := range repeats { + if r.TIndex > 0x7F { + return nil, fmt.Errorf("tindex out of range: %d", r.TIndex) + } + out.WriteByte(byte(r.TIndex&0x7F) | 0x80) + _ = binary.Write(&out, binary.BigEndian, r.CIndex) + _ = binary.Write(&out, binary.BigEndian, r.Size) + out.WriteByte(r.TIndex2) + _ = binary.Write(&out, binary.BigEndian, r.CIndex2) + } + + return out.Bytes(), nil +} + +func DecodeSignature(sig []byte) ([]Section, error) { + var sections []Section + i := 0 + for i < len(sig) { + if i+5 > len(sig) { + return nil, fmt.Errorf("signature truncated at %d", i) + } + tRaw := sig[i] + i++ + cindex := binary.BigEndian.Uint16(sig[i : i+2]) + i += 2 + size := binary.BigEndian.Uint16(sig[i : i+2]) + i += 2 + + if tRaw&0x80 != 0 { + if i+3 > len(sig) { + return nil, fmt.Errorf("repeat truncated at %d", i) + } + t2 := sig[i] + i++ + c2 := binary.BigEndian.Uint16(sig[i : i+2]) + i += 2 + sections = append(sections, RepeatSection{ + TIndex: tRaw & 0x7F, + CIndex: cindex, + Size: size, + TIndex2: t2, + CIndex2: c2, + }) + } else { + sections = append(sections, StaticSection{ + TIndex: tRaw & 0x7F, + CIndex: cindex, + Size: size, + }) + } + } + return sections, nil +} diff --git a/lib/sapient/malleable/spans.go b/lib/sapient/malleable/spans.go new file mode 100644 index 00000000..2446b529 --- /dev/null +++ b/lib/sapient/malleable/spans.go @@ -0,0 +1,57 @@ +package malleable + +import ( + "fmt" + + v3 "github.com/0xsequence/go-sequence/core/v3" +) + +type Span struct { + Start int + Len int +} + +type ByteRange struct { + CallIndex int + Offset int + Size int +} + +func (r ByteRange) Slice(payload *v3.CallsPayload) ([]byte, error) { + if payload == nil { + return nil, fmt.Errorf("payload is nil") + } + if r.CallIndex < 0 || r.CallIndex >= len(payload.Calls) { + return nil, fmt.Errorf("call index out of range: %d", r.CallIndex) + } + data := payload.Calls[r.CallIndex].Data + if r.Offset < 0 || r.Size < 0 || r.Offset+r.Size > len(data) { + return nil, fmt.Errorf("range out of bounds: [%d,%d) with len %d", r.Offset, r.Offset+r.Size, len(data)) + } + return data[r.Offset : r.Offset+r.Size], nil +} + +type RangeSelector struct { + Range ByteRange +} + +func NewRangeSelector(callIndex, offset, size int) RangeSelector { + return RangeSelector{ + Range: ByteRange{ + CallIndex: callIndex, + Offset: offset, + Size: size, + }, + } +} + +func (r RangeSelector) Resolve(payload *v3.CallsPayload) ([]ByteRange, error) { + if _, err := r.Range.Slice(payload); err != nil { + return nil, err + } + return []ByteRange{r.Range}, nil +} + +func (r RangeSelector) String() string { + return fmt.Sprintf("range(call=%d,offset=%d,size=%d)", r.Range.CallIndex, r.Range.Offset, r.Range.Size) +}