diff --git a/evm/vm/code_stream.py b/evm/vm/code_stream.py index f98eaffc12..a6703da38f 100644 --- a/evm/vm/code_stream.py +++ b/evm/vm/code_stream.py @@ -10,13 +10,15 @@ class CodeStream(object): stream = None + depth_processed = None logger = logging.getLogger('evm.vm.CodeStream') def __init__(self, code_bytes): validate_is_bytes(code_bytes) self.stream = io.BytesIO(code_bytes) - self._validity_cache = {} + self.invalid_positions = set() + self.depth_processed = 0 def read(self, size): return self.stream.read(size) @@ -30,6 +32,9 @@ def __iter__(self): def __next__(self): return self.next() + def __getitem__(self, i): + return self.stream.getvalue()[i] + def next(self): next_opcode_as_byte = self.read(1) @@ -63,31 +68,30 @@ def seek(self, pc): finally: self.pc = anchor_pc - _validity_cache = None + invalid_positions = None def is_valid_opcode(self, position): if position >= len(self): return False - - if position not in self._validity_cache: - with self.seek(max(0, position - 32)): - prefix = self.read(min(position, 32)) - - for offset, opcode in enumerate(reversed(prefix)): - if opcode < opcode_values.PUSH1 or opcode > opcode_values.PUSH32: - continue - - push_size = 1 + opcode - opcode_values.PUSH1 - if push_size <= offset: - continue - - opcode_position = position - 1 - offset - if not self.is_valid_opcode(opcode_position): - continue - - self._validity_cache[position] = False - break + if position in self.invalid_positions: + return False + if position <= self.depth_processed: + return True + else: + i = self.depth_processed + while i <= position: + opcode = self.__getitem__(i) + if opcode >= opcode_values.PUSH1 and opcode <= opcode_values.PUSH32: + left_bound = (i + 1) + right_bound = left_bound + (opcode - 95) + invalid_range = range(left_bound, right_bound) + self.invalid_positions.update(invalid_range) + i = right_bound + else: + self.depth_processed = i + i += 1 + + if position in self.invalid_positions: + return False else: - self._validity_cache[position] = True - - return self._validity_cache[position] + return True diff --git a/tests/core/code-stream/test_code_stream.py b/tests/core/code-stream/test_code_stream.py index efa5889bee..73ee9ba03d 100644 --- a/tests/core/code-stream/test_code_stream.py +++ b/tests/core/code-stream/test_code_stream.py @@ -55,8 +55,65 @@ def test_seek_reverts_to_original_stream_position_when_context_exits(): assert code_stream.peek() == opcode_values.ADD +def test_get_item_returns_correct_opcode(): + code_stream = CodeStream(b'\x01\x02\x30') + assert code_stream.__getitem__(0) == opcode_values.ADD + assert code_stream.__getitem__(1) == opcode_values.MUL + assert code_stream.__getitem__(2) == opcode_values.ADDRESS + + def test_is_valid_opcode_invalidates_bytes_after_PUSHXX_opcodes(): - code_stream = CodeStream(b'\x01\x60\x02') + code_stream = CodeStream(b'\x02\x60\x02\x04') assert code_stream.is_valid_opcode(0) is True assert code_stream.is_valid_opcode(1) is True assert code_stream.is_valid_opcode(2) is False + assert code_stream.is_valid_opcode(3) is True + assert code_stream.is_valid_opcode(4) is False + + +def test_harder_is_valid_opcode(): + code_stream = CodeStream(b'\x02\x03\x72' + (b'\x04' * 32) + b'\x05') + # valid: 0 - 2 :: 22 - 35 + # invalid: 3-21 (PUSH19) :: 36+ (too long) + assert code_stream.is_valid_opcode(0) is True + assert code_stream.is_valid_opcode(1) is True + assert code_stream.is_valid_opcode(2) is True + assert code_stream.is_valid_opcode(3) is False + assert code_stream.is_valid_opcode(21) is False + assert code_stream.is_valid_opcode(22) is True + assert code_stream.is_valid_opcode(35) is True + assert code_stream.is_valid_opcode(36) is False + + +def test_even_harder_is_valid_opcode(): + test = b'\x02\x03\x7d' + (b'\x04' * 32) + b'\x05\x7e' + (b'\x04' * 35) + b'\x01\x61\x01\x01\x01' + code_stream = CodeStream(test) + # valid: 0 - 2 :: 33 - 36 :: 68 - 73 :: 76 + # invalid: 3 - 32 (PUSH30) :: 37 - 67 (PUSH31) :: 74, 75 (PUSH2) :: 77+ (too long) + assert code_stream.is_valid_opcode(0) is True + assert code_stream.is_valid_opcode(1) is True + assert code_stream.is_valid_opcode(2) is True + assert code_stream.is_valid_opcode(3) is False + assert code_stream.is_valid_opcode(32) is False + assert code_stream.is_valid_opcode(33) is True + assert code_stream.is_valid_opcode(36) is True + assert code_stream.is_valid_opcode(37) is False + assert code_stream.is_valid_opcode(67) is False + assert code_stream.is_valid_opcode(68) is True + assert code_stream.is_valid_opcode(71) is True + assert code_stream.is_valid_opcode(72) is True + assert code_stream.is_valid_opcode(73) is True + assert code_stream.is_valid_opcode(74) is False + assert code_stream.is_valid_opcode(75) is False + assert code_stream.is_valid_opcode(76) is True + assert code_stream.is_valid_opcode(77) is False + + +def test_right_number_of_bytes_invalidated_after_pushxx(): + code_stream = CodeStream(b'\x02\x03\x60\x02\x02') + assert code_stream.is_valid_opcode(0) is True + assert code_stream.is_valid_opcode(1) is True + assert code_stream.is_valid_opcode(2) is True + assert code_stream.is_valid_opcode(3) is False + assert code_stream.is_valid_opcode(4) is True + assert code_stream.is_valid_opcode(5) is False