From 34776a139655e1cde8ff6d529ff6ab697c7634f4 Mon Sep 17 00:00:00 2001 From: Robbie Ginsburg Date: Wed, 23 Sep 2026 20:20:42 -0400 Subject: [PATCH] Validate authorization response state Co-authored-by: Copilot App <223556219+Copilot@users.noreply.github.com> --- .../middleware/handlers/redirectHandler.ts | 11 +- .../handlers/redirectHandler.spec.ts | 100 ++++++++++++++++++ 2 files changed, 106 insertions(+), 5 deletions(-) create mode 100644 Common/msal-node-wrapper/test/middleware/handlers/redirectHandler.spec.ts diff --git a/Common/msal-node-wrapper/src/middleware/handlers/redirectHandler.ts b/Common/msal-node-wrapper/src/middleware/handlers/redirectHandler.ts index 73a4c49..78be397 100644 --- a/Common/msal-node-wrapper/src/middleware/handlers/redirectHandler.ts +++ b/Common/msal-node-wrapper/src/middleware/handlers/redirectHandler.ts @@ -5,7 +5,7 @@ import { Request, Response, NextFunction, RequestHandler } from "express"; import { StringUtils } from "@azure/msal-common"; -import { AuthorizationCodePayload, AuthorizationCodeRequest } from "@azure/msal-node"; +import { AuthorizationCodeRequest } from "@azure/msal-node"; import { WebAppAuthProvider } from "../../provider/WebAppAuthProvider"; import { AppState } from "../MiddlewareOptions"; import { EMPTY_STRING, ErrorMessages } from "../../utils/Constants"; @@ -18,6 +18,10 @@ function redirectHandler(this: WebAppAuthProvider): RequestHandler { return next(new Error(ErrorMessages.AUTH_CODE_RESPONSE_NOT_FOUND)); } + if (!req.body.state || req.body.state !== req.session.tokenRequestParams?.state) { + return next(new Error(ErrorMessages.CSRF_TOKEN_MISMATCH)); + } + const tokenRequest = { ...req.session.tokenRequestParams, code: req.body.code as string @@ -30,10 +34,7 @@ function redirectHandler(this: WebAppAuthProvider): RequestHandler { msalInstance.getTokenCache().deserialize(req.session.tokenCache); } - const tokenResponse = await msalInstance.acquireTokenByCode( - tokenRequest, - req.body as AuthorizationCodePayload - ); + const tokenResponse = await msalInstance.acquireTokenByCode(tokenRequest); req.session.tokenCache = msalInstance.getTokenCache().serialize(); // eslint-disable-next-line @typescript-eslint/no-non-null-assertion diff --git a/Common/msal-node-wrapper/test/middleware/handlers/redirectHandler.spec.ts b/Common/msal-node-wrapper/test/middleware/handlers/redirectHandler.spec.ts new file mode 100644 index 0000000..68aa4d0 --- /dev/null +++ b/Common/msal-node-wrapper/test/middleware/handlers/redirectHandler.spec.ts @@ -0,0 +1,100 @@ +/* + * Copyright (c) Microsoft Corporation. All rights reserved. + * Licensed under the MIT License. + */ + +import { NextFunction, Request, Response } from "express"; +import { WebAppAuthProvider } from "../../../src/provider/WebAppAuthProvider"; +import redirectHandler from "../../../src/middleware/handlers/redirectHandler"; +import { ErrorMessages } from "../../../src/utils/Constants"; + +describe("redirectHandler", () => { + const state = "encoded-state"; + const account = { homeAccountId: "home-account-id" }; + + const createHandler = () => { + const acquireTokenByCode = jest.fn().mockResolvedValue({ account }); + const tokenCache = { + deserialize: jest.fn(), + serialize: jest.fn().mockReturnValue("serialized-cache"), + }; + const provider = { + getLogger: () => ({ trace: jest.fn() }), + getMsalClient: () => ({ + acquireTokenByCode, + getTokenCache: () => tokenCache, + }), + getCryptoProvider: () => ({ + base64Decode: jest.fn().mockReturnValue(JSON.stringify({ redirectTo: "/target" })), + }), + } as unknown as WebAppAuthProvider; + + return { + acquireTokenByCode, + handler: redirectHandler.call(provider), + }; + }; + + it("redeems the code with the stored request when state matches", async () => { + const { acquireTokenByCode, handler } = createHandler(); + const req = { + body: { + code: "authorization-code", + state, + unexpected: "authorization-response-data", + }, + session: { + tokenRequestParams: { + code: "", + scopes: ["openid"], + state, + redirectUri: "http://localhost/redirect", + }, + }, + } as unknown as Request; + const res = { + redirect: jest.fn(), + } as unknown as Response; + const next = jest.fn() as NextFunction; + + await handler(req, res, next); + + expect(acquireTokenByCode).toHaveBeenCalledTimes(1); + expect(acquireTokenByCode).toHaveBeenCalledWith({ + code: "authorization-code", + scopes: ["openid"], + state, + redirectUri: "http://localhost/redirect", + }); + expect(req.session.tokenCache).toBe("serialized-cache"); + expect(req.session.account).toBe(account); + expect(req.session.isAuthenticated).toBe(true); + expect(res.redirect).toHaveBeenCalledWith("/target"); + expect(next).not.toHaveBeenCalled(); + }); + + it.each([ + ["missing", undefined], + ["mismatched", "different-state"], + ])("rejects a %s state before redeeming the code", async (_description, responseState) => { + const { acquireTokenByCode, handler } = createHandler(); + const req = { + body: { + code: "authorization-code", + state: responseState, + }, + session: { + tokenRequestParams: { + state, + }, + }, + } as unknown as Request; + const res = {} as Response; + const next = jest.fn() as NextFunction; + + await handler(req, res, next); + + expect(next).toHaveBeenCalledWith(new Error(ErrorMessages.CSRF_TOKEN_MISMATCH)); + expect(acquireTokenByCode).not.toHaveBeenCalled(); + }); +});