Skip to content

Commit eb7b283

Browse files
committed
Fix directory expansion bug
1 parent 8bbb991 commit eb7b283

1 file changed

Lines changed: 36 additions & 20 deletions

File tree

hail/python/hailtop/aiocloud/aiogoogle/client/storage_client.py

Lines changed: 36 additions & 20 deletions
Original file line numberDiff line numberDiff line change
@@ -30,7 +30,13 @@
3030

3131
from hailtop import timex
3232
from hailtop.aiocloud.common import AnonymousCloudCredentials
33-
from hailtop.aiotools import FeedableAsyncIterable, Transfer, WeightedSemaphore, WriteBuffer, weighted_bounded_gather2
33+
from hailtop.aiotools import (
34+
FeedableAsyncIterable,
35+
Transfer,
36+
WeightedSemaphore,
37+
WriteBuffer,
38+
weighted_bounded_gather2,
39+
)
3440
from hailtop.aiotools.fs import (
3541
AsyncFS,
3642
AsyncFSFactory,
@@ -51,6 +57,8 @@
5157
blocking_to_async,
5258
retry_transient_errors,
5359
secret_alnum_string,
60+
url_basename,
61+
url_join,
5462
)
5563

5664
from ...common.session import BaseSession
@@ -476,16 +484,13 @@ async def download_single_file(self, bucket: str, filename: str, dest: str) -> N
476484
user_project = self._get_user_project_for_bucket(bucket)
477485
bucket_instance = self._client.bucket(bucket, user_project=user_project)
478486
blob = bucket_instance.blob(filename)
479-
if dest.startswith('file://'):
480-
local_dest = urllib.parse.urlparse(dest).path
481-
else:
482-
local_dest = dest
483487

484-
os.makedirs(os.path.dirname(local_dest), exist_ok=True)
488+
if not os.path.exists(os.path.dirname(dest)):
489+
os.makedirs(os.path.dirname(dest), exist_ok=True)
485490
await blocking_to_async(
486491
self._thread_pool,
487492
blob.download_to_filename,
488-
local_dest,
493+
dest,
489494
single_shot_download=True,
490495
timeout=self._timeout,
491496
)
@@ -494,17 +499,14 @@ async def download_large_file(self, bucket: str, src: str, dest: str) -> None:
494499
user_project = self._get_user_project_for_bucket(bucket)
495500
bucket_instance = self._client.bucket(bucket, user_project=user_project)
496501
blob = bucket_instance.blob(src)
497-
if dest.startswith('file://'):
498-
local_dest = urllib.parse.urlparse(dest).path
499-
else:
500-
local_dest = dest
501502

502-
os.makedirs(os.path.dirname(local_dest), exist_ok=True)
503+
if not os.path.exists(os.path.dirname(dest)):
504+
os.makedirs(os.path.dirname(dest), exist_ok=True)
503505
await blocking_to_async(
504506
self._thread_pool,
505507
transfer_manager.download_chunks_concurrently,
506508
blob=blob,
507-
filename=local_dest,
509+
filename=dest,
508510
chunk_size=self.CHUNK_SIZE,
509511
download_kwargs={'timeout': self._timeout},
510512
max_workers=self.MAX_WORKERS,
@@ -999,9 +1001,13 @@ async def copy_to_local(
9991001
if await self.isfile(transfer.src):
10001002
stat = await self.statfile(transfer.src)
10011003
size = await stat.size()
1002-
filename = self.get_bucket_and_name(transfer.src)[1]
1003-
if transfer.dest.endswith('/') or transfer.treat_dest_as == Transfer.DEST_DIR:
1004-
target_dest = (transfer.dest if transfer.dest.endswith('/') else transfer.dest + '/') + filename
1004+
filename = url_basename(transfer.src)
1005+
if transfer.treat_dest_as == Transfer.DEST_DIR or (
1006+
transfer.treat_dest_as == Transfer.INFER_DEST and transfer.dest.endswith('/')
1007+
):
1008+
target_dest = url_join(transfer.dest, filename)
1009+
elif transfer.treat_dest_as == Transfer.DEST_IS_TARGET and os.path.isdir(transfer.dest):
1010+
raise IsADirectoryError(transfer.dest)
10051011
else:
10061012
target_dest = transfer.dest
10071013
if size > self._storage_client.CHUNK_SIZE:
@@ -1034,14 +1040,24 @@ async def copy_to_local(
10341040
async for file in await self.listfiles(transfer.src, recursive=True):
10351041
stat = await file.status()
10361042
size = await stat.size()
1037-
filename = file.basename()
1038-
target_dest = (transfer.dest if transfer.dest.endswith('/') else transfer.dest + '/') + filename
1043+
filename = await file.url_full()
1044+
relfilename = filename[len(transfer.src) :].lstrip('/')
1045+
if transfer.treat_dest_as == Transfer.DEST_DIR or (
1046+
transfer.treat_dest_as == Transfer.INFER_DEST and transfer.dest.endswith('/')
1047+
):
1048+
target_dest = url_join(transfer.dest, relfilename)
1049+
elif os.path.isfile(transfer.dest):
1050+
raise NotADirectoryError(transfer.dest)
1051+
else:
1052+
target_dest = url_join(
1053+
url_join(transfer.dest, url_basename(transfer.src.rstrip('/'))), relfilename
1054+
)
10391055
if size > self._storage_client.CHUNK_SIZE:
10401056
copy_operations.append(
10411057
functools.partial(
10421058
self._copy_single_large_file,
10431059
xfer_sema,
1044-
await file.url(),
1060+
filename,
10451061
target_dest,
10461062
size,
10471063
report,
@@ -1054,7 +1070,7 @@ async def copy_to_local(
10541070
functools.partial(
10551071
self._copy_single_local_file,
10561072
xfer_sema,
1057-
await file.url(),
1073+
filename,
10581074
target_dest,
10591075
size,
10601076
report,

0 commit comments

Comments
 (0)