Add Github SSO users to default organization on signup

This commit is contained in:
carlosmonastyrski
2025-05-01 20:41:30 -03:00
parent 296493484f
commit 63279280fd
12 changed files with 112 additions and 38 deletions
+20
View File
@@ -48,3 +48,23 @@ export const SecretNameSchema = BaseSecretNameSchema.refine(
) )
.refine((el) => !el.includes(":"), "Secret name cannot contain colon.") .refine((el) => !el.includes(":"), "Secret name cannot contain colon.")
.refine((el) => !el.includes("/"), "Secret name cannot contain forward slash."); .refine((el) => !el.includes("/"), "Secret name cannot contain forward slash.");
const DefaultOrgSchema = z.object({
useDefaultOrg: z.literal(true),
organizationName: z.string().trim().optional()
});
const CustomOrgSchema = z.object({
useDefaultOrg: z.literal(false),
organizationName: GenericResourceNameSchema
});
export const OrganizationInputSchema = z.preprocess(
(data) => {
if (typeof data === "object" && data && "useDefaultOrg" in data === false) {
return { ...data, useDefaultOrg: false };
}
return data;
},
z.discriminatedUnion("useDefaultOrg", [DefaultOrgSchema, CustomOrgSchema])
);
+1
View File
@@ -623,6 +623,7 @@ export const registerRoutes = async (
tokenService, tokenService,
orgDAL, orgDAL,
totpService, totpService,
orgMembershipDAL,
auditLogService auditLogService
}); });
const passwordService = authPaswordServiceFactory({ const passwordService = authPaswordServiceFactory({
+6 -1
View File
@@ -23,6 +23,7 @@ import { fetchGithubEmails, fetchGithubUser } from "@app/lib/requests/github";
import { authRateLimit } from "@app/server/config/rateLimiter"; import { authRateLimit } from "@app/server/config/rateLimiter";
import { AuthMethod } from "@app/services/auth/auth-type"; import { AuthMethod } from "@app/services/auth/auth-type";
import { OrgAuthMethod } from "@app/services/org/org-types"; import { OrgAuthMethod } from "@app/services/org/org-types";
import { getServerCfg } from "@app/services/super-admin/super-admin-service";
export const registerSsoRouter = async (server: FastifyZodProvider) => { export const registerSsoRouter = async (server: FastifyZodProvider) => {
const appCfg = getConfig(); const appCfg = getConfig();
@@ -342,8 +343,12 @@ export const registerSsoRouter = async (server: FastifyZodProvider) => {
}` }`
); );
} }
const serverCfg = await getServerCfg();
return res.redirect( return res.redirect(
`${appCfg.SITE_URL}/signup/sso?token=${encodeURIComponent(req.passportUser.providerAuthToken)}` `${appCfg.SITE_URL}/signup/sso?token=${encodeURIComponent(req.passportUser.providerAuthToken)}${
serverCfg.defaultAuthOrgId ? `&defaultOrgAllowed=true` : ""
}`
); );
} }
}); });
+20 -19
View File
@@ -4,7 +4,7 @@ import { UsersSchema } from "@app/db/schemas";
import { getConfig } from "@app/lib/config/env"; import { getConfig } from "@app/lib/config/env";
import { ForbiddenRequestError } from "@app/lib/errors"; import { ForbiddenRequestError } from "@app/lib/errors";
import { authRateLimit } from "@app/server/config/rateLimiter"; import { authRateLimit } from "@app/server/config/rateLimiter";
import { GenericResourceNameSchema } from "@app/server/lib/schemas"; import { OrganizationInputSchema } from "@app/server/lib/schemas";
import { getServerCfg } from "@app/services/super-admin/super-admin-service"; import { getServerCfg } from "@app/services/super-admin/super-admin-service";
import { PostHogEventTypes } from "@app/services/telemetry/telemetry-types"; import { PostHogEventTypes } from "@app/services/telemetry/telemetry-types";
@@ -88,24 +88,25 @@ export const registerSignupRouter = async (server: FastifyZodProvider) => {
rateLimit: authRateLimit rateLimit: authRateLimit
}, },
schema: { schema: {
body: z.object({ body: z
email: z.string().trim(), .object({
firstName: z.string().trim(), email: z.string().trim(),
lastName: z.string().trim().optional(), firstName: z.string().trim(),
protectedKey: z.string().trim(), lastName: z.string().trim().optional(),
protectedKeyIV: z.string().trim(), protectedKey: z.string().trim(),
protectedKeyTag: z.string().trim(), protectedKeyIV: z.string().trim(),
publicKey: z.string().trim(), protectedKeyTag: z.string().trim(),
encryptedPrivateKey: z.string().trim(), publicKey: z.string().trim(),
encryptedPrivateKeyIV: z.string().trim(), encryptedPrivateKey: z.string().trim(),
encryptedPrivateKeyTag: z.string().trim(), encryptedPrivateKeyIV: z.string().trim(),
salt: z.string().trim(), encryptedPrivateKeyTag: z.string().trim(),
verifier: z.string().trim(), salt: z.string().trim(),
organizationName: GenericResourceNameSchema, verifier: z.string().trim(),
providerAuthToken: z.string().trim().optional().nullish(), providerAuthToken: z.string().trim().optional().nullish(),
attributionSource: z.string().trim().optional(), attributionSource: z.string().trim().optional(),
password: z.string() password: z.string()
}), })
.and(OrganizationInputSchema),
response: { response: {
200: z.object({ 200: z.object({
message: z.string(), message: z.string(),
@@ -2,7 +2,7 @@ import bcrypt from "bcrypt";
import jwt from "jsonwebtoken"; import jwt from "jsonwebtoken";
import { Knex } from "knex"; import { Knex } from "knex";
import { OrgMembershipRole, TUsers, UserDeviceSchema } from "@app/db/schemas"; import { OrgMembershipRole, OrgMembershipStatus, TableName, TUsers, UserDeviceSchema } from "@app/db/schemas";
import { TAuditLogServiceFactory } from "@app/ee/services/audit-log/audit-log-service"; import { TAuditLogServiceFactory } from "@app/ee/services/audit-log/audit-log-service";
import { EventType } from "@app/ee/services/audit-log/audit-log-types"; import { EventType } from "@app/ee/services/audit-log/audit-log-types";
import { isAuthMethodSaml } from "@app/ee/services/permission/permission-fns"; import { isAuthMethodSaml } from "@app/ee/services/permission/permission-fns";
@@ -20,6 +20,8 @@ import { getServerCfg } from "@app/services/super-admin/super-admin-service";
import { TAuthTokenServiceFactory } from "../auth-token/auth-token-service"; import { TAuthTokenServiceFactory } from "../auth-token/auth-token-service";
import { TokenType } from "../auth-token/auth-token-types"; import { TokenType } from "../auth-token/auth-token-types";
import { TOrgDALFactory } from "../org/org-dal"; import { TOrgDALFactory } from "../org/org-dal";
import { getDefaultOrgMembershipRole } from "../org/org-role-fns";
import { TOrgMembershipDALFactory } from "../org-membership/org-membership-dal";
import { SmtpTemplates, TSmtpService } from "../smtp/smtp-service"; import { SmtpTemplates, TSmtpService } from "../smtp/smtp-service";
import { LoginMethod } from "../super-admin/super-admin-types"; import { LoginMethod } from "../super-admin/super-admin-types";
import { TTotpServiceFactory } from "../totp/totp-service"; import { TTotpServiceFactory } from "../totp/totp-service";
@@ -48,6 +50,7 @@ type TAuthLoginServiceFactoryDep = {
smtpService: TSmtpService; smtpService: TSmtpService;
totpService: Pick<TTotpServiceFactory, "verifyUserTotp" | "verifyWithUserRecoveryCode">; totpService: Pick<TTotpServiceFactory, "verifyUserTotp" | "verifyWithUserRecoveryCode">;
auditLogService: Pick<TAuditLogServiceFactory, "createAuditLog">; auditLogService: Pick<TAuditLogServiceFactory, "createAuditLog">;
orgMembershipDAL: TOrgMembershipDALFactory;
}; };
export type TAuthLoginFactory = ReturnType<typeof authLoginServiceFactory>; export type TAuthLoginFactory = ReturnType<typeof authLoginServiceFactory>;
@@ -56,6 +59,7 @@ export const authLoginServiceFactory = ({
tokenService, tokenService,
smtpService, smtpService,
orgDAL, orgDAL,
orgMembershipDAL,
totpService, totpService,
auditLogService auditLogService
}: TAuthLoginServiceFactoryDep) => { }: TAuthLoginServiceFactoryDep) => {
@@ -719,6 +723,31 @@ export const authLoginServiceFactory = ({
authMethods: [authMethod], authMethods: [authMethod],
isGhost: false isGhost: false
}); });
if (authMethod === AuthMethod.GITHUB && serverCfg.defaultAuthOrgId) {
let orgId = "";
const defaultOrg = await orgDAL.findOrgById(serverCfg.defaultAuthOrgId);
if (!defaultOrg) throw new BadRequestError({ message: "Failed to find default organization" });
orgId = defaultOrg.id;
const [orgMembership] = await orgDAL.findMembership({
[`${TableName.OrgMembership}.userId` as "userId"]: user.id,
[`${TableName.OrgMembership}.orgId` as "id"]: orgId
});
if (!orgMembership) {
const { role, roleId } = await getDefaultOrgMembershipRole(defaultOrg.defaultMembershipRole);
await orgMembershipDAL.create({
userId: user.id,
inviteEmail: email,
orgId,
role,
roleId,
status: OrgMembershipStatus.Accepted,
isActive: true
});
}
}
} else { } else {
const isLinkingRequired = !user?.authMethods?.includes(authMethod); const isLinkingRequired = !user?.authMethods?.includes(authMethod);
if (isLinkingRequired) { if (isLinkingRequired) {
@@ -9,7 +9,7 @@ import { isAuthMethodSaml } from "@app/ee/services/permission/permission-fns";
import { getConfig } from "@app/lib/config/env"; import { getConfig } from "@app/lib/config/env";
import { infisicalSymmetricDecrypt, infisicalSymmetricEncypt } from "@app/lib/crypto/encryption"; import { infisicalSymmetricDecrypt, infisicalSymmetricEncypt } from "@app/lib/crypto/encryption";
import { generateUserSrpKeys, getUserPrivateKey } from "@app/lib/crypto/srp"; import { generateUserSrpKeys, getUserPrivateKey } from "@app/lib/crypto/srp";
import { ForbiddenRequestError, NotFoundError } from "@app/lib/errors"; import { BadRequestError, ForbiddenRequestError, NotFoundError } from "@app/lib/errors";
import { getMinExpiresIn } from "@app/lib/fn"; import { getMinExpiresIn } from "@app/lib/fn";
import { isDisposableEmail } from "@app/lib/validator"; import { isDisposableEmail } from "@app/lib/validator";
import { TGroupProjectDALFactory } from "@app/services/group-project/group-project-dal"; import { TGroupProjectDALFactory } from "@app/services/group-project/group-project-dal";
@@ -150,7 +150,8 @@ export const authSignupServiceFactory = ({
encryptedPrivateKeyTag, encryptedPrivateKeyTag,
ip, ip,
userAgent, userAgent,
authorization authorization,
useDefaultOrg
}: TCompleteAccountSignupDTO) => { }: TCompleteAccountSignupDTO) => {
const appCfg = getConfig(); const appCfg = getConfig();
const serverCfg = await getServerCfg(); const serverCfg = await getServerCfg();
@@ -293,15 +294,24 @@ export const authSignupServiceFactory = ({
}); });
if (!organizationId) { if (!organizationId) {
const newOrganization = await orgService.createOrganization({ let orgId = "";
userId: user.id, if (useDefaultOrg && serverCfg.defaultAuthOrgId) {
userEmail: user.email ?? user.username, const defaultOrg = await orgDAL.findOrgById(serverCfg.defaultAuthOrgId);
orgName: organizationName if (!defaultOrg) throw new BadRequestError({ message: "Failed to find default organization" });
}); orgId = defaultOrg.id;
} else {
if (!organizationName) throw new BadRequestError({ message: "Organization name is required" });
const newOrganization = await orgService.createOrganization({
userId: user.id,
userEmail: user.email ?? user.username,
orgName: organizationName
});
if (!newOrganization) throw new Error("Failed to create organization"); if (!newOrganization) throw new Error("Failed to create organization");
orgId = newOrganization.id;
}
organizationId = newOrganization.id; organizationId = orgId;
} }
const updatedMembersips = await orgDAL.updateMembership( const updatedMembersips = await orgDAL.updateMembership(
@@ -12,12 +12,13 @@ export type TCompleteAccountSignupDTO = {
encryptedPrivateKeyTag: string; encryptedPrivateKeyTag: string;
salt: string; salt: string;
verifier: string; verifier: string;
organizationName: string; organizationName?: string;
providerAuthToken?: string | null; providerAuthToken?: string | null;
attributionSource?: string | undefined; attributionSource?: string | undefined;
ip: string; ip: string;
userAgent: string; userAgent: string;
authorization: string; authorization: string;
useDefaultOrg?: boolean;
}; };
export type TCompleteAccountInviteDTO = { export type TCompleteAccountInviteDTO = {
+1
View File
@@ -107,6 +107,7 @@ export type CompleteAccountSignupDTO = CompleteAccountDTO & {
providerAuthToken?: string; providerAuthToken?: string;
attributionSource?: string; attributionSource?: string;
organizationName: string; organizationName: string;
useDefaultOrg?: boolean;
}; };
export type VerifySignupInviteDTO = { export type VerifySignupInviteDTO = {
@@ -235,7 +235,7 @@ export const OverviewPage = () => {
Default organization Default organization
</div> </div>
<div className="mb-4 max-w-sm text-sm text-mineshaft-400"> <div className="mb-4 max-w-sm text-sm text-mineshaft-400">
Select the default organization you want to set for SAML/LDAP/OIDC based Select the default organization you want to set for SAML/LDAP/OIDC/Github
logins. When selected, user logins will be automatically scoped to the logins. When selected, user logins will be automatically scoped to the
selected organization. selected organization.
</div> </div>
@@ -13,6 +13,7 @@ export const SignupSsoPage = () => {
const { t } = useTranslation(); const { t } = useTranslation();
const search = useSearch({ from: ROUTE_PATHS.Auth.SignUpSsoPage.id }); const search = useSearch({ from: ROUTE_PATHS.Auth.SignUpSsoPage.id });
const token = search.token as string; const token = search.token as string;
const defaultOrgAllowed = search.defaultOrgAllowed as boolean | undefined;
const [step, setStep] = useState(0); const [step, setStep] = useState(0);
const [password, setPassword] = useState(""); const [password, setPassword] = useState("");
@@ -57,6 +58,7 @@ export const SignupSsoPage = () => {
password={password} password={password}
setPassword={setPassword} setPassword={setPassword}
providerAuthToken={token} providerAuthToken={token}
forceDefaultOrg={defaultOrgAllowed}
/> />
); );
default: default:
@@ -30,6 +30,7 @@ type Props = {
name: string; name: string;
providerOrganizationName: string; providerOrganizationName: string;
providerAuthToken?: string; providerAuthToken?: string;
forceDefaultOrg?: boolean;
}; };
/** /**
@@ -51,7 +52,8 @@ export const UserInfoSSOStep = ({
providerOrganizationName, providerOrganizationName,
password, password,
setPassword, setPassword,
providerAuthToken providerAuthToken,
forceDefaultOrg
}: Props) => { }: Props) => {
const [nameError, setNameError] = useState(false); const [nameError, setNameError] = useState(false);
const [organizationName, setOrganizationName] = useState(""); const [organizationName, setOrganizationName] = useState("");
@@ -84,7 +86,7 @@ export const UserInfoSSOStep = ({
} else { } else {
setNameError(false); setNameError(false);
} }
if (!organizationName) { if (!organizationName && !forceDefaultOrg) {
setOrganizationNameError(true); setOrganizationNameError(true);
errorCheck = true; errorCheck = true;
} else { } else {
@@ -160,7 +162,8 @@ export const UserInfoSSOStep = ({
salt: result.salt, salt: result.salt,
verifier: result.verifier, verifier: result.verifier,
organizationName, organizationName,
attributionSource attributionSource,
useDefaultOrg: forceDefaultOrg
}); });
// unset signup JWT token and set JWT token // unset signup JWT token and set JWT token
@@ -267,7 +270,7 @@ export const UserInfoSSOStep = ({
</p> </p>
)} )}
</div> </div>
{providerOrganizationName === undefined && ( {!forceDefaultOrg && providerOrganizationName === undefined && (
<div className="relative z-0 flex w-full min-w-[20rem] flex-col items-center justify-end rounded-lg py-2 lg:w-1/6"> <div className="relative z-0 flex w-full min-w-[20rem] flex-col items-center justify-end rounded-lg py-2 lg:w-1/6">
<p className="mb-1 ml-1 w-full text-left text-sm font-medium text-bunker-300"> <p className="mb-1 ml-1 w-full text-left text-sm font-medium text-bunker-300">
Organization Name Organization Name
@@ -279,7 +282,7 @@ export const UserInfoSSOStep = ({
isRequired isRequired
className="h-12" className="h-12"
maxLength={64} maxLength={64}
disabled isDisabled={forceDefaultOrg}
/> />
{organizationNameError && ( {organizationNameError && (
<p className="ml-1 mt-1 w-full text-left text-xs text-red-600"> <p className="ml-1 mt-1 w-full text-left text-xs text-red-600">
@@ -5,7 +5,8 @@ import { z } from "zod";
import { SignupSsoPage } from "./SignUpSsoPage"; import { SignupSsoPage } from "./SignUpSsoPage";
const SignupSSOPageQueryParamsSchema = z.object({ const SignupSSOPageQueryParamsSchema = z.object({
token: z.string() token: z.string(),
defaultOrgAllowed: z.boolean().optional()
}); });
export const Route = createFileRoute("/_restrict-login-signup/signup/sso")({ export const Route = createFileRoute("/_restrict-login-signup/signup/sso")({