Skip to content
Draft
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
Original file line number Diff line number Diff line change
Expand Up @@ -19,6 +19,8 @@

import static org.apache.beam.sdk.values.Row.toRow;

import com.fasterxml.jackson.core.JsonProcessingException;
import com.fasterxml.jackson.databind.ObjectMapper;
import java.io.InputStream;
import java.math.BigDecimal;
import java.util.ArrayList;
Expand All @@ -40,6 +42,8 @@
import org.yaml.snakeyaml.error.YAMLException;

public class YamlUtils {
private static final ObjectMapper OBJECT_MAPPER = new ObjectMapper();

private static final Map<Schema.TypeName, Function<String, @Nullable Object>> YAML_VALUE_PARSERS =
ImmutableMap
.<Schema.TypeName,
Expand Down Expand Up @@ -116,6 +120,9 @@ public static Row toBeamRow(
}

if (yamlValue instanceof List) {
if (fieldType.getTypeName() == Schema.TypeName.STRING) {
return toJsonString(field, yamlValue);
}
FieldType innerType =
Preconditions.checkNotNull(
fieldType.getCollectionElementType(),
Expand All @@ -142,6 +149,8 @@ public static Row toBeamRow(
return toBeamRow((Map<String, Object>) yamlValue, nestedSchema, convertNamesToCamelCase);
} else if (fieldType.getTypeName() == Schema.TypeName.MAP) {
return yamlValue;
} else if (fieldType.getTypeName() == Schema.TypeName.STRING) {
return toJsonString(field, yamlValue);
}
}

Expand Down Expand Up @@ -198,6 +207,18 @@ public static String yamlStringFromMap(@Nullable Map<String, Object> map) {
}
}

private static String toJsonString(Field field, Object yamlValue) {
try {
return OBJECT_MAPPER.writeValueAsString(yamlValue);
} catch (JsonProcessingException e) {
throw new IllegalArgumentException(
String.format(
"Failed to serialize YAML %s to JSON string for field '%s': %s",
yamlValue instanceof List ? "list" : "map", field.getName(), e.getMessage()),
e);
}
}

private static List<String> findNonSerializableKeys(Map<String, Object> map) {
List<String> problematicKeys = new ArrayList<>();
Yaml yaml = new Yaml();
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -288,4 +288,38 @@ public void testYamlStringFromMapWithValidMap() {
org.junit.Assert.assertNotNull(yaml);
org.junit.Assert.assertTrue(yaml.contains("string_key"));
}

@Test
public void testMapAndListToStringField() {
Schema schema =
Schema.builder().addStringField("map_field").addStringField("list_field").build();

String yaml =
"map_field:\n"
+ " key: value\n"
+ " nested:\n"
+ " foo: 123\n"
+ "list_field:\n"
+ " - a\n"
+ " - b\n";

Row row = YamlUtils.toBeamRow(yaml, schema);
assertEquals("{\"key\":\"value\",\"nested\":{\"foo\":123}}", row.getString("map_field"));
assertEquals("[\"a\",\"b\"]", row.getString("list_field"));
}

@Test
public void testMapAndListToStringFieldWithCamelCase() {
Schema schema = Schema.builder().addStringField("mapField").addStringField("listField").build();

Map<String, Object> map = new java.util.HashMap<>();
Map<String, Object> nestedMap = new java.util.HashMap<>();
nestedMap.put("key", "value");
map.put("map_field", nestedMap);
map.put("list_field", Arrays.asList("a", "b"));

Row row = YamlUtils.toBeamRow(map, schema, true);
assertEquals("{\"key\":\"value\"}", row.getString("mapField"));
assertEquals("[\"a\",\"b\"]", row.getString("listField"));
}
}
Original file line number Diff line number Diff line change
Expand Up @@ -361,7 +361,24 @@ public void testBuildTransformWithManaged() {
+ "schema: '"
+ PROTO_SCHEMA
+ "'\n"
+ "message_name: MyMessage");
+ "message_name: MyMessage",
"topic: topic_6\n"
+ "bootstrap_servers: some bootstrap\n"
+ "format: AVRO\n"
+ "schema:\n"
+ " type: record\n"
+ " name: my_record\n"
+ " fields:\n"
+ " - name: bool\n"
+ " type: boolean",
"topic: topic_7\n"
+ "bootstrap_servers: some bootstrap\n"
+ "format: JSON\n"
+ "schema:\n"
+ " type: object\n"
+ " properties:\n"
+ " name:\n"
+ " type: string");

for (String config : configs) {
// Kafka Read SchemaTransform gets built in ManagedSchemaTransformProvider's expand
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -268,7 +268,16 @@ public void testBuildTransformWithManaged() {
+ "schema: '"
+ PROTO_SCHEMA
+ "'\n"
+ "message_name: MyMessage");
+ "message_name: MyMessage",
"topic: topic_4\n"
+ "bootstrap_servers: some bootstrap\n"
+ "format: AVRO\n"
+ "schema:\n"
+ " type: record\n"
+ " name: my_record\n"
+ " fields:\n"
+ " - name: str\n"
+ " type: string");

for (String config : configs) {
// Kafka Write SchemaTransform gets built in ManagedSchemaTransformProvider's expand
Expand Down
5 changes: 5 additions & 0 deletions sdks/python/apache_beam/transforms/external.py
Original file line number Diff line number Diff line change
Expand Up @@ -22,6 +22,7 @@
import copy
import functools
import glob
import json
import logging
import re
import subprocess
Expand Down Expand Up @@ -256,6 +257,10 @@ def dict_to_row_recursive(field_type, py_value):
key: dict_to_row_recursive(field_type.map_type.value_type, value)
for key, value in py_value.items()
}
elif (type_info == 'atomic_type' and
field_type.atomic_type == schema_pb2.STRING and
isinstance(py_value, (dict, list))):
return json.dumps(py_value)
else:
return py_value

Expand Down
27 changes: 27 additions & 0 deletions sdks/python/apache_beam/transforms/external_test.py
Original file line number Diff line number Diff line change
Expand Up @@ -45,6 +45,7 @@
from apache_beam.transforms import external
from apache_beam.transforms.external import MANAGED_SCHEMA_TRANSFORM_IDENTIFIER
from apache_beam.transforms.external import AnnotationBasedPayloadBuilder
from apache_beam.transforms.external import ExplicitSchemaTransformPayloadBuilder
from apache_beam.transforms.external import ImplicitSchemaPayloadBuilder
from apache_beam.transforms.external import JavaClassLookupPayloadBuilder
from apache_beam.transforms.external import JavaExternalTransform
Expand Down Expand Up @@ -524,6 +525,32 @@ def test_build_payload(self):
self.assertEqual('bbb', schema_transform_config.object_field.str_sub_field)
self.assertEqual(456, schema_transform_config.object_field.int_sub_field)

def test_explicit_payload_builder_with_dict_and_list_to_string_field(self):
schema = schema_pb2.Schema(
fields=[
schema_pb2.Field(
name='str_field',
type=schema_pb2.FieldType(atomic_type=schema_pb2.STRING)),
schema_pb2.Field(
name='list_field',
type=schema_pb2.FieldType(atomic_type=schema_pb2.STRING)),
])

payload_builder = ExplicitSchemaTransformPayloadBuilder(
identifier='dummy_id',
schema_proto=schema,
str_field={'foo': 'bar'},
list_field=[1, 2, 'baz'])
payload = payload_builder.build()

self.assertEqual('dummy_id', payload.identifier)

coder = RowCoder(payload.configuration_schema)
config = coder.decode(payload.configuration_row)

self.assertEqual('{"foo": "bar"}', config.str_field)
self.assertEqual('[1, 2, "baz"]', config.list_field)


class SchemaAwareExternalTransformTest(unittest.TestCase):
class MockDiscoveryService:
Expand Down
44 changes: 44 additions & 0 deletions sdks/python/apache_beam/yaml/yaml_io_test.py
Original file line number Diff line number Diff line change
Expand Up @@ -983,6 +983,50 @@ def test_dicom_search_without_error_handling_raises(self):
'''))


class YamlKafkaTest(unittest.TestCase):
def test_read_from_kafka_json_schema_expansion(self):
# Regression test for https://github.com/apache/beam/issues/35186.
# Verifies that ReadFromKafka expands with a nested JSON schema map without
# UnsupportedOperationException.
p = beam.Pipeline(
options=beam.options.pipeline_options.PipelineOptions(
pickle_library='cloudpickle'))
_ = p | YamlTransform(
'''
type: ReadFromKafka
config:
topic: my-topic
bootstrap_servers: kafka:9092
format: JSON
schema:
type: object
properties:
name:
type: string
''')

def test_read_from_kafka_avro_schema_expansion(self):
# Verifies that ReadFromKafka expands with a nested AVRO schema map without
# UnsupportedOperationException.
p = beam.Pipeline(
options=beam.options.pipeline_options.PipelineOptions(
pickle_library='cloudpickle'))
_ = p | YamlTransform(
'''
type: ReadFromKafka
config:
topic: my-topic
bootstrap_servers: kafka:9092
format: AVRO
schema:
type: record
name: my_record
fields:
- name: bool
type: boolean
''')


if __name__ == '__main__':
logging.getLogger().setLevel(logging.INFO)
unittest.main()
Loading