|
7 | 7 | import math |
8 | 8 | import sys |
9 | 9 | from enum import auto, Enum, IntFlag |
10 | | -from io import BytesIO, IOBase, UnsupportedOperation, SEEK_CUR, SEEK_END, SEEK_SET |
| 10 | +from io import BytesIO, IOBase, UnsupportedOperation, SEEK_SET |
11 | 11 | from typing import IO, Iterator, Optional |
12 | 12 |
|
13 | 13 | from .validation import calculate_crc64 |
@@ -88,9 +88,6 @@ class StructuredMessageEncodeStream(IOBase): # pylint: disable=too-many-instanc |
88 | 88 | _current_region_length: int |
89 | 89 | _current_region_offset: int |
90 | 90 |
|
91 | | - _checksum_offset: int |
92 | | - """Tracks the offset the checksum has been calculated up to for seeking purposes""" |
93 | | - |
94 | 91 | _message_crc64: int |
95 | 92 | _segment_crc64s: dict[int, int] |
96 | 93 |
|
@@ -121,7 +118,6 @@ def __init__( |
121 | 118 | self._current_region_length = self._message_header_length |
122 | 119 | self._current_region_offset = 0 |
123 | 120 |
|
124 | | - self._checksum_offset = 0 |
125 | 121 | self._message_crc64 = 0 |
126 | 122 | self._segment_crc64s = {} |
127 | 123 |
|
@@ -171,9 +167,15 @@ def _update_current_region_length(self) -> None: |
171 | 167 | def __len__(self): |
172 | 168 | return self.message_length |
173 | 169 |
|
| 170 | + @property |
| 171 | + def closed(self) -> bool: |
| 172 | + return self._inner_stream.closed |
| 173 | + |
174 | 174 | def close(self) -> None: |
175 | | - self._inner_stream.close() |
176 | | - super().close() |
| 175 | + # Do not close the inner stream or this stream. |
| 176 | + # The inner stream is caller-owned and must survive for retries. |
| 177 | + # This stream may be re-read after a seek(0) on retry. |
| 178 | + pass |
177 | 179 |
|
178 | 180 | def readable(self) -> bool: |
179 | 181 | return True |
@@ -224,66 +226,23 @@ def seek(self, offset: int, whence: int = SEEK_SET) -> int: |
224 | 226 | if not self.seekable(): |
225 | 227 | raise UnsupportedOperation("Inner stream is not seekable.") |
226 | 228 |
|
227 | | - if whence == SEEK_SET: |
228 | | - position = offset |
229 | | - elif whence == SEEK_CUR: |
230 | | - position = self.tell() + offset |
231 | | - elif whence == SEEK_END: |
232 | | - position = self.message_length + offset |
233 | | - else: |
234 | | - raise ValueError(f"Invalid value for whence: {whence}") |
235 | | - |
236 | | - if position < 0: |
237 | | - raise ValueError(f"Cannot seek to negative position: {position}") |
238 | | - if position > self.tell(): |
239 | | - raise UnsupportedOperation("This stream only supports seeking backwards.") |
240 | | - |
241 | | - # MESSAGE_HEADER |
242 | | - if position < self._message_header_length: |
243 | | - self._current_region = SMRegion.MESSAGE_HEADER |
244 | | - self._current_region_offset = position |
245 | | - self._content_offset = 0 |
246 | | - self._current_segment_number = 0 |
247 | | - # MESSAGE_FOOTER |
248 | | - elif position >= self.message_length - self._message_footer_length: |
249 | | - self._current_region = SMRegion.MESSAGE_FOOTER |
250 | | - self._current_region_offset = position - (self.message_length - self._message_footer_length) |
251 | | - self._content_offset = self.content_length |
252 | | - self._current_segment_number = self._num_segments |
253 | | - else: |
254 | | - # The size of a "full" segment. Fine to use for calculating new segment number and pos |
255 | | - full_segment_size = self._segment_header_length + self._segment_size + self._segment_footer_length |
256 | | - new_segment_num = 1 + (position - self._message_header_length) // full_segment_size |
257 | | - segment_pos = (position - self._message_header_length) % full_segment_size |
258 | | - previous_segments_total_content_size = (new_segment_num - 1) * self._segment_size |
259 | | - |
260 | | - # We need the size of the segment we are seeking to for some of the calculations below |
261 | | - new_segment_size = self._segment_size |
262 | | - if new_segment_num == self._num_segments: |
263 | | - # The last segment size is the remaining content length |
264 | | - new_segment_size = self.content_length - previous_segments_total_content_size |
265 | | - |
266 | | - # SEGMENT_HEADER |
267 | | - if segment_pos < self._segment_header_length: |
268 | | - self._current_region = SMRegion.SEGMENT_HEADER |
269 | | - self._current_region_offset = segment_pos |
270 | | - self._content_offset = previous_segments_total_content_size |
271 | | - # SEGMENT_CONTENT |
272 | | - elif segment_pos < self._segment_header_length + new_segment_size: |
273 | | - self._current_region = SMRegion.SEGMENT_CONTENT |
274 | | - self._current_region_offset = segment_pos - self._segment_header_length |
275 | | - self._content_offset = previous_segments_total_content_size + self._current_region_offset |
276 | | - # SEGMENT_FOOTER |
277 | | - else: |
278 | | - self._current_region = SMRegion.SEGMENT_FOOTER |
279 | | - self._current_region_offset = segment_pos - self._segment_header_length - new_segment_size |
280 | | - self._content_offset = previous_segments_total_content_size + new_segment_size |
| 229 | + if whence != SEEK_SET: |
| 230 | + raise UnsupportedOperation("This stream only supports SEEK_SET.") |
281 | 231 |
|
282 | | - self._current_segment_number = new_segment_num |
| 232 | + if offset != 0: |
| 233 | + raise UnsupportedOperation("This stream only supports seeking to position 0.") |
283 | 234 |
|
284 | | - self._update_current_region_length() |
285 | | - self._inner_stream.seek((self._initial_content_position or 0) + self._content_offset) |
286 | | - return position |
| 235 | + # Reset to initial state |
| 236 | + self._content_offset = 0 |
| 237 | + self._current_segment_number = 0 |
| 238 | + self._current_region = SMRegion.MESSAGE_HEADER |
| 239 | + self._current_region_length = self._message_header_length |
| 240 | + self._current_region_offset = 0 |
| 241 | + self._message_crc64 = 0 |
| 242 | + self._segment_crc64s = {} |
| 243 | + |
| 244 | + self._inner_stream.seek(self._initial_content_position or 0) |
| 245 | + return 0 |
287 | 246 |
|
288 | 247 | def read(self, size: int = -1) -> bytes: |
289 | 248 | if self.closed: # pylint: disable=using-constant-test |
@@ -386,31 +345,20 @@ def _read_metadata_region(self, region: SMRegion, size: int, output: BytesIO) -> |
386 | 345 | return read_size |
387 | 346 |
|
388 | 347 | def _read_content(self, size: int, output: BytesIO) -> int: |
389 | | - # Will be non-zero if there is data to read that does not need to have checksum calculated. |
390 | | - # Will always be positive as stream can only seek backwards. |
391 | | - checksum_offset = self._checksum_offset - self._content_offset |
392 | | - |
393 | 348 | read_size = min(size, self._current_region_length - self._current_region_offset) |
394 | | - if checksum_offset != 0: |
395 | | - # Only read up to checksum offset this iteration |
396 | | - read_size = min(read_size, checksum_offset) |
397 | 349 |
|
398 | 350 | content = self._inner_stream.read(read_size) |
399 | 351 | if len(content) != read_size: |
400 | 352 | raise ValueError("Content ended early when encoding structured message.") |
401 | 353 | output.write(content) |
402 | 354 |
|
403 | 355 | if StructuredMessageProperties.CRC64 in self.flags: |
404 | | - if checksum_offset == 0: |
405 | | - self._segment_crc64s[self._current_segment_number] = calculate_crc64( |
406 | | - content, self._segment_crc64s[self._current_segment_number] |
407 | | - ) |
408 | | - self._message_crc64 = calculate_crc64(content, self._message_crc64) |
| 356 | + self._segment_crc64s[self._current_segment_number] = calculate_crc64( |
| 357 | + content, self._segment_crc64s[self._current_segment_number] |
| 358 | + ) |
| 359 | + self._message_crc64 = calculate_crc64(content, self._message_crc64) |
409 | 360 |
|
410 | 361 | self._content_offset += read_size |
411 | | - # Only update the checksum offset if we've read new data |
412 | | - if self._content_offset > self._checksum_offset: |
413 | | - self._checksum_offset += read_size |
414 | 362 | self._current_region_offset += read_size |
415 | 363 | if self._current_region_offset == self._current_region_length: |
416 | 364 | self._advance_region(SMRegion.SEGMENT_CONTENT) |
|
0 commit comments