diff --git a/pyiceberg/io/fileformat.py b/pyiceberg/io/fileformat.py index 337e698605..65c0cbeff9 100644 --- a/pyiceberg/io/fileformat.py +++ b/pyiceberg/io/fileformat.py @@ -21,7 +21,7 @@ from abc import ABC, abstractmethod from dataclasses import dataclass -from typing import TYPE_CHECKING, Any +from typing import TYPE_CHECKING, Any, Protocol, runtime_checkable from pyiceberg.io import OutputFile from pyiceberg.manifest import FileFormat @@ -143,18 +143,17 @@ def __exit__(self, exc_type: Any, exc_val: Any, exc_tb: Any) -> None: self._result = self.close() -class FileFormatModel(ABC): +@runtime_checkable +class FileFormatModel(Protocol): """Represents a file format's capabilities. Creates writers.""" @property - @abstractmethod def format(self) -> FileFormat: ... - @abstractmethod def file_extension(self) -> str: """Return file extension without dot, e.g. 'parquet', 'orc'.""" + ... - @abstractmethod def create_writer( self, output_file: OutputFile, @@ -162,9 +161,9 @@ def create_writer( properties: Properties, ) -> FileFormatWriter: ... - @abstractmethod def add_field_metadata(self, field: NestedField, metadata: dict[bytes, bytes], include_field_ids: bool) -> None: """Add format-specific Arrow field metadata.""" + ... class FileFormatFactory: diff --git a/tests/io/test_fileformat.py b/tests/io/test_fileformat.py index 328fee9274..d5d487fa6d 100644 --- a/tests/io/test_fileformat.py +++ b/tests/io/test_fileformat.py @@ -77,3 +77,36 @@ def close(self) -> DataFileStatistics: writer = _DummyWriter() with pytest.raises(RuntimeError, match="Writer has not been closed yet"): writer.result() + + +class _StructuralModel: + """Non-inheriting class that structurally conforms to the FileFormatModel Protocol.""" + + @property + def format(self) -> FileFormat: + return FileFormat.ORC + + def file_extension(self) -> str: + return "orc" + + def create_writer(self, output_file: Any, file_schema: Any, properties: Any) -> Any: + raise NotImplementedError + + def add_field_metadata(self, field: Any, metadata: Any, include_field_ids: bool) -> None: + pass + + +def test_file_format_model_is_protocol() -> None: + """A structurally-conforming class (no inheritance) passes isinstance() against FileFormatModel.""" + assert isinstance(_StructuralModel(), FileFormatModel) + + +def test_structural_model_works_with_factory() -> None: + """A structurally-conforming class (no inheritance) can be registered and retrieved via FileFormatFactory.""" + original = dict(FileFormatFactory._registry) + try: + model = _StructuralModel() + FileFormatFactory.register(model) + assert FileFormatFactory.get(FileFormat.ORC) is model + finally: + FileFormatFactory._registry = original