diff --git a/streaming/base/storage/upload.py b/streaming/base/storage/upload.py index 57f678948..3047cd827 100644 --- a/streaming/base/storage/upload.py +++ b/streaming/base/storage/upload.py @@ -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 @@ -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) @@ -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]], @@ -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 diff --git a/tests/test_upload.py b/tests/test_upload.py index d3b2997bf..4c0488102 100644 --- a/tests/test_upload.py +++ b/tests/test_upload.py @@ -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') @@ -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.*'):