diff --git a/src/main/java/org/springframework/data/repository/query/ExtensionAwareEvaluationContextProvider.java b/src/main/java/org/springframework/data/repository/query/ExtensionAwareEvaluationContextProvider.java index fe7078399..9b6be4f15 100644 --- a/src/main/java/org/springframework/data/repository/query/ExtensionAwareEvaluationContextProvider.java +++ b/src/main/java/org/springframework/data/repository/query/ExtensionAwareEvaluationContextProvider.java @@ -42,6 +42,7 @@ import org.springframework.expression.spel.SpelMessage; import org.springframework.expression.spel.support.ReflectivePropertyAccessor; import org.springframework.expression.spel.support.StandardEvaluationContext; import org.springframework.util.Assert; +import org.springframework.util.StringUtils; import org.springframework.util.TypeUtils; /** @@ -109,16 +110,53 @@ public class ExtensionAwareEvaluationContextProvider implements EvaluationContex // Add parameters for indexed access ec.setRootObject(parameterValues); - // Add parameters as variables if named - for (Parameter param : parameters) { - if (param.isNamedParameter()) { - ec.setVariable(param.getName(), parameterValues[param.getIndex()]); - } - } + Map variables = collectVariables(parameters, parameterValues); + + ec.setVariables(variables); return ec; } + protected > Map collectVariables(T parameters, + Object[] parameterValues) { + + Map variables = new HashMap(); + + registerSpecialParameterVariablesIfPresent(parameters, parameterValues, variables); + registerNamedMethodParameterVariables(parameters, parameterValues, variables); + + return variables; + } + + protected > void registerNamedMethodParameterVariables(T parameters, + Object[] parameterValues, Map variables) { + + for (Parameter param : parameters) { + if (param.isNamedParameter()) { + variables.put(param.getName(), parameterValues[param.getIndex()]); + } + } + } + + protected > void registerSpecialParameterVariablesIfPresent( + T parameters, Object[] parameterValues, Map variables) { + + if (parameters.hasSpecialParameter()) { + + for (Parameter param : parameters) { + if (param.isSpecialParameter()) { + + Object paramValue = parameterValues[param.getIndex()]; + if (paramValue == null) { + continue; + } + + variables.put(StringUtils.uncapitalize(param.getType().getSimpleName()), paramValue); + } + } + } + } + /** * Returns the {@link EvaluationContextExtension} to be used. Either from the current configuration or the configured * {@link BeanFactory}. @@ -251,6 +289,10 @@ public class ExtensionAwareEvaluationContextProvider implements EvaluationContex final Method function = targetObject instanceof Map && ((Map) targetObject).containsKey(name) ? (Method) ((Map) targetObject) .get(name) : (Method) functions.get(name); + if (function == null) { + return null; + } + Class[] parameterTypes = function.getParameterTypes(); if (parameterTypes.length != argumentTypes.size()) { return null; diff --git a/src/main/java/org/springframework/data/repository/query/Parameters.java b/src/main/java/org/springframework/data/repository/query/Parameters.java index 95cc6234b..d73c79886 100644 --- a/src/main/java/org/springframework/data/repository/query/Parameters.java +++ b/src/main/java/org/springframework/data/repository/query/Parameters.java @@ -37,7 +37,7 @@ import org.springframework.util.Assert; */ public abstract class Parameters, T extends Parameter> implements Iterable { - @SuppressWarnings("unchecked") public static final List> TYPES = Arrays.asList(Pageable.class, Sort.class); + public static final List> TYPES = Arrays.asList(Pageable.class, Sort.class); private static final String PARAM_ON_SPECIAL = format("You must not user @%s on a parameter typed %s or %s", Param.class.getSimpleName(), Pageable.class.getSimpleName(), Sort.class.getSimpleName()); diff --git a/src/test/java/org/springframework/data/repository/query/ExtensibleEvaluationContextProviderUnitTests.java b/src/test/java/org/springframework/data/repository/query/ExtensibleEvaluationContextProviderUnitTests.java index 85bd0e87f..7026ad850 100644 --- a/src/test/java/org/springframework/data/repository/query/ExtensibleEvaluationContextProviderUnitTests.java +++ b/src/test/java/org/springframework/data/repository/query/ExtensibleEvaluationContextProviderUnitTests.java @@ -27,6 +27,10 @@ import java.util.Map; import org.junit.Before; import org.junit.Test; +import org.springframework.data.domain.PageRequest; +import org.springframework.data.domain.Pageable; +import org.springframework.data.domain.Sort; +import org.springframework.data.domain.Sort.Direction; import org.springframework.data.repository.query.spi.EvaluationContextExtension; import org.springframework.data.repository.query.spi.EvaluationContextExtensionSupport; import org.springframework.expression.EvaluationContext; @@ -134,6 +138,40 @@ public class ExtensibleEvaluationContextProviderUnitTests { assertThat(evaluateExpression("_first.DUMMY_KEY", provider), is((Object) "dummy")); } + /** + * @see DATACMNS-533 + */ + @Test + public void exposesPageableParameter() throws Exception { + + this.parameters = new DefaultParameters(SampleRepo.class.getMethod("findByFirstname", String.class, Pageable.class)); + ExtensionAwareEvaluationContextProvider provider = new ExtensionAwareEvaluationContextProvider( + Collections. emptyList()); + + PageRequest pageable = new PageRequest(2, 3, new Sort(Direction.DESC, "lastname")); + + assertThat(evaluateExpression("#pageable.offset", provider, new Object[] { "test", pageable }), is((Object) 6)); + assertThat(evaluateExpression("#pageable.pageSize", provider, new Object[] { "test", pageable }), is((Object) 3)); + assertThat(evaluateExpression("#pageable.sort.toString()", provider, new Object[] { "test", pageable }), + is((Object) "lastname: DESC")); + } + + /** + * @see DATACMNS-533 + */ + @Test + public void exposesSortParameter() throws Exception { + + this.parameters = new DefaultParameters(SampleRepo.class.getMethod("findByFirstname", String.class, Sort.class)); + ExtensionAwareEvaluationContextProvider provider = new ExtensionAwareEvaluationContextProvider( + Collections. emptyList()); + + Sort sort = new Sort(Direction.DESC, "lastname"); + + assertThat(evaluateExpression("#sort.toString()", provider, new Object[] { "test", sort }), + is((Object) "lastname: DESC")); + } + public static class DummyExtension extends EvaluationContextExtensionSupport { public static String DUMMY_KEY = "dummy"; @@ -196,13 +234,21 @@ public class ExtensibleEvaluationContextProviderUnitTests { } private Object evaluateExpression(String expression, EvaluationContextProvider provider) { + return evaluateExpression(expression, provider, new Object[] { "parameterValue" }); + } - EvaluationContext evaluationContext = provider.getEvaluationContext(parameters, new Object[] { "parameterValue" }); + private Object evaluateExpression(String expression, EvaluationContextProvider provider, Object[] args) { + + EvaluationContext evaluationContext = provider.getEvaluationContext(parameters, args); return new SpelExpressionParser().parseExpression(expression).getValue(evaluationContext); } interface SampleRepo { List findByFirstname(@Param("firstname") String firstname); + + List findByFirstname(@Param("firstname") String firstname, Pageable pageable); + + List findByFirstname(@Param("firstname") String firstname, Sort sort); } }