From d38ffeae6b1f8ff405c0835b05f888bb7ce8c65e Mon Sep 17 00:00:00 2001 From: rstoyanchev Date: Wed, 19 Oct 2022 07:16:22 +0100 Subject: [PATCH] Consistent raw value handling in GraphQlArgumentBinder This commit ensures a single method is used to bind raw values, whether those a top level argument value, or a value nested within a collection or map. See gh-449 --- .../graphql/data/GraphQlArgumentBinder.java | 299 ++++++++---------- .../data/GraphQlArgumentBinderTests.java | 16 +- 2 files changed, 133 insertions(+), 182 deletions(-) diff --git a/spring-graphql/src/main/java/org/springframework/graphql/data/GraphQlArgumentBinder.java b/spring-graphql/src/main/java/org/springframework/graphql/data/GraphQlArgumentBinder.java index 16163bc2..6abf540d 100644 --- a/spring-graphql/src/main/java/org/springframework/graphql/data/GraphQlArgumentBinder.java +++ b/spring-graphql/src/main/java/org/springframework/graphql/data/GraphQlArgumentBinder.java @@ -34,7 +34,6 @@ import org.springframework.beans.MutablePropertyValues; import org.springframework.beans.SimpleTypeConverter; import org.springframework.beans.TypeMismatchException; import org.springframework.core.CollectionFactory; -import org.springframework.core.MethodParameter; import org.springframework.core.ResolvableType; import org.springframework.core.convert.ConversionService; import org.springframework.core.convert.TypeDescriptor; @@ -113,7 +112,7 @@ public class GraphQlArgumentBinder { * Bind a single argument, or the full arguments map, onto an object of the * given target type. * @param environment for access to the arguments - * @param argumentName the name of the argument to bind, or {@code null} to + * @param name the name of the argument to bind, or {@code null} to * use the full arguments map * @param targetType the type of Object to create * @return the created Object, possibly {@code null} @@ -124,58 +123,32 @@ public class GraphQlArgumentBinder { * is the argument path where the issue occurred. */ @Nullable - @SuppressWarnings("unchecked") public Object bind( - DataFetchingEnvironment environment, @Nullable String argumentName, ResolvableType targetType) + DataFetchingEnvironment environment, @Nullable String name, ResolvableType targetType) throws BindException { - Object rawValue = (argumentName != null ? - environment.getArgument(argumentName) : environment.getArguments()); + Object rawValue = (name != null ? + environment.getArgument(name) : environment.getArguments()); - if (rawValue == null) { - return wrapAsOptionalIfNecessary(null, targetType); + DataBinder binder = new DataBinder(null, name != null ? ("Arguments[" + name + "]") : "Arguments"); + initDataBinder(binder); + BindingResult bindingResult = binder.getBindingResult(); + + Stack segments = new Stack<>(); + if (name != null) { + segments.push(name); } Class targetClass = targetType.resolve(); Assert.notNull(targetClass, "Could not determine target type from " + targetType); - DataBinder binder = new DataBinder(null, argumentName != null ? argumentName : "arguments"); - initDataBinder(binder); - BindingResult bindingResult = binder.getBindingResult(); - Stack segments = new Stack<>(); + Object targetValue = bindRawValue(rawValue, targetType, targetClass, bindingResult, segments); - try { - // From Collection - - if (isApproximableCollectionType(rawValue)) { - segments.push(argumentName); - return bindCollection((Collection) rawValue, targetType, bindingResult, segments); - } - - if (targetClass == Optional.class) { - targetClass = targetType.getNested(2).resolve(); - Assert.notNull(targetClass, "Could not determine Optional type from " + targetType); - } - - // From Map - - if (rawValue instanceof Map) { - Object target = bindMap((Map) rawValue, targetType, bindingResult, segments); - return wrapAsOptionalIfNecessary(target, targetType); - } - - // From Scalar - - if (targetClass.isInstance(rawValue)) { - return wrapAsOptionalIfNecessary(rawValue, targetType); - } - - Object target = convertValue(rawValue, targetClass, bindingResult, segments); - return wrapAsOptionalIfNecessary(target, targetType); - } - finally { - checkBindingResult(bindingResult); + if (bindingResult.hasErrors()) { + throw new BindException(bindingResult); } + + return targetValue; } private void initDataBinder(DataBinder binder) { @@ -183,23 +156,45 @@ public class GraphQlArgumentBinder { this.dataBinderInitializers.forEach(initializer -> initializer.accept(binder)); } - @Nullable - private Object wrapAsOptionalIfNecessary(@Nullable Object value, ResolvableType type) { - return (type.resolve(Object.class).equals(Optional.class) ? Optional.ofNullable(value) : value); - } - - private boolean isApproximableCollectionType(@Nullable Object rawValue) { - return (rawValue != null && - (CollectionFactory.isApproximableCollectionType(rawValue.getClass()) || - rawValue instanceof List)); // it may be SingletonList - } - @SuppressWarnings({"ConstantConditions", "unchecked"}) - private Collection bindCollection( + @Nullable + private Object bindRawValue( + Object rawValue, ResolvableType targetValueType, Class targetValueClass, + BindingResult bindingResult, Stack segments) { + + boolean isOptional = targetValueClass == Optional.class; + + if (isOptional) { + targetValueType = targetValueType.getNested(2); + targetValueClass = targetValueType.resolve(); + } + + Object targetValue; + if (rawValue == null || targetValueClass == Object.class) { + targetValue = rawValue; + } + else if (rawValue instanceof Collection) { + targetValue = bindCollection((Collection) rawValue, targetValueType, bindingResult, segments); + } + else if (rawValue instanceof Map) { + targetValue = bindMap((Map) rawValue, targetValueType, bindingResult, segments); + } + else { + targetValue = (targetValueClass.isAssignableFrom(rawValue.getClass()) ? + rawValue : convertValue(rawValue, targetValueClass, bindingResult, segments)); + } + + return (isOptional ? Optional.ofNullable(targetValue) : targetValue); + } + + private Collection bindCollection( Collection rawCollection, ResolvableType collectionType, BindingResult bindingResult, Stack segments) { - if (!Collection.class.isAssignableFrom(collectionType.resolve())) { + Class collectionClass = collectionType.resolve(); + Assert.notNull(collectionClass, "Unknown Collection class"); + + if (!Collection.class.isAssignableFrom(collectionClass)) { bindingResult.rejectValue(toArgumentPath(segments), "typeMismatch", "Expected collection: " + collectionType); return Collections.emptyList(); } @@ -211,133 +206,95 @@ public class GraphQlArgumentBinder { return Collections.emptyList(); } - Collection collection = CollectionFactory.createCollection(collectionType.getRawClass(), elementClass, rawCollection.size()); - int i = 0; + Collection collection = + CollectionFactory.createCollection(collectionClass, elementClass, rawCollection.size()); + + int index = 0; for (Object rawValue : rawCollection) { - segments.push("[" + i++ + "]"); - if (rawValue == null || elementClass.isAssignableFrom(rawValue.getClass())) { - collection.add((T) rawValue); - } - else if (rawValue instanceof Map) { - collection.add((T) bindMap((Map) rawValue, elementType, bindingResult, segments)); - } - else { - collection.add((T) convertValue(rawValue, elementClass, bindingResult, segments)); - } + segments.push("[" + index++ + "]"); + collection.add(bindRawValue(rawValue, elementType, elementClass, bindingResult, segments)); segments.pop(); } return collection; } @Nullable - @SuppressWarnings("unchecked") private Object bindMap( Map rawMap, ResolvableType targetType, BindingResult bindingResult, Stack segments) { - try { - Class targetClass = targetType.resolve(); - Assert.notNull(targetClass, "Unknown target class"); + Class targetClass = targetType.resolve(); + Assert.notNull(targetClass, "Unknown target class"); - if (Map.class.isAssignableFrom(targetClass)) { - ResolvableType valueType = targetType.asMap().getGeneric(1); - Class valueClass = valueType.resolve(); - if (valueClass == null) { - bindingResult.rejectValue(toArgumentPath(segments), "unknownMapValueType", "Unknown Map value type"); - return Collections.emptyMap(); - } - Map map = CollectionFactory.createMap(targetClass, rawMap.size()); - for (Map.Entry entry : rawMap.entrySet()) { - Object rawValue = entry.getValue(); - segments.push("[" + entry.getKey() + "]"); - if (rawValue == null || valueType.isAssignableFrom(rawValue.getClass())) { - map.put(entry.getKey(), entry.getValue()); - } - else if (rawValue instanceof Map) { - map.put(entry.getKey(), bindMap( - (Map) rawValue, valueType, bindingResult, segments)); - } - else { - map.put(entry.getKey(), convertValue(rawValue, valueClass, bindingResult, segments)); - } - segments.pop(); - } - return map; + if (Map.class.isAssignableFrom(targetClass)) { + ResolvableType valueType = targetType.asMap().getGeneric(1); + Class valueClass = valueType.resolve(); + if (valueClass == null) { + bindingResult.rejectValue(toArgumentPath(segments), "unknownMapValueType", "Unknown Map value type"); + return Collections.emptyMap(); } - - Object target; - Constructor ctor = BeanUtils.getResolvableConstructor(targetClass); - - // Default constructor + data binding via properties - - if (ctor.getParameterCount() == 0) { - target = BeanUtils.instantiateClass(ctor); - DataBinder dataBinder = new DataBinder(target); - initDataBinder(dataBinder); - dataBinder.getBindingResult().setNestedPath(toArgumentPath(segments)); - dataBinder.setConversionService(getConversionService()); - dataBinder.bind(toPropertyValues(rawMap)); - - if (dataBinder.getBindingResult().hasErrors()) { - addDataBinderErrors(dataBinder, bindingResult, segments); - throw new BindException(bindingResult); - } - - return target; - } - - // Data class constructor - - if (!segments.isEmpty()) { - segments.push("."); - } - - String[] paramNames = BeanUtils.getParameterNames(ctor); - Class[] paramTypes = ctor.getParameterTypes(); - Object[] args = new Object[paramTypes.length]; - - for (int i = 0; i < paramNames.length; i++) { - String paramName = paramNames[i]; - Object rawValue = rawMap.get(paramName); - segments.push(paramName); - MethodParameter methodParam = new MethodParameter(ctor, i); - if (rawValue == null && methodParam.isOptional()) { - args[i] = (paramTypes[i] == Optional.class ? Optional.empty() : null); - } - else if (paramTypes[i] == Object.class) { - args[i] = rawValue; - } - else if (isApproximableCollectionType(rawValue)) { - ResolvableType elementType = ResolvableType.forMethodParameter(methodParam); - args[i] = bindCollection((Collection) rawValue, elementType, bindingResult, segments); - } - else if (rawValue instanceof Map) { - boolean isOptional = (paramTypes[i] == Optional.class); - ResolvableType type = ResolvableType.forMethodParameter(methodParam.nestedIfOptional()); - Object value = bindMap((Map) rawValue, type, bindingResult, segments); - args[i] = (isOptional ? Optional.ofNullable(value) : value); - } - else { - args[i] = convertValue(rawValue, paramTypes[i], new TypeDescriptor(methodParam), bindingResult, segments); - } + Map map = CollectionFactory.createMap(targetClass, rawMap.size()); + for (Map.Entry entry : rawMap.entrySet()) { + String key = entry.getKey(); + segments.push("[" + key + "]"); + map.put(key, bindRawValue(entry.getValue(), valueType, valueClass, bindingResult, segments)); segments.pop(); } - - if (segments.size() > 1) { - segments.pop(); - } - - try { - return BeanUtils.instantiateClass(ctor, args); - } - catch (BeanInstantiationException ex) { - // Swallow if we had binding errors, it's as far as we could go - checkBindingResult(bindingResult); - throw ex; - } + return map; } - catch (BindException ex) { - return null; + + Object target; + Constructor constructor = BeanUtils.getResolvableConstructor(targetClass); + + // Default constructor + data binding via properties + + if (constructor.getParameterCount() == 0) { + target = BeanUtils.instantiateClass(constructor); + DataBinder dataBinder = new DataBinder(target); + initDataBinder(dataBinder); + dataBinder.getBindingResult().setNestedPath(toArgumentPath(segments)); + dataBinder.setConversionService(getConversionService()); + dataBinder.bind(toPropertyValues(rawMap)); + + if (dataBinder.getBindingResult().hasErrors()) { + copyBindingErrors(dataBinder, bindingResult, segments); + return null; + } + + return target; + } + + // Data class constructor + + if (!segments.isEmpty()) { + segments.push("."); + } + + String[] paramNames = BeanUtils.getParameterNames(constructor); + Class[] paramTypes = constructor.getParameterTypes(); + Object[] args = new Object[paramTypes.length]; + + for (int i = 0; i < paramNames.length; i++) { + String name = paramNames[i]; + segments.push(name); + ResolvableType paramType = ResolvableType.forConstructorParameter(constructor, i); + args[i] = bindRawValue(rawMap.get(name), paramType, paramTypes[i], bindingResult, segments); + segments.pop(); + } + + if (segments.size() > 1) { + segments.pop(); + } + + try { + return BeanUtils.instantiateClass(constructor, args); + } + catch (BeanInstantiationException ex) { + // Ignore if we had binding errors already + if (bindingResult.hasErrors()) { + return null; + } + throw ex; } } @@ -409,7 +366,7 @@ public class GraphQlArgumentBinder { return null; } - private static void addDataBinderErrors(DataBinder binder, BindingResult bindingResult, Stack segments) { + private static void copyBindingErrors(DataBinder binder, BindingResult bindingResult, Stack segments) { String path = (!segments.isEmpty() ? toArgumentPath(segments) + "." : ""); binder.getBindingResult().getFieldErrors().forEach(error -> bindingResult.addError( new FieldError(bindingResult.getObjectName(), path + error.getField(), @@ -417,10 +374,4 @@ public class GraphQlArgumentBinder { error.getArguments(), error.getDefaultMessage()))); } - private void checkBindingResult(BindingResult bindingResult) throws BindException { - if (bindingResult.hasErrors()) { - throw new BindException(bindingResult); - } - } - } diff --git a/spring-graphql/src/test/java/org/springframework/graphql/data/GraphQlArgumentBinderTests.java b/spring-graphql/src/test/java/org/springframework/graphql/data/GraphQlArgumentBinderTests.java index dead81ad..6f575324 100644 --- a/spring-graphql/src/test/java/org/springframework/graphql/data/GraphQlArgumentBinderTests.java +++ b/spring-graphql/src/test/java/org/springframework/graphql/data/GraphQlArgumentBinderTests.java @@ -120,8 +120,8 @@ class GraphQlArgumentBinderTests { .extracting(ex -> ((BindException) ex).getFieldErrors()) .satisfies(errors -> { assertThat(errors).hasSize(1); - assertThat(errors.get(0).getObjectName()).isEqualTo("key"); - assertThat(errors.get(0).getField()).isEqualTo("age"); + assertThat(errors.get(0).getObjectName()).isEqualTo("Arguments[key]"); + assertThat(errors.get(0).getField()).isEqualTo("key.age"); assertThat(errors.get(0).getRejectedValue()).isEqualTo("invalid"); }); } @@ -246,12 +246,12 @@ class GraphQlArgumentBinderTests { .satisfies(errors -> { assertThat(errors).hasSize(2); - assertThat(errors.get(0).getObjectName()).isEqualTo("key"); - assertThat(errors.get(0).getField()).isEqualTo("age"); + assertThat(errors.get(0).getObjectName()).isEqualTo("Arguments[key]"); + assertThat(errors.get(0).getField()).isEqualTo("key.age"); assertThat(errors.get(0).getRejectedValue()).isEqualTo("invalid"); - assertThat(errors.get(1).getObjectName()).isEqualTo("key"); - assertThat(errors.get(1).getField()).isEqualTo("item.age"); + assertThat(errors.get(0).getObjectName()).isEqualTo("Arguments[key]"); + assertThat(errors.get(1).getField()).isEqualTo("key.item.age"); assertThat(errors.get(1).getRejectedValue()).isEqualTo("invalid"); }); } @@ -272,8 +272,8 @@ class GraphQlArgumentBinderTests { assertThat(errors).hasSize(2); for (int i = 0; i < errors.size(); i++) { FieldError error = errors.get(i); - assertThat(error.getObjectName()).isEqualTo("key"); - assertThat(error.getField()).isEqualTo("items[" + i + "].age"); + assertThat(error.getObjectName()).isEqualTo("Arguments[key]"); + assertThat(error.getField()).isEqualTo("key.items[" + i + "].age"); assertThat(error.getRejectedValue()).isEqualTo("invalid"); assertThat(error.getDefaultMessage()).startsWith("Failed to convert property value"); }