diff --git a/spring-session-core/src/main/java/org/springframework/session/web/http/HttpSessionAdapter.java b/spring-session-core/src/main/java/org/springframework/session/web/http/HttpSessionAdapter.java index becdbda9..39a779e2 100644 --- a/spring-session-core/src/main/java/org/springframework/session/web/http/HttpSessionAdapter.java +++ b/spring-session-core/src/main/java/org/springframework/session/web/http/HttpSessionAdapter.java @@ -1,5 +1,5 @@ /* - * Copyright 2014-2017 the original author or authors. + * Copyright 2014-2018 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. @@ -24,8 +24,13 @@ import java.util.Set; import javax.servlet.ServletContext; import javax.servlet.http.HttpSession; +import javax.servlet.http.HttpSessionBindingEvent; +import javax.servlet.http.HttpSessionBindingListener; import javax.servlet.http.HttpSessionContext; +import org.apache.commons.logging.Log; +import org.apache.commons.logging.LogFactory; + import org.springframework.session.Session; /** @@ -33,11 +38,14 @@ import org.springframework.session.Session; * * @param the {@link Session} type * @author Rob Winch + * @author Vedran Pavic * @since 1.1 */ @SuppressWarnings("deprecation") class HttpSessionAdapter implements HttpSession { + private static final Log logger = LogFactory.getLog(HttpSessionAdapter.class); + private S session; private final ServletContext servletContext; @@ -129,7 +137,28 @@ class HttpSessionAdapter implements HttpSession { @Override public void setAttribute(String name, Object value) { checkState(); + Object oldValue = this.session.getAttribute(name); this.session.setAttribute(name, value); + if (value != oldValue) { + if (oldValue instanceof HttpSessionBindingListener) { + try { + ((HttpSessionBindingListener) oldValue).valueUnbound( + new HttpSessionBindingEvent(this, name, oldValue)); + } + catch (Throwable th) { + logger.error("Error invoking session binding event listener", th); + } + } + if (value instanceof HttpSessionBindingListener) { + try { + ((HttpSessionBindingListener) value) + .valueBound(new HttpSessionBindingEvent(this, name, value)); + } + catch (Throwable th) { + logger.error("Error invoking session binding event listener", th); + } + } + } } @Override @@ -140,7 +169,17 @@ class HttpSessionAdapter implements HttpSession { @Override public void removeAttribute(String name) { checkState(); + Object oldValue = this.session.getAttribute(name); this.session.removeAttribute(name); + if (oldValue instanceof HttpSessionBindingListener) { + try { + ((HttpSessionBindingListener) oldValue) + .valueUnbound(new HttpSessionBindingEvent(this, name, oldValue)); + } + catch (Throwable th) { + logger.error("Error invoking session binding event listener", th); + } + } } @Override diff --git a/spring-session-core/src/test/java/org/springframework/session/web/http/SessionRepositoryFilterTests.java b/spring-session-core/src/test/java/org/springframework/session/web/http/SessionRepositoryFilterTests.java index 225e2d5b..cd1fad20 100644 --- a/spring-session-core/src/test/java/org/springframework/session/web/http/SessionRepositoryFilterTests.java +++ b/spring-session-core/src/test/java/org/springframework/session/web/http/SessionRepositoryFilterTests.java @@ -27,6 +27,8 @@ import java.util.Map; import java.util.NoSuchElementException; import java.util.concurrent.ConcurrentHashMap; import java.util.concurrent.TimeUnit; +import java.util.concurrent.atomic.AtomicBoolean; +import java.util.concurrent.atomic.AtomicInteger; import javax.servlet.FilterChain; import javax.servlet.ServletContext; @@ -36,6 +38,8 @@ import javax.servlet.http.HttpServlet; import javax.servlet.http.HttpServletRequest; import javax.servlet.http.HttpServletResponse; import javax.servlet.http.HttpSession; +import javax.servlet.http.HttpSessionBindingEvent; +import javax.servlet.http.HttpSessionBindingListener; import javax.servlet.http.HttpSessionContext; import org.assertj.core.data.Offset; @@ -1386,6 +1390,122 @@ public class SessionRepositoryFilterTests { .hasMessage("httpSessionIdResolver cannot be null"); } + @Test + public void bindingListenerBindListener() throws Exception { + String bindingListenerName = "bindingListener"; + CountingHttpSessionBindingListener bindingListener = new CountingHttpSessionBindingListener(); + + doFilter(new DoInFilter() { + + @Override + public void doFilter(HttpServletRequest wrappedRequest) { + HttpSession session = wrappedRequest.getSession(); + session.setAttribute(bindingListenerName, bindingListener); + } + + }); + + assertThat(bindingListener.getCounter()).isEqualTo(1); + } + + @Test + public void bindingListenerBindListenerThenUnbind() throws Exception { + String bindingListenerName = "bindingListener"; + CountingHttpSessionBindingListener bindingListener = new CountingHttpSessionBindingListener(); + + doFilter(new DoInFilter() { + + @Override + public void doFilter(HttpServletRequest wrappedRequest) { + HttpSession session = wrappedRequest.getSession(); + session.setAttribute(bindingListenerName, bindingListener); + session.removeAttribute(bindingListenerName); + } + + }); + + assertThat(bindingListener.getCounter()).isEqualTo(0); + } + + @Test + public void bindingListenerBindSameListenerTwice() throws Exception { + String bindingListenerName = "bindingListener"; + CountingHttpSessionBindingListener bindingListener = new CountingHttpSessionBindingListener(); + + doFilter(new DoInFilter() { + + @Override + public void doFilter(HttpServletRequest wrappedRequest) { + HttpSession session = wrappedRequest.getSession(); + session.setAttribute(bindingListenerName, bindingListener); + session.setAttribute(bindingListenerName, bindingListener); + } + + }); + + assertThat(bindingListener.getCounter()).isEqualTo(1); + } + + @Test + public void bindingListenerBindListenerOverwrite() throws Exception { + String bindingListenerName = "bindingListener"; + CountingHttpSessionBindingListener bindingListener1 = new CountingHttpSessionBindingListener(); + CountingHttpSessionBindingListener bindingListener2 = new CountingHttpSessionBindingListener(); + + doFilter(new DoInFilter() { + + @Override + public void doFilter(HttpServletRequest wrappedRequest) { + HttpSession session = wrappedRequest.getSession(); + session.setAttribute(bindingListenerName, bindingListener1); + session.setAttribute(bindingListenerName, bindingListener2); + } + + }); + + assertThat(bindingListener1.getCounter()).isEqualTo(0); + assertThat(bindingListener2.getCounter()).isEqualTo(1); + } + + @Test + public void bindingListenerBindThrowsException() throws Exception { + String bindingListenerName = "bindingListener"; + CountingHttpSessionBindingListener bindingListener = new CountingHttpSessionBindingListener(); + + doFilter(new DoInFilter() { + + @Override + public void doFilter(HttpServletRequest wrappedRequest) { + HttpSession session = wrappedRequest.getSession(); + bindingListener.setThrowException(); + session.setAttribute(bindingListenerName, bindingListener); + } + + }); + + assertThat(bindingListener.getCounter()).isEqualTo(0); + } + + @Test + public void bindingListenerBindListenerThenUnbindThrowsException() throws Exception { + String bindingListenerName = "bindingListener"; + CountingHttpSessionBindingListener bindingListener = new CountingHttpSessionBindingListener(); + + doFilter(new DoInFilter() { + + @Override + public void doFilter(HttpServletRequest wrappedRequest) { + HttpSession session = wrappedRequest.getSession(); + session.setAttribute(bindingListenerName, bindingListener); + bindingListener.setThrowException(); + session.removeAttribute(bindingListenerName); + } + + }); + + assertThat(bindingListener.getCounter()).isEqualTo(1); + } + // --- helper methods private void assertNewSession() { @@ -1488,4 +1608,39 @@ public class SessionRepositoryFilterTests { } + private static class CountingHttpSessionBindingListener + implements HttpSessionBindingListener { + + private final AtomicInteger counter = new AtomicInteger(0); + + private final AtomicBoolean throwException = new AtomicBoolean(false); + + @Override + public void valueBound(HttpSessionBindingEvent event) { + if (this.throwException.get()) { + this.throwException.compareAndSet(true, false); + throw new RuntimeException("bind exception"); + } + this.counter.incrementAndGet(); + } + + @Override + public void valueUnbound(HttpSessionBindingEvent event) { + if (this.throwException.get()) { + this.throwException.compareAndSet(true, false); + throw new RuntimeException("unbind exception"); + } + this.counter.decrementAndGet(); + } + + int getCounter() { + return this.counter.get(); + } + + void setThrowException() { + this.throwException.compareAndSet(false, true); + } + + } + }