diff --git a/app/composables/useAuthRefresh.ts b/app/composables/useAuthRefresh.ts index 78b354e..a655b75 100644 --- a/app/composables/useAuthRefresh.ts +++ b/app/composables/useAuthRefresh.ts @@ -10,29 +10,20 @@ export const useAuthRefresh = () => { const { status, data, refresh } = useAuth(); const isRefreshing = ref(false); const refreshError = ref(undefined); + let activeRefresh: Promise | undefined; - const refreshToken = async (): Promise => { - // Prevent multiple simultaneous refresh attempts - if (isRefreshing.value) { - return { success: false, error: "Refresh in progress" }; - } - - // Check if user is authenticated - if (status.value !== "authenticated") { - return { success: false, error: "No active session" }; - } - + const doRefresh = async (): Promise => { isRefreshing.value = true; refreshError.value = undefined; try { await refresh(); // From sidebase methods - // Verify the refresh was successful - if (status.value === "authenticated") { + const sessionError = (data.value as Session | null)?.error; + if (status.value === "authenticated" && !sessionError) { return { success: true }; } else { - refreshError.value = "Session refresh failed"; + refreshError.value = sessionError ?? "Session refresh failed"; return { success: false, error: refreshError.value }; } } catch (error) { @@ -43,7 +34,21 @@ export const useAuthRefresh = () => { return { success: false, error: errorMessage }; } finally { isRefreshing.value = false; + activeRefresh = undefined; + } + }; + + const refreshToken = (): Promise => { + if (activeRefresh) { + return activeRefresh; + } + + if (status.value !== "authenticated") { + return Promise.resolve({ success: false, error: "No active session" }); } + + activeRefresh = doRefresh(); + return activeRefresh; }; const shouldRefreshToken = (bufferSeconds: number = 120): boolean => { diff --git a/app/plugins/api.ts b/app/plugins/api.ts index 7b42a47..6edd3fa 100644 --- a/app/plugins/api.ts +++ b/app/plugins/api.ts @@ -31,7 +31,7 @@ declare module "vue" { } export default defineNuxtPlugin(() => { - const { signIn, getSession } = useAuth(); + const { signIn, getSession, data } = useAuth(); const { shouldRefreshToken, refreshToken } = useAuthRefresh(); const toast = useToast(); @@ -43,13 +43,18 @@ export default defineNuxtPlugin(() => { baseURL: baseUrl, timeout: 60000, // matches HA async onRequest({ options }) { - const sessionData = await getSession(); + let sessionData = await getSession(); - if (shouldRefreshToken(120)) { + if (sessionData?.error) { + await signIn(idpProvider); + return; + } else if (shouldRefreshToken(120)) { const refreshStatus = await refreshToken(); if (!refreshStatus.success) { await signIn(idpProvider); // Force sign in again if auto refresh fails + return; } + sessionData = data.value; } // Annoying workaround to avoid typescript from complaining - cast to Headers then set explicitly diff --git a/nuxt-auth.d.ts b/nuxt-auth.d.ts index acaa68a..d270f8f 100644 --- a/nuxt-auth.d.ts +++ b/nuxt-auth.d.ts @@ -1,7 +1,7 @@ +// next-auth module definitions import type { DefaultSession } from "next-auth"; declare module "next-auth" { - /* Returned by `useAuth`, `getSession` and `getServerSession` */ interface Session extends DefaultSession { user: { name: string; @@ -10,5 +10,6 @@ declare module "next-auth" { accessToken?: string; expires?: Date; expiresAt?: number; + error?: string; } } diff --git a/release-please-config.json b/release-please-config.json index 42b7a9c..fbafd00 100644 --- a/release-please-config.json +++ b/release-please-config.json @@ -5,7 +5,8 @@ "bump-patch-for-minor-pre-major": true, "packages": { ".": { - "release-type": "node" + "release-type": "node", + "release-as": "1.0.0" } } } diff --git a/server/routes/flame/api/auth/[...].ts b/server/routes/flame/api/auth/[...].ts index fdae7ab..3633def 100644 --- a/server/routes/flame/api/auth/[...].ts +++ b/server/routes/flame/api/auth/[...].ts @@ -144,7 +144,62 @@ function buildProvider() { return providers; } -async function refreshAccessToken(token: JWT) { +const REFRESH_BUFFER_SECONDS = 120; +const REFRESH_RESULT_TTL_MS = 10_000; +// Need to track them for edge cases +const activeRefreshes = new Map>(); + +let tokenEndpointPromise: Promise | undefined; + +function nowSeconds() { + return Math.floor(Date.now() / 1000); +} + +function jwtExpiry(accessToken: unknown): number | undefined { + if (typeof accessToken !== "string") return undefined; + try { + const payload = JSON.parse( + Buffer.from(accessToken.split(".")[1] ?? "", "base64url").toString(), + ); + return typeof payload.exp === "number" ? payload.exp : undefined; + } catch { + return undefined; + } +} + +function getTokenEndpoint(issuer: string, proxy: RequestInit) { + tokenEndpointPromise ??= fetch( + `${issuer}/.well-known/openid-configuration`, + proxy, + ) + .then(async (r) => { + if (!r.ok) throw new Error(`OIDC discovery failed with ${r.status}`); + return (await r.json()).token_endpoint as string; + }) + .catch((error) => { + tokenEndpointPromise = undefined; + throw error; + }); + return tokenEndpointPromise; +} + +// stop parallel requests from using the same refresh token +function refreshAccessTokenOnce(token: JWT) { + const key = token.refresh_token as string; + let refresh = activeRefreshes.get(key); + if (!refresh) { + refresh = refreshAccessToken(token); + activeRefreshes.set(key, refresh); + refresh + .catch(() => undefined) + .finally(() => + setTimeout(() => activeRefreshes.delete(key), REFRESH_RESULT_TTL_MS), + ); + } + return refresh; +} + +async function refreshAccessToken(token: JWT): Promise { const clientId = process.env.NUXT_IDP_CLIENT_ID ?? "node-ui"; const clientSecret = process.env.NUXT_IDP_CLIENT_SECRET ?? ""; const clientIssuer = @@ -152,11 +207,10 @@ async function refreshAccessToken(token: JWT) { const proxy = createProxy(); - const discovery = await fetch( - `${clientIssuer}/.well-known/openid-configuration`, + const tokenEndpoint = await getTokenEndpoint( + clientIssuer, proxy as RequestInit, - ).then((r) => r.json()); - const tokenEndpoint: string = discovery.token_endpoint; + ); const response = await fetch(tokenEndpoint, { ...(proxy as RequestInit), @@ -177,8 +231,12 @@ async function refreshAccessToken(token: JWT) { return { ...token, access_token: refreshedTokens.access_token, - expires_at: Date.now() + refreshedTokens.expires_in * 1000, + expires_at: + typeof refreshedTokens.expires_in === "number" + ? nowSeconds() + refreshedTokens.expires_in + : jwtExpiry(refreshedTokens.access_token), refresh_token: refreshedTokens.refresh_token ?? token.refresh_token, + error: undefined, }; } @@ -223,6 +281,7 @@ export default NuxtAuthHandler({ ...session, accessToken: token.access_token as string | undefined, expiresAt: token.expires_at as number | undefined, + error: token.error as string | undefined, }; }, /* on JWT token creation or mutation */ @@ -232,12 +291,11 @@ export default NuxtAuthHandler({ user, }: { token: JWT; - account: Account | null; - user: User; + account?: Account | null; + user?: User; }) { if (account && user) { if (account.type === "credentials") { - // CredentialsProvider (Hub password grant): tokens live on the user object const u = user as { access_token?: string; refresh_token?: string; @@ -253,18 +311,21 @@ export default NuxtAuthHandler({ return { ...token, access_token: account.access_token, - expires_at: account.expires_at as number, + expires_at: account.expires_at ?? jwtExpiry(account.access_token), refresh_token: account.refresh_token, }; } - if (Date.now() < (token.expires_at as number)) { - // Subsequent logins, but the `access_token` is still valid + + const expiresAt = token.expires_at as number | undefined; + if (!expiresAt || nowSeconds() < expiresAt - REFRESH_BUFFER_SECONDS) { return token; } - if (!token.refresh_token) throw new TypeError("Missing refresh_token"); + if (!token.refresh_token) { + return { ...token, error: "RefreshAccessTokenError" }; + } try { - return refreshAccessToken(token); + return await refreshAccessTokenOnce(token); } catch (error) { console.error("Error refreshing access_token", error); // If we fail to refresh the token, return an error so we can handle it on the page diff --git a/server/routes/flame/api/token.get.ts b/server/routes/flame/api/token.get.ts deleted file mode 100644 index 15803b9..0000000 --- a/server/routes/flame/api/token.get.ts +++ /dev/null @@ -1,6 +0,0 @@ -import { getToken } from "#auth"; - -export default eventHandler(async (event) => { - const token = await getToken({ event }); - return token || "No token found"; -}); diff --git a/test/composables/useAuthRefresh.test.ts b/test/composables/useAuthRefresh.test.ts index 191207a..3f51ab5 100644 --- a/test/composables/useAuthRefresh.test.ts +++ b/test/composables/useAuthRefresh.test.ts @@ -112,22 +112,60 @@ describe("useAuthRefresh", () => { expect(result).toEqual({ success: true }); }); - it("returns 'Refresh in progress' when called concurrently", async () => { + it("shares a single in-flight refresh between concurrent callers", async () => { const expiry = Math.floor(Date.now() / 1000) + 3600; let resolveRefresh!: () => void; - const slowRefresh = () => - new Promise((res) => { - resolveRefresh = res; - }); + const slowRefresh = vi.fn( + () => + new Promise((res) => { + resolveRefresh = res; + }), + ); mockAuthAs("authenticated", expiry, slowRefresh); const { refreshToken } = useAuthRefresh(); - const first = refreshToken(); // starts, isRefreshing → true - const second = refreshToken(); // should be blocked + const first = refreshToken(); + const second = refreshToken(); resolveRefresh(); const [r1, r2] = await Promise.all([first, second]); + expect(slowRefresh).toHaveBeenCalledTimes(1); expect(r1).toEqual({ success: true }); - expect(r2).toEqual({ success: false, error: "Refresh in progress" }); + expect(r2).toEqual({ success: true }); + }); + + it("starts a new refresh once the previous one has settled", async () => { + const expiry = Math.floor(Date.now() / 1000) + 3600; + const refreshFn = vi.fn().mockResolvedValue(undefined); + mockAuthAs("authenticated", expiry, refreshFn); + const { refreshToken } = useAuthRefresh(); + + await refreshToken(); + await refreshToken(); + expect(refreshFn).toHaveBeenCalledTimes(2); + }); + + it("returns the session error when the server failed to refresh", async () => { + const expiry = Math.floor(Date.now() / 1000) - 10; + const data = ref<{ expiresAt: number; error?: string }>({ + expiresAt: expiry, + }); + vi.mocked(useAuth).mockReturnValue({ + status: ref("authenticated"), + data, + refresh: vi.fn(async () => { + data.value = { expiresAt: expiry, error: "RefreshAccessTokenError" }; + }), + signIn: vi.fn(), + signOut: vi.fn(), + }); + const { refreshToken, refreshError } = useAuthRefresh(); + + const result = await refreshToken(); + expect(result).toEqual({ + success: false, + error: "RefreshAccessTokenError", + }); + expect(refreshError.value).toBe("RefreshAccessTokenError"); }); it("returns error when refresh() throws", async () => {