diff --git a/spring-session/src/main/java/org/springframework/session/jdbc/JdbcOperationsSessionRepository.java b/spring-session/src/main/java/org/springframework/session/jdbc/JdbcOperationsSessionRepository.java index 8ba0b18..7664f62 100644 --- a/spring-session/src/main/java/org/springframework/session/jdbc/JdbcOperationsSessionRepository.java +++ b/spring-session/src/main/java/org/springframework/session/jdbc/JdbcOperationsSessionRepository.java @@ -16,7 +16,6 @@ package org.springframework.session.jdbc; -import java.sql.Connection; import java.sql.PreparedStatement; import java.sql.ResultSet; import java.sql.SQLException; @@ -38,14 +37,14 @@ import org.springframework.core.convert.TypeDescriptor; import org.springframework.core.convert.support.GenericConversionService; import org.springframework.core.serializer.support.DeserializingConverter; import org.springframework.core.serializer.support.SerializingConverter; +import org.springframework.dao.DataAccessException; import org.springframework.expression.Expression; import org.springframework.expression.spel.standard.SpelExpressionParser; import org.springframework.jdbc.core.BatchPreparedStatementSetter; import org.springframework.jdbc.core.JdbcOperations; import org.springframework.jdbc.core.JdbcTemplate; -import org.springframework.jdbc.core.PreparedStatementCreator; import org.springframework.jdbc.core.PreparedStatementSetter; -import org.springframework.jdbc.core.RowMapper; +import org.springframework.jdbc.core.ResultSetExtractor; import org.springframework.jdbc.support.lob.DefaultLobHandler; import org.springframework.jdbc.support.lob.LobHandler; import org.springframework.scheduling.annotation.Scheduled; @@ -186,7 +185,8 @@ public class JdbcOperationsSessionRepository implements private final TransactionOperations transactionOperations; - private final RowMapper mapper = new ExpiringSessionMapper(); + private final ResultSetExtractor> extractor = + new ExpiringSessionResultSetExtractor(); /** * The name of database table used by Spring Session to store sessions. @@ -382,17 +382,15 @@ public class JdbcOperationsSessionRepository implements public ExpiringSession doInTransaction(TransactionStatus status) { List sessions = JdbcOperationsSessionRepository.this.jdbcOperations.query( - new PreparedStatementCreator() { + getQuery(GET_SESSION_QUERY), + new PreparedStatementSetter() { - public PreparedStatement createPreparedStatement(Connection con) throws SQLException { - PreparedStatement ps = con.prepareStatement(getQuery(GET_SESSION_QUERY), - ResultSet.TYPE_SCROLL_INSENSITIVE, ResultSet.CONCUR_READ_ONLY); + public void setValues(PreparedStatement ps) throws SQLException { ps.setString(1, id); - return ps; } }, - JdbcOperationsSessionRepository.this.mapper + JdbcOperationsSessionRepository.this.extractor ); if (sessions.isEmpty()) { return null; @@ -434,18 +432,15 @@ public class JdbcOperationsSessionRepository implements public List doInTransaction(TransactionStatus status) { return JdbcOperationsSessionRepository.this.jdbcOperations.query( - new PreparedStatementCreator() { + getQuery(LIST_SESSIONS_BY_PRINCIPAL_NAME_QUERY), + new PreparedStatementSetter() { - public PreparedStatement createPreparedStatement(Connection con) throws SQLException { - PreparedStatement ps = con.prepareStatement( - getQuery(LIST_SESSIONS_BY_PRINCIPAL_NAME_QUERY), - ResultSet.TYPE_SCROLL_INSENSITIVE, ResultSet.CONCUR_READ_ONLY); + public void setValues(PreparedStatement ps) throws SQLException { ps.setString(1, indexValue); - return ps; } }, - JdbcOperationsSessionRepository.this.mapper + JdbcOperationsSessionRepository.this.extractor ); } @@ -661,23 +656,34 @@ public class JdbcOperationsSessionRepository implements } - private class ExpiringSessionMapper implements RowMapper { + private class ExpiringSessionResultSetExtractor + implements ResultSetExtractor> { - public ExpiringSession mapRow(ResultSet rs, int rowNum) throws SQLException { - MapSession session = new MapSession(rs.getString("SESSION_ID")); - session.setCreationTime(rs.getLong("CREATION_TIME")); - session.setLastAccessedTime(rs.getLong("LAST_ACCESS_TIME")); - session.setMaxInactiveIntervalInSeconds(rs.getInt("MAX_INACTIVE_INTERVAL")); - String attributeName = rs.getString("ATTRIBUTE_NAME"); - if (attributeName != null) { - session.setAttribute(attributeName, deserialize(rs, "ATTRIBUTE_BYTES")); - while (rs.next() && session.getId().equals(rs.getString("SESSION_ID"))) { - session.setAttribute(rs.getString("ATTRIBUTE_NAME"), - deserialize(rs, "ATTRIBUTE_BYTES")); + public List extractData(ResultSet rs) throws SQLException, DataAccessException { + List sessions = new ArrayList(); + while (rs.next()) { + String id = rs.getString("SESSION_ID"); + MapSession session; + if (sessions.size() > 0 && getLast(sessions).getId().equals(id)) { + session = (MapSession) getLast(sessions); } - rs.previous(); + else { + session = new MapSession(id); + session.setCreationTime(rs.getLong("CREATION_TIME")); + session.setLastAccessedTime(rs.getLong("LAST_ACCESS_TIME")); + session.setMaxInactiveIntervalInSeconds(rs.getInt("MAX_INACTIVE_INTERVAL")); + } + String attributeName = rs.getString("ATTRIBUTE_NAME"); + if (attributeName != null) { + session.setAttribute(attributeName, deserialize(rs, "ATTRIBUTE_BYTES")); + } + sessions.add(session); } - return session; + return sessions; + } + + private ExpiringSession getLast(List sessions) { + return sessions.get(sessions.size() - 1); } } diff --git a/spring-session/src/test/java/org/springframework/session/jdbc/JdbcOperationsSessionRepositoryTests.java b/spring-session/src/test/java/org/springframework/session/jdbc/JdbcOperationsSessionRepositoryTests.java index 8ec47f0..f56cd73 100644 --- a/spring-session/src/test/java/org/springframework/session/jdbc/JdbcOperationsSessionRepositoryTests.java +++ b/spring-session/src/test/java/org/springframework/session/jdbc/JdbcOperationsSessionRepositoryTests.java @@ -33,9 +33,8 @@ import org.mockito.Mock; import org.mockito.runners.MockitoJUnitRunner; import org.springframework.jdbc.core.JdbcOperations; -import org.springframework.jdbc.core.PreparedStatementCreator; import org.springframework.jdbc.core.PreparedStatementSetter; -import org.springframework.jdbc.core.RowMapper; +import org.springframework.jdbc.core.ResultSetExtractor; import org.springframework.security.authentication.UsernamePasswordAuthenticationToken; import org.springframework.security.core.Authentication; import org.springframework.security.core.authority.AuthorityUtils; @@ -235,14 +234,17 @@ public class JdbcOperationsSessionRepositoryTests { @Test public void getSessionNotFound() { String sessionId = "testSessionId"; + given(this.jdbcOperations.query(isA(String.class), + isA(PreparedStatementSetter.class), isA(ResultSetExtractor.class))) + .willReturn(Collections.emptyList()); JdbcOperationsSessionRepository.JdbcSession session = this.repository .getSession(sessionId); assertThat(session).isNull(); assertPropagationRequiresNew(); - verify(this.jdbcOperations, times(1)).query( - isA(PreparedStatementCreator.class), isA(RowMapper.class)); + verify(this.jdbcOperations, times(1)).query(isA(String.class), + isA(PreparedStatementSetter.class), isA(ResultSetExtractor.class)); } @Test @@ -250,16 +252,17 @@ public class JdbcOperationsSessionRepositoryTests { MapSession expired = new MapSession(); expired.setLastAccessedTime(System.currentTimeMillis() - (MapSession.DEFAULT_MAX_INACTIVE_INTERVAL_SECONDS * 1000 + 1000)); - given(this.jdbcOperations.query(isA(PreparedStatementCreator.class), - isA(RowMapper.class))).willReturn(Collections.singletonList(expired)); + given(this.jdbcOperations.query(isA(String.class), + isA(PreparedStatementSetter.class), isA(ResultSetExtractor.class))) + .willReturn(Collections.singletonList(expired)); JdbcOperationsSessionRepository.JdbcSession session = this.repository .getSession(expired.getId()); assertThat(session).isNull(); assertPropagationRequiresNew(); - verify(this.jdbcOperations, times(1)).query( - isA(PreparedStatementCreator.class), isA(RowMapper.class)); + verify(this.jdbcOperations, times(1)).query(isA(String.class), + isA(PreparedStatementSetter.class), isA(ResultSetExtractor.class)); verify(this.jdbcOperations, times(1)).update(startsWith("DELETE"), eq(expired.getId())); } @@ -268,8 +271,9 @@ public class JdbcOperationsSessionRepositoryTests { public void getSessionFound() { MapSession saved = new MapSession(); saved.setAttribute("savedName", "savedValue"); - given(this.jdbcOperations.query(isA(PreparedStatementCreator.class), - isA(RowMapper.class))).willReturn(Collections.singletonList(saved)); + given(this.jdbcOperations.query(isA(String.class), + isA(PreparedStatementSetter.class), isA(ResultSetExtractor.class))) + .willReturn(Collections.singletonList(saved)); JdbcOperationsSessionRepository.JdbcSession session = this.repository .getSession(saved.getId()); @@ -278,8 +282,8 @@ public class JdbcOperationsSessionRepositoryTests { assertThat(session.isNew()).isFalse(); assertThat(session.getAttribute("savedName")).isEqualTo("savedValue"); assertPropagationRequiresNew(); - verify(this.jdbcOperations, times(1)).query( - isA(PreparedStatementCreator.class), isA(RowMapper.class)); + verify(this.jdbcOperations, times(1)).query(isA(String.class), + isA(PreparedStatementSetter.class), isA(ResultSetExtractor.class)); } @Test @@ -306,6 +310,9 @@ public class JdbcOperationsSessionRepositoryTests { @Test public void findByIndexNameAndIndexValuePrincipalIndexNameNotFound() { String principal = "username"; + given(this.jdbcOperations.query(isA(String.class), + isA(PreparedStatementSetter.class), isA(ResultSetExtractor.class))) + .willReturn(Collections.emptyList()); Map sessions = this.repository .findByIndexNameAndIndexValue( @@ -314,8 +321,8 @@ public class JdbcOperationsSessionRepositoryTests { assertThat(sessions).isEmpty(); assertPropagationRequiresNew(); - verify(this.jdbcOperations, times(1)).query( - isA(PreparedStatementCreator.class), isA(RowMapper.class)); + verify(this.jdbcOperations, times(1)).query(isA(String.class), + isA(PreparedStatementSetter.class), isA(ResultSetExtractor.class)); } @Test @@ -330,8 +337,9 @@ public class JdbcOperationsSessionRepositoryTests { MapSession saved2 = new MapSession(); saved2.setAttribute(SPRING_SECURITY_CONTEXT, authentication); saved.add(saved2); - given(this.jdbcOperations.query(isA(PreparedStatementCreator.class), - isA(RowMapper.class))).willReturn(saved); + given(this.jdbcOperations.query(isA(String.class), + isA(PreparedStatementSetter.class), isA(ResultSetExtractor.class))) + .willReturn(saved); Map sessions = this.repository .findByIndexNameAndIndexValue( @@ -340,8 +348,8 @@ public class JdbcOperationsSessionRepositoryTests { assertThat(sessions).hasSize(2); assertPropagationRequiresNew(); - verify(this.jdbcOperations, times(1)).query( - isA(PreparedStatementCreator.class), isA(RowMapper.class)); + verify(this.jdbcOperations, times(1)).query(isA(String.class), + isA(PreparedStatementSetter.class), isA(ResultSetExtractor.class)); } @Test