From 4c99a735e19b2c15cb374cf126ce93a74a086553 Mon Sep 17 00:00:00 2001 From: goutamadwant Date: Mon, 28 Sep 2026 21:51:01 -0700 Subject: [PATCH] Validate FileBasedSink num_shards type --- sdks/python/apache_beam/io/filebasedsink.py | 6 ++++-- sdks/python/apache_beam/io/filebasedsink_test.py | 11 +++++++++++ 2 files changed, 15 insertions(+), 2 deletions(-) diff --git a/sdks/python/apache_beam/io/filebasedsink.py b/sdks/python/apache_beam/io/filebasedsink.py index 8e0b39e1ac38..e9bc8e92a12e 100644 --- a/sdks/python/apache_beam/io/filebasedsink.py +++ b/sdks/python/apache_beam/io/filebasedsink.py @@ -81,8 +81,8 @@ def __init__( Raises: TypeError: if file path parameters are not a :class:`str` or :class:`~apache_beam.options.value_provider.ValueProvider`, or if - **compression_type** is not member of - :class:`~apache_beam.io.filesystem.CompressionTypes`. + **num_shards** is not an :class:`int`, or if **compression_type** is not + member of :class:`~apache_beam.io.filesystem.CompressionTypes`. ValueError: if **shard_name_template** is not of expected format. """ @@ -94,6 +94,8 @@ def __init__( raise TypeError( 'file_name_suffix must be a string or ValueProvider;' 'got %r instead' % file_name_suffix) + if not isinstance(num_shards, int): + raise TypeError('num_shards must be an int; got %r instead' % num_shards) if not CompressionTypes.is_valid_compression_type(compression_type): raise TypeError( diff --git a/sdks/python/apache_beam/io/filebasedsink_test.py b/sdks/python/apache_beam/io/filebasedsink_test.py index 121bc479200f..0372b16c14fc 100644 --- a/sdks/python/apache_beam/io/filebasedsink_test.py +++ b/sdks/python/apache_beam/io/filebasedsink_test.py @@ -155,6 +155,17 @@ def test_file_sink_display_data(self): ] hc.assert_that(dd.items, hc.contains_inanyorder(*expected_items)) + def test_file_sink_rejects_as_singleton_num_shards(self): + pipeline = beam.Pipeline() + count = (pipeline | beam.Create([1, 2]) | beam.combiners.Count.Globally()) + num_shards = beam.pvalue.AsSingleton(count) + + with self.assertRaisesRegex(TypeError, 'num_shards must be an int'): + MyFileBasedSink( + os.path.join(self._new_tempdir(), 'output'), + coder=coders.ToBytesCoder(), + num_shards=num_shards) + def test_empty_write(self): temp_path = tempfile.NamedTemporaryFile().name sink = MyFileBasedSink(