Skip to content

Commit e3e0ea7

Browse files
committed
support generics in from row and to row conversions
1 parent 9fba823 commit e3e0ea7

14 files changed

Lines changed: 952 additions & 69 deletions

File tree

sdks/java/core/src/main/java/org/apache/beam/sdk/schemas/FieldValueTypeInformation.java

Lines changed: 9 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -105,7 +105,11 @@ public abstract static class Builder {
105105

106106
public abstract Builder setDescription(@Nullable String fieldDescription);
107107

108-
abstract FieldValueTypeInformation build();
108+
public abstract FieldValueTypeInformation build();
109+
}
110+
111+
public static Builder builder() {
112+
return new AutoValue_FieldValueTypeInformation.Builder();
109113
}
110114

111115
public static FieldValueTypeInformation forOneOf(
@@ -311,7 +315,8 @@ public FieldValueTypeInformation withName(String name) {
311315
return toBuilder().setName(name).build();
312316
}
313317

314-
static @Nullable FieldValueTypeInformation getIterableComponentType(TypeDescriptor<?> valueType) {
318+
public static @Nullable FieldValueTypeInformation getIterableComponentType(
319+
TypeDescriptor<?> valueType) {
315320
// TODO: Figure out nullable elements.
316321
TypeDescriptor<?> componentType = ReflectUtils.getIterableComponentType(valueType);
317322
if (componentType == null) {
@@ -331,13 +336,13 @@ public FieldValueTypeInformation withName(String name) {
331336
}
332337

333338
// If the type is a map type, returns the key type, otherwise returns a null reference.
334-
private static @Nullable FieldValueTypeInformation getMapKeyType(
339+
public static @Nullable FieldValueTypeInformation getMapKeyType(
335340
TypeDescriptor<?> typeDescriptor) {
336341
return getMapType(typeDescriptor, 0);
337342
}
338343

339344
// If the type is a map type, returns the value type, otherwise returns a null reference.
340-
private static @Nullable FieldValueTypeInformation getMapValueType(
345+
public static @Nullable FieldValueTypeInformation getMapValueType(
341346
TypeDescriptor<?> typeDescriptor) {
342347
return getMapType(typeDescriptor, 1);
343348
}

sdks/java/core/src/main/java/org/apache/beam/sdk/schemas/FromRowUsingCreator.java

Lines changed: 3 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -82,10 +82,10 @@ public T apply(Row row) {
8282
return null;
8383
}
8484
if (row instanceof RowWithGetters) {
85-
Object target = ((RowWithGetters) row).getGetterTarget();
86-
if (target.getClass().equals(typeDescriptor.getRawType())) {
85+
RowWithGetters rowWithGetters = (RowWithGetters) row;
86+
if (rowWithGetters.getGetterTargetType().equals(typeDescriptor)) {
8787
// Efficient path: simply extract the underlying object instead of creating a new one.
88-
return (T) target;
88+
return (T) rowWithGetters.getGetterTarget();
8989
}
9090
}
9191
if (fieldConverters == null) {

sdks/java/core/src/main/java/org/apache/beam/sdk/schemas/GetterBasedSchemaProvider.java

Lines changed: 74 additions & 24 deletions
Original file line numberDiff line numberDiff line change
@@ -20,16 +20,20 @@
2020
import static org.apache.beam.sdk.util.Preconditions.checkStateNotNull;
2121

2222
import java.util.ArrayList;
23+
import java.util.Arrays;
2324
import java.util.Collection;
2425
import java.util.List;
2526
import java.util.Map;
2627
import java.util.Objects;
2728
import java.util.Optional;
29+
import java.util.function.Function;
30+
import java.util.stream.Collectors;
2831
import org.apache.beam.sdk.schemas.Schema.FieldType;
2932
import org.apache.beam.sdk.schemas.Schema.LogicalType;
3033
import org.apache.beam.sdk.schemas.Schema.TypeName;
3134
import org.apache.beam.sdk.schemas.logicaltypes.EnumerationType;
3235
import org.apache.beam.sdk.schemas.logicaltypes.OneOfType;
36+
import org.apache.beam.sdk.schemas.utils.ReflectUtils;
3337
import org.apache.beam.sdk.transforms.SerializableFunction;
3438
import org.apache.beam.sdk.values.Row;
3539
import org.apache.beam.sdk.values.TypeDescriptor;
@@ -117,9 +121,11 @@ private class ToRowWithValueGetters<T extends @NonNull Object>
117121
implements SerializableFunction<T, Row> {
118122
private final Schema schema;
119123
private final Factory<List<FieldValueGetter<T, Object>>> getterFactory;
124+
private final TypeDescriptor getterTargetType;
120125

121-
public ToRowWithValueGetters(Schema schema) {
126+
public ToRowWithValueGetters(Schema schema, TypeDescriptor getterTargetType) {
122127
this.schema = schema;
128+
this.getterTargetType = getterTargetType;
123129
// Since we know that this factory is always called from inside the lambda with the same
124130
// schema, return a caching factory that caches the first value seen for each class. This
125131
// prevents having to lookup the getter list each time createGetters is called.
@@ -128,13 +134,13 @@ public ToRowWithValueGetters(Schema schema) {
128134
(Factory<List<FieldValueGetter<T, Object>>>)
129135
(typeDescriptor, schema1) ->
130136
(List)
131-
GetterBasedSchemaProvider.this.fieldValueGetters(
132-
typeDescriptor, schema1));
137+
GetterBasedSchemaProvider.this.fieldValueGetters(typeDescriptor, schema1),
138+
GetterBasedSchemaProvider.this::fieldValueTypeInformations);
133139
}
134140

135141
@Override
136142
public Row apply(T input) {
137-
return Row.withSchema(schema).withFieldValueGetters(getterFactory, input);
143+
return Row.withSchema(schema).withFieldValueGetters(getterFactory, input, getterTargetType);
138144
}
139145

140146
private GetterBasedSchemaProvider getOuter() {
@@ -172,7 +178,7 @@ public <T> SerializableFunction<T, Row> toRowFunction(TypeDescriptor<T> typeDesc
172178
Verify.verifyNotNull(
173179
schemaFor(typeDescriptor), "can't create a ToRowFunction with null schema");
174180

175-
return new ToRowWithValueGetters<>(schema);
181+
return new ToRowWithValueGetters<>(schema, typeDescriptor);
176182
}
177183

178184
@Override
@@ -193,26 +199,39 @@ public boolean equals(@Nullable Object obj) {
193199
private static class RowValueGettersFactory<T extends @NonNull Object>
194200
implements Factory<List<FieldValueGetter<T, Object>>> {
195201
private final Factory<List<FieldValueGetter<T, Object>>> gettersFactory;
202+
private final Factory<List<FieldValueTypeInformation>> typeInfoFactory;
196203
private final @NotOnlyInitialized Factory<List<FieldValueGetter<T, Object>>>
197204
cachingGettersFactory;
198205

199206
static <T extends @NonNull Object> Factory<List<FieldValueGetter<T, Object>>> of(
200-
Factory<List<FieldValueGetter<T, Object>>> gettersFactory) {
201-
return new RowValueGettersFactory<>(gettersFactory).cachingGettersFactory;
207+
Factory<List<FieldValueGetter<T, Object>>> gettersFactory,
208+
Factory<List<FieldValueTypeInformation>> typeInfoFactory) {
209+
return new RowValueGettersFactory(gettersFactory, typeInfoFactory).cachingGettersFactory;
202210
}
203211

204-
RowValueGettersFactory(Factory<List<FieldValueGetter<T, Object>>> gettersFactory) {
212+
RowValueGettersFactory(
213+
Factory<List<FieldValueGetter<T, Object>>> gettersFactory,
214+
Factory<List<FieldValueTypeInformation>> typeInfoFactory) {
205215
this.gettersFactory = gettersFactory;
216+
this.typeInfoFactory = typeInfoFactory;
206217
this.cachingGettersFactory = new CachingFactory<>(this);
207218
}
208219

209220
@Override
210221
public List<FieldValueGetter<T, Object>> create(
211222
TypeDescriptor<?> typeDescriptor, Schema schema) {
212223
List<FieldValueGetter<T, Object>> getters = gettersFactory.create(typeDescriptor, schema);
224+
Map<String, FieldValueTypeInformation> typeInfoByName =
225+
typeInfoFactory.create(typeDescriptor, schema).stream()
226+
.collect(Collectors.toMap(FieldValueTypeInformation::getName, Function.identity()));
213227
List<FieldValueGetter<T, Object>> rowGetters = new ArrayList<>(getters.size());
214228
for (int i = 0; i < getters.size(); i++) {
215-
rowGetters.add(rowValueGetter(getters.get(i), schema.getField(i).getType()));
229+
FieldValueGetter getter = Verify.verifyNotNull(getters.get(i));
230+
rowGetters.add(
231+
rowValueGetter(
232+
getter,
233+
schema.getField(i).getType(),
234+
Verify.verifyNotNull(typeInfoByName.get(getter.name())).getType()));
216235
}
217236
return rowGetters;
218237
}
@@ -228,26 +247,49 @@ && needsConversion(Verify.verifyNotNull(type.getCollectionElementType())))
228247
|| needsConversion(Verify.verifyNotNull(type.getMapValueType()))));
229248
}
230249

231-
FieldValueGetter<T, Object> rowValueGetter(FieldValueGetter base, FieldType type) {
250+
FieldValueGetter<T, Object> rowValueGetter(
251+
FieldValueGetter base, FieldType type, @Nullable TypeDescriptor<?> getterReturnType) {
232252
TypeName typeName = type.getTypeName();
233253
if (!needsConversion(type)) {
234254
return base;
235255
}
236256
if (typeName.equals(TypeName.ROW)) {
237-
return new GetRow(base, Verify.verifyNotNull(type.getRowSchema()), cachingGettersFactory);
238-
} else if (typeName.equals(TypeName.ARRAY)) {
257+
return new GetRow(
258+
base,
259+
getterReturnType,
260+
Verify.verifyNotNull(type.getRowSchema()),
261+
cachingGettersFactory);
262+
} else if (typeName.equals(TypeName.ARRAY) || typeName.equals(TypeName.ITERABLE)) {
239263
FieldType elementType = Verify.verifyNotNull(type.getCollectionElementType());
240-
return elementType.getTypeName().equals(TypeName.ROW)
241-
? new GetEagerCollection(base, converter(elementType))
242-
: new GetCollection(base, converter(elementType));
243-
} else if (typeName.equals(TypeName.ITERABLE)) {
244-
return new GetIterable(
245-
base, converter(Verify.verifyNotNull(type.getCollectionElementType())));
264+
TypeDescriptor<?> elementTypeDescriptor =
265+
Optional.ofNullable(getterReturnType)
266+
.map(ReflectUtils::getIterableComponentType)
267+
.orElse(null);
268+
if (TypeName.ARRAY == typeName) {
269+
return TypeName.ROW == elementType.getTypeName()
270+
? new GetEagerCollection(base, converter(elementType, elementTypeDescriptor))
271+
: new GetCollection(base, converter(elementType, elementTypeDescriptor));
272+
} else { // TypeName.ITERABLE
273+
return new GetIterable(base, converter(elementType, elementTypeDescriptor));
274+
}
246275
} else if (typeName.equals(TypeName.MAP)) {
276+
@Nullable
277+
TypeDescriptor[] resolvedKeyValueTypes =
278+
Optional.ofNullable(getterReturnType)
279+
.<@Nullable TypeDescriptor[]>map(
280+
getterType ->
281+
Arrays.stream(Map.class.getTypeParameters())
282+
.<@Nullable TypeDescriptor>map(
283+
typeVar -> {
284+
TypeDescriptor resolved = getterType.resolveType(typeVar);
285+
return resolved.hasUnresolvedParameters() ? null : resolved;
286+
})
287+
.<@Nullable TypeDescriptor>toArray(TypeDescriptor[]::new))
288+
.orElse(new TypeDescriptor[] {null, null});
247289
return new GetMap(
248290
base,
249-
converter(Verify.verifyNotNull(type.getMapKeyType())),
250-
converter(Verify.verifyNotNull(type.getMapValueType())));
291+
converter(Verify.verifyNotNull(type.getMapKeyType()), resolvedKeyValueTypes[0]),
292+
converter(Verify.verifyNotNull(type.getMapValueType()), resolvedKeyValueTypes[1]));
251293
} else if (type.isLogicalType(OneOfType.IDENTIFIER)) {
252294
OneOfType oneOfType = type.getLogicalType(OneOfType.class);
253295
Schema oneOfSchema = oneOfType.getOneOfSchema();
@@ -257,7 +299,7 @@ FieldValueGetter<T, Object> rowValueGetter(FieldValueGetter base, FieldType type
257299
Maps.newHashMapWithExpectedSize(values.size());
258300
for (Map.Entry<String, Integer> kv : values.entrySet()) {
259301
FieldType fieldType = oneOfSchema.getField(kv.getKey()).getType();
260-
FieldValueGetter<?, ?> converter = converter(fieldType);
302+
FieldValueGetter<?, ?> converter = converter(fieldType, null);
261303
converters.put(kv.getValue(), converter);
262304
}
263305

@@ -268,27 +310,35 @@ FieldValueGetter<T, Object> rowValueGetter(FieldValueGetter base, FieldType type
268310
return base;
269311
}
270312

271-
FieldValueGetter<?, ?> converter(FieldType type) {
272-
return rowValueGetter(IDENTITY, type);
313+
FieldValueGetter<?, ?> converter(FieldType type, @Nullable TypeDescriptor<?> getterReturnType) {
314+
return rowValueGetter(IDENTITY, type, getterReturnType);
273315
}
274316

275317
static class GetRow<T extends @NonNull Object, V extends @NonNull Object>
276318
extends Converter<T, V> {
277319
final Schema schema;
278320
final Factory<List<FieldValueGetter<V, Object>>> factory;
321+
final @Nullable TypeDescriptor<?> valueType;
279322

280323
GetRow(
281324
FieldValueGetter<T, V> getter,
325+
@Nullable TypeDescriptor<?> getterReturnType,
282326
Schema schema,
283327
Factory<List<FieldValueGetter<V, Object>>> factory) {
284328
super(getter);
285329
this.schema = schema;
286330
this.factory = factory;
331+
this.valueType = getterReturnType;
287332
}
288333

289334
@Override
290335
Object convert(V value) {
291-
return Row.withSchema(schema).withFieldValueGetters(factory, value);
336+
return Row.withSchema(schema)
337+
.withFieldValueGetters(
338+
factory,
339+
value,
340+
Optional.ofNullable(valueType)
341+
.orElse((TypeDescriptor) TypeDescriptor.of(value.getClass())));
292342
}
293343
}
294344

sdks/java/core/src/main/java/org/apache/beam/sdk/schemas/utils/AutoValueUtils.java

Lines changed: 6 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -162,7 +162,7 @@ private static String getAutoValueGeneratedName(String baseClass) {
162162
Optional<Constructor<?>> constructor =
163163
Arrays.stream(generatedTypeDescriptor.getRawType().getDeclaredConstructors())
164164
.filter(c -> !Modifier.isPrivate(c.getModifiers()))
165-
.filter(c -> matchConstructor(c, schemaTypes))
165+
.filter(c -> matchConstructor(generatedTypeDescriptor, c, schemaTypes))
166166
.findAny();
167167
return constructor
168168
.map(
@@ -177,7 +177,9 @@ private static String getAutoValueGeneratedName(String baseClass) {
177177
}
178178

179179
private static boolean matchConstructor(
180-
Constructor<?> constructor, List<FieldValueTypeInformation> getterTypes) {
180+
TypeDescriptor typeDescriptor,
181+
Constructor<?> constructor,
182+
List<FieldValueTypeInformation> getterTypes) {
181183
if (constructor.getParameters().length != getterTypes.size()) {
182184
return false;
183185
}
@@ -197,7 +199,8 @@ private static boolean matchConstructor(
197199
// Verify that constructor parameters match (name and type) the inferred schema.
198200
for (Parameter parameter : constructor.getParameters()) {
199201
FieldValueTypeInformation type = typeMap.get(parameter.getName());
200-
if (type == null || type.getRawType() != parameter.getType()) {
202+
if (type == null
203+
|| !type.getType().equals(typeDescriptor.resolveType(parameter.getParameterizedType()))) {
201204
valid = false;
202205
break;
203206
}

sdks/java/core/src/main/java/org/apache/beam/sdk/values/Row.java

Lines changed: 4 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -838,9 +838,11 @@ public int nextFieldId() {
838838

839839
@Internal
840840
public <T> Row withFieldValueGetters(
841-
Factory<List<FieldValueGetter<T, Object>>> fieldValueGetterFactory, T getterTarget) {
841+
Factory<List<FieldValueGetter<T, Object>>> fieldValueGetterFactory,
842+
T getterTarget,
843+
TypeDescriptor<?> getterTargetType) {
842844
checkState(getterTarget != null, "getters require withGetterTarget.");
843-
return new RowWithGetters<>(schema, fieldValueGetterFactory, getterTarget);
845+
return new RowWithGetters<>(schema, fieldValueGetterFactory, getterTarget, getterTargetType);
844846
}
845847

846848
public Row build() {

sdks/java/core/src/main/java/org/apache/beam/sdk/values/RowWithGetters.java

Lines changed: 11 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -44,14 +44,19 @@
4444
@SuppressWarnings("rawtypes")
4545
public class RowWithGetters<T extends @NonNull Object> extends Row {
4646
private final T getterTarget;
47+
private final TypeDescriptor<?> getterTargetType;
4748
private final List<FieldValueGetter<T, Object>> getters;
4849
private @Nullable Map<Integer, @Nullable Object> cache = null;
4950

5051
RowWithGetters(
51-
Schema schema, Factory<List<FieldValueGetter<T, Object>>> getterFactory, T getterTarget) {
52+
Schema schema,
53+
Factory<List<FieldValueGetter<T, Object>>> getterFactory,
54+
T getterTarget,
55+
TypeDescriptor<?> getterTargetType) {
5256
super(schema);
5357
this.getterTarget = getterTarget;
54-
this.getters = getterFactory.create(TypeDescriptor.of(getterTarget.getClass()), schema);
58+
this.getterTargetType = getterTargetType;
59+
this.getters = getterFactory.create(getterTargetType, schema);
5560
}
5661

5762
@Override
@@ -90,6 +95,10 @@ public <W> W getValue(int fieldIdx) {
9095
return (W) fieldValue;
9196
}
9297

98+
public TypeDescriptor<?> getGetterTargetType() {
99+
return getterTargetType;
100+
}
101+
93102
private boolean cacheFieldType(Field field) {
94103
TypeName typeName = field.getType().getTypeName();
95104
return typeName.equals(TypeName.MAP)

0 commit comments

Comments
 (0)