Skip to content

Commit cda094f

Browse files
committed
Add FileDownloadConfig annotation for FlyteFile inputs
Port new BlobType fields file_extension and enable_legacy_filename to flytekit. FlyteFile inputs can be annotated with the FileDownloadConfig annotation to configure the file extension to use during the copilot download phase. e.g. ```python def t1(file: Annotated[FlyteFile, FileDownloadConfig(file_extension="csv")]): ... # copilot downloads the file to e.g. /inputs/file.csv versus... def t1(file: FlyteFile["csv"]): ... # copilot downloads the file to e.g. /inputs/file ```
1 parent 17991b6 commit cda094f

5 files changed

Lines changed: 167 additions & 8 deletions

File tree

flytekit/core/type_engine.py

Lines changed: 48 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -10,6 +10,7 @@
1010
import json
1111
import mimetypes
1212
import os
13+
import re
1314
import sys
1415
import textwrap
1516
import threading
@@ -102,6 +103,53 @@ def get_batch_size(t: Type) -> Optional[int]:
102103
return None
103104

104105

106+
class FileDownloadConfig:
107+
"""
108+
This is used to annotate a FlyteFile when we want to download the file with a specific extension. For example,
109+
110+
```python
111+
# ContainerTask
112+
def t1(file: Annotated[FlyteFile, FileDownloadConfig(file_extension="csv")]):
113+
... # copilot downloads the file to e.g. /inputs/file.csv
114+
115+
versus...
116+
117+
def t1(file: FlyteFile["csv"]):
118+
... # copilot downloads the file to e.g. /inputs/file
119+
```
120+
121+
file_extension: (Default is "") The file extension (e.g. "csv", "parquet") to use during copilot download.
122+
enable_legacy_filename: (Default is False) When true and file_extension is non-empty, the copilot download phase
123+
writes the blob to both the full path (with extension) and the old path (without extension), preserving backward compatibility for
124+
workflows with tasks that may read from both.
125+
"""
126+
127+
def __init__(self, file_extension: str = "", enable_legacy_filename: bool = False):
128+
self._file_extension = file_extension
129+
self._enable_legacy_filename = enable_legacy_filename
130+
131+
if self._file_extension is not "":
132+
pattern = r"^[a-zA-Z0-9]+(\.[a-zA-Z0-9]+)*$"
133+
if not re.match(pattern, self._file_extension):
134+
raise ValueError(f"Invalid file extension: {self._file_extension}")
135+
136+
@property
137+
def file_extension(self) -> str:
138+
return self._file_extension
139+
140+
@property
141+
def enable_legacy_filename(self) -> bool:
142+
return self._enable_legacy_filename
143+
144+
145+
def get_file_download_config(t: Type) -> Optional[FileDownloadConfig]:
146+
if is_annotated(t):
147+
for arg in get_args(t):
148+
if isinstance(arg, FileDownloadConfig):
149+
return arg
150+
return None
151+
152+
105153
def modify_literal_uris(lit: Literal):
106154
"""
107155
Modifies the literal object recursively to replace the URIs with the native paths in case they are of

flytekit/models/core/types.py

Lines changed: 37 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -38,13 +38,19 @@ class BlobDimensionality(object):
3838
SINGLE = _types_pb2.BlobType.SINGLE
3939
MULTIPART = _types_pb2.BlobType.MULTIPART
4040

41-
def __init__(self, format, dimensionality):
41+
def __init__(self, format, dimensionality, file_extension="", enable_legacy_filename=False):
4242
"""
4343
:param Text format: A string describing the format of the underlying blob data.
4444
:param int dimensionality: An integer from BlobType.BlobDimensionality enum
45+
:param Text file_extension: The file extension (e.g. "csv", "parquet") to use
46+
during copilot download, e.g. "csv", "parquet". Empty by default.
47+
:param bool enable_legacy_filename: When True and file_extension is set, the copilot
48+
download phase writes the blob to both the extended path and the base path.
4549
"""
4650
self._format = format
4751
self._dimensionality = dimensionality
52+
self._file_extension = file_extension
53+
self._enable_legacy_filename = enable_legacy_filename
4854

4955
@property
5056
def format(self):
@@ -62,16 +68,44 @@ def dimensionality(self):
6268
"""
6369
return self._dimensionality
6470

71+
@property
72+
def file_extension(self):
73+
"""
74+
The file extension (e.g. "csv", "parquet") to use during copilot download.
75+
Default is "", which means no extension is appended.
76+
:rtype: Text
77+
"""
78+
return self._file_extension
79+
80+
@property
81+
def enable_legacy_filename(self):
82+
"""
83+
When True and file_extension is set, the copilot download writes the blob to
84+
both the full path (with extension) and the old path (without extension).
85+
:rtype: bool
86+
"""
87+
return self._enable_legacy_filename
88+
6589
def to_flyte_idl(self):
6690
"""
6791
:rtype: flyteidl.core.types_pb2.BlobType
6892
"""
69-
return _types_pb2.BlobType(format=self.format, dimensionality=self.dimensionality)
93+
return _types_pb2.BlobType(
94+
format=self.format,
95+
dimensionality=self.dimensionality,
96+
file_extension=self._file_extension,
97+
enable_legacy_filename=self._enable_legacy_filename,
98+
)
7099

71100
@classmethod
72101
def from_flyte_idl(cls, proto):
73102
"""
74103
:param flyteidl.core.types_pb2.BlobType proto:
75104
:rtype: BlobType
76105
"""
77-
return cls(format=proto.format, dimensionality=proto.dimensionality)
106+
return cls(
107+
format=proto.format,
108+
dimensionality=proto.dimensionality,
109+
file_extension=proto.file_extension,
110+
enable_legacy_filename=proto.enable_legacy_filename,
111+
)

flytekit/types/file/file.py

Lines changed: 31 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -25,6 +25,7 @@
2525
AsyncTypeTransformer,
2626
TypeEngine,
2727
TypeTransformerFailedError,
28+
get_file_download_config,
2829
get_underlying_type,
2930
)
3031
from flytekit.exceptions.user import FlyteAssertion
@@ -449,8 +450,26 @@ def get_format(t: typing.Union[typing.Type[FlyteFile], os.PathLike]) -> str:
449450
return ""
450451
return cast(FlyteFile, t).extension()
451452

452-
def _blob_type(self, format: str) -> BlobType:
453-
return BlobType(format=format, dimensionality=BlobType.BlobDimensionality.SINGLE)
453+
@staticmethod
454+
def get_file_extension(t: typing.Union[typing.Type[FlyteFile], os.PathLike]) -> str:
455+
if t is os.PathLike:
456+
return ""
457+
file_download_config = get_file_download_config(t)
458+
if file_download_config is None:
459+
return ""
460+
return file_download_config.file_extension or ""
461+
462+
@staticmethod
463+
def get_enable_legacy_filename(t: typing.Union[typing.Type[FlyteFile], os.PathLike]) -> str:
464+
if t is os.PathLike:
465+
return False
466+
file_download_config = get_file_download_config(t)
467+
if file_download_config is None:
468+
return False
469+
return file_download_config.enable_legacy_filename or False
470+
471+
def _blob_type(self, format: str, file_extension: str = "", enable_legacy_filename: bool = False) -> BlobType:
472+
return BlobType(format=format, dimensionality=BlobType.BlobDimensionality.SINGLE, file_extension=file_extension, enable_legacy_filename=enable_legacy_filename)
454473

455474
def assert_type(
456475
self, t: typing.Union[typing.Type[FlyteFile], os.PathLike], v: typing.Union[FlyteFile, os.PathLike, str]
@@ -463,7 +482,11 @@ def assert_type(
463482
)
464483

465484
def get_literal_type(self, t: typing.Union[typing.Type[FlyteFile], os.PathLike]) -> LiteralType:
466-
return LiteralType(blob=self._blob_type(format=FlyteFilePathTransformer.get_format(t)))
485+
return LiteralType(blob=self._blob_type(
486+
format=FlyteFilePathTransformer.get_format(t),
487+
file_extension=FlyteFilePathTransformer.get_file_extension(t),
488+
enable_legacy_filename=FlyteFilePathTransformer.get_enable_legacy_filename(t),
489+
))
467490

468491
def get_mime_type_from_extension(self, extension: str) -> typing.Union[str, typing.Sequence[str]]:
469492
extension_to_mime_type = {
@@ -537,7 +560,11 @@ async def async_to_literal(
537560
raise ValueError(f"Incorrect type {python_type}, must be either a FlyteFile or os.PathLike")
538561

539562
# information used by all cases
540-
meta = BlobMetadata(type=self._blob_type(format=FlyteFilePathTransformer.get_format(python_type)))
563+
meta = BlobMetadata(type=self._blob_type(
564+
format=FlyteFilePathTransformer.get_format(python_type),
565+
file_extension=FlyteFilePathTransformer.get_file_extension(python_type),
566+
enable_legacy_filename=FlyteFilePathTransformer.get_enable_legacy_filename(python_type),
567+
))
541568

542569
if isinstance(python_val, FlyteFile):
543570
# Cast the source path to str type to avoid error raised when the source path is used as the blob uri,

tests/flytekit/unit/core/test_flyte_file.py

Lines changed: 29 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -17,7 +17,7 @@
1717
from flytekit.core.hash import HashMethod
1818
from flytekit.core.launch_plan import LaunchPlan
1919
from flytekit.core.task import task
20-
from flytekit.core.type_engine import TypeEngine
20+
from flytekit.core.type_engine import FileDownloadConfig, TypeEngine
2121
from flytekit.core.workflow import workflow
2222
from flytekit.models.core.types import BlobType
2323
from flytekit.models.literals import LiteralMap, Blob, BlobMetadata
@@ -764,6 +764,34 @@ def test_headers():
764764
assert len(FlyteFilePathTransformer.get_additional_headers(".gz")) == 1
765765

766766

767+
def test_transform_flytefile_with_file_download_config():
768+
csv_file_no_config = FlyteFile["csv"]
769+
lt = FlyteFilePathTransformer().get_literal_type(csv_file_no_config)
770+
assert lt.blob.file_extension == ""
771+
assert lt.blob.enable_legacy_filename == False
772+
773+
legacy_file = Annotated[FlyteFile["csv"], FileDownloadConfig(file_extension="csv", enable_legacy_filename=True)]
774+
lt = FlyteFilePathTransformer().get_literal_type(legacy_file)
775+
assert lt.blob.file_extension == "csv"
776+
assert lt.blob.enable_legacy_filename == True
777+
778+
779+
def test_file_download_config_valid_compound_extension():
780+
config = FileDownloadConfig(file_extension="tar.gz")
781+
assert config.file_extension == "tar.gz"
782+
783+
784+
@pytest.mark.parametrize("bad_ext", [
785+
".csv",
786+
"my file",
787+
"../../escape",
788+
"csv!",
789+
])
790+
def test_file_download_config_rejects_invalid_extensions(bad_ext):
791+
with pytest.raises(ValueError, match="Invalid file extension"):
792+
FileDownloadConfig(file_extension=bad_ext)
793+
794+
767795
def test_new_remote_file():
768796
nf = FlyteFile.new_remote_file(name="foo.txt")
769797
assert isinstance(nf, FlyteFile)

tests/flytekit/unit/models/core/test_types.py

Lines changed: 22 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -15,11 +15,33 @@ def test_blob_type():
1515
)
1616
assert o.format == "csv"
1717
assert o.dimensionality == _types.BlobType.BlobDimensionality.SINGLE
18+
assert o.file_extension == ""
19+
assert o.enable_legacy_filename == False
1820

1921
o2 = _types.BlobType.from_flyte_idl(o.to_flyte_idl())
2022
assert o == o2
2123
assert o2.format == "csv"
2224
assert o2.dimensionality == _types.BlobType.BlobDimensionality.SINGLE
25+
assert o2.file_extension == ""
26+
assert o2.enable_legacy_filename == False
27+
28+
o = _types.BlobType(
29+
format="csv",
30+
dimensionality=_types.BlobType.BlobDimensionality.SINGLE,
31+
file_extension="csv",
32+
enable_legacy_filename=True,
33+
)
34+
assert o.format == "csv"
35+
assert o.dimensionality == _types.BlobType.BlobDimensionality.SINGLE
36+
assert o.file_extension == "csv"
37+
assert o.enable_legacy_filename == True
38+
39+
o2 = _types.BlobType.from_flyte_idl(o.to_flyte_idl())
40+
assert o == o2
41+
assert o2.format == "csv"
42+
assert o2.dimensionality == _types.BlobType.BlobDimensionality.SINGLE
43+
assert o2.file_extension == "csv"
44+
assert o2.enable_legacy_filename == True
2345

2446

2547
def test_enum_type():

0 commit comments

Comments
 (0)