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 b99aef13..654e1618 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,12 +34,10 @@ 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; import org.springframework.lang.Nullable; -import org.springframework.util.Assert; import org.springframework.validation.BindException; import org.springframework.validation.BindingErrorProcessor; import org.springframework.validation.BindingResult; @@ -66,12 +64,13 @@ public class GraphQlArgumentBinder { */ private static final int DEFAULT_AUTO_GROW_COLLECTION_LIMIT = 1024; + @Nullable private final SimpleTypeConverter typeConverter; private final BindingErrorProcessor bindingErrorProcessor = new DefaultBindingErrorProcessor(); - private List> dataBinderInitializers = new ArrayList<>(); + private final List> dataBinderInitializers = new ArrayList<>(); public GraphQlArgumentBinder() { @@ -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,30 @@ 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); - } - - Class targetClass = targetType.resolve(); - Assert.notNull(targetClass, "Could not determine target type from " + targetType); - - DataBinder binder = new DataBinder(null, argumentName != null ? argumentName : "arguments"); + DataBinder binder = new DataBinder(null, name != null ? ("Arguments[" + name + "]") : "Arguments"); initDataBinder(binder); BindingResult bindingResult = binder.getBindingResult(); + Stack segments = new Stack<>(); - - try { - // From Collection - - if (isApproximableCollectionType(rawValue)) { - segments.push(argumentName); - return createCollection((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 = createValue((Map) rawValue, targetClass, 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); + if (name != null) { + segments.push(name); } - finally { - checkBindingResult(bindingResult); + + Object targetValue = bindRawValue( + rawValue, targetType, targetType.resolve(Object.class), bindingResult, segments); + + if (bindingResult.hasErrors()) { + throw new BindException(bindingResult); } + + return targetValue; } private void initDataBinder(DataBinder binder) { @@ -183,123 +154,143 @@ 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 createCollection( - Collection rawCollection, ResolvableType collectionType, + @Nullable + private Object bindRawValue( + Object rawValue, ResolvableType targetType, Class targetClass, BindingResult bindingResult, Stack segments) { - if (!Collection.class.isAssignableFrom(collectionType.resolve())) { - bindingResult.rejectValue(toArgumentPath(segments), "typeMismatch", "Expected collection: " + collectionType); - return Collections.emptyList(); + boolean isOptional = (targetClass == Optional.class); + + if (isOptional) { + targetType = targetType.getNested(2); + targetClass = targetType.resolve(); } + Object value; + if (rawValue == null || targetClass == Object.class) { + value = rawValue; + } + else if (rawValue instanceof Collection) { + value = bindCollection((Collection) rawValue, targetType, targetClass, bindingResult, segments); + } + else if (rawValue instanceof Map) { + value = bindMap((Map) rawValue, targetType, targetClass, bindingResult, segments); + } + else { + value = (targetClass.isAssignableFrom(rawValue.getClass()) ? + rawValue : convertValue(rawValue, targetClass, bindingResult, segments)); + } + + return (isOptional ? Optional.ofNullable(value) : value); + } + + private Collection bindCollection( + Collection rawCollection, ResolvableType collectionType, Class collectionClass, + BindingResult bindingResult, Stack segments) { + + ResolvableType elementType = collectionType.asCollection().getGeneric(0); Class elementClass = collectionType.asCollection().getGeneric(0).resolve(); if (elementClass == null) { - bindingResult.rejectValue(toArgumentPath(segments), "unknownElementType", "Unknown element type"); - return Collections.emptyList(); + bindingResult.rejectValue(toArgumentPath(segments), "unknownTargetType", "Unknown target type"); + return Collections.emptyList(); // Keep going, report as many errors as we can } - 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) createValueOrNull((Map) rawValue, elementClass, 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 - private Object createValueOrNull( - Map rawMap, Class targetType, BindingResult result, Stack segments) { - - try { - return createValue(rawMap, targetType, result, segments); - } - catch (BindException ex) { - return null; - } + private static String toArgumentPath(Stack path) { + StringBuilder sb = new StringBuilder(); + path.forEach(sb::append); + return sb.toString(); } - @SuppressWarnings("unchecked") - private Object createValue( - Map rawMap, Class targetType, BindingResult bindingResult, - Stack segments) throws BindException { + @Nullable + private Object bindMap( + Map rawMap, ResolvableType targetType, Class targetClass, + BindingResult bindingResult, Stack segments) { - Object target; - Constructor ctor = BeanUtils.getResolvableConstructor(targetType); - - // 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(initBindValues(rawMap)); - - if (dataBinder.getBindingResult().hasErrors()) { - addErrors(dataBinder, bindingResult, segments); - throw new BindException(bindingResult); - } - - return target; + if (Map.class.isAssignableFrom(targetClass)) { + return bindMapToMap(rawMap, targetType, bindingResult, segments, targetClass); } - // Data class constructor + Constructor constructor = BeanUtils.getResolvableConstructor(targetClass); + if (constructor.getParameterCount() > 0) { + return bindMapToObjectViaConstructor(rawMap, constructor, bindingResult, segments); + } - if (!segments.isEmpty()) { + Object target = BeanUtils.instantiateClass(constructor); + DataBinder dataBinder = new DataBinder(target); + initDataBinder(dataBinder); + dataBinder.getBindingResult().setNestedPath(toArgumentPath(segments)); + dataBinder.setConversionService(getConversionService()); + dataBinder.bind(createPropertyValues(rawMap)); + + if (dataBinder.getBindingResult().hasErrors()) { + String nestedPath = dataBinder.getBindingResult().getNestedPath(); + for (FieldError error : dataBinder.getBindingResult().getFieldErrors()) { + bindingResult.addError( + new FieldError(bindingResult.getObjectName(), nestedPath + error.getField(), + error.getRejectedValue(), error.isBindingFailure(), error.getCodes(), + error.getArguments(), error.getDefaultMessage())); + } + return null; + } + + return target; + } + + private Map bindMapToMap( + Map rawMap, ResolvableType targetType, BindingResult bindingResult, + Stack segments, Class targetClass) { + + ResolvableType valueType = targetType.asMap().getGeneric(1); + Class valueClass = valueType.resolve(); + if (valueClass == null) { + bindingResult.rejectValue(toArgumentPath(segments), "unknownTargetType", "Unknown target type"); + return Collections.emptyMap(); // Keep going, report as many errors as we can + } + + 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(); + } + + return map; + } + + @Nullable + private Object bindMapToObjectViaConstructor( + Map rawMap, Constructor constructor, BindingResult bindingResult, + Stack segments) { + + if (segments.size() > 0) { segments.push("."); } - String[] paramNames = BeanUtils.getParameterNames(ctor); - Class[] paramTypes = ctor.getParameterTypes(); + String[] paramNames = BeanUtils.getParameterNames(constructor); + Class[] paramTypes = constructor.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] = createCollection((Collection) rawValue, elementType, bindingResult, segments); - } - else if (rawValue instanceof Map) { - boolean isOptional = (paramTypes[i] == Optional.class); - Class type = (isOptional ? methodParam.nestedIfOptional().getNestedParameterType() : paramTypes[i]); - Object value = createValueOrNull((Map) rawValue, type, bindingResult, segments); - args[i] = (isOptional ? Optional.ofNullable(value) : value); - } - else { - args[i] = convertValue(rawValue, paramTypes[i], new TypeDescriptor(methodParam), bindingResult, segments); - } + 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(); } @@ -308,26 +299,28 @@ public class GraphQlArgumentBinder { } try { - return BeanUtils.instantiateClass(ctor, args); + return BeanUtils.instantiateClass(constructor, args); } catch (BeanInstantiationException ex) { - // Swallow if we had binding errors, it's as far as we could go - checkBindingResult(bindingResult); + // Ignore: we had binding errors to begin with + if (bindingResult.hasErrors()) { + return null; + } throw ex; } } - private MutablePropertyValues initBindValues(Map rawMap) { + private static MutablePropertyValues createPropertyValues(Map rawMap) { MutablePropertyValues mpvs = new MutablePropertyValues(); Stack segments = new Stack<>(); for (String key : rawMap.keySet()) { - addBindValues(mpvs, key, rawMap.get(key), segments); + addPropertyValue(mpvs, key, rawMap.get(key), segments); } return mpvs; } @SuppressWarnings("unchecked") - private void addBindValues(MutablePropertyValues mpvs, String name, Object value, Stack segments) { + private static void addPropertyValue(MutablePropertyValues mpvs, String name, Object value, Stack segments) { if (value instanceof List) { List items = (List) value; if (items.isEmpty()) { @@ -337,7 +330,7 @@ public class GraphQlArgumentBinder { } else { for (int i = 0; i < items.size(); i++) { - addBindValues(mpvs, name + "[" + i + "]", items.get(i), segments); + addPropertyValue(mpvs, name + "[" + i + "]", items.get(i), segments); } } } @@ -345,7 +338,7 @@ public class GraphQlArgumentBinder { segments.push(name + "."); Map map = (Map) value; for (String key : map.keySet()) { - addBindValues(mpvs, key, map.get(key), segments); + addPropertyValue(mpvs, key, map.get(key), segments); } segments.pop(); } @@ -356,25 +349,14 @@ public class GraphQlArgumentBinder { } } - private String toArgumentPath(Stack path) { - StringBuilder sb = new StringBuilder(); - path.forEach(sb::append); - return sb.toString(); - } - @SuppressWarnings("unchecked") @Nullable - private T convertValue(@Nullable Object rawValue, Class type, BindingResult result, Stack segments) { - return (T) convertValue(rawValue, type, TypeDescriptor.valueOf(type), result, segments); - } - - @Nullable - private Object convertValue( - @Nullable Object rawValue, Class type, TypeDescriptor descriptor, - BindingResult bindingResult, Stack segments) { + private T convertValue( + @Nullable Object rawValue, Class type, BindingResult bindingResult, Stack segments) { + Object value = null; try { - return getTypeConverter().convertIfNecessary(rawValue, type, descriptor); + value = getTypeConverter().convertIfNecessary(rawValue, (Class) type, TypeDescriptor.valueOf(type)); } catch (TypeMismatchException ex) { String name = toArgumentPath(segments); @@ -382,21 +364,7 @@ public class GraphQlArgumentBinder { bindingResult.recordFieldValue(name, type, rawValue); this.bindingErrorProcessor.processPropertyAccessException(ex, bindingResult); } - return null; - } - - private void addErrors(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(), - error.getRejectedValue(), error.isBindingFailure(), error.getCodes(), - error.getArguments(), error.getDefaultMessage()))); - } - - private void checkBindingResult(BindingResult bindingResult) throws BindException { - if (bindingResult.hasErrors()) { - throw new BindException(bindingResult); - } + return (T) value; } } 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 60bd142c..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,14 +272,45 @@ 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"); } }); } + @Test + void primaryConstructorWithMapArgument() throws Exception { + + Object result = this.binder.bind( + environment( + "{\"key\":{" + + "\"map\":{" + + "\"item1\":{" + + "\"name\":\"Jason\"," + + "\"age\":\"21\"" + + "}," + + "\"item2\":{" + + "\"name\":\"James\"," + + "\"age\":\"22\"" + + "}" + + "}}}"), + "key", + ResolvableType.forClass(PrimaryConstructorItemMapBean.class)); + + assertThat(result).isNotNull().isInstanceOf(PrimaryConstructorItemMapBean.class); + Map map = ((PrimaryConstructorItemMapBean) result).getMap(); + + Item item1 = map.get("item1"); + assertThat(item1.getName()).isEqualTo("Jason"); + assertThat(item1.getAge()).isEqualTo(21); + + Item item2 = map.get("item2"); + assertThat(item2.getName()).isEqualTo("James"); + assertThat(item2.getAge()).isEqualTo(22); + } + @Test // gh-447 @SuppressWarnings("unchecked") void primaryConstructorWithGenericObject() throws Exception { @@ -428,6 +459,20 @@ class GraphQlArgumentBinderTests { } + static class PrimaryConstructorItemMapBean { + + private final Map map; + + public PrimaryConstructorItemMapBean(Map map) { + this.map = map; + } + + public Map getMap() { + return this.map; + } + } + + @SuppressWarnings("OptionalUsedAsFieldOrParameterType") static class PrimaryConstructorOptionalItemBean {