Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
33 changes: 19 additions & 14 deletions app/composables/useAuthRefresh.ts
Original file line number Diff line number Diff line change
Expand Up @@ -10,29 +10,20 @@ export const useAuthRefresh = () => {
const { status, data, refresh } = useAuth();
const isRefreshing = ref(false);
const refreshError = ref<string | undefined>(undefined);
let activeRefresh: Promise<RefreshResponse> | undefined;

const refreshToken = async (): Promise<RefreshResponse> => {
// 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<RefreshResponse> => {
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) {
Expand All @@ -43,7 +34,21 @@ export const useAuthRefresh = () => {
return { success: false, error: errorMessage };
} finally {
isRefreshing.value = false;
activeRefresh = undefined;
}
};

const refreshToken = (): Promise<RefreshResponse> => {
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 => {
Expand Down
11 changes: 8 additions & 3 deletions app/plugins/api.ts
Original file line number Diff line number Diff line change
Expand Up @@ -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();

Expand All @@ -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);
Comment thread
brucetony marked this conversation as resolved.
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
Expand Down
3 changes: 2 additions & 1 deletion nuxt-auth.d.ts
Original file line number Diff line number Diff line change
@@ -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;
Expand All @@ -10,5 +10,6 @@ declare module "next-auth" {
accessToken?: string;
expires?: Date;
expiresAt?: number;
error?: string;
}
}
3 changes: 2 additions & 1 deletion release-please-config.json
Original file line number Diff line number Diff line change
Expand Up @@ -5,7 +5,8 @@
"bump-patch-for-minor-pre-major": true,
"packages": {
".": {
"release-type": "node"
"release-type": "node",
"release-as": "1.0.0"
Comment thread
brucetony marked this conversation as resolved.
}
}
}
89 changes: 75 additions & 14 deletions server/routes/flame/api/auth/[...].ts
Original file line number Diff line number Diff line change
Expand Up @@ -144,19 +144,73 @@ 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<string, Promise<JWT>>();

let tokenEndpointPromise: Promise<string> | 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<JWT> {
const clientId = process.env.NUXT_IDP_CLIENT_ID ?? "node-ui";
const clientSecret = process.env.NUXT_IDP_CLIENT_SECRET ?? "";
const clientIssuer =
process.env.NUXT_PUBLIC_IDP_ISSUER ?? "http://localhost:8080/realms/flame";

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),
Expand All @@ -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,
};
}

Expand Down Expand Up @@ -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 */
Expand All @@ -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;
Expand All @@ -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
Expand Down
6 changes: 0 additions & 6 deletions server/routes/flame/api/token.get.ts

This file was deleted.

54 changes: 46 additions & 8 deletions test/composables/useAuthRefresh.test.ts
Original file line number Diff line number Diff line change
Expand Up @@ -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<void>((res) => {
resolveRefresh = res;
});
const slowRefresh = vi.fn(
() =>
new Promise<void>((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 () => {
Expand Down