From 83a3db608954c5b2c926c3bdc91886e199746fca Mon Sep 17 00:00:00 2001 From: Oliver Gierke Date: Tue, 7 Jul 2015 12:01:54 +0200 Subject: [PATCH] DATAJPA-757 - Count queries for String-based native queries are now executed as native ones. MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit AbstractStringBasedJpaQuery now considers the native flag of the underlying JpaQueryMethod when creating count queries. Had to widen the return type of AbstractJpaQuery.createJpaQuery(…) to Query to accommodate the return value EntityManager.createNativeQuery(…). Tweaked the implementation of PagedExecution to handle non-Long return values potentially converting them to longs to retain API. --- .../repository/query/AbstractJpaQuery.java | 9 ++-- .../query/AbstractStringBasedJpaQuery.java | 10 +++-- .../repository/query/JpaQueryExecution.java | 19 ++++---- .../query/SimpleJpaQueryUnitTests.java | 43 +++++++++++++------ 4 files changed, 51 insertions(+), 30 deletions(-) diff --git a/src/main/java/org/springframework/data/jpa/repository/query/AbstractJpaQuery.java b/src/main/java/org/springframework/data/jpa/repository/query/AbstractJpaQuery.java index d87ee193b..8d47fe3e2 100644 --- a/src/main/java/org/springframework/data/jpa/repository/query/AbstractJpaQuery.java +++ b/src/main/java/org/springframework/data/jpa/repository/query/AbstractJpaQuery.java @@ -182,7 +182,8 @@ public abstract class AbstractJpaQuery implements RepositoryQuery { Assert.notNull(query, "Query must not be null!"); Assert.notNull(method, "JpaQueryMethod must not be null!"); - Map hints = Jpa21Utils.tryGetFetchGraphHints(em, method.getEntityGraph(), getQueryMethod().getEntityInformation().getJavaType()); + Map hints = Jpa21Utils.tryGetFetchGraphHints(em, method.getEntityGraph(), + getQueryMethod().getEntityInformation().getJavaType()); for (Map.Entry hint : hints.entrySet()) { query.setHint(hint.getKey(), hint.getValue()); @@ -191,8 +192,8 @@ public abstract class AbstractJpaQuery implements RepositoryQuery { return query; } - protected TypedQuery createCountQuery(Object[] values) { - TypedQuery countQuery = doCreateCountQuery(values); + protected Query createCountQuery(Object[] values) { + Query countQuery = doCreateCountQuery(values); return method.applyHintsToCountQuery() ? applyHints(countQuery, method) : countQuery; } @@ -210,5 +211,5 @@ public abstract class AbstractJpaQuery implements RepositoryQuery { * @param values must not be {@literal null}. * @return */ - protected abstract TypedQuery doCreateCountQuery(Object[] values); + protected abstract Query doCreateCountQuery(Object[] values); } diff --git a/src/main/java/org/springframework/data/jpa/repository/query/AbstractStringBasedJpaQuery.java b/src/main/java/org/springframework/data/jpa/repository/query/AbstractStringBasedJpaQuery.java index 697bfde1b..f466ee6c5 100644 --- a/src/main/java/org/springframework/data/jpa/repository/query/AbstractStringBasedJpaQuery.java +++ b/src/main/java/org/springframework/data/jpa/repository/query/AbstractStringBasedJpaQuery.java @@ -17,7 +17,6 @@ package org.springframework.data.jpa.repository.query; import javax.persistence.EntityManager; import javax.persistence.Query; -import javax.persistence.TypedQuery; import org.springframework.data.repository.query.EvaluationContextProvider; import org.springframework.data.repository.query.ParameterAccessor; @@ -105,8 +104,13 @@ abstract class AbstractStringBasedJpaQuery extends AbstractJpaQuery { * @see org.springframework.data.jpa.repository.query.AbstractJpaQuery#doCreateCountQuery(java.lang.Object[]) */ @Override - protected TypedQuery doCreateCountQuery(Object[] values) { - return createBinder(values).bind(getEntityManager().createQuery(countQuery.getQueryString(), Long.class)); + protected Query doCreateCountQuery(Object[] values) { + + String queryString = countQuery.getQueryString(); + EntityManager em = getEntityManager(); + + return createBinder(values).bind( + getQueryMethod().isNativeQuery() ? em.createNativeQuery(queryString) : em.createQuery(queryString, Long.class)); } /** diff --git a/src/main/java/org/springframework/data/jpa/repository/query/JpaQueryExecution.java b/src/main/java/org/springframework/data/jpa/repository/query/JpaQueryExecution.java index 394946a19..521079803 100644 --- a/src/main/java/org/springframework/data/jpa/repository/query/JpaQueryExecution.java +++ b/src/main/java/org/springframework/data/jpa/repository/query/JpaQueryExecution.java @@ -22,11 +22,10 @@ import javax.persistence.EntityManager; import javax.persistence.NoResultException; import javax.persistence.Query; import javax.persistence.StoredProcedureQuery; -import javax.persistence.TypedQuery; import org.springframework.core.convert.ConversionService; import org.springframework.core.convert.support.ConfigurableConversionService; -import org.springframework.core.convert.support.GenericConversionService; +import org.springframework.core.convert.support.DefaultConversionService; import org.springframework.data.domain.PageImpl; import org.springframework.data.domain.Pageable; import org.springframework.data.domain.Slice; @@ -53,7 +52,7 @@ public abstract class JpaQueryExecution { static { - ConfigurableConversionService conversionService = new GenericConversionService(); + ConfigurableConversionService conversionService = new DefaultConversionService(); conversionService.addConverter(JpaResultConverters.BlobToByteArrayConverter.INSTANCE); CONVERSION_SERVICE = conversionService; @@ -90,8 +89,8 @@ public abstract class JpaQueryExecution { return result; } - return CONVERSION_SERVICE.canConvert(result.getClass(), requiredType) ? CONVERSION_SERVICE.convert(result, - requiredType) : result; + return CONVERSION_SERVICE.canConvert(result.getClass(), requiredType) + ? CONVERSION_SERVICE.convert(result, requiredType) : result; } /** @@ -173,10 +172,10 @@ public abstract class JpaQueryExecution { protected Object doExecute(AbstractJpaQuery repositoryQuery, Object[] values) { // Execute query to compute total - TypedQuery projection = repositoryQuery.createCountQuery(values); + Query projection = repositoryQuery.createCountQuery(values); - List totals = projection.getResultList(); - Long total = totals.size() == 1 ? totals.get(0) : totals.size(); + List totals = projection.getResultList(); + Long total = totals.size() == 1 ? CONVERSION_SERVICE.convert(totals.get(0), Long.class) : totals.size(); ParameterAccessor accessor = new ParametersParameterAccessor(parameters, values); Pageable pageable = accessor.getPageable(); @@ -187,8 +186,8 @@ public abstract class JpaQueryExecution { Query query = repositoryQuery.createQuery(values); - List content = pageable == null || total > pageable.getOffset() ? query.getResultList() : Collections - .emptyList(); + List content = pageable == null || total > pageable.getOffset() ? query.getResultList() + : Collections.emptyList(); return new PageImpl(content, pageable, total); } diff --git a/src/test/java/org/springframework/data/jpa/repository/query/SimpleJpaQueryUnitTests.java b/src/test/java/org/springframework/data/jpa/repository/query/SimpleJpaQueryUnitTests.java index b47df59fd..c42ef877d 100644 --- a/src/test/java/org/springframework/data/jpa/repository/query/SimpleJpaQueryUnitTests.java +++ b/src/test/java/org/springframework/data/jpa/repository/query/SimpleJpaQueryUnitTests.java @@ -1,5 +1,5 @@ /* - * Copyright 2008-2014 the original author or authors. + * Copyright 2008-2015 the original author or authors. * * Licensed under the Apache License, Version 2.0 (the "License"); * you may not use this file except in compliance with the License. @@ -44,8 +44,8 @@ import org.springframework.data.jpa.provider.QueryExtractor; import org.springframework.data.jpa.repository.Query; import org.springframework.data.jpa.repository.sample.UserRepository; import org.springframework.data.repository.core.RepositoryMetadata; -import org.springframework.data.repository.query.DefaultEvaluationContextProvider; import org.springframework.data.repository.query.EvaluationContextProvider; +import org.springframework.data.repository.query.ExtensionAwareEvaluationContextProvider; import org.springframework.data.repository.query.RepositoryQuery; import org.springframework.expression.spel.standard.SpelExpressionParser; @@ -60,14 +60,15 @@ public class SimpleJpaQueryUnitTests { static final String USER_QUERY = "select u from User u"; static final SpelExpressionParser PARSER = new SpelExpressionParser(); - private static final EvaluationContextProvider EVALUATION_CONTEXT_PROVIDER = DefaultEvaluationContextProvider.INSTANCE; + private static final EvaluationContextProvider EVALUATION_CONTEXT_PROVIDER = new ExtensionAwareEvaluationContextProvider(); JpaQueryMethod method; @Mock EntityManager em; @Mock EntityManagerFactory emf; @Mock QueryExtractor extractor; - @Mock TypedQuery query; + @Mock javax.persistence.Query query; + @Mock TypedQuery typedQuery; @Mock RepositoryMetadata metadata; @Mock ParameterBinder binder; @@ -78,7 +79,7 @@ public class SimpleJpaQueryUnitTests { public void setUp() throws SecurityException, NoSuchMethodException { when(em.createQuery(anyString())).thenReturn(query); - when(em.createQuery(anyString(), eq(Long.class))).thenReturn(query); + when(em.createQuery(anyString(), eq(Long.class))).thenReturn(typedQuery); when(em.getEntityManagerFactory()).thenReturn(emf); when(emf.createEntityManager()).thenReturn(em); when(metadata.getDomainType()).thenReturn((Class) User.class); @@ -92,13 +93,13 @@ public class SimpleJpaQueryUnitTests { public void prefersDeclaredCountQueryOverCreatingOne() throws Exception { method = new JpaQueryMethod(SimpleJpaQueryUnitTests.class.getMethod("prefersDeclaredCountQueryOverCreatingOne"), - metadata, extractor); // mock(JpaQueryMethod.class); - when(em.createQuery("foo", Long.class)).thenReturn(query); + metadata, extractor); + when(em.createQuery("foo", Long.class)).thenReturn(typedQuery); SimpleJpaQuery jpaQuery = new SimpleJpaQuery(method, em, "select u from User u", EVALUATION_CONTEXT_PROVIDER, PARSER); - assertThat(jpaQuery.createCountQuery(new Object[] {}), is(query)); + assertThat(jpaQuery.createCountQuery(new Object[] {}), is((javax.persistence.Query) typedQuery)); } /** @@ -112,8 +113,8 @@ public class SimpleJpaQueryUnitTests { Method method = UserRepository.class.getMethod("findAllPaged", Pageable.class); JpaQueryMethod queryMethod = new JpaQueryMethod(method, metadata, extractor); - AbstractJpaQuery jpaQuery = new SimpleJpaQuery(queryMethod, em, "select u from User u", - EVALUATION_CONTEXT_PROVIDER, PARSER); + AbstractJpaQuery jpaQuery = new SimpleJpaQuery(queryMethod, em, "select u from User u", EVALUATION_CONTEXT_PROVIDER, + PARSER); jpaQuery.createCountQuery(new Object[] { new PageRequest(1, 10) }); verify(query, times(0)).setFirstResult(anyInt()); @@ -169,7 +170,7 @@ public class SimpleJpaQueryUnitTests { public void doesNotValidateCountQueryIfNotPagingMethod() throws Exception { Method method = SampleRepository.class.getMethod("findByAnnotatedQuery"); - when(em.createQuery(contains("count"))).thenThrow(IllegalArgumentException.class); + when(em.createQuery(Mockito.contains("count"))).thenThrow(IllegalArgumentException.class); createJpaQuery(method); } @@ -183,7 +184,7 @@ public class SimpleJpaQueryUnitTests { Method method = SampleRepository.class.getMethod("pageByAnnotatedQuery", Pageable.class); - when(em.createQuery(contains("count"))).thenThrow(IllegalArgumentException.class); + when(em.createQuery(Mockito.contains("count"))).thenThrow(IllegalArgumentException.class); exception.expect(IllegalArgumentException.class); exception.expectMessage("Count"); exception.expectMessage(method.getName()); @@ -205,7 +206,23 @@ public class SimpleJpaQueryUnitTests { assertThat(query instanceof NativeJpaQuery, is(true)); } - private RepositoryQuery createJpaQuery(Method method) { + /** + * @see DATAJPA-757 + */ + @Test + public void createsNativeCountQuery() throws Exception { + + when(em.createNativeQuery(anyString())).thenReturn(query); + + AbstractJpaQuery jpaQuery = createJpaQuery( + UserRepository.class.getMethod("findUsersInNativeQueryWithPagination", Pageable.class)); + + jpaQuery.doCreateCountQuery(new Object[] { new PageRequest(0, 10) }); + + verify(em, times(1)).createNativeQuery(anyString()); + } + + private AbstractJpaQuery createJpaQuery(Method method) { JpaQueryMethod queryMethod = new JpaQueryMethod(method, metadata, extractor); return JpaQueryFactory.INSTANCE.fromQueryAnnotation(queryMethod, em, EVALUATION_CONTEXT_PROVIDER);