diff --git a/spring-data-rest-webmvc/src/main/java/org/springframework/data/rest/webmvc/json/patch/SpelPath.java b/spring-data-rest-webmvc/src/main/java/org/springframework/data/rest/webmvc/json/patch/SpelPath.java index 469300bce..b1e40e5d2 100644 --- a/spring-data-rest-webmvc/src/main/java/org/springframework/data/rest/webmvc/json/patch/SpelPath.java +++ b/spring-data-rest-webmvc/src/main/java/org/springframework/data/rest/webmvc/json/patch/SpelPath.java @@ -36,9 +36,11 @@ import org.springframework.expression.EvaluationContext; import org.springframework.expression.Expression; import org.springframework.expression.ExpressionException; import org.springframework.expression.spel.SpelEvaluationException; +import org.springframework.expression.spel.SpelMessage; import org.springframework.expression.spel.standard.SpelExpressionParser; import org.springframework.expression.spel.support.SimpleEvaluationContext; import org.springframework.util.Assert; +import org.springframework.util.CollectionUtils; import org.springframework.util.ConcurrentReferenceHashMap; /** @@ -210,6 +212,7 @@ class SpelPath { static class TypedSpelPath extends SpelPath { private static final String INVALID_PATH_REFERENCE = "Invalid path reference %s on type %s (from source %s)!"; + private static final String INVALID_COLLECTION_INDEX = "Invalid collection index %s for collection of size %s. Use '…/-' or the collection's actual size as index to append to it!"; private static final Map TYPED_PATHS = new ConcurrentReferenceHashMap<>(32); private static final EvaluationContext CONTEXT = SimpleEvaluationContext.forReadWriteDataBinding().build(); @@ -286,7 +289,24 @@ class SpelPath { Assert.notNull(root, "Root object must not be null!"); - return expression.getValueType(CONTEXT, root); + try { + + return expression.getValueType(CONTEXT, root); + + } catch (SpelEvaluationException o_O) { + + if (!SpelMessage.COLLECTION_INDEX_OUT_OF_BOUNDS.equals(o_O.getMessageCode())) { + throw o_O; + } + + Object collectionOrArray = getParent().getValue(root); + + if (Collection.class.isInstance(collectionOrArray)) { + return CollectionUtils.findCommonElementType(Collection.class.cast(collectionOrArray)); + } + } + + throw new IllegalArgumentException(String.format("Cannot obtain type for path %s on %s!", path, root)); } /** @@ -384,6 +404,11 @@ class SpelPath { } else { List list = parentPath.getValue(target); + + if (listIndex > list.size()) { + throw new PatchException(String.format(INVALID_COLLECTION_INDEX, listIndex, list.size())); + } + list.add(listIndex >= 0 ? listIndex.intValue() : list.size(), value); } } diff --git a/spring-data-rest-webmvc/src/test/java/org/springframework/data/rest/webmvc/json/patch/AddOperationUnitTests.java b/spring-data-rest-webmvc/src/test/java/org/springframework/data/rest/webmvc/json/patch/AddOperationUnitTests.java index 8474f8f3f..6cf5a1ed5 100755 --- a/spring-data-rest-webmvc/src/test/java/org/springframework/data/rest/webmvc/json/patch/AddOperationUnitTests.java +++ b/spring-data-rest-webmvc/src/test/java/org/springframework/data/rest/webmvc/json/patch/AddOperationUnitTests.java @@ -15,6 +15,7 @@ */ package org.springframework.data.rest.webmvc.json.patch; +import static org.assertj.core.api.Assertions.*; import static org.assertj.core.api.Assertions.assertThat; import static org.junit.Assert.*; @@ -111,4 +112,29 @@ public class AddOperationUnitTests { assertThat(todo.getUninitialized()).containsExactly("Text"); } + + @Test // DATAREST-1273 + public void addsItemToTheEndOfACollectionViaIndex() { + + List todos = new ArrayList(); + todos.add(new Todo(1L, "A", false)); + + Todo todo = new Todo(2L, "B", true); + AddOperation.of("/1", todo).perform(todos, Todo.class); + + assertThat(todos).element(1).isEqualTo(todo); + } + + @Test // DATAREST-1273 + public void rejectsAdditionBeyondEndOfList() { + + List todos = new ArrayList(); + todos.add(new Todo(1L, "A", false)); + + assertThatExceptionOfType(PatchException.class) // + .isThrownBy(() -> AddOperation.of("/2", new Todo(2L, "B", true)).perform(todos, Todo.class)) // + .withMessageContaining("index") // + .withMessageContaining("2") // + .withMessageContaining("1"); + } }