Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
40 changes: 27 additions & 13 deletions streaming/base/storage/upload.py
Original file line number Diff line number Diff line change
Expand Up @@ -48,6 +48,26 @@
}


def _provider_prefix(uri: str) -> str:
"""Return the cloud provider scheme for ``uri``, or ``''`` for local paths.

``urllib.parse.urlparse`` treats Windows drive letters as URL schemes
(``D:/path`` → ``scheme='d'``). Those are local filesystem paths, not cloud
providers, so map single alphabetic schemes to the local uploader.
"""
obj = urllib.parse.urlparse(uri)
scheme = obj.scheme
if len(scheme) == 1 and scheme.isalpha():
return ''
if scheme == 'dbfs':
path = pathlib.Path(uri)
if len(path.parts) >= 2:
prefix = os.path.join(path.parts[0], path.parts[1])
if prefix == 'dbfs:/Volumes':
return prefix
return scheme


class GCSAuthentication(Enum):
HMAC = 1
SERVICE_ACCOUNT = 2
Expand Down Expand Up @@ -86,13 +106,7 @@ def get(cls,
CloudUploader: An instance of sub-class.
"""
cls._validate(cls, out)
obj = urllib.parse.urlparse(out) if isinstance(out, str) else urllib.parse.urlparse(out[1])
provider_prefix = obj.scheme
if obj.scheme == 'dbfs':
path = pathlib.Path(out) if isinstance(out, str) else pathlib.Path(out[1])
prefix = os.path.join(path.parts[0], path.parts[1])
if prefix == 'dbfs:/Volumes':
provider_prefix = prefix
provider_prefix = _provider_prefix(out if isinstance(out, str) else out[1])
return getattr(sys.modules[__name__],
UPLOADERS[provider_prefix])(out, keep_local, progress_bar, retry, exist_ok)

Expand All @@ -114,15 +128,15 @@ def _validate(self, out: Union[str, tuple[str, str]]) -> None:
ValueError: Invalid Cloud provider prefix.
"""
if isinstance(out, str):
obj = urllib.parse.urlparse(out)
provider_prefix = _provider_prefix(out)
else:
if len(out) != 2:
raise ValueError(f'Invalid `out` argument. It is either a string of ' +
f'local/remote directory or a list of two strings with ' +
f'[local, remote].')
obj = urllib.parse.urlparse(out[1])
if obj.scheme not in UPLOADERS:
raise ValueError(f'Invalid Cloud provider prefix: {obj.scheme}.')
provider_prefix = _provider_prefix(out[1])
if provider_prefix not in UPLOADERS:
raise ValueError(f'Invalid Cloud provider prefix: {provider_prefix}.')

def __init__(self,
out: Union[str, tuple[str, str]],
Expand Down Expand Up @@ -158,8 +172,8 @@ def __init__(self,
self.retry = retry

if isinstance(out, str):
# It is a remote directory
if urllib.parse.urlparse(out).scheme != '':
# It is a remote directory (cloud scheme). Windows drive letters are local.
if _provider_prefix(out) != '':
self.local = mkdtemp()
self.remote = out
# It is a local directory
Expand Down
17 changes: 17 additions & 0 deletions tests/test_upload.py
Original file line number Diff line number Diff line change
Expand Up @@ -49,6 +49,9 @@ class TestCloudUploader:
[None, 'gs://bucket/dir/file', GCSUploader],
['/tmp/dir/filepath', LocalUploader],
['./relative/dir/filepath', LocalUploader],
# Windows absolute paths parse as scheme='d' / 'c' — must stay local (#960).
['D:/datasets/train', LocalUploader],
['C:/Users/data/out', LocalUploader],
],
)
@pytest.mark.usefixtures('gcs_hmac_credentials')
Expand All @@ -70,6 +73,20 @@ def test_instantiation_type(
cw = CloudUploader.get(out_root)
assert isinstance(cw, mapping[-1])

def test_windows_drive_letter_not_cloud_scheme(self, tmp_path: Any):
"""MDSWriter-style absolute Windows paths must not raise Invalid Cloud provider (#960)."""
from streaming.base.storage.upload import _provider_prefix
assert _provider_prefix('D:/test') == ''
assert _provider_prefix('C:\\train\\out') == ''
assert _provider_prefix('s3://bucket/key') == 's3'
assert _provider_prefix('/unix/abs') == ''
# Instantiation uses a real temp dir so LocalUploader can mkdir.
win_style = str(tmp_path / 'win_out')
# Simulate urlparse drive-letter behavior without requiring Windows.
with patch('streaming.base.storage.upload._provider_prefix', return_value=''):
cw = CloudUploader.get(out=win_style)
assert isinstance(cw, LocalUploader)

@pytest.mark.parametrize('out', [(), ('s3://bucket/dir',), ('./dir1', './dir2', './dir3')])
def test_invalid_out_parameter_length(self, out: Any):
with pytest.raises(ValueError, match=f'Invalid `out` argument.*'):
Expand Down