22
33import static com .github .tomakehurst .wiremock .client .WireMock .*;
44import static org .assertj .core .api .AssertionsForClassTypes .assertThat ;
5+ import static org .assertj .core .api .AssertionsForClassTypes .assertThatThrownBy ;
56
67import com .github .tomakehurst .wiremock .junit5 .WireMockRuntimeInfo ;
78import com .github .tomakehurst .wiremock .junit5 .WireMockTest ;
9+ import io .github .hikingc .matrixsdk .api .auth .BrowserLauncher ;
810import io .github .hikingc .matrixsdk .api .auth .TokenMetadata ;
11+ import io .github .hikingc .matrixsdk .exceptions .MatrixException ;
12+ import java .io .IOException ;
13+ import java .net .ServerSocket ;
914import java .net .URI ;
15+ import java .net .URLDecoder ;
1016import java .net .http .HttpClient ;
17+ import java .net .http .HttpRequest ;
18+ import java .net .http .HttpResponse ;
19+ import java .nio .charset .StandardCharsets ;
20+ import java .security .MessageDigest ;
21+ import java .util .Arrays ;
22+ import java .util .Base64 ;
23+ import java .util .concurrent .atomic .AtomicReference ;
1124import org .junit .jupiter .api .BeforeEach ;
1225import org .junit .jupiter .api .Test ;
1326
1427@ WireMockTest
15- class MatrixAuthTest {
28+ class MatrixOAuthLoginTest {
1629
1730 private static MatrixAuth matrixAuth ;
31+ private static String baseUrl ;
32+ private int callbackPort ;
1833 private TokenMetadata tokens = new TokenMetadata ("ABCD" , null , null , null , null );
1934
2035 @ BeforeEach
21- void setupAuth (WireMockRuntimeInfo wireMockRuntimeInfo ) {
22- matrixAuth =
23- new MatrixAuth (
24- URI .create (wireMockRuntimeInfo .getHttpBaseUrl ()), HttpClient .newBuilder ().build ());
36+ void setupAuth (WireMockRuntimeInfo wireMockRuntimeInfo ) throws IOException {
37+ baseUrl = wireMockRuntimeInfo .getHttpBaseUrl ();
38+ matrixAuth = new MatrixAuth (URI .create (baseUrl ), HttpClient .newBuilder ().build ());
39+ callbackPort = findFreePort ();
40+
2541 stubFor (
2642 get (urlEqualTo ("/.well-known/matrix/client" ))
43+ .willReturn (
44+ okJson (
45+ """
46+ {"m.homeserver": {"base_url": "%s"}}
47+ """
48+ .formatted (baseUrl ))));
49+
50+ stubFor (
51+ get (urlEqualTo ("/_matrix/client/v1/auth_metadata" ))
52+ .willReturn (
53+ okJson (
54+ """
55+ {
56+ "issuer": "%1$s/",
57+ "authorization_endpoint": "%1$s/oauth2/auth",
58+ "token_endpoint": "%1$s/oauth2/token",
59+ "registration_endpoint": "%1$s/oauth2/clients/register",
60+ "revocation_endpoint": "%1$s/oauth2/revoke",
61+ "grant_types_supported": ["authorization_code", "refresh_token"],
62+ "response_types_supported": ["code"],
63+ "response_modes_supported": ["query", "fragment"],
64+ "code_challenge_methods_supported": ["S256"]
65+ }
66+ """
67+ .formatted (baseUrl ))));
68+
69+ stubFor (
70+ post (urlEqualTo ("/oauth2/clients/register" ))
2771 .willReturn (
2872 aResponse ()
29- .withStatus (200 )
73+ .withStatus (201 )
3074 .withHeader ("Content-Type" , "application/json" )
3175 .withBody (
32- "{\" m.homeserver\" : {\" base_url\" : \" "
33- + wireMockRuntimeInfo .getHttpBaseUrl ()
34- + "\" }}" )));
76+ """
77+ {"client_id": "test-client-id"}""" )));
78+ }
79+
80+ @ Test
81+ void performOAuthLogin_happyPath_returnsTokens () {
82+ stubFor (
83+ post (urlEqualTo ("/oauth2/token" ))
84+ .willReturn (
85+ okJson (
86+ """
87+ {
88+ "access_token": "syt_test_token",
89+ "token_type": "Bearer",
90+ "expires_in": 300,
91+ "refresh_token": "test_refresh",
92+ "scope": "urn:matrix:client:api:*"
93+ }
94+ """ )));
95+
96+ AtomicReference <String > capturedCodeChallenge = new AtomicReference <>();
97+ BrowserLauncher fakeBrowser =
98+ fakeBrowserCompleting (callbackPort , "fake-auth-code" , capturedCodeChallenge );
99+
100+ TokenMetadata token =
101+ matrixAuth .performOAuthLogin ("TestClient" , callbackPort , "TESTDEVICE01" , fakeBrowser );
102+
103+ assertThat (token .accessToken ()).isEqualTo ("syt_test_token" );
104+ assertThat (token .refreshToken ()).isEqualTo ("test_refresh" );
105+
106+ verify (
107+ postRequestedFor (urlEqualTo ("/oauth2/token" ))
108+ .withRequestBody (containing ("grant_type=authorization_code" ))
109+ .withRequestBody (containing ("code=fake-auth-code" ))
110+ .withRequestBody (containing ("client_id=test-client-id" )));
111+ }
112+
113+ @ Test
114+ void performOAuthLogin_verifiesPkceChallengeMatchesVerifier () throws Exception {
115+ stubFor (
116+ post (urlEqualTo ("/oauth2/token" ))
117+ .willReturn (
118+ okJson (
119+ """
120+ {
121+ "access_token": "syt_test_token",
122+ "token_type": "Bearer",
123+ "expires_in": 300,
124+ "refresh_token": "test_refresh",
125+ "scope": "urn:matrix:client:api:*"
126+ }
127+ """ )));
128+
129+ AtomicReference <String > capturedCodeChallenge = new AtomicReference <>();
130+ BrowserLauncher fakeBrowser =
131+ fakeBrowserCompleting (callbackPort , "fake-auth-code" , capturedCodeChallenge );
132+
133+ matrixAuth .performOAuthLogin ("TestClient" , callbackPort , "TESTDEVICE01" , fakeBrowser );
134+
135+ String sentVerifier = extractFormParam (lastTokenRequestBody (), "code_verifier" );
136+ String expectedChallenge = sha256Base64Url (sentVerifier );
137+
138+ assertThat (capturedCodeChallenge .get ())
139+ .as (
140+ "code_challenge sent in auth URI should equal SHA256(code_verifier) sent to token endpoint" )
141+ .isEqualTo (expectedChallenge );
142+ }
143+
144+ @ Test
145+ void performOAuthLogin_stateMismatch_throwsMatrixException () {
146+ BrowserLauncher maliciousBrowser =
147+ uri -> sendCallback (callbackPort , "state=not-the-real-state&code=whatever" );
148+
149+ assertThatThrownBy (
150+ () ->
151+ matrixAuth .performOAuthLogin (
152+ "TestClient" , callbackPort , "TESTDEVICE01" , maliciousBrowser ))
153+ .isInstanceOf (MatrixException .class )
154+ .hasMessageContaining ("A fatal error has ceased authorization flow." );
155+ }
156+
157+ @ Test
158+ void performOAuthLogin_authorizationDenied_throwsMatrixException () {
159+ BrowserLauncher denyingBrowser =
160+ uri -> {
161+ String state = extractQueryParam (uri .getRawQuery (), "state" );
162+ sendCallback (
163+ callbackPort ,
164+ "state=" + state + "&error=access_denied&error_description=User+declined" );
165+ };
166+
167+ assertThatThrownBy (
168+ () ->
169+ matrixAuth .performOAuthLogin (
170+ "TestClient" , callbackPort , "TESTDEVICE01" , denyingBrowser ))
171+ .isInstanceOf (MatrixException .class )
172+ .hasMessageContaining ("A fatal error has ceased authorization flow." );
173+ }
174+
175+ // ---------------------------------------------------------------------
176+ // Helpers
177+ // ---------------------------------------------------------------------
178+
179+ private BrowserLauncher fakeBrowserCompleting (
180+ int port , String fakeCode , AtomicReference <String > capturedCodeChallenge ) {
181+ return uri -> {
182+ String query = uri .getRawQuery ();
183+ String state = extractQueryParam (query , "state" );
184+ capturedCodeChallenge .set (extractQueryParam (query , "code_challenge" ));
185+ sendCallback (port , "state=" + state + "&code=" + fakeCode );
186+ };
187+ }
188+
189+ private static void sendCallback (int port , String rawQuery ) {
190+ try {
191+ HttpClient .newHttpClient ()
192+ .send (
193+ HttpRequest .newBuilder (
194+ URI .create ("http://127.0.0.1:" + port + "/callback?" + rawQuery ))
195+ .GET ()
196+ .build (),
197+ HttpResponse .BodyHandlers .discarding ());
198+ } catch (Exception e ) {
199+ throw new RuntimeException ("Fake browser failed to hit callback" , e );
200+ }
201+ }
202+
203+ private static int findFreePort () throws IOException {
204+ try (ServerSocket socket = new ServerSocket (0 )) {
205+ return socket .getLocalPort ();
206+ }
207+ }
208+
209+ private static String extractQueryParam (String query , String key ) {
210+ if (query == null ) return null ;
211+ return Arrays .stream (query .split ("&" ))
212+ .map (pair -> pair .split ("=" , 2 ))
213+ .filter (kv -> kv [0 ].equals (key ))
214+ .map (kv -> kv .length > 1 ? URLDecoder .decode (kv [1 ], StandardCharsets .UTF_8 ) : "" )
215+ .findFirst ()
216+ .orElse (null );
217+ }
218+
219+ private static String extractFormParam (String body , String key ) {
220+ return extractQueryParam (body , key );
221+ }
222+
223+ private String lastTokenRequestBody () {
224+ return findAll (postRequestedFor (urlEqualTo ("/oauth2/token" ))).get (0 ).getBodyAsString ();
225+ }
226+
227+ private static String sha256Base64Url (String input ) throws Exception {
228+ MessageDigest digest = MessageDigest .getInstance ("SHA-256" );
229+ byte [] hash = digest .digest (input .getBytes (StandardCharsets .UTF_8 ));
230+ return Base64 .getUrlEncoder ().withoutPadding ().encodeToString (hash );
35231 }
36232
37233 @ Test
@@ -42,11 +238,11 @@ void getCurrentAccountInformation_WithACorrectPayload_ReturnAnObject() {
42238 .willReturn (
43239 okJson (
44240 """
45- {
46- "device_id": "ABC1234",
47- "user_id": "@joe:example.org"
48- }
49- """ )));
241+ {
242+ "device_id": "ABC1234",
243+ "user_id": "@joe:example.org"
244+ }
245+ """ )));
50246 var response = matrixAuth .getCurrentAccountInformation ("token" );
51247 assertThat (response ).isNotNull ();
52248 }
@@ -59,22 +255,22 @@ void getFetchWellKnown() {
59255 .willReturn (
60256 okJson (
61257 """
62- {
63- "contacts": [
64- {
65- "email_address": "admin@example.org",
66- "matrix_id": "@admin:example.org",
67- "role": "m.role.admin"
68- },
69- {
70- "email_address": "security@example.org",
71- "role": "m.role.security"
72- }
73- ],
74- "support_page": "https://example.org/support.html"
75- }
76-
77- """ )));
258+ {
259+ "contacts": [
260+ {
261+ "email_address": "admin@example.org",
262+ "matrix_id": "@admin:example.org",
263+ "role": "m.role.admin"
264+ },
265+ {
266+ "email_address": "security@example.org",
267+ "role": "m.role.security"
268+ }
269+ ],
270+ "support_page": "https://example.org/support.html"
271+ }
272+
273+ """ )));
78274 var response = matrixAuth .fetchWellKnown ();
79275 assertThat (response ).isNotNull ();
80276 }
@@ -87,16 +283,16 @@ void getVersions_WithACorrectPayload_ReturnAnObject() {
87283 .willReturn (
88284 okJson (
89285 """
90- {
91- "unstable_features": {
92- "org.example.my_feature": true
93- },
94- "versions": [
95- "r0.0.1",
96- "v1.1"
97- ]
98- }
99- """ )));
286+ {
287+ "unstable_features": {
288+ "org.example.my_feature": true
289+ },
290+ "versions": [
291+ "r0.0.1",
292+ "v1.1"
293+ ]
294+ }
295+ """ )));
100296 assertThat (tokens .accessToken ()).isNotNull ();
101297 var response = matrixAuth .getVersions (tokens .accessToken ());
102298 assertThat (response ).isNotNull ();
0 commit comments