@@ -64,6 +64,25 @@ def fetch(self, data_range, stream=False):
6464 return buff
6565
6666
67+ class TrackingResponse (io .BytesIO ):
68+ """BytesIO test double that records reads and connection release."""
69+ def __init__ (self , data , fail_read = False ):
70+ super (TrackingResponse , self ).__init__ (data )
71+ self .bytes_read = 0
72+ self .fail_read = fail_read
73+ self .release_count = 0
74+
75+ def read (self , size = - 1 ):
76+ if self .fail_read :
77+ raise IOError ("simulated read failure" )
78+ content = super (TrackingResponse , self ).read (size )
79+ self .bytes_read += len (content )
80+ return content
81+
82+ def release_conn (self ):
83+ self .release_count += 1
84+
85+
6786class TestPartialBuffer (unittest .TestCase ):
6887 def test_handles_short_reads_from_the_stream (self ):
6988 """A socket-backed response may return less than asked for per read."""
@@ -89,11 +108,29 @@ def test_handles_a_server_sending_less_than_declared(self):
89108
90109 def test_does_not_buffer_more_than_declared_size (self ):
91110 """A server sending more than it declared must not enlarge the buffer."""
92- oversized = io . BytesIO (b'x' * 10000 )
111+ oversized = TrackingResponse (b'x' * 10000 )
93112 pb = rz .PartialBuffer (oversized , 0 , 100 , stream = False )
94113 self .assertEqual (len (pb .read (0 )), 100 )
95114 # the rest of the response was never pulled into memory
96- self .assertEqual (oversized .tell (), 100 )
115+ self .assertEqual (oversized .bytes_read , 100 )
116+ self .assertTrue (oversized .closed )
117+ self .assertEqual (oversized .release_count , 1 )
118+
119+ def test_closes_source_when_buffering_fails (self ):
120+ source = TrackingResponse (b'x' * 100 , fail_read = True )
121+ with self .assertRaises (IOError ):
122+ rz .PartialBuffer (source , 0 , 100 , stream = False )
123+ self .assertTrue (source .closed )
124+ self .assertEqual (source .release_count , 1 )
125+
126+ def test_stream_owns_source_until_closed (self ):
127+ source = TrackingResponse (b'x' * 100 )
128+ pb = rz .PartialBuffer (source , 0 , 100 , stream = True )
129+ self .assertFalse (source .closed )
130+ self .assertEqual (source .release_count , 0 )
131+ pb .close ()
132+ self .assertTrue (source .closed )
133+ self .assertEqual (source .release_count , 1 )
97134
98135 def setUp (self ):
99136 if not hasattr (self , 'assertRaisesRegex' ):
@@ -242,26 +279,39 @@ class Fetcher(rz.RemoteFetcher):
242279 def __init__ (self , header ):
243280 super (Fetcher , self ).__init__ ('http://test.com/file.zip' )
244281 self .header = header
282+ self .response = TrackingResponse (b'x' * 100 )
245283
246284 def _request (self , kwargs ):
247- return io . BytesIO ( b'x' * 100 ) , self .header
285+ return self . response , self .header
248286
287+ invalid = Fetcher ('bytes 100-50/1000' )
249288 with self .assertRaises (rz .RemoteZipError ):
250- Fetcher ('bytes 100-50/1000' ).fetch ((0 , 99 ))
289+ invalid .fetch ((0 , 99 ))
290+ self .assertTrue (invalid .response .closed )
291+ self .assertEqual (invalid .response .release_count , 1 )
251292
293+ invalid = Fetcher ('bytes -500/1000' )
252294 with self .assertRaises (rz .RemoteZipError ):
253- Fetcher ('bytes -500/1000' ).fetch ((0 , 99 ))
295+ invalid .fetch ((0 , 99 ))
296+ self .assertTrue (invalid .response .closed )
297+ self .assertEqual (invalid .response .release_count , 1 )
254298
255299 # a malformed header must not leak a bare ValueError to the caller.
256300 # 'bytes */1000' is the RFC 7233 unsatisfied-range form, so this is not
257301 # only about hostile input.
258302 for bad in ('bytes abc-def/1000' , 'bytes -500-100/1000' , 'bytes /1000' ,
259303 'bytes */1000' , '' ):
304+ malformed = Fetcher (bad )
260305 with self .assertRaises (rz .RemoteZipError ):
261- Fetcher (bad ).fetch ((0 , 99 ))
306+ malformed .fetch ((0 , 99 ))
307+ self .assertTrue (malformed .response .closed )
308+ self .assertEqual (malformed .response .release_count , 1 )
262309
263310 # an unknown total length is legitimate and must still be accepted
264- Fetcher ('bytes 0-99/*' ).fetch ((0 , 99 ))
311+ valid = Fetcher ('bytes 0-99/*' )
312+ valid .fetch ((0 , 99 ))
313+ self .assertTrue (valid .response .closed )
314+ self .assertEqual (valid .response .release_count , 1 )
265315
266316 def test_parse_range_header (self ):
267317 range_min , range_max = rz .RemoteFetcher .parse_range_header ('bytes 0-11/12' )
0 commit comments