Skip to content

Commit c8a3b57

Browse files
Validate authorization response state (#139)
Co-authored-by: Copilot App <223556219+Copilot@users.noreply.github.com>
1 parent 589c62e commit c8a3b57

2 files changed

Lines changed: 106 additions & 5 deletions

File tree

‎Common/msal-node-wrapper/src/middleware/handlers/redirectHandler.ts‎

Lines changed: 6 additions & 5 deletions
Original file line numberDiff line numberDiff line change
@@ -5,7 +5,7 @@
55

66
import { Request, Response, NextFunction, RequestHandler } from "express";
77
import { StringUtils } from "@azure/msal-common";
8-
import { AuthorizationCodePayload, AuthorizationCodeRequest } from "@azure/msal-node";
8+
import { AuthorizationCodeRequest } from "@azure/msal-node";
99
import { WebAppAuthProvider } from "../../provider/WebAppAuthProvider";
1010
import { AppState } from "../MiddlewareOptions";
1111
import { EMPTY_STRING, ErrorMessages } from "../../utils/Constants";
@@ -18,6 +18,10 @@ function redirectHandler(this: WebAppAuthProvider): RequestHandler {
1818
return next(new Error(ErrorMessages.AUTH_CODE_RESPONSE_NOT_FOUND));
1919
}
2020

21+
if (!req.body.state || req.body.state !== req.session.tokenRequestParams?.state) {
22+
return next(new Error(ErrorMessages.CSRF_TOKEN_MISMATCH));
23+
}
24+
2125
const tokenRequest = {
2226
...req.session.tokenRequestParams,
2327
code: req.body.code as string
@@ -30,10 +34,7 @@ function redirectHandler(this: WebAppAuthProvider): RequestHandler {
3034
msalInstance.getTokenCache().deserialize(req.session.tokenCache);
3135
}
3236

33-
const tokenResponse = await msalInstance.acquireTokenByCode(
34-
tokenRequest,
35-
req.body as AuthorizationCodePayload
36-
);
37+
const tokenResponse = await msalInstance.acquireTokenByCode(tokenRequest);
3738

3839
req.session.tokenCache = msalInstance.getTokenCache().serialize();
3940
// eslint-disable-next-line @typescript-eslint/no-non-null-assertion
Lines changed: 100 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,100 @@
1+
/*
2+
* Copyright (c) Microsoft Corporation. All rights reserved.
3+
* Licensed under the MIT License.
4+
*/
5+
6+
import { NextFunction, Request, Response } from "express";
7+
import { WebAppAuthProvider } from "../../../src/provider/WebAppAuthProvider";
8+
import redirectHandler from "../../../src/middleware/handlers/redirectHandler";
9+
import { ErrorMessages } from "../../../src/utils/Constants";
10+
11+
describe("redirectHandler", () => {
12+
const state = "encoded-state";
13+
const account = { homeAccountId: "home-account-id" };
14+
15+
const createHandler = () => {
16+
const acquireTokenByCode = jest.fn().mockResolvedValue({ account });
17+
const tokenCache = {
18+
deserialize: jest.fn(),
19+
serialize: jest.fn().mockReturnValue("serialized-cache"),
20+
};
21+
const provider = {
22+
getLogger: () => ({ trace: jest.fn() }),
23+
getMsalClient: () => ({
24+
acquireTokenByCode,
25+
getTokenCache: () => tokenCache,
26+
}),
27+
getCryptoProvider: () => ({
28+
base64Decode: jest.fn().mockReturnValue(JSON.stringify({ redirectTo: "/target" })),
29+
}),
30+
} as unknown as WebAppAuthProvider;
31+
32+
return {
33+
acquireTokenByCode,
34+
handler: redirectHandler.call(provider),
35+
};
36+
};
37+
38+
it("redeems the code with the stored request when state matches", async () => {
39+
const { acquireTokenByCode, handler } = createHandler();
40+
const req = {
41+
body: {
42+
code: "authorization-code",
43+
state,
44+
unexpected: "authorization-response-data",
45+
},
46+
session: {
47+
tokenRequestParams: {
48+
code: "",
49+
scopes: ["openid"],
50+
state,
51+
redirectUri: "http://localhost/redirect",
52+
},
53+
},
54+
} as unknown as Request;
55+
const res = {
56+
redirect: jest.fn(),
57+
} as unknown as Response;
58+
const next = jest.fn() as NextFunction;
59+
60+
await handler(req, res, next);
61+
62+
expect(acquireTokenByCode).toHaveBeenCalledTimes(1);
63+
expect(acquireTokenByCode).toHaveBeenCalledWith({
64+
code: "authorization-code",
65+
scopes: ["openid"],
66+
state,
67+
redirectUri: "http://localhost/redirect",
68+
});
69+
expect(req.session.tokenCache).toBe("serialized-cache");
70+
expect(req.session.account).toBe(account);
71+
expect(req.session.isAuthenticated).toBe(true);
72+
expect(res.redirect).toHaveBeenCalledWith("/target");
73+
expect(next).not.toHaveBeenCalled();
74+
});
75+
76+
it.each([
77+
["missing", undefined],
78+
["mismatched", "different-state"],
79+
])("rejects a %s state before redeeming the code", async (_description, responseState) => {
80+
const { acquireTokenByCode, handler } = createHandler();
81+
const req = {
82+
body: {
83+
code: "authorization-code",
84+
state: responseState,
85+
},
86+
session: {
87+
tokenRequestParams: {
88+
state,
89+
},
90+
},
91+
} as unknown as Request;
92+
const res = {} as Response;
93+
const next = jest.fn() as NextFunction;
94+
95+
await handler(req, res, next);
96+
97+
expect(next).toHaveBeenCalledWith(new Error(ErrorMessages.CSRF_TOKEN_MISMATCH));
98+
expect(acquireTokenByCode).not.toHaveBeenCalled();
99+
});
100+
});

0 commit comments

Comments
 (0)