Skip to content

Commit ead1b66

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

1 file changed

Lines changed: 40 additions & 20 deletions

File tree

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

Lines changed: 40 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,11 +1001,17 @@ 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)
10051009
else:
10061010
target_dest = transfer.dest
1011+
1012+
if os.path.isdir(target_dest):
1013+
raise IsADirectoryError(transfer.dest)
1014+
10071015
if size > self._storage_client.CHUNK_SIZE:
10081016
copy_operations.append(
10091017
functools.partial(
@@ -1034,14 +1042,26 @@ async def copy_to_local(
10341042
async for file in await self.listfiles(transfer.src, recursive=True):
10351043
stat = await file.status()
10361044
size = await stat.size()
1037-
filename = file.basename()
1038-
target_dest = (transfer.dest if transfer.dest.endswith('/') else transfer.dest + '/') + filename
1045+
filename = await file.url_full()
1046+
relfilename = filename[len(transfer.src) :].lstrip('/')
1047+
if transfer.treat_dest_as == Transfer.DEST_DIR or (
1048+
transfer.treat_dest_as == Transfer.INFER_DEST and transfer.dest.endswith('/')
1049+
):
1050+
target_dest = url_join(
1051+
url_join(transfer.dest, url_basename(transfer.src.rstrip('/'))), relfilename
1052+
)
1053+
else:
1054+
target_dest = url_join(transfer.dest, relfilename)
1055+
1056+
if os.path.isfile(target_dest):
1057+
raise NotADirectoryError(transfer.dest)
1058+
10391059
if size > self._storage_client.CHUNK_SIZE:
10401060
copy_operations.append(
10411061
functools.partial(
10421062
self._copy_single_large_file,
10431063
xfer_sema,
1044-
await file.url(),
1064+
filename,
10451065
target_dest,
10461066
size,
10471067
report,
@@ -1054,7 +1074,7 @@ async def copy_to_local(
10541074
functools.partial(
10551075
self._copy_single_local_file,
10561076
xfer_sema,
1057-
await file.url(),
1077+
filename,
10581078
target_dest,
10591079
size,
10601080
report,

0 commit comments

Comments
 (0)