CsrfWebFilter places Mono<CsrfToken>

Fixes: gh-4855
This commit is contained in:
Rob Winch
2017-11-20 14:16:49 -06:00
parent edccafca84
commit d55db837e1
11 changed files with 73 additions and 114 deletions

View File

@@ -30,6 +30,10 @@ import java.util.regex.Pattern;
* @since 5.0
*/
public class CsrfRequestDataValueProcessor implements RequestDataValueProcessor {
/**
* The default request attribute to look for a {@link CsrfToken}.
*/
public static final String DEFAULT_CSRF_ATTR_NAME = "_csrf";
private static final Pattern DISABLE_CSRF_TOKEN_PATTERN = Pattern
.compile("(?i)^(GET|HEAD|TRACE|OPTIONS)$");
@@ -62,7 +66,7 @@ public class CsrfRequestDataValueProcessor implements RequestDataValueProcessor
exchange.getAttributes().remove(DISABLE_CSRF_TOKEN_ATTR);
return Collections.emptyMap();
}
CsrfToken token = exchange.getAttribute(CsrfToken.class.getName());
CsrfToken token = exchange.getAttribute(DEFAULT_CSRF_ATTR_NAME);
if(token == null) {
return Collections.emptyMap();
}

View File

@@ -47,12 +47,16 @@ import java.util.Set;
* {@link WebSessionServerCsrfTokenRepository}. This is preferred to storing the token in
* a cookie which can be modified by a client application.
* </p>
* <p>
* The {@code Mono&lt;CsrfToken&gt;} is exposes as a request attribute with the name of
* {@code CsrfToken.class.getName()}. If the token is new it will automatically be saved
* at the time it is subscribed.
* </p>
*
* @author Rob Winch
* @since 5.0
*/
public class CsrfWebFilter implements WebFilter {
private ServerWebExchangeMatcher requireCsrfProtectionMatcher = new DefaultRequireCsrfProtectionMatcher();
private ServerCsrfTokenRepository csrfTokenRepository = new WebSessionServerCsrfTokenRepository();
@@ -105,11 +109,11 @@ public class CsrfWebFilter implements WebFilter {
}
private Mono<Void> continueFilterChain(ServerWebExchange exchange, WebFilterChain chain) {
return csrfToken(exchange)
.doOnSuccess(csrfToken -> exchange.getAttributes().put(CsrfToken.class.getName(), csrfToken))
.doOnSuccess(csrfToken -> exchange.getAttributes().put(csrfToken.getParameterName(), csrfToken))
.flatMap( t -> chain.filter(exchange))
.then();
return Mono.defer(() ->{
Mono<CsrfToken> csrfToken = csrfToken(exchange);
exchange.getAttributes().put(CsrfToken.class.getName(), csrfToken);
return chain.filter(exchange);
});
}
private Mono<CsrfToken> csrfToken(ServerWebExchange exchange) {

View File

@@ -17,7 +17,6 @@ package org.springframework.security.web.server.csrf;
import org.springframework.util.Assert;
import org.springframework.web.server.ServerWebExchange;
import org.springframework.web.server.WebSession;
import reactor.core.publisher.Mono;
import javax.servlet.http.HttpServletRequest;
@@ -49,20 +48,15 @@ public class WebSessionServerCsrfTokenRepository
@Override
public Mono<CsrfToken> generateToken(ServerWebExchange exchange) {
return exchange.getSession()
.map(WebSession::getAttributes)
.map(this::createCsrfToken);
return Mono.fromCallable(() -> createCsrfToken());
}
@Override
public Mono<CsrfToken> saveToken(ServerWebExchange exchange, CsrfToken token) {
if(token != null) {
return Mono.just(token);
}
return exchange.getSession()
.doOnSuccess(session -> putToken(session.getAttributes(), token))
.doOnNext(session -> putToken(session.getAttributes(), token))
.flatMap(session -> session.changeSessionId())
.flatMap(r -> Mono.justOrEmpty(token));
.then(Mono.justOrEmpty(token));
}
private void putToken(Map<String, Object> attributes, CsrfToken token) {
@@ -111,11 +105,6 @@ public class WebSessionServerCsrfTokenRepository
this.sessionAttributeName = sessionAttributeName;
}
private CsrfToken createCsrfToken(Map<String, Object> attributes) {
return new LazyCsrfToken(attributes, createCsrfToken());
}
private CsrfToken createCsrfToken() {
return new DefaultCsrfToken(this.headerName, this.parameterName, createNewToken());
}
@@ -124,58 +113,4 @@ public class WebSessionServerCsrfTokenRepository
return UUID.randomUUID().toString();
}
private class LazyCsrfToken implements CsrfToken {
private final Map<String, Object> attributes;
private final CsrfToken delegate;
private LazyCsrfToken(Map<String, Object> attributes, CsrfToken delegate) {
this.attributes = attributes;
this.delegate = delegate;
}
@Override
public String getHeaderName() {
return this.delegate.getHeaderName();
}
@Override
public String getParameterName() {
return this.delegate.getParameterName();
}
@Override
public String getToken() {
putToken(this.attributes, this.delegate);
return this.delegate.getToken();
}
@Override
public boolean equals(Object o) {
if (this == o)
return true;
if (o == null || !(o instanceof CsrfToken))
return false;
CsrfToken that = (CsrfToken) o;
if (!getToken().equals(that.getToken()))
return false;
if (!getParameterName().equals(that.getParameterName()))
return false;
return getHeaderName().equals(that.getHeaderName());
}
@Override
public int hashCode() {
int result = getToken().hashCode();
result = 31 * result + getParameterName().hashCode();
result = 31 * result + getHeaderName().hashCode();
return result;
}
@Override
public String toString() {
return "LazyCsrfToken{" + "delegate=" + this.delegate + '}';
}
}
}

View File

@@ -60,8 +60,8 @@ public class LoginPageGeneratingWebFilter implements WebFilter {
private Mono<DataBuffer> createBuffer(ServerWebExchange exchange) {
MultiValueMap<String, String> queryParams = exchange.getRequest()
.getQueryParams();
CsrfToken token = exchange.getAttribute(CsrfToken.class.getName());
return Mono.justOrEmpty(token)
Mono<CsrfToken> token = exchange.getAttributeOrDefault(CsrfToken.class.getName(), Mono.empty());
return token
.map(LoginPageGeneratingWebFilter::csrfToken)
.defaultIfEmpty("")
.map(csrfTokenHtmlInput -> {

View File

@@ -57,8 +57,8 @@ public class LogoutPageGeneratingWebFilter implements WebFilter {
}
private Mono<DataBuffer> createBuffer(ServerWebExchange exchange) {
CsrfToken token = exchange.getAttribute(CsrfToken.class.getName());
return Mono.justOrEmpty(token)
Mono<CsrfToken> token = exchange.getAttributeOrDefault(CsrfToken.class.getName(), Mono.empty());
return token
.map(LogoutPageGeneratingWebFilter::csrfToken)
.defaultIfEmpty("")
.map(csrfTokenHtmlInput -> {