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
1 change: 1 addition & 0 deletions packages/core/src/@types/config.ts
Original file line number Diff line number Diff line change
Expand Up @@ -486,6 +486,7 @@ export interface RouterGlobalContext<DefaultUser extends User = User, SignUpSche
signUp?: SignUpConfig<DefaultUser, SignUpSchema>
jwtManager: JWTManager<DefaultUser>
rateLimiters: InferRules<Required<RateLimiterConfig>>
sessionStrategyMode: "jwt" | "database"
}

export interface SchemaRegistryContext {
Expand Down
12 changes: 11 additions & 1 deletion packages/core/src/@types/session.ts
Original file line number Diff line number Diff line change
Expand Up @@ -11,6 +11,7 @@ import type {
FromShapeToObject,
Prettify,
Identities,
OAuthProviderRecord,
} from "@/@types/index.ts"
import type { DatabaseAdapter } from "@/@types/adapter.ts"

Expand Down Expand Up @@ -213,6 +214,10 @@ export interface GetStatelessSessionReturn<DefaultUser extends User = User> {

export type GetStatefulSessionReturn<DefaultUser extends User = User> = GetStatelessSessionReturn<DefaultUser>

export type GetProviderTokensStatefulReturn =
| { success: true; tokens: OAuthTokenPayload; headers: Headers }
| { success: false; tokens: null; error: { code: string; message: string }; headers: Headers; statusCode: number }

/**
* Abstraction layer for session management.
*/
Expand All @@ -229,14 +234,16 @@ export interface SessionStrategy<DefaultUser extends User = User> {
*/
createSession(session: User): Promise<string>

getProviderTokens(oauth: string, request: Request): Promise<GetProviderTokensStatefulReturn>

/**
* Attempt to refresh using the refresh token cookie.
* Returns null session + cookie-clearing response on any failure.
*/
refreshSession(
headers: Headers,
session: DeepPartial<Session<DefaultUser>>,
skipCSRFCheck?: boolean
skipCSRFCheck: boolean
): Promise<{
session: Session<DefaultUser> | null
headers: Headers
Expand All @@ -263,6 +270,7 @@ export interface CreateSessionStrategyOptions<Identity extends Identities> {
cookies: () => InternalCookieStoreConfig
logger?: InternalLogger
identity: SchemaRegistryContext
oauth: OAuthProviderRecord
}

/** Options specialized for the JWT-backed session strategy. */
Expand All @@ -272,6 +280,7 @@ export interface JWTStrategyOptions<DefaultUser extends User = User> {
logger?: InternalLogger
cookies: () => InternalCookieStoreConfig
identity: SchemaRegistryContext
oauth: OAuthProviderRecord
}

export interface DatabaseStrategyOptions<DefaultUser extends User = User> {
Expand All @@ -280,6 +289,7 @@ export interface DatabaseStrategyOptions<DefaultUser extends User = User> {
logger?: InternalLogger
cookies: () => InternalCookieStoreConfig
identity: SchemaRegistryContext
oauth: OAuthProviderRecord
}

/** Minimal token issue/verify surface used by session code paths. */
Expand Down
106 changes: 21 additions & 85 deletions packages/core/src/api/getProviderTokens.ts
Original file line number Diff line number Diff line change
@@ -1,78 +1,21 @@
import { getCookie } from "@/cookie.ts"
import { HeadersBuilder } from "@aura-stack/router"
import { fetchAsync } from "@/shared/fetch-async.ts"
import { toUnionHeaders } from "@/shared/utils.ts"
import { secureApiHeaders } from "@/shared/headers.ts"
import { createBasicAuthHeader, shouldRefresh, toUnionHeaders } from "@/shared/utils.ts"
import { AuraAuthError } from "@/shared/errors.ts"
import { createValidation, handleApiError } from "@/shared/utils/api.ts"
import type {
FunctionAPIContext,
GetProviderTokensAPIOptions,
GetProviderTokensAPIReturn,
LiteralUnion,
OAuthTokenPayload,
BuiltInOAuthProvider,
RuntimeOAuthProvider,
} from "@/@types/index.ts"
import { isObject, isRefreshTokenObject } from "@/shared/assert.ts"

export const refreshProviderToken = async (
payload: OAuthTokenPayload,
provider: RuntimeOAuthProvider
): Promise<OAuthTokenPayload> => {
if (!provider.refreshToken || (isObject(provider.refreshToken) && !isRefreshTokenObject(provider.refreshToken))) {
throw new AuraAuthError({ code: "OAUTH_INVALID_REFRESH_TOKEN_CONFIG" })
}
if (!payload.refreshToken) {
throw new AuraAuthError({ code: "OAUTH_INVALID_REFRESH_TOKEN_CONFIG" })
}
const url = isRefreshTokenObject(provider.refreshToken) ? provider.refreshToken.url : provider.refreshToken

const isCredentialsAuth = isRefreshTokenObject(provider.refreshToken)
? provider.refreshToken.authorization?.type === "credentials"
: false
const response = await fetchAsync(url, {
method: "POST",
headers: {
"Content-Type": "application/x-www-form-urlencoded",
...(isCredentialsAuth ? {} : { Authorization: createBasicAuthHeader(provider.clientId!, provider.clientSecret!) }),
...(typeof provider.refreshToken === "object" && provider.refreshToken.headers ? provider.refreshToken.headers : {}),
},
body: new URLSearchParams({
grant_type: "refresh_token",
refresh_token: payload.refreshToken!,
...(isCredentialsAuth ? { client_id: provider.clientId!, client_secret: provider.clientSecret! } : {}),
...(typeof provider.refreshToken === "object" && provider.refreshToken.params ? provider.refreshToken.params : {}),
}),
})
if (!response.ok) {
throw new AuraAuthError({ code: "OAUTH_INVALID_REFRESH_TOKEN_RESPONSE" })
}
const data = await response.json()
const now = Math.floor(Date.now() / 1000)

return {
accessToken: data.access_token ?? payload.accessToken,
expiresAt: now + (data.expires_in ?? 3600),
refreshToken: data.refresh_token ?? payload.refreshToken,
refreshTokenExpiresAt: data.refresh_token_expires_in
? now + data.refresh_token_expires_in
: payload.refreshTokenExpiresAt,
scopes: typeof data.scope === "string" ? data.scope.split(" ") : payload.scopes,
tokenType: data.token_type ?? payload.tokenType,
idToken: data.id_token ?? payload.idToken,
issuedAt: now,
}
}

export const getProviderTokens = async (
oauth: LiteralUnion<BuiltInOAuthProvider>,
{ ctx, request: requestInit, headers: headersInit, skipCSRFCheck = false }: FunctionAPIContext<GetProviderTokensAPIOptions>
): Promise<GetProviderTokensAPIReturn> => {
const { cookies, identity, jwtManager } = ctx
const initialHeaders = new Headers(headersInit ?? requestInit?.headers)
try {
const { provider, headers, request, rateLimit } = await createValidation(ctx, initialHeaders)
const { request, rateLimit } = await createValidation(ctx, initialHeaders)
.verifyOAuthProvider(oauth)
.verifySession()
.verifyCSRFToken(skipCSRFCheck)
Expand All @@ -83,36 +26,29 @@ export const getProviderTokens = async (
if (rateLimit) {
return rateLimit as unknown as GetProviderTokensAPIReturn
}

const cookieName = `${cookies.accessToken.name}.${oauth}`
const cookie = getCookie(request, cookieName)

const decodedToken = await jwtManager.verifyToken(cookie)
const tokens = await identity.schemaRegistry.parseOAuthTokens(decodedToken)

const refreshWindow = provider.refreshWindow ?? 300
const refreshed = shouldRefresh(tokens, refreshWindow)

if (refreshed) {
const refreshedTokens = await refreshProviderToken(tokens, provider!)
const encodedTokens = await jwtManager.createToken(refreshedTokens as unknown as Record<string, unknown>)
const builder = new HeadersBuilder(secureApiHeaders)
.setCookie(cookieName, encodedTokens, cookies.accessToken.attributes)
.toHeaders()
const newHeaders = toUnionHeaders(builder, headers)
const getTokens = await ctx.sessionStrategy.getProviderTokens(oauth, request)
if (getTokens.success) {
const { success, tokens, headers } = getTokens
return {
success: true,
tokens: refreshedTokens,
headers: newHeaders,
toResponse: () => Response.json({ success: true, tokens: refreshedTokens }, { status: 200, headers: newHeaders }),
success,
tokens,
headers,
toResponse: () => {
return Response.json({ success, tokens }, { status: success ? 200 : 400, headers })
},
}
}

return {
success: true,
tokens,
headers,
toResponse: () => Response.json({ success: true, tokens }, { status: 200, headers }),
success: false,
tokens: null,
headers: getTokens.headers,
error: getTokens.error,
toResponse: () => {
return Response.json(
{ success: false, tokens: null },
{ status: getTokens.statusCode, headers: getTokens.headers }
)
},
}
} catch (error) {
const { code, message, statusCode } = handleApiError(error, "PROVIDER_TOKENS_ERROR", "Failed to get provider tokens")
Expand Down
2 changes: 2 additions & 0 deletions packages/core/src/router/context.ts
Original file line number Diff line number Diff line change
Expand Up @@ -63,13 +63,15 @@ export const createContext = <Identity extends Identities, SignUpSchema extends
signUp: config?.signUp,
jwtManager: createJoseManager(isStatelessStrategy(config?.session) ? config?.session?.jwt : undefined, jose),
rateLimiters: createRateLimiterInstance(config?.rateLimiter),
sessionStrategyMode: isStatelessStrategy(config?.session) ? "jwt" : "database",
} as InternalContext<Identity, SignUpSchema>
ctx.sessionStrategy = createSessionStrategy<Identity>({
cookies: () => ctx.cookies,
jose: ctx.jose,
config: config?.session,
logger: ctx.logger,
identity: ctx.identity,
oauth: ctx.oauth,
})
return ctx
}
Loading
Loading