Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
30 changes: 22 additions & 8 deletions sdks/python/apache_beam/runners/pipeline_utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -93,16 +93,30 @@ def validate_transform(transform_id):
if output_coder.spec.urn != common_urns.coders.KV.urn:
raise ValueError(
"Bad coder for output of %s: %s" % (transform_id, output_coder))
input_key_coder_id = input_coder.component_coder_ids[0]
output_key_coder_id = output_coder.component_coder_ids[0]
if input_key_coder_id != output_key_coder_id:
raise ValueError(
"Input key coder %s does not match output key coder %s for "
"transform %s" %
(input_key_coder_id, output_key_coder_id, transform_id))
output_values_coder_id = output_coder.component_coder_ids[1]
output_values_coder = pipeline_proto.components.coders[
output_coder.component_coder_ids[1]]
if (input_coder.component_coder_ids[0]
!= output_coder.component_coder_ids[0] or
output_values_coder.spec.urn != common_urns.coders.ITERABLE.urn or
output_values_coder.component_coder_ids[0]
!= input_coder.component_coder_ids[1]):
output_values_coder_id]
if output_values_coder.spec.urn != common_urns.coders.ITERABLE.urn:
raise ValueError(
"Output value coder %s for transform %s must be an iterable "
"coder, but uses URN %s" % (
output_values_coder_id,
transform_id,
output_values_coder.spec.urn))
input_value_coder_id = input_coder.component_coder_ids[1]
output_value_coder_id = output_values_coder.component_coder_ids[0]
if output_value_coder_id != input_value_coder_id:
raise ValueError(
"Incompatible input coder %s and output coder %s for transform %s" %
(transform_id, input_coder, output_coder))
"Input value coder %s does not match output value coder %s for "
"transform %s" %
(input_value_coder_id, output_value_coder_id, transform_id))
elif transform_proto.spec.urn == common_urns.primitives.ASSIGN_WINDOWS.urn:
if not transform_proto.inputs:
raise ValueError("Missing input for transform: %s" % transform_proto)
Expand Down
93 changes: 93 additions & 0 deletions sdks/python/apache_beam/runners/pipeline_utils_test.py
Original file line number Diff line number Diff line change
Expand Up @@ -28,10 +28,103 @@
from apache_beam.portability.api import beam_runner_api_pb2
from apache_beam.runners.pipeline_utils import merge_common_environments
from apache_beam.runners.pipeline_utils import merge_superset_dep_environments
from apache_beam.runners.pipeline_utils import validate_pipeline_graph
from apache_beam.runners.portability.expansion_service_test import FibTransform


class PipelineUtilitiesTest(unittest.TestCase):
@staticmethod
def _pipeline_with_gbk_coders(
input_key_coder='key_coder',
output_key_coder='key_coder',
output_values_coder_urn=common_urns.coders.ITERABLE.urn,
input_value_coder='input_value_coder',
output_value_coder='input_value_coder'):
leaf_coders = {
coder_id: beam_runner_api_pb2.Coder(
spec=beam_runner_api_pb2.FunctionSpec(
urn=common_urns.coders.BYTES.urn))
for coder_id in {
input_key_coder,
output_key_coder,
input_value_coder,
output_value_coder, }
}
output_values_coder = beam_runner_api_pb2.Coder(
spec=beam_runner_api_pb2.FunctionSpec(urn=output_values_coder_urn),
component_coder_ids=([output_value_coder] if output_values_coder_urn
== common_urns.coders.ITERABLE.urn else []))
return beam_runner_api_pb2.Pipeline(
components=beam_runner_api_pb2.Components(
coders={
**leaf_coders,
'input_coder': beam_runner_api_pb2.Coder(
spec=beam_runner_api_pb2.FunctionSpec(
urn=common_urns.coders.KV.urn),
component_coder_ids=[input_key_coder, input_value_coder]),
'output_coder': beam_runner_api_pb2.Coder(
spec=beam_runner_api_pb2.FunctionSpec(
urn=common_urns.coders.KV.urn),
component_coder_ids=[
output_key_coder, 'output_values_coder'
]),
'output_values_coder': output_values_coder,
},
pcollections={
'input': beam_runner_api_pb2.PCollection(
coder_id='input_coder'),
'output': beam_runner_api_pb2.PCollection(
coder_id='output_coder'),
},
transforms={
'gbk': beam_runner_api_pb2.PTransform(
spec=beam_runner_api_pb2.FunctionSpec(
urn=common_urns.primitives.GROUP_BY_KEY.urn),
inputs={'input': 'input'},
outputs={'output': 'output'}),
}),
root_transform_ids=['gbk'])

def test_validate_pipeline_graph_accepts_valid_gbk_coders(self):
validate_pipeline_graph(self._pipeline_with_gbk_coders())

def test_validate_pipeline_graph_reports_gbk_key_coder_mismatch(self):
pipeline = self._pipeline_with_gbk_coders(
input_key_coder='input_key_coder', output_key_coder='output_key_coder')

with self.assertRaises(ValueError) as error:
validate_pipeline_graph(pipeline)

self.assertEqual(
str(error.exception),
'Input key coder input_key_coder does not match output key coder '
'output_key_coder for transform gbk')

def test_validate_pipeline_graph_reports_non_iterable_gbk_output(self):
pipeline = self._pipeline_with_gbk_coders(
output_values_coder_urn=common_urns.coders.BYTES.urn)

with self.assertRaises(ValueError) as error:
validate_pipeline_graph(pipeline)

self.assertEqual(
str(error.exception),
'Output value coder output_values_coder for transform gbk must be an '
'iterable coder, but uses URN beam:coder:bytes:v1')

def test_validate_pipeline_graph_reports_gbk_value_coder_mismatch(self):
pipeline = self._pipeline_with_gbk_coders(
input_value_coder='input_value_coder',
output_value_coder='output_value_coder')

with self.assertRaises(ValueError) as error:
validate_pipeline_graph(pipeline)

self.assertEqual(
str(error.exception),
'Input value coder input_value_coder does not match output value coder '
'output_value_coder for transform gbk')

def test_equal_environments_merged(self):
pipeline_proto = merge_common_environments(
beam_runner_api_pb2.Pipeline(
Expand Down
Loading