fix: check url first for sso provider identification
This commit is contained in:
parent
27ac83613d
commit
3e73b7cc8a
2 changed files with 14 additions and 12 deletions
|
|
@ -2,14 +2,14 @@ import { APIError } from "better-auth/api";
|
||||||
import type { GenericEndpointContext } from "@better-auth/core";
|
import type { GenericEndpointContext } from "@better-auth/core";
|
||||||
import { db } from "~/server/db/db";
|
import { db } from "~/server/db/db";
|
||||||
import { logger } from "~/server/utils/logger";
|
import { logger } from "~/server/utils/logger";
|
||||||
import { extractProviderIdFromContext, normalizeEmail } from "../utils/sso-context";
|
import { extractProviderIdFromContext, extractProviderIdFromUrl, normalizeEmail } from "../utils/sso-context";
|
||||||
|
|
||||||
export function isSsoCallbackRequest(ctx: GenericEndpointContext | null) {
|
export function isSsoCallbackRequest(ctx: GenericEndpointContext | null) {
|
||||||
if (!ctx) {
|
if (!ctx?.request?.url) {
|
||||||
return false;
|
return false;
|
||||||
}
|
}
|
||||||
|
|
||||||
return extractProviderIdFromContext(ctx) !== null;
|
return extractProviderIdFromUrl(ctx.request.url) !== null;
|
||||||
}
|
}
|
||||||
|
|
||||||
export const requireSsoInvitation = async (userEmail: string, ctx: GenericEndpointContext | null) => {
|
export const requireSsoInvitation = async (userEmail: string, ctx: GenericEndpointContext | null) => {
|
||||||
|
|
|
||||||
|
|
@ -4,6 +4,16 @@ export function normalizeEmail(email: string): string {
|
||||||
return email.trim().toLowerCase();
|
return email.trim().toLowerCase();
|
||||||
}
|
}
|
||||||
|
|
||||||
|
export function extractProviderIdFromUrl(url: string): string | null {
|
||||||
|
try {
|
||||||
|
const pathname = new URL(url, "http://localhost").pathname;
|
||||||
|
const match = pathname.match(/\/sso\/(?:saml2\/)?callback\/([^/]+)$/);
|
||||||
|
return match?.[1] ?? null;
|
||||||
|
} catch {
|
||||||
|
return null;
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
export function extractProviderIdFromContext(ctx?: GenericEndpointContext | null) {
|
export function extractProviderIdFromContext(ctx?: GenericEndpointContext | null) {
|
||||||
if (!ctx) {
|
if (!ctx) {
|
||||||
return null;
|
return null;
|
||||||
|
|
@ -14,15 +24,7 @@ export function extractProviderIdFromContext(ctx?: GenericEndpointContext | null
|
||||||
}
|
}
|
||||||
|
|
||||||
if (ctx.request?.url) {
|
if (ctx.request?.url) {
|
||||||
try {
|
return extractProviderIdFromUrl(ctx.request.url);
|
||||||
const pathname = new URL(ctx.request.url, "http://localhost").pathname;
|
|
||||||
const ssoCallbackMatch = pathname.match(/\/sso\/(?:saml2\/)?callback\/([^/]+)$/);
|
|
||||||
if (ssoCallbackMatch) {
|
|
||||||
return ssoCallbackMatch[1];
|
|
||||||
}
|
|
||||||
} catch {
|
|
||||||
return null;
|
|
||||||
}
|
|
||||||
}
|
}
|
||||||
|
|
||||||
return null;
|
return null;
|
||||||
|
|
|
||||||
Loading…
Reference in a new issue