test_download_driver.py 11 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302303304305306
  1. """
  2. Provider-agnostic unit tests for the base ranged download driver
  3. (``BaseBucketObject.download_to_file`` / ``_download_ranged``).
  4. The driver is the engine behind transparent large downloads on providers that
  5. do not override it (GCP, OpenStack Swift). Because the mock provider is
  6. AWS-backed and AWS overrides the driver with boto3's native downloader, the
  7. driver is exercised here directly against in-memory fakes so it has coverage
  8. in CI without cloud credentials.
  9. """
  10. import os
  11. import shutil
  12. import tempfile
  13. import threading
  14. import unittest
  15. from cloudbridge.base.resources import BaseBucketObject
  16. from cloudbridge.interfaces.exceptions import InvalidValueException
  17. from cloudbridge.interfaces.resources import TransferConfig
  18. class _Recorder:
  19. """Thread-safe range log shared by the original and cloned fake
  20. services."""
  21. def __init__(self, content):
  22. self.content = content
  23. self._lock = threading.Lock()
  24. self.ranges = [] # (offset, length) served
  25. self.services_used = set() # id() of each service that served a range
  26. self.clone_count = 0
  27. self.single_shot = False
  28. self.active = 0
  29. self.max_active = 0
  30. self.fail_on_offset = None # offset that should raise
  31. self.on_serve = None # hook called as each range is served
  32. def serve_range(self, service, offset, length):
  33. with self._lock:
  34. self.active += 1
  35. self.max_active = max(self.max_active, self.active)
  36. try:
  37. if self.on_serve:
  38. self.on_serve()
  39. if self.fail_on_offset == offset:
  40. raise RuntimeError("boom at offset %d" % offset)
  41. # Hold briefly so concurrent fetches genuinely overlap.
  42. threading.Event().wait(0.02)
  43. with self._lock:
  44. self.ranges.append((offset, length))
  45. self.services_used.add(id(service))
  46. return self.content[offset:offset + length]
  47. finally:
  48. with self._lock:
  49. self.active -= 1
  50. class _FakeService:
  51. def __init__(self, recorder, provider):
  52. self._recorder = recorder
  53. self._provider = provider
  54. def download_range(self, bucket, object_name, offset, length):
  55. return self._recorder.serve_range(self, offset, length)
  56. class _FakeStorage:
  57. def __init__(self, service):
  58. self._bucket_objects = service
  59. class _FakeProvider:
  60. def __init__(self, recorder):
  61. self._recorder = recorder
  62. self.storage = _FakeStorage(_FakeService(recorder, self))
  63. def clone(self, zone=None):
  64. self._recorder.clone_count += 1
  65. return _FakeProvider(self._recorder)
  66. def _get_config_value(self, key, default_value=None):
  67. return default_value
  68. class _DriverObject(BaseBucketObject):
  69. """A BaseBucketObject wired to fakes with tiny transfer sizes."""
  70. def __init__(self, provider, threshold, part_size, concurrency):
  71. super(_DriverObject, self).__init__(provider)
  72. self._threshold = threshold
  73. self._part_size = part_size
  74. self._concurrency = concurrency
  75. @property
  76. def id(self):
  77. return "obj"
  78. @property
  79. def name(self):
  80. return "obj"
  81. @property
  82. def size(self):
  83. return len(self._provider._recorder.content)
  84. @property
  85. def bucket(self):
  86. return "BUCKET"
  87. def save_content(self, target_stream):
  88. self._provider._recorder.single_shot = True
  89. target_stream.write(self._provider._recorder.content)
  90. def _multipart_threshold(self, config=None):
  91. if config is not None and config.threshold is not None:
  92. return config.threshold
  93. return self._threshold
  94. def _multipart_part_size(self, config=None):
  95. if config is not None and config.part_size is not None:
  96. return config.part_size
  97. return self._part_size
  98. def _multipart_max_concurrency(self, config=None):
  99. if config is not None and config.max_concurrency is not None:
  100. return config.max_concurrency
  101. return self._concurrency
  102. class DownloadDriverTestCase(unittest.TestCase):
  103. def _driver(self, recorder, threshold, part_size, concurrency):
  104. return _DriverObject(
  105. _FakeProvider(recorder), threshold, part_size, concurrency)
  106. def _download(self, driver, config=None):
  107. fd, path = tempfile.mkstemp()
  108. os.close(fd)
  109. os.remove(path)
  110. try:
  111. driver.download_to_file(path, config)
  112. with open(path, 'rb') as f:
  113. return f.read()
  114. finally:
  115. if os.path.exists(path):
  116. os.remove(path)
  117. def test_reassembles_content_in_order(self):
  118. content = b"abcdefghijABCDEFGHIJ0123456789x" # 31 bytes -> 8 ranges
  119. recorder = _Recorder(content)
  120. driver = self._driver(
  121. recorder, threshold=10, part_size=4, concurrency=3)
  122. self.assertEqual(self._download(driver), content)
  123. self.assertFalse(recorder.single_shot)
  124. # Ranges tile the object exactly: no gaps, no overlap, short tail.
  125. self.assertEqual(
  126. sorted(recorder.ranges),
  127. [(offset, min(4, 31 - offset)) for offset in range(0, 31, 4)])
  128. def test_below_threshold_uses_single_shot(self):
  129. content = b"tiny content"
  130. recorder = _Recorder(content)
  131. driver = self._driver(
  132. recorder, threshold=100, part_size=4, concurrency=3)
  133. self.assertEqual(self._download(driver), content)
  134. self.assertTrue(recorder.single_shot)
  135. self.assertEqual(recorder.ranges, [])
  136. def test_downloads_ranges_concurrently_via_cloned_services(self):
  137. concurrency = 4
  138. content = bytes(range(12)) # 12 ranges of one byte each
  139. recorder = _Recorder(content)
  140. driver = self._driver(
  141. recorder, threshold=1, part_size=1, concurrency=concurrency)
  142. self.assertEqual(self._download(driver), content)
  143. # A clone per worker, reused across ranges.
  144. self.assertEqual(recorder.clone_count, concurrency)
  145. self.assertEqual(len(recorder.services_used), concurrency)
  146. # Real parallelism happened, bounded by the configured concurrency.
  147. self.assertGreater(recorder.max_active, 1)
  148. self.assertLessEqual(recorder.max_active, concurrency)
  149. def test_single_concurrency_does_not_clone(self):
  150. content = b"abcdefghij"
  151. recorder = _Recorder(content)
  152. driver = self._driver(
  153. recorder, threshold=1, part_size=4, concurrency=1)
  154. self.assertEqual(self._download(driver), content)
  155. self.assertEqual(recorder.clone_count, 0)
  156. self.assertEqual(recorder.max_active, 1)
  157. def test_per_call_config_overrides_concurrency(self):
  158. content = bytes(range(12))
  159. recorder = _Recorder(content)
  160. driver = self._driver(
  161. recorder, threshold=1, part_size=1, concurrency=1)
  162. result = None
  163. fd, path = tempfile.mkstemp()
  164. os.close(fd)
  165. try:
  166. driver.download_to_file(path, TransferConfig(max_concurrency=3))
  167. with open(path, 'rb') as f:
  168. result = f.read()
  169. finally:
  170. os.remove(path)
  171. self.assertEqual(result, content)
  172. self.assertEqual(recorder.clone_count, 3)
  173. self.assertGreater(recorder.max_active, 1)
  174. self.assertLessEqual(recorder.max_active, 3)
  175. def test_removes_partial_file_and_raises_on_range_failure(self):
  176. content = bytes(range(16))
  177. recorder = _Recorder(content)
  178. recorder.fail_on_offset = 8
  179. driver = self._driver(
  180. recorder, threshold=1, part_size=4, concurrency=2)
  181. fd, path = tempfile.mkstemp()
  182. os.close(fd)
  183. os.remove(path)
  184. try:
  185. with self.assertRaises(Exception):
  186. driver.download_to_file(path)
  187. self.assertFalse(os.path.exists(path))
  188. finally:
  189. if os.path.exists(path):
  190. os.remove(path)
  191. def test_destination_only_appears_once_complete(self):
  192. content = bytes(range(256))
  193. recorder = _Recorder(content)
  194. driver = self._driver(
  195. recorder, threshold=1, part_size=16, concurrency=3)
  196. fd, path = tempfile.mkstemp()
  197. os.close(fd)
  198. os.remove(path)
  199. seen_early = []
  200. recorder.on_serve = lambda: seen_early.append(os.path.exists(path))
  201. try:
  202. driver.download_to_file(path)
  203. with open(path, 'rb') as f:
  204. self.assertEqual(f.read(), content)
  205. finally:
  206. if os.path.exists(path):
  207. os.remove(path)
  208. # A partially written object is never visible at the destination.
  209. self.assertTrue(seen_early)
  210. self.assertNotIn(True, seen_early)
  211. def test_survives_concurrent_downloader_taking_the_destination(self):
  212. # Galaxy gives every download of a dataset the same cache .tmp path,
  213. # so a second download of the same dataset can rename the destination
  214. # away while this one is still fetching ranges.
  215. content = bytes(range(256))
  216. recorder = _Recorder(content)
  217. driver = self._driver(
  218. recorder, threshold=1, part_size=16, concurrency=3)
  219. directory = tempfile.mkdtemp()
  220. path = os.path.join(directory, 'dataset.dat')
  221. taken = os.path.join(directory, 'taken.dat')
  222. def steal_destination():
  223. if os.path.exists(path):
  224. os.replace(path, taken)
  225. recorder.on_serve = steal_destination
  226. try:
  227. driver.download_to_file(path)
  228. with open(path, 'rb') as f:
  229. self.assertEqual(f.read(), content)
  230. finally:
  231. shutil.rmtree(directory)
  232. def test_failed_download_leaves_an_existing_destination_intact(self):
  233. content = bytes(range(64))
  234. recorder = _Recorder(content)
  235. recorder.fail_on_offset = 16
  236. driver = self._driver(
  237. recorder, threshold=1, part_size=16, concurrency=2)
  238. directory = tempfile.mkdtemp()
  239. path = os.path.join(directory, 'dataset.dat')
  240. with open(path, 'wb') as f:
  241. f.write(b'previously cached')
  242. try:
  243. with self.assertRaises(Exception):
  244. driver.download_to_file(path)
  245. # The cached copy survives a failed refetch, and no scratch file
  246. # is left behind next to it.
  247. with open(path, 'rb') as f:
  248. self.assertEqual(f.read(), b'previously cached')
  249. self.assertEqual(os.listdir(directory), ['dataset.dat'])
  250. finally:
  251. shutil.rmtree(directory)
  252. def test_part_size_must_be_positive(self):
  253. content = bytes(range(16))
  254. recorder = _Recorder(content)
  255. driver = self._driver(
  256. recorder, threshold=1, part_size=4, concurrency=2)
  257. with self.assertRaises(InvalidValueException):
  258. self._download(driver, TransferConfig(part_size=0))
  259. self.assertEqual(recorder.ranges, [])
  260. if __name__ == "__main__":
  261. unittest.main()