diff --git a/pyproject.toml b/pyproject.toml index 23f53f4..5903c69 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -4,7 +4,7 @@ build-backend = "setuptools.build_meta" [project] name = "remotezip" -version = "0.12.5" +version = "0.12.6" description = "Access zip file content hosted remotely without downloading the full file." readme = "README.md" requires-python = ">=2.7" @@ -45,4 +45,4 @@ Homepage = "https://github.com/gtsystem/python-remotezip" remotezip = "remotezip:main" [tool.setuptools] -py-modules = ["remotezip"] \ No newline at end of file +py-modules = ["remotezip"] diff --git a/remotezip.py b/remotezip.py index 03d1e11..0e507fe 100755 --- a/remotezip.py +++ b/remotezip.py @@ -30,8 +30,45 @@ class PartialBuffer: however, any attempt to read data outside the partial data is going to fail with OutOfBound error. """ + @staticmethod + def _read_up_to(buffer, size): + """Read at most `size` bytes into a new buffer. + + A single read() is not enough: a socket-backed response can return + fewer bytes than requested while more are still coming. Data is written + straight into the result so that no intermediate copy of the whole + range is held. + """ + result = io.BytesIO() + remaining = size + while remaining > 0: + chunk = buffer.read(remaining) + if not chunk: + break + result.write(chunk) + remaining -= len(chunk) + result.seek(0) + return result + + @staticmethod + def _close_buffer(buffer): + """Close a response buffer and release its connection, if any.""" + try: + buffer.close() + finally: + if hasattr(buffer, 'release_conn'): + buffer.release_conn() + def __init__(self, buffer, offset, size, stream): - self.buffer = buffer if stream else io.BytesIO(buffer.read()) + # Read at most `size` bytes: the declared range is what this buffer + # represents, and a server may send more than it announced. + if stream: + self.buffer = buffer + else: + try: + self.buffer = self._read_up_to(buffer, size) + finally: + self._close_buffer(buffer) self._offset = offset self._size = size self._position = offset @@ -56,9 +93,7 @@ def read(self, size=0): def close(self): """Ensure memory and connections are closed""" if not self.buffer.closed: - self.buffer.close() - if hasattr(self.buffer, 'release_conn'): - self.buffer.release_conn() + self._close_buffer(self.buffer) def tell(self): """Returns the current position on the virtual buffer""" @@ -223,7 +258,18 @@ def fetch(self, data_range, stream=False): kwargs = self.prepare_request(data_range) try: res, range_header = self._request(kwargs) - range_min, range_max = self.parse_range_header(range_header) + try: + try: + range_min, range_max = self.parse_range_header(range_header) + except ValueError: + raise RemoteZipError( + "Malformed Content-Range returned by the server: %s" % range_header) + if range_max is None or range_max < range_min: + raise RemoteZipError( + "Invalid Content-Range returned by the server: %s" % range_header) + except RemoteZipError: + PartialBuffer._close_buffer(res) + raise return PartialBuffer(res, range_min, range_max - range_min + 1, stream) except IOError as e: raise RemoteIOError(str(e)) diff --git a/test_remotezip.py b/test_remotezip.py index 186a055..1127123 100644 --- a/test_remotezip.py +++ b/test_remotezip.py @@ -64,7 +64,74 @@ def fetch(self, data_range, stream=False): return buff +class TrackingResponse(io.BytesIO): + """BytesIO test double that records reads and connection release.""" + def __init__(self, data, fail_read=False): + super(TrackingResponse, self).__init__(data) + self.bytes_read = 0 + self.fail_read = fail_read + self.release_count = 0 + + def read(self, size=-1): + if self.fail_read: + raise IOError("simulated read failure") + content = super(TrackingResponse, self).read(size) + self.bytes_read += len(content) + return content + + def release_conn(self): + self.release_count += 1 + + class TestPartialBuffer(unittest.TestCase): + def test_handles_short_reads_from_the_stream(self): + """A socket-backed response may return less than asked for per read.""" + class ShortReader(io.RawIOBase): + def __init__(self, data, chunk): + self._b = io.BytesIO(data) + self._chunk = chunk + + def read(self, n=-1): + if n is None or n < 0: + return self._b.read() + return self._b.read(min(n, self._chunk)) + + data = b'z' * 1000 + for chunk in (1000, 512, 100, 1): + pb = rz.PartialBuffer(ShortReader(data, chunk), 0, len(data), stream=False) + self.assertEqual(pb.read(0), data) + + def test_handles_a_server_sending_less_than_declared(self): + """A truncated response must not hang or raise; it yields what arrived.""" + pb = rz.PartialBuffer(io.BytesIO(b'z' * 40), 0, 1000, stream=False) + self.assertEqual(pb.read(0), b'z' * 40) + + def test_does_not_buffer_more_than_declared_size(self): + """A server sending more than it declared must not enlarge the buffer.""" + oversized = TrackingResponse(b'x' * 10000) + pb = rz.PartialBuffer(oversized, 0, 100, stream=False) + self.assertEqual(len(pb.read(0)), 100) + # the rest of the response was never pulled into memory + self.assertEqual(oversized.bytes_read, 100) + self.assertTrue(oversized.closed) + self.assertEqual(oversized.release_count, 1) + + def test_closes_source_when_buffering_fails(self): + source = TrackingResponse(b'x' * 100, fail_read=True) + with self.assertRaises(IOError): + rz.PartialBuffer(source, 0, 100, stream=False) + self.assertTrue(source.closed) + self.assertEqual(source.release_count, 1) + + def test_stream_owns_source_until_closed(self): + source = TrackingResponse(b'x' * 100) + pb = rz.PartialBuffer(source, 0, 100, stream=True) + self.assertFalse(source.closed) + self.assertEqual(source.release_count, 0) + pb.close() + self.assertTrue(source.closed) + self.assertEqual(source.release_count, 1) + def setUp(self): if not hasattr(self, 'assertRaisesRegex'): self.assertRaisesRegex = self.assertRaisesRegexp @@ -206,6 +273,46 @@ def test_build_range_header(self): header = rz.RemoteFetcher.build_range_header(-123, None) self.assertEqual(header, 'bytes=-123') + def test_fetch_rejects_invalid_content_range(self): + """A server must not be able to declare a range that ends before it starts.""" + class Fetcher(rz.RemoteFetcher): + def __init__(self, header): + super(Fetcher, self).__init__('http://test.com/file.zip') + self.header = header + self.response = TrackingResponse(b'x' * 100) + + def _request(self, kwargs): + return self.response, self.header + + invalid = Fetcher('bytes 100-50/1000') + with self.assertRaises(rz.RemoteZipError): + invalid.fetch((0, 99)) + self.assertTrue(invalid.response.closed) + self.assertEqual(invalid.response.release_count, 1) + + invalid = Fetcher('bytes -500/1000') + with self.assertRaises(rz.RemoteZipError): + invalid.fetch((0, 99)) + self.assertTrue(invalid.response.closed) + self.assertEqual(invalid.response.release_count, 1) + + # a malformed header must not leak a bare ValueError to the caller. + # 'bytes */1000' is the RFC 7233 unsatisfied-range form, so this is not + # only about hostile input. + for bad in ('bytes abc-def/1000', 'bytes -500-100/1000', 'bytes /1000', + 'bytes */1000', ''): + malformed = Fetcher(bad) + with self.assertRaises(rz.RemoteZipError): + malformed.fetch((0, 99)) + self.assertTrue(malformed.response.closed) + self.assertEqual(malformed.response.release_count, 1) + + # an unknown total length is legitimate and must still be accepted + valid = Fetcher('bytes 0-99/*') + valid.fetch((0, 99)) + self.assertTrue(valid.response.closed) + self.assertEqual(valid.response.release_count, 1) + def test_parse_range_header(self): range_min, range_max = rz.RemoteFetcher.parse_range_header('bytes 0-11/12') self.assertEqual(range_min, 0) diff --git a/uv.lock b/uv.lock index d694115..a7d505a 100644 --- a/uv.lock +++ b/uv.lock @@ -211,7 +211,7 @@ wheels = [ [[package]] name = "remotezip" -version = "0.12.5" +version = "0.12.6" source = { editable = "." } dependencies = [ { name = "requests", version = "2.15.1", source = { registry = "https://pypi.org/simple" }, marker = "python_full_version < '3.10'" },