2020import static org .apache .beam .sdk .util .Preconditions .checkStateNotNull ;
2121
2222import java .util .ArrayList ;
23+ import java .util .Arrays ;
2324import java .util .Collection ;
2425import java .util .List ;
2526import java .util .Map ;
2627import java .util .Objects ;
2728import java .util .Optional ;
29+ import java .util .function .Function ;
30+ import java .util .stream .Collectors ;
2831import org .apache .beam .sdk .schemas .Schema .FieldType ;
2932import org .apache .beam .sdk .schemas .Schema .LogicalType ;
3033import org .apache .beam .sdk .schemas .Schema .TypeName ;
3134import org .apache .beam .sdk .schemas .logicaltypes .EnumerationType ;
3235import org .apache .beam .sdk .schemas .logicaltypes .OneOfType ;
36+ import org .apache .beam .sdk .schemas .utils .ReflectUtils ;
3337import org .apache .beam .sdk .transforms .SerializableFunction ;
3438import org .apache .beam .sdk .values .Row ;
3539import 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
0 commit comments