Harden cheap block parser bounds (#21021)

This commit is contained in:
Zachary Brown
2026-08-05 08:14:44 -07:00
committed by GitHub
parent f9a8b77c27
commit 5152b699ba
2 changed files with 69 additions and 1 deletions
+57
View File
@@ -34,6 +34,8 @@ from chia.full_node.full_block_utils import (
generator_from_block,
get_height_and_tx_status_from_block,
header_block_from_block,
skip_bytes,
skip_list,
skip_reward_chain_block,
)
from chia.types.blockchain_format.serialized_program import SerializedProgram
@@ -138,6 +140,26 @@ def test_skip_reward_chain_block_handles_combined_optional_tag(has_icc: bool, ha
assert len(skip_reward_chain_block(memoryview(bytes(reward_chain_block)))) == 0
def test_skip_list_rejects_count_exceeding_remaining_buffer() -> None:
with pytest.raises(ValueError, match="list count 2 exceeds remaining buffer 1"):
skip_list(memoryview((2).to_bytes(4, "big") + b"\x00"), skip_bytes)
def test_skip_list_rejects_short_count_prefix() -> None:
with pytest.raises(ValueError, match="list count prefix requires 4 bytes, remaining buffer 1"):
skip_list(memoryview(b"\x00"), skip_bytes)
def test_skip_bytes_rejects_length_exceeding_remaining_buffer() -> None:
with pytest.raises(ValueError, match="byte length 4 exceeds remaining buffer 3"):
skip_bytes(memoryview((4).to_bytes(4, "big") + b"abc"))
def test_skip_bytes_rejects_short_length_prefix() -> None:
with pytest.raises(ValueError, match="byte length prefix requires 4 bytes, remaining buffer 1"):
skip_bytes(memoryview(b"\x00"))
def get_foliage_block_data() -> Iterator[FoliageBlockData]:
for pool_signature in [g2(), None]:
pool_target = PoolTarget(
@@ -305,6 +327,41 @@ def get_full_blocks(shard: int) -> Iterator[FullBlock]:
)
def test_block_info_from_block_rejects_refs_count_exceeding_remaining_buffer() -> None:
block_bytes = bytearray(bytes(next(get_full_blocks(0))))
block_bytes[-4:] = (1).to_bytes(4, "big")
with pytest.raises(ValueError, match="refs count 1 exceeds remaining buffer 0"):
block_info_from_block(memoryview(block_bytes))
def test_block_info_from_block_rejects_short_refs_count_prefix() -> None:
block_bytes = bytearray(bytes(next(get_full_blocks(0))))
del block_bytes[-3:]
with pytest.raises(ValueError, match="refs count prefix requires 4 bytes, remaining buffer 1"):
block_info_from_block(memoryview(block_bytes))
def test_cheap_parser_matches_round_tripped_block() -> None:
block = next(get_full_blocks(0))
block_bytes = memoryview(bytes(block))
round_tripped = FullBlock.from_bytes(block_bytes)
height, is_tx_block = get_height_and_tx_status_from_block(block_bytes)
assert height == round_tripped.height
assert is_tx_block == round_tripped.is_transaction_block()
assert generator_from_block(block_bytes) is None
block_info = block_info_from_block(block_bytes)
assert block_info.prev_header_hash == round_tripped.prev_header_hash
assert block_info.transactions_generator == round_tripped.transactions_generator
assert block_info.transactions_generator_ref_list == round_tripped.transactions_generator_ref_list
header_block = HeaderBlock.from_bytes(header_block_from_block(block_bytes))
assert header_block == get_block_header(round_tripped, None)
@pytest.mark.anyio
@pytest.mark.parametrize("shard", [0, 1, 2, 3])
@pytest.mark.skipif(_is_macos_intel(), reason="Very slow on macOS Intel")
+12 -1
View File
@@ -13,17 +13,24 @@ from chia.types.blockchain_format.serialized_program import SerializedProgram
def skip_list(buf: memoryview, skip_item: Callable[[memoryview], memoryview]) -> memoryview:
if len(buf) < 4:
raise ValueError(f"list count prefix requires 4 bytes, remaining buffer {len(buf)}")
n = int.from_bytes(buf[:4], "big", signed=False)
buf = buf[4:]
if n > len(buf):
raise ValueError(f"list count {n} exceeds remaining buffer {len(buf)}")
for _ in range(n):
buf = skip_item(buf)
return buf
def skip_bytes(buf: memoryview) -> memoryview:
if len(buf) < 4:
raise ValueError(f"byte length prefix requires 4 bytes, remaining buffer {len(buf)}")
n = int.from_bytes(buf[:4], "big", signed=False)
buf = buf[4:]
assert n >= 0
if n > len(buf):
raise ValueError(f"byte length {n} exceeds remaining buffer {len(buf)}")
return buf[n:]
@@ -279,8 +286,12 @@ def block_info_from_block(buf: memoryview) -> GeneratorBlockInfo:
else:
buf = buf[1:]
if len(buf) < 4:
raise ValueError(f"refs count prefix requires 4 bytes, remaining buffer {len(buf)}")
refs_length = uint32.from_bytes(buf[:4])
buf = buf[4:]
if refs_length * 4 > len(buf):
raise ValueError(f"refs count {refs_length} exceeds remaining buffer {len(buf)}")
refs = []
for i in range(refs_length):