Skip to content

Commit 47f57a8

Browse files
committed
AB#69392 Key the LTI launch session off the state.
To allow multiple launches to happen concurrently we key the session storage off the state value for the launch (which should be unique), and this way multiple tools can launch concurrently.
1 parent 02a7811 commit 47f57a8

3 files changed

Lines changed: 233 additions & 0 deletions

File tree

src/main/java/uk/ac/ox/ctl/ltiauth/CustomLti13Configurer.java

Lines changed: 14 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -1,13 +1,18 @@
11
package uk.ac.ox.ctl.ltiauth;
22

33
import org.springframework.security.oauth2.client.registration.ClientRegistrationRepository;
4+
import org.springframework.security.oauth2.client.web.AuthorizationRequestRepository;
5+
import org.springframework.security.oauth2.core.endpoint.OAuth2AuthorizationRequest;
46
import uk.ac.ox.ctl.lti13.Lti13Configurer;
57
import uk.ac.ox.ctl.lti13.security.oauth2.client.lti.authentication.OidcLaunchFlowAuthenticationProvider;
68
import uk.ac.ox.ctl.lti13.security.oauth2.client.lti.web.OAuth2AuthorizationRequestRedirectFilter;
79
import uk.ac.ox.ctl.lti13.security.oauth2.client.lti.web.OAuth2LoginAuthenticationFilter;
810
import uk.ac.ox.ctl.lti13.security.oauth2.client.lti.web.OIDCInitiatingLoginRequestResolver;
911
import uk.ac.ox.ctl.lti13.security.oauth2.client.lti.web.OptimisticAuthorizationRequestRepository;
1012
import uk.ac.ox.ctl.lti13.security.oauth2.client.lti.web.PathOIDCInitiationRegistrationResolver;
13+
import uk.ac.ox.ctl.lti13.security.oauth2.client.lti.web.StateAuthorizationRequestRepository;
14+
15+
import java.time.Duration;
1116

1217
/**
1318
* This overrides the standard configurer to add our token passing redirect along with allowing client
@@ -25,6 +30,15 @@ public CustomLti13Configurer(JWTService jwtService, ClientRegistrationService cl
2530
this.clientRegistrationService = clientRegistrationService;
2631
}
2732

33+
@Override
34+
protected OptimisticAuthorizationRequestRepository configureRequestRepository() {
35+
AuthorizationRequestRepository<OAuth2AuthorizationRequest> sessionRepository =
36+
new MultiStateHttpSessionOAuth2AuthorizationRequestRepository();
37+
StateAuthorizationRequestRepository stateRepository = new StateAuthorizationRequestRepository(Duration.ofMinutes(1));
38+
stateRepository.setLimitIpAddress(limitIpAddresses);
39+
return new OptimisticAuthorizationRequestRepository(sessionRepository, stateRepository);
40+
}
41+
2842
@Override
2943
protected OAuth2LoginAuthenticationFilter configureLoginFilter(ClientRegistrationRepository clientRegistrationRepository, OidcLaunchFlowAuthenticationProvider oidcLaunchFlowAuthenticationProvider, OptimisticAuthorizationRequestRepository authorizationRequestRepository) {
3044
OAuth2LoginAuthenticationFilter loginFilter = super.configureLoginFilter(clientRegistrationRepository, oidcLaunchFlowAuthenticationProvider, authorizationRequestRepository);
Lines changed: 118 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,118 @@
1+
package uk.ac.ox.ctl.ltiauth;
2+
3+
import jakarta.servlet.http.HttpServletRequest;
4+
import jakarta.servlet.http.HttpServletResponse;
5+
import jakarta.servlet.http.HttpSession;
6+
import org.springframework.security.oauth2.client.web.AuthorizationRequestRepository;
7+
import org.springframework.security.oauth2.client.web.HttpSessionOAuth2AuthorizationRequestRepository;
8+
import org.springframework.security.oauth2.core.endpoint.OAuth2AuthorizationRequest;
9+
import org.springframework.security.oauth2.core.endpoint.OAuth2ParameterNames;
10+
import org.springframework.util.Assert;
11+
12+
import java.util.LinkedHashMap;
13+
import java.util.Map;
14+
15+
/**
16+
* Stores authorization requests in the HTTP session by OAuth state so concurrent LTI launches in the same browser
17+
* session do not overwrite each other.
18+
*/
19+
public class MultiStateHttpSessionOAuth2AuthorizationRequestRepository implements AuthorizationRequestRepository<OAuth2AuthorizationRequest> {
20+
21+
static final int MAX_AUTHORIZATION_REQUESTS = 100;
22+
static final String AUTHORIZATION_REQUESTS_ATTR_NAME =
23+
MultiStateHttpSessionOAuth2AuthorizationRequestRepository.class.getName() + ".AUTHORIZATION_REQUESTS";
24+
private static final String LEGACY_AUTHORIZATION_REQUEST_ATTR_NAME =
25+
HttpSessionOAuth2AuthorizationRequestRepository.class.getName() + ".AUTHORIZATION_REQUEST";
26+
27+
@Override
28+
public OAuth2AuthorizationRequest loadAuthorizationRequest(HttpServletRequest request) {
29+
Assert.notNull(request, "request cannot be null");
30+
String state = request.getParameter(OAuth2ParameterNames.STATE);
31+
if (state == null) {
32+
return null;
33+
}
34+
HttpSession session = request.getSession(false);
35+
if (session == null) {
36+
return null;
37+
}
38+
39+
OAuth2AuthorizationRequest authorizationRequest = getOrCreateAuthorizationRequests(session).get(state);
40+
if (authorizationRequest != null) {
41+
return authorizationRequest;
42+
}
43+
44+
Object legacyAuthorizationRequest = session.getAttribute(LEGACY_AUTHORIZATION_REQUEST_ATTR_NAME);
45+
if (legacyAuthorizationRequest instanceof OAuth2AuthorizationRequest legacyRequest && state.equals(legacyRequest.getState())) {
46+
return legacyRequest;
47+
}
48+
return null;
49+
}
50+
51+
@Override
52+
public void saveAuthorizationRequest(OAuth2AuthorizationRequest authorizationRequest, HttpServletRequest request, HttpServletResponse response) {
53+
Assert.notNull(request, "request cannot be null");
54+
Assert.notNull(response, "response cannot be null");
55+
56+
if (authorizationRequest == null) {
57+
removeAuthorizationRequest(request, response);
58+
return;
59+
}
60+
61+
String state = authorizationRequest.getState();
62+
Assert.hasText(state, "authorizationRequest.state cannot be empty");
63+
64+
HttpSession session = request.getSession();
65+
Map<String, OAuth2AuthorizationRequest> authorizationRequests = getOrCreateAuthorizationRequests(session);
66+
authorizationRequests.put(state, authorizationRequest);
67+
// Remove any legacy single-entry value so new launches always use the state-keyed storage.
68+
session.removeAttribute(LEGACY_AUTHORIZATION_REQUEST_ATTR_NAME);
69+
}
70+
71+
@Override
72+
public OAuth2AuthorizationRequest removeAuthorizationRequest(HttpServletRequest request, HttpServletResponse response) {
73+
Assert.notNull(request, "request cannot be null");
74+
Assert.notNull(response, "response cannot be null");
75+
76+
String state = request.getParameter(OAuth2ParameterNames.STATE);
77+
if (state == null) {
78+
return null;
79+
}
80+
81+
HttpSession session = request.getSession(false);
82+
if (session == null) {
83+
return null;
84+
}
85+
86+
Map<String, OAuth2AuthorizationRequest> authorizationRequests = getOrCreateAuthorizationRequests(session);
87+
OAuth2AuthorizationRequest authorizationRequest = authorizationRequests.remove(state);
88+
if (authorizationRequests.isEmpty()) {
89+
session.removeAttribute(AUTHORIZATION_REQUESTS_ATTR_NAME);
90+
}
91+
92+
if (authorizationRequest != null) {
93+
return authorizationRequest;
94+
}
95+
96+
Object legacyAuthorizationRequest = session.getAttribute(LEGACY_AUTHORIZATION_REQUEST_ATTR_NAME);
97+
if (legacyAuthorizationRequest instanceof OAuth2AuthorizationRequest legacyRequest && state.equals(legacyRequest.getState())) {
98+
session.removeAttribute(LEGACY_AUTHORIZATION_REQUEST_ATTR_NAME);
99+
return legacyRequest;
100+
}
101+
return null;
102+
}
103+
104+
@SuppressWarnings("unchecked")
105+
private Map<String, OAuth2AuthorizationRequest> getOrCreateAuthorizationRequests(HttpSession session) {
106+
Object attribute = session.getAttribute(AUTHORIZATION_REQUESTS_ATTR_NAME);
107+
if (attribute instanceof Map<?, ?> map) {
108+
return (Map<String, OAuth2AuthorizationRequest>) map;
109+
}
110+
session.setAttribute(AUTHORIZATION_REQUESTS_ATTR_NAME, new LinkedHashMap<String, OAuth2AuthorizationRequest>() {
111+
@Override
112+
protected boolean removeEldestEntry(Map.Entry<String, OAuth2AuthorizationRequest> eldest) {
113+
return size() > MAX_AUTHORIZATION_REQUESTS;
114+
}
115+
});
116+
return (Map<String, OAuth2AuthorizationRequest>) session.getAttribute(AUTHORIZATION_REQUESTS_ATTR_NAME);
117+
}
118+
}
Lines changed: 101 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,101 @@
1+
package uk.ac.ox.ctl.ltiauth;
2+
3+
import org.junit.jupiter.api.Test;
4+
import org.springframework.mock.web.MockHttpServletRequest;
5+
import org.springframework.mock.web.MockHttpServletResponse;
6+
import org.springframework.mock.web.MockHttpSession;
7+
import org.springframework.security.oauth2.client.web.HttpSessionOAuth2AuthorizationRequestRepository;
8+
import org.springframework.security.oauth2.core.endpoint.OAuth2AuthorizationRequest;
9+
import org.springframework.security.oauth2.core.endpoint.OAuth2ParameterNames;
10+
11+
import java.util.Map;
12+
13+
import static org.junit.jupiter.api.Assertions.assertEquals;
14+
import static org.junit.jupiter.api.Assertions.assertNotNull;
15+
import static org.junit.jupiter.api.Assertions.assertNull;
16+
17+
class MultiStateHttpSessionOAuth2AuthorizationRequestRepositoryTest {
18+
19+
private final MultiStateHttpSessionOAuth2AuthorizationRequestRepository repository =
20+
new MultiStateHttpSessionOAuth2AuthorizationRequestRepository();
21+
22+
@Test
23+
void storesConcurrentRequestsByState() {
24+
MockHttpSession session = new MockHttpSession();
25+
MockHttpServletResponse response = new MockHttpServletResponse();
26+
27+
repository.saveAuthorizationRequest(authorizationRequest("state-1"), requestWithSession(session), response);
28+
repository.saveAuthorizationRequest(authorizationRequest("state-2"), requestWithSession(session), response);
29+
30+
assertEquals("state-1", repository.loadAuthorizationRequest(callbackRequest(session, "state-1")).getState());
31+
assertEquals("state-2", repository.loadAuthorizationRequest(callbackRequest(session, "state-2")).getState());
32+
}
33+
34+
@Test
35+
void removesOnlyMatchingState() {
36+
MockHttpSession session = new MockHttpSession();
37+
MockHttpServletResponse response = new MockHttpServletResponse();
38+
39+
repository.saveAuthorizationRequest(authorizationRequest("state-1"), requestWithSession(session), response);
40+
repository.saveAuthorizationRequest(authorizationRequest("state-2"), requestWithSession(session), response);
41+
42+
OAuth2AuthorizationRequest removed = repository.removeAuthorizationRequest(callbackRequest(session, "state-1"), response);
43+
44+
assertNotNull(removed);
45+
assertEquals("state-1", removed.getState());
46+
assertNull(repository.loadAuthorizationRequest(callbackRequest(session, "state-1")));
47+
assertEquals("state-2", repository.loadAuthorizationRequest(callbackRequest(session, "state-2")).getState());
48+
}
49+
50+
@Test
51+
void canReadLegacySingleRequestEntry() {
52+
MockHttpSession session = new MockHttpSession();
53+
OAuth2AuthorizationRequest authorizationRequest = authorizationRequest("legacy-state");
54+
session.setAttribute(HttpSessionOAuth2AuthorizationRequestRepository.class.getName() + ".AUTHORIZATION_REQUEST", authorizationRequest);
55+
56+
assertEquals("legacy-state", repository.loadAuthorizationRequest(callbackRequest(session, "legacy-state")).getState());
57+
assertEquals("legacy-state", repository.removeAuthorizationRequest(callbackRequest(session, "legacy-state"), new MockHttpServletResponse()).getState());
58+
assertNull(session.getAttribute(HttpSessionOAuth2AuthorizationRequestRepository.class.getName() + ".AUTHORIZATION_REQUEST"));
59+
}
60+
61+
@Test
62+
void storesAtMostHundredRequests() {
63+
MockHttpSession session = new MockHttpSession();
64+
MockHttpServletResponse response = new MockHttpServletResponse();
65+
66+
for (int i = 1; i <= 101; i++) {
67+
repository.saveAuthorizationRequest(authorizationRequest("state-" + i), requestWithSession(session), response);
68+
}
69+
70+
@SuppressWarnings("unchecked")
71+
Map<String, OAuth2AuthorizationRequest> authorizationRequests =
72+
(Map<String, OAuth2AuthorizationRequest>) session.getAttribute(MultiStateHttpSessionOAuth2AuthorizationRequestRepository.AUTHORIZATION_REQUESTS_ATTR_NAME);
73+
74+
assertNotNull(authorizationRequests);
75+
assertEquals(100, authorizationRequests.size());
76+
assertNull(repository.loadAuthorizationRequest(callbackRequest(session, "state-1")));
77+
assertEquals("state-101", repository.loadAuthorizationRequest(callbackRequest(session, "state-101")).getState());
78+
}
79+
80+
private MockHttpServletRequest requestWithSession(MockHttpSession session) {
81+
MockHttpServletRequest request = new MockHttpServletRequest();
82+
request.setSession(session);
83+
return request;
84+
}
85+
86+
private MockHttpServletRequest callbackRequest(MockHttpSession session, String state) {
87+
MockHttpServletRequest request = requestWithSession(session);
88+
request.setParameter(OAuth2ParameterNames.STATE, state);
89+
return request;
90+
}
91+
92+
private OAuth2AuthorizationRequest authorizationRequest(String state) {
93+
return OAuth2AuthorizationRequest.authorizationCode()
94+
.authorizationUri("https://canvas.example/authorize")
95+
.clientId("client-id")
96+
.redirectUri("https://tool.example/lti/login")
97+
.state(state)
98+
.authorizationRequestUri("https://canvas.example/authorize?state=" + state)
99+
.build();
100+
}
101+
}

0 commit comments

Comments
 (0)