diff --git a/sdks/python/apache_beam/io/gcp/bigquery_file_loads.py b/sdks/python/apache_beam/io/gcp/bigquery_file_loads.py index 4e45d0324ee2..2f396a591f55 100644 --- a/sdks/python/apache_beam/io/gcp/bigquery_file_loads.py +++ b/sdks/python/apache_beam/io/gcp/bigquery_file_loads.py @@ -82,6 +82,32 @@ _SLEEP_DURATION_BETWEEN_POLLS = 10 +def _has_partitioning_load_parameters(additional_parameters): + return ( + 'timePartitioning' in additional_parameters or + 'rangePartitioning' in additional_parameters) + + +def _add_destination_partitioning_load_parameters( + additional_parameters, destination_table): + if destination_table is None: + return additional_parameters + + additional_parameters = dict(additional_parameters) + time_partitioning = getattr(destination_table, 'timePartitioning', None) + range_partitioning = getattr(destination_table, 'rangePartitioning', None) + + if ('timePartitioning' not in additional_parameters and + isinstance(time_partitioning, bigquery_tools.bigquery.TimePartitioning)): + additional_parameters['timePartitioning'] = time_partitioning + + if ('rangePartitioning' not in additional_parameters and isinstance( + range_partitioning, bigquery_tools.bigquery.RangePartitioning)): + additional_parameters['rangePartitioning'] = range_partitioning + + return additional_parameters + + def _generate_job_name(job_name, job_type, step_name): return bigquery_tools.generate_bq_job_name( job_name=job_name, @@ -688,6 +714,7 @@ def start_bundle(self): self.bq_io_metadata = create_bigquery_io_metadata(self._step_name) self.pending_jobs = [] self.schema_cache = {} + self.destination_table_cache = {} def process( self, @@ -716,6 +743,7 @@ def process( additional_parameters = self.additional_bq_parameters.get() else: additional_parameters = self.additional_bq_parameters + additional_parameters = dict(additional_parameters or {}) table_reference = bigquery_tools.parse_table_reference(destination) if table_reference.projectId is None: @@ -735,28 +763,55 @@ def process( create_disposition = self.create_disposition if self.temporary_tables: + destination_table = None + hashed_dest = bigquery_tools.get_hashable_destination(table_reference) + need_schema = schema is None and hashed_dest not in self.schema_cache + need_partitioning = not _has_partitioning_load_parameters( + additional_parameters) + if need_schema or need_partitioning: + try: + if hashed_dest in self.destination_table_cache: + destination_table = self.destination_table_cache[hashed_dest] + else: + destination_table = self.bq_wrapper.get_table( + project_id=table_reference.projectId, + dataset_id=table_reference.datasetId, + table_id=table_reference.tableId) + self.destination_table_cache[hashed_dest] = destination_table + except Exception as e: + if need_schema: + _LOGGER.warning( + "Input schema is absent and could not fetch the final " + "destination table's schema [%s]. Creating temp table [%s] " + "will likely fail: %s", + hashed_dest, + job_name, + e) + destination_table = None + # we need to create temp tables, so we need a schema. # if there is no input schema, fetch the destination table's schema if schema is None: - hashed_dest = bigquery_tools.get_hashable_destination(table_reference) if hashed_dest in self.schema_cache: schema = self.schema_cache[hashed_dest] - else: - try: - schema = bigquery_tools.table_schema_to_dict( - bigquery_tools.BigQueryWrapper().get_table( - project_id=table_reference.projectId, - dataset_id=table_reference.datasetId, - table_id=table_reference.tableId).schema) + elif destination_table is not None: + destination_schema = getattr(destination_table, 'schema', None) + if isinstance(destination_schema, + bigquery_tools.bigquery.TableSchema): + schema = bigquery_tools.table_schema_to_dict(destination_schema) self.schema_cache[hashed_dest] = schema - except Exception as e: + else: _LOGGER.warning( - "Input schema is absent and could not fetch the final " - "destination table's schema [%s]. Creating temp table [%s] " - "will likely fail: %s", + "Input schema is absent and the final destination table [%s] " + "does not have a usable schema. Creating temp table [%s] will " + "likely fail.", hashed_dest, - job_name, - e) + job_name) + + if (destination_table is not None and + not _has_partitioning_load_parameters(additional_parameters)): + additional_parameters = _add_destination_partitioning_load_parameters( + additional_parameters, destination_table) # If we are using temporary tables, then we must always create the # temporary tables, so we replace the create_disposition. diff --git a/sdks/python/apache_beam/io/gcp/bigquery_file_loads_test.py b/sdks/python/apache_beam/io/gcp/bigquery_file_loads_test.py index 191719e6a208..fda4a3e9d520 100644 --- a/sdks/python/apache_beam/io/gcp/bigquery_file_loads_test.py +++ b/sdks/python/apache_beam/io/gcp/bigquery_file_loads_test.py @@ -703,6 +703,157 @@ def test_one_load_job_failed_after_waiting(self, sleep_mock): sleep_mock.assert_called_once() + def test_temporary_table_load_inherits_destination_time_partitioning(self): + destination = 'project1:dataset1.table1' + partition = (destination, (0, ['gs://bucket/file1'])) + job_reference = bigquery_api.JobReference( + projectId='project1', jobId='job_name1') + destination_table = bigquery_api.Table( + timePartitioning=bigquery_api.TimePartitioning(type='DAY')) + + dofn = bqfl.TriggerLoadJobs( + schema=_ELEMENTS_SCHEMA, test_client=mock.Mock(), temporary_tables=True) + dofn.start_bundle() + dofn.bq_wrapper.get_table = mock.Mock(return_value=destination_table) + dofn.bq_wrapper.perform_load_job = mock.Mock(return_value=job_reference) + + list(dofn.process(partition, 'test_job', pane_info=mock.Mock(index=0))) + + load_call = dofn.bq_wrapper.perform_load_job.call_args.kwargs + self.assertEqual( + load_call['additional_load_parameters']['timePartitioning'], + destination_table.timePartitioning) + dofn.bq_wrapper.get_table.assert_called_once_with( + project_id='project1', dataset_id='dataset1', table_id='table1') + + def test_temporary_table_load_inherits_destination_range_partitioning(self): + destination = 'project1:dataset1.table1' + partition = (destination, (0, ['gs://bucket/file1'])) + job_reference = bigquery_api.JobReference( + projectId='project1', jobId='job_name1') + destination_table = bigquery_api.Table( + rangePartitioning=bigquery_api.RangePartitioning()) + + dofn = bqfl.TriggerLoadJobs( + schema=_ELEMENTS_SCHEMA, test_client=mock.Mock(), temporary_tables=True) + dofn.start_bundle() + dofn.bq_wrapper.get_table = mock.Mock(return_value=destination_table) + dofn.bq_wrapper.perform_load_job = mock.Mock(return_value=job_reference) + + list(dofn.process(partition, 'test_job', pane_info=mock.Mock(index=0))) + + load_call = dofn.bq_wrapper.perform_load_job.call_args.kwargs + self.assertEqual( + load_call['additional_load_parameters']['rangePartitioning'], + destination_table.rangePartitioning) + dofn.bq_wrapper.get_table.assert_called_once_with( + project_id='project1', dataset_id='dataset1', table_id='table1') + + def test_temporary_table_load_keeps_explicit_partitioning_parameters(self): + destination = 'project1:dataset1.table1' + partition = (destination, (0, ['gs://bucket/file1'])) + explicit_partitioning = {'timePartitioning': {'type': 'DAY'}} + job_reference = bigquery_api.JobReference( + projectId='project1', jobId='job_name1') + + dofn = bqfl.TriggerLoadJobs( + schema=_ELEMENTS_SCHEMA, + test_client=mock.Mock(), + temporary_tables=True, + additional_bq_parameters=explicit_partitioning) + dofn.start_bundle() + dofn.bq_wrapper.get_table = mock.Mock() + dofn.bq_wrapper.perform_load_job = mock.Mock(return_value=job_reference) + + list(dofn.process(partition, 'test_job', pane_info=mock.Mock(index=0))) + + load_call = dofn.bq_wrapper.perform_load_job.call_args.kwargs + self.assertEqual( + load_call['additional_load_parameters'], explicit_partitioning) + dofn.bq_wrapper.get_table.assert_not_called() + + def test_temporary_table_load_uses_cached_schema_with_explicit_partitioning( + self): + destination = 'project1:dataset1.table1' + partition = (destination, (0, ['gs://bucket/file1'])) + explicit_partitioning = {'timePartitioning': {'type': 'DAY'}} + job_reference = bigquery_api.JobReference( + projectId='project1', jobId='job_name1') + table_reference = bigquery_tools.parse_table_reference(destination) + hashed_dest = bigquery_tools.get_hashable_destination(table_reference) + + dofn = bqfl.TriggerLoadJobs( + schema=None, + test_client=mock.Mock(), + temporary_tables=True, + additional_bq_parameters=explicit_partitioning) + dofn.start_bundle() + dofn.schema_cache[hashed_dest] = _ELEMENTS_SCHEMA + dofn.bq_wrapper.get_table = mock.Mock() + dofn.bq_wrapper.perform_load_job = mock.Mock(return_value=job_reference) + + list(dofn.process(partition, 'test_job', pane_info=mock.Mock(index=0))) + + load_call = dofn.bq_wrapper.perform_load_job.call_args.kwargs + self.assertEqual(load_call['schema'], _ELEMENTS_SCHEMA) + self.assertEqual( + load_call['additional_load_parameters'], explicit_partitioning) + dofn.bq_wrapper.get_table.assert_not_called() + + def test_temporary_table_load_caches_destination_table_per_bundle(self): + destination = 'project1:dataset1.table1' + first_partition = (destination, (0, ['gs://bucket/file1'])) + second_partition = (destination, (1, ['gs://bucket/file2'])) + job_reference = bigquery_api.JobReference( + projectId='project1', jobId='job_name1') + destination_table = bigquery_api.Table( + timePartitioning=bigquery_api.TimePartitioning(type='DAY')) + + dofn = bqfl.TriggerLoadJobs( + schema=_ELEMENTS_SCHEMA, test_client=mock.Mock(), temporary_tables=True) + dofn.start_bundle() + dofn.bq_wrapper.get_table = mock.Mock(return_value=destination_table) + dofn.bq_wrapper.perform_load_job = mock.Mock(return_value=job_reference) + + list( + dofn.process(first_partition, 'test_job', pane_info=mock.Mock(index=0))) + list( + dofn.process( + second_partition, 'test_job', pane_info=mock.Mock(index=1))) + + dofn.bq_wrapper.get_table.assert_called_once_with( + project_id='project1', dataset_id='dataset1', table_id='table1') + load_call = dofn.bq_wrapper.perform_load_job.call_args.kwargs + self.assertEqual( + load_call['additional_load_parameters']['timePartitioning'], + destination_table.timePartitioning) + + def test_temporary_table_load_ignores_invalid_mock_partitioning_metadata( + self): + destination = 'project1:dataset1.table1' + partition = (destination, (0, ['gs://bucket/file1'])) + job_reference = bigquery_api.JobReference( + projectId='project1', jobId='job_name1') + destination_table = mock.Mock() + destination_table.timePartitioning = mock.Mock() + destination_table.rangePartitioning = mock.Mock() + + dofn = bqfl.TriggerLoadJobs( + schema=_ELEMENTS_SCHEMA, test_client=mock.Mock(), temporary_tables=True) + dofn.start_bundle() + dofn.bq_wrapper.get_table = mock.Mock(return_value=destination_table) + dofn.bq_wrapper.perform_load_job = mock.Mock(return_value=job_reference) + + list(dofn.process(partition, 'test_job', pane_info=mock.Mock(index=0))) + + load_call = dofn.bq_wrapper.perform_load_job.call_args.kwargs + self.assertNotIn( + 'timePartitioning', load_call['additional_load_parameters']) + self.assertNotIn( + 'rangePartitioning', load_call['additional_load_parameters']) + dofn.bq_wrapper.get_table.assert_called_once_with( + project_id='project1', dataset_id='dataset1', table_id='table1') + def test_multiple_partition_files(self): destination = 'project1:dataset1.table1' diff --git a/sdks/python/apache_beam/io/gcp/bigquery_write_it_test.py b/sdks/python/apache_beam/io/gcp/bigquery_write_it_test.py index c694383dcf9b..194d5243fcaa 100644 --- a/sdks/python/apache_beam/io/gcp/bigquery_write_it_test.py +++ b/sdks/python/apache_beam/io/gcp/bigquery_write_it_test.py @@ -42,6 +42,7 @@ from apache_beam.io.gcp.bigquery_tools import BigQueryWrapper from apache_beam.io.gcp.bigquery_tools import FileFormat from apache_beam.io.gcp.internal.clients import bigquery +from apache_beam.io.gcp.tests.bigquery_matcher import BigQueryTableMatcher from apache_beam.io.gcp.tests.bigquery_matcher import BigqueryFullResultMatcher from apache_beam.testing.test_pipeline import TestPipeline from apache_beam.testing.util import assert_that @@ -87,7 +88,7 @@ def tearDown(self): self.dataset_id, self.project) - def create_table(self, table_name): + def create_table(self, table_name, time_partitioning=None): table_schema = bigquery.TableSchema() table_field = bigquery.TableFieldSchema() table_field.name = 'int64' @@ -112,6 +113,8 @@ def create_table(self, table_name): datasetId=self.dataset_id, tableId=table_name), schema=table_schema) + if time_partitioning is not None: + table.timePartitioning = time_partitioning request = bigquery.BigqueryTablesInsertRequest( projectId=self.project, datasetId=self.dataset_id, table=table) self.bigquery_client.client.tables.Insert(request) @@ -377,6 +380,72 @@ def test_big_query_write_without_schema(self): write_disposition=beam.io.BigQueryDisposition.WRITE_APPEND, temp_file_format=FileFormat.JSON)) + @pytest.mark.it_postcommit + @mock.patch( + "apache_beam.io.gcp.bigquery_file_loads._MAXIMUM_SOURCE_URIS", new=1) + @retry(reraise=True, stop=stop_after_attempt(3)) + def test_big_query_file_loads_existing_partitioned_table(self): + table_name = 'python_file_loads_partitioned_table' + self.create_table( + table_name, + time_partitioning=bigquery.TimePartitioning(field='date', type='DAY')) + table_id = '{}.{}'.format(self.dataset_id, table_name) + + input_data = [{ + 'int64': 1, 'bytes': b'abc', 'date': '2026-01-01', 'time': '00:00:00' + }, + { + 'int64': 2, + 'bytes': b'xyz', + 'date': '2026-01-02', + 'time': '12:34:56' + }] + for row in input_data: + row['bytes'] = base64.b64encode(row['bytes']) + + args = self.test_pipeline.get_full_options_as_args( + on_success_matcher=hc.all_of( + BigqueryFullResultMatcher( + project=self.project, + query=( + "SELECT int64, bytes, date, time FROM %s ORDER BY int64" % + table_id), + data=[ + ( + 1, + b'abc', + datetime.date(2026, 1, 1), + datetime.time(0, 0, 0), + ), + ( + 2, + b'xyz', + datetime.date(2026, 1, 2), + datetime.time(12, 34, 56), + ), + ]), + BigQueryTableMatcher( + project=self.project, + dataset=self.dataset_id, + table=table_name, + expected_properties={ + 'timePartitioning': { + 'field': 'date', + 'type': 'DAY', + } + }))) + + with beam.Pipeline(argv=args) as p: + # pylint: disable=expression-not-assigned + ( + p | 'create' >> beam.Create(input_data) + | 'write' >> beam.io.WriteToBigQuery( + table_id, + write_disposition=beam.io.BigQueryDisposition.WRITE_APPEND, + max_file_size=1, # bytes + method=beam.io.WriteToBigQuery.Method.FILE_LOADS, + temp_file_format=FileFormat.JSON)) + @pytest.mark.it_postcommit def test_big_query_write_insert_errors_reporting(self): """