Skip to content

Commit dda5f3d

Browse files
committed
[Storage] Simplify Encoder seek, fix SM streaming retry (#46564)
1 parent 660572e commit dda5f3d

8 files changed

Lines changed: 217 additions & 425 deletions

File tree

sdk/storage/azure-storage-blob/assets.json

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -2,5 +2,5 @@
22
"AssetsRepo": "Azure/azure-sdk-assets",
33
"AssetsRepoPrefixPath": "python",
44
"TagPrefix": "python/storage/azure-storage-blob",
5-
"Tag": "python/storage/azure-storage-blob_e0a670a6a4"
5+
"Tag": "python/storage/azure-storage-blob_89384f00b6"
66
}

sdk/storage/azure-storage-blob/azure/storage/blob/_shared/streams.py

Lines changed: 28 additions & 80 deletions
Original file line numberDiff line numberDiff line change
@@ -7,7 +7,7 @@
77
import math
88
import sys
99
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
1111
from typing import IO, Iterator, Optional
1212

1313
from .validation import calculate_crc64
@@ -88,9 +88,6 @@ class StructuredMessageEncodeStream(IOBase): # pylint: disable=too-many-instanc
8888
_current_region_length: int
8989
_current_region_offset: int
9090

91-
_checksum_offset: int
92-
"""Tracks the offset the checksum has been calculated up to for seeking purposes"""
93-
9491
_message_crc64: int
9592
_segment_crc64s: dict[int, int]
9693

@@ -121,7 +118,6 @@ def __init__(
121118
self._current_region_length = self._message_header_length
122119
self._current_region_offset = 0
123120

124-
self._checksum_offset = 0
125121
self._message_crc64 = 0
126122
self._segment_crc64s = {}
127123

@@ -171,9 +167,15 @@ def _update_current_region_length(self) -> None:
171167
def __len__(self):
172168
return self.message_length
173169

170+
@property
171+
def closed(self) -> bool:
172+
return self._inner_stream.closed
173+
174174
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
177179

178180
def readable(self) -> bool:
179181
return True
@@ -224,66 +226,23 @@ def seek(self, offset: int, whence: int = SEEK_SET) -> int:
224226
if not self.seekable():
225227
raise UnsupportedOperation("Inner stream is not seekable.")
226228

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.")
281231

282-
self._current_segment_number = new_segment_num
232+
if offset != 0:
233+
raise UnsupportedOperation("This stream only supports seeking to position 0.")
283234

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
287246

288247
def read(self, size: int = -1) -> bytes:
289248
if self.closed: # pylint: disable=using-constant-test
@@ -386,31 +345,20 @@ def _read_metadata_region(self, region: SMRegion, size: int, output: BytesIO) ->
386345
return read_size
387346

388347
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-
393348
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)
397349

398350
content = self._inner_stream.read(read_size)
399351
if len(content) != read_size:
400352
raise ValueError("Content ended early when encoding structured message.")
401353
output.write(content)
402354

403355
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)
409360

410361
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
414362
self._current_region_offset += read_size
415363
if self._current_region_offset == self._current_region_length:
416364
self._advance_region(SMRegion.SEGMENT_CONTENT)

sdk/storage/azure-storage-blob/tests/test_content_validation.py

Lines changed: 44 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -567,3 +567,47 @@ def download_hook_fail_once(response):
567567
content = downloader.read()
568568
assert download_call_count == 2 # Original + retry
569569
assert content == data
570+
571+
@BlobPreparer()
572+
@pytest.mark.parametrize('a', [True, 'md5', 'crc64']) # a: validate_content
573+
@GenericTestProxyParametrize1()
574+
@recorded_by_proxy
575+
def test_streaming_with_retry(self, a, **kwargs):
576+
storage_account_name = kwargs.pop("storage_account_name")
577+
578+
# Setup with retry enabled
579+
token_credential = self.get_credential(BlobServiceClient)
580+
self.bsc = BlobServiceClient(
581+
self.account_url(storage_account_name, "blob"),
582+
token_credential,
583+
retry_total=1,
584+
initial_backoff=0.1,
585+
increment_base=0.1,
586+
logging_enable=True
587+
)
588+
self.container = self.bsc.get_container_client(self.get_resource_name('utcontainer'))
589+
try:
590+
self.container.create_container()
591+
except ResourceExistsError:
592+
pass
593+
blob = self.container.get_blob_client(self._get_blob_reference())
594+
595+
content = b'abc' * 512
596+
assert_method = assert_structured_message if a == 'crc64' else assert_content_md5
597+
598+
call_count = 0
599+
def hook_fail_once(response):
600+
nonlocal call_count
601+
call_count += 1
602+
# Assert content validation headers are present on both attempts
603+
assert_method(response)
604+
if call_count == 1:
605+
response.http_response.status_code = 408 # Request Timeout - triggers retry
606+
607+
# Use stage_block to test structured message streaming
608+
blob.stage_block('1', BytesIO(content), validate_content=a, raw_response_hook=hook_fail_once)
609+
assert call_count == 2 # Original + retry
610+
611+
blob.commit_block_list([BlobBlock('1')])
612+
result = blob.download_blob()
613+
assert result.read() == content

sdk/storage/azure-storage-blob/tests/test_content_validation_async.py

Lines changed: 45 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -542,3 +542,48 @@ def download_hook_fail_once(response):
542542
content = await downloader.read()
543543
assert download_call_count == 2 # Original + retry
544544
assert content == data
545+
546+
@BlobPreparer()
547+
@pytest.mark.parametrize('a', [True, 'md5', 'crc64']) # a: validate_content
548+
@GenericTestProxyParametrize1()
549+
@recorded_by_proxy_async
550+
async def test_streaming_with_retry(self, a, **kwargs):
551+
storage_account_name = kwargs.pop("storage_account_name")
552+
553+
# Setup with retry enabled
554+
token_credential = self.get_credential(BlobServiceClient, is_async=True)
555+
self.bsc = BlobServiceClient(
556+
self.account_url(storage_account_name, "blob"),
557+
token_credential,
558+
retry_total=1,
559+
initial_backoff=0.1,
560+
increment_base=0.1,
561+
logging_enable=True
562+
)
563+
self.container = self.bsc.get_container_client(self.get_resource_name('utcontainer'))
564+
try:
565+
await self.container.create_container()
566+
except ResourceExistsError:
567+
pass
568+
blob = self.container.get_blob_client(self._get_blob_reference())
569+
570+
content = b'abc' * 512
571+
assert_method = assert_structured_message if a == 'crc64' else assert_content_md5
572+
573+
# Test stage_block streaming with retry
574+
call_count = 0
575+
def hook_fail_once(response):
576+
nonlocal call_count
577+
call_count += 1
578+
# Assert content validation headers are present on both attempts
579+
assert_method(response)
580+
if call_count == 1:
581+
response.http_response.status_code = 408 # Request Timeout - triggers retry
582+
583+
# Use stage_block to test structured message streaming
584+
await blob.stage_block('1', BytesIO(content), validate_content=a, raw_response_hook=hook_fail_once)
585+
assert call_count == 2 # Original + retry
586+
587+
await blob.commit_block_list([BlobBlock('1')])
588+
result = await blob.download_blob()
589+
assert await result.read() == content

0 commit comments

Comments
 (0)