|
22 | 22 | from cassandra.cluster import Cluster |
23 | 23 | from cassandra.connection import (Connection, HEADER_DIRECTION_TO_CLIENT, ProtocolError, |
24 | 24 | locally_supported_compressions, ConnectionHeartbeat, _Frame, Timer, TimerManager, |
25 | | - ConnectionException, ConnectionShutdown, DefaultEndPoint, ShardAwarePortGenerator) |
| 25 | + ConnectionException, ConnectionShutdown, DefaultEndPoint, ShardAwarePortGenerator, |
| 26 | + _ConnectionIOBuffer) |
26 | 27 | from cassandra.marshal import uint8_pack, uint32_pack, int32_pack |
27 | 28 | from cassandra.protocol import (write_stringmultimap, write_int, write_string, |
28 | 29 | SupportedMessage, ProtocolHandler) |
@@ -571,3 +572,80 @@ def test_generate_is_repeatable_with_same_mock(self, mock_randrange): |
571 | 572 | second_run = list(itertools.islice(gen.generate(0, 2), 5)) |
572 | 573 |
|
573 | 574 | assert first_run == second_run |
| 575 | + |
| 576 | + |
| 577 | +class TestConnectionIOBufferReset(unittest.TestCase): |
| 578 | + """Verify _reset_buffer and reset_cql_frame_buffer position semantics.""" |
| 579 | + |
| 580 | + def test_reset_buffer_discards_consumed_data(self): |
| 581 | + buf = BytesIO(b'\x01\x02\x03\x04\x05') |
| 582 | + buf.seek(0) |
| 583 | + # Consume first 3 bytes |
| 584 | + assert buf.read(3) == b'\x01\x02\x03' |
| 585 | + new_buf = _ConnectionIOBuffer._reset_buffer(buf) |
| 586 | + # New buffer should contain only unconsumed data |
| 587 | + new_buf.seek(0) |
| 588 | + assert new_buf.read() == b'\x04\x05' |
| 589 | + |
| 590 | + def test_reset_buffer_position_at_end(self): |
| 591 | + buf = BytesIO(b'\x01\x02\x03') |
| 592 | + buf.seek(0) |
| 593 | + buf.read(1) |
| 594 | + new_buf = _ConnectionIOBuffer._reset_buffer(buf) |
| 595 | + # Position should be at end (ready for appending) |
| 596 | + assert new_buf.tell() == 2 |
| 597 | + |
| 598 | + def test_reset_buffer_fully_consumed(self): |
| 599 | + buf = BytesIO(b'\x01\x02') |
| 600 | + buf.seek(0) |
| 601 | + buf.read(2) |
| 602 | + new_buf = _ConnectionIOBuffer._reset_buffer(buf) |
| 603 | + new_buf.seek(0) |
| 604 | + assert new_buf.read() == b'' |
| 605 | + |
| 606 | + def test_reset_buffer_nothing_consumed(self): |
| 607 | + buf = BytesIO(b'\x01\x02\x03') |
| 608 | + buf.seek(0) |
| 609 | + new_buf = _ConnectionIOBuffer._reset_buffer(buf) |
| 610 | + new_buf.seek(0) |
| 611 | + assert new_buf.read() == b'\x01\x02\x03' |
| 612 | + |
| 613 | + @staticmethod |
| 614 | + def _make_iobuf(checksumming=False): |
| 615 | + conn = Mock() |
| 616 | + conn._is_checksumming_enabled = checksumming |
| 617 | + iobuf = _ConnectionIOBuffer(conn) |
| 618 | + # Keep a strong reference so the weakref.proxy inside iobuf stays valid |
| 619 | + iobuf._conn_ref = conn |
| 620 | + if checksumming: |
| 621 | + iobuf.set_checksumming_buffer() |
| 622 | + return iobuf |
| 623 | + |
| 624 | + def test_reset_cql_frame_buffer_checksumming_uses_tell_position(self): |
| 625 | + """ |
| 626 | + When checksumming is enabled, reset_cql_frame_buffer delegates to |
| 627 | + _reset_buffer which relies on tell() to determine consumed data. |
| 628 | + Verify that seeking to an arbitrary position before reset correctly |
| 629 | + preserves only the unconsumed tail. |
| 630 | + """ |
| 631 | + iobuf = self._make_iobuf(checksumming=True) |
| 632 | + # Write some data into the cql_frame_buffer |
| 633 | + iobuf.cql_frame_buffer.write(b'\xAA\xBB\xCC\xDD\xEE') |
| 634 | + # Seek to position 3, simulating that first 3 bytes were consumed |
| 635 | + iobuf.cql_frame_buffer.seek(3) |
| 636 | + iobuf.reset_cql_frame_buffer() |
| 637 | + # After reset, only the unconsumed tail should remain |
| 638 | + iobuf.cql_frame_buffer.seek(0) |
| 639 | + assert iobuf.cql_frame_buffer.read() == b'\xDD\xEE' |
| 640 | + |
| 641 | + def test_reset_cql_frame_buffer_no_checksumming_resets_io_buffer(self): |
| 642 | + """ |
| 643 | + Without checksumming, reset_cql_frame_buffer delegates to |
| 644 | + reset_io_buffer (since cql_frame_buffer IS the io_buffer). |
| 645 | + """ |
| 646 | + iobuf = self._make_iobuf(checksumming=False) |
| 647 | + iobuf.io_buffer.write(b'\x01\x02\x03\x04') |
| 648 | + iobuf.io_buffer.seek(2) |
| 649 | + iobuf.reset_cql_frame_buffer() |
| 650 | + iobuf.io_buffer.seek(0) |
| 651 | + assert iobuf.io_buffer.read() == b'\x03\x04' |
0 commit comments