Adjust SCIM and SAML impl to use username / nameID, patch LDAP edge-cases

This commit is contained in:
Tuan Dang
2024-02-25 18:16:26 -08:00
parent 8c31566e17
commit a33c50b75a
15 changed files with 136 additions and 82 deletions
+2 -2
View File
@@ -27,7 +27,7 @@ export const registerLdapRouter = async (server: FastifyZodProvider) => {
passport.use( passport.use(
new LdapStrategy( new LdapStrategy(
server.services.ldap.getLDAPConfiguration, server.services.ldap.getLdapPassportOpts,
// eslint-disable-next-line // eslint-disable-next-line
async (req, user, cb) => { async (req, user, cb) => {
try { try {
@@ -98,7 +98,7 @@ export const registerLdapRouter = async (server: FastifyZodProvider) => {
} }
}, },
handler: async (req) => { handler: async (req) => {
const ldap = await server.services.ldap.getLdapCfg({ const ldap = await server.services.ldap.getLdapCfgWithPermissionCheck({
actor: req.permission.type, actor: req.permission.type,
actorId: req.permission.id, actorId: req.permission.id,
orgId: req.query.organizationId, orgId: req.query.organizationId,
+3 -4
View File
@@ -94,15 +94,14 @@ export const registerSamlRouter = async (server: FastifyZodProvider) => {
async (req, profile, cb) => { async (req, profile, cb) => {
try { try {
if (!profile) throw new BadRequestError({ message: "Missing profile" }); if (!profile) throw new BadRequestError({ message: "Missing profile" });
const { firstName } = profile;
const email = profile?.email ?? (profile?.emailAddress as string); // emailRippling is added because in Rippling the field `email` reserved
if (!email || !firstName) { if (!profile.nameID || !profile.firstName) {
throw new BadRequestError({ message: "Invalid request. Missing email or first name" }); throw new BadRequestError({ message: "Invalid request. Missing email or first name" });
} }
const { isUserCompleted, providerAuthToken } = await server.services.saml.samlLogin({ const { isUserCompleted, providerAuthToken } = await server.services.saml.samlLogin({
email, username: profile.nameID,
email: profile.email,
firstName: profile.firstName as string, firstName: profile.firstName as string,
lastName: profile.lastName as string, lastName: profile.lastName as string,
relayState: (req.body as { RelayState?: string }).RelayState, relayState: (req.body as { RelayState?: string }).RelayState,
+15 -10
View File
@@ -122,7 +122,7 @@ export const registerScimRouter = async (server: FastifyZodProvider) => {
emails: z.array( emails: z.array(
z.object({ z.object({
primary: z.boolean(), primary: z.boolean(),
value: z.string().email(), value: z.string(),
type: z.string().trim() type: z.string().trim()
}) })
), ),
@@ -168,7 +168,7 @@ export const registerScimRouter = async (server: FastifyZodProvider) => {
emails: z.array( emails: z.array(
z.object({ z.object({
primary: z.boolean(), primary: z.boolean(),
value: z.string().email(), value: z.string(),
type: z.string().trim() type: z.string().trim()
}) })
), ),
@@ -198,13 +198,15 @@ export const registerScimRouter = async (server: FastifyZodProvider) => {
familyName: z.string().trim(), familyName: z.string().trim(),
givenName: z.string().trim() givenName: z.string().trim()
}), }),
// emails: z.array( // optional? emails: z
// z.object({ .array(
// primary: z.boolean(), z.object({
// value: z.string().email(), primary: z.boolean(),
// type: z.string().trim() value: z.string().email(),
// }) type: z.string().trim()
// ), })
)
.optional(),
// displayName: z.string().trim(), // displayName: z.string().trim(),
active: z.boolean() active: z.boolean()
}), }),
@@ -231,8 +233,11 @@ export const registerScimRouter = async (server: FastifyZodProvider) => {
}, },
onRequest: verifyAuth([AuthMode.SCIM_TOKEN]), onRequest: verifyAuth([AuthMode.SCIM_TOKEN]),
handler: async (req) => { handler: async (req) => {
const primaryEmail = req.body.emails?.find((email) => email.primary)?.value;
const user = await req.server.services.scim.createScimUser({ const user = await req.server.services.scim.createScimUser({
email: req.body.userName, username: req.body.userName,
email: primaryEmail as string,
firstName: req.body.name.givenName, firstName: req.body.name.givenName,
lastName: req.body.name.familyName, lastName: req.body.name.familyName,
orgId: req.permission.orgId as string orgId: req.permission.orgId as string
@@ -13,6 +13,7 @@ import {
infisicalSymmetricEncypt infisicalSymmetricEncypt
} from "@app/lib/crypto/encryption"; } from "@app/lib/crypto/encryption";
import { BadRequestError } from "@app/lib/errors"; import { BadRequestError } from "@app/lib/errors";
import { logger } from "@app/lib/logger";
import { TOrgPermission } from "@app/lib/types"; import { TOrgPermission } from "@app/lib/types";
import { AuthMethod, AuthTokenType } from "@app/services/auth/auth-type"; import { AuthMethod, AuthTokenType } from "@app/services/auth/auth-type";
import { TOrgBotDALFactory } from "@app/services/org/org-bot-dal"; import { TOrgBotDALFactory } from "@app/services/org/org-bot-dal";
@@ -204,8 +205,8 @@ export const ldapConfigServiceFactory = ({
return ldapConfig; return ldapConfig;
}; };
const getLdapCfg2 = async (orgId: string) => { const getLdapCfg = async (filter: { orgId: string; isActive?: boolean }) => {
const ldapConfig = await ldapConfigDAL.findOne({ orgId }); const ldapConfig = await ldapConfigDAL.findOne(filter);
if (!ldapConfig) throw new BadRequestError({ message: "Failed to find organization LDAP data" }); if (!ldapConfig) throw new BadRequestError({ message: "Failed to find organization LDAP data" });
const orgBot = await orgBotDAL.findOne({ orgId: ldapConfig.orgId }); const orgBot = await orgBotDAL.findOne({ orgId: ldapConfig.orgId });
@@ -272,44 +273,55 @@ export const ldapConfigServiceFactory = ({
}; };
}; };
const getLdapCfg = async ({ actor, actorId, orgId, actorOrgId }: TOrgPermission) => { const getLdapCfgWithPermissionCheck = async ({ actor, actorId, orgId, actorOrgId }: TOrgPermission) => {
const { permission } = await permissionService.getOrgPermission(actor, actorId, orgId, actorOrgId); const { permission } = await permissionService.getOrgPermission(actor, actorId, orgId, actorOrgId);
ForbiddenError.from(permission).throwUnlessCan(OrgPermissionActions.Read, OrgPermissionSubjects.Sso); ForbiddenError.from(permission).throwUnlessCan(OrgPermissionActions.Read, OrgPermissionSubjects.Sso);
return getLdapCfg2(orgId); return getLdapCfg({
orgId
});
}; };
// eslint-disable-next-line // eslint-disable-next-line
const getLDAPConfiguration = (req: FastifyRequest, callback: any) => { const getLdapPassportOpts = (req: FastifyRequest, done: any) => {
const { organizationSlug } = req.body as { const { organizationSlug } = req.body as {
organizationSlug: string; organizationSlug: string;
}; };
const boot = async () => { const boot = async () => {
const organization = await orgDAL.findOne({ slug: organizationSlug }); try {
const ldapConfig = await getLdapCfg2(organization.id); // repeat? const organization = await orgDAL.findOne({ slug: organizationSlug });
req.ldapConfig = ldapConfig; const ldapConfig = await getLdapCfg({
orgId: organization.id,
isActive: true
});
req.ldapConfig = ldapConfig;
const opts = { const opts = {
server: { server: {
url: ldapConfig.url, url: ldapConfig.url,
bindDN: ldapConfig.bindDN, bindDN: ldapConfig.bindDN,
bindCredentials: ldapConfig.bindPass, bindCredentials: ldapConfig.bindPass,
searchBase: ldapConfig.searchBase, searchBase: ldapConfig.searchBase,
searchFilter: "(uid={{username}})", searchFilter: "(uid={{username}})",
searchAttributes: ["uid", "givenName", "sn"], searchAttributes: ["uid", "givenName", "sn"],
...(ldapConfig.caCert !== "" ...(ldapConfig.caCert !== ""
? { ? {
tlsOptions: { tlsOptions: {
ca: [ldapConfig.caCert] ca: [ldapConfig.caCert]
}
} }
} : {})
: {}) },
}, passReqToCallback: true
passReqToCallback: true };
};
// eslint-disable-next-line // eslint-disable-next-line
callback(null, opts); done(null, opts);
} catch (err) {
logger.error(err);
// eslint-disable-next-line
done(err as Error);
}
}; };
process.nextTick(async () => { process.nextTick(async () => {
@@ -403,8 +415,9 @@ export const ldapConfigServiceFactory = ({
return { return {
createLdapCfg, createLdapCfg,
updateLdapCfg, updateLdapCfg,
getLdapCfgWithPermissionCheck,
getLdapCfg, getLdapCfg,
getLDAPConfiguration, getLdapPassportOpts,
ldapLogin ldapLogin
}; };
}; };
@@ -17,7 +17,8 @@ export const getDefaultOnPremFeatures = () => {
customAlerts: false, customAlerts: false,
auditLogs: false, auditLogs: false,
auditLogsRetentionDays: 0, auditLogsRetentionDays: 0,
samlSSO: false, samlSSO: true,
scim: true,
status: null, status: null,
trial_end: null, trial_end: null,
has_used_trial: true, has_used_trial: true,
@@ -23,8 +23,8 @@ export const getDefaultOnPremFeatures = (): TFeatureSet => ({
customAlerts: false, customAlerts: false,
auditLogs: false, auditLogs: false,
auditLogsRetentionDays: 0, auditLogsRetentionDays: 0,
samlSSO: false, samlSSO: true,
scim: false, scim: true,
ldap: true, ldap: true,
status: null, status: null,
trial_end: null, trial_end: null,
@@ -24,8 +24,8 @@ export type TFeatureSet = {
customAlerts: false; customAlerts: false;
auditLogs: false; auditLogs: false;
auditLogsRetentionDays: 0; auditLogsRetentionDays: 0;
samlSSO: false; samlSSO: true;
scim: false; scim: true;
ldap: true; ldap: true;
status: null; status: null;
trial_end: null; trial_end: null;
@@ -5,6 +5,7 @@ import {
OrgMembershipRole, OrgMembershipRole,
OrgMembershipStatus, OrgMembershipStatus,
SecretKeyEncoding, SecretKeyEncoding,
TableName,
TSamlConfigs, TSamlConfigs,
TSamlConfigsUpdate TSamlConfigsUpdate
} from "@app/db/schemas"; } from "@app/db/schemas";
@@ -31,7 +32,7 @@ import { TCreateSamlCfgDTO, TGetSamlCfgDTO, TSamlLoginDTO, TUpdateSamlCfgDTO } f
type TSamlConfigServiceFactoryDep = { type TSamlConfigServiceFactoryDep = {
samlConfigDAL: TSamlConfigDALFactory; samlConfigDAL: TSamlConfigDALFactory;
userDAL: Pick<TUserDALFactory, "create" | "findUserByEmail" | "transaction" | "updateById">; userDAL: Pick<TUserDALFactory, "create" | "findOne" | "transaction" | "updateById">;
orgDAL: Pick< orgDAL: Pick<
TOrgDALFactory, TOrgDALFactory,
"createMembership" | "updateMembershipById" | "findMembership" | "findOrgById" | "findOne" | "updateById" "createMembership" | "updateMembershipById" | "findMembership" | "findOrgById" | "findOne" | "updateById"
@@ -299,16 +300,30 @@ export const samlConfigServiceFactory = ({
}; };
}; };
const samlLogin = async ({ firstName, email, lastName, authProvider, orgId, relayState }: TSamlLoginDTO) => { const samlLogin = async ({
username,
email,
firstName,
lastName,
authProvider,
orgId,
relayState
}: TSamlLoginDTO) => {
const appCfg = getConfig(); const appCfg = getConfig();
let user = await userDAL.findUserByEmail(email); let user = await userDAL.findOne({ username });
const organization = await orgDAL.findOrgById(orgId); const organization = await orgDAL.findOrgById(orgId);
if (!organization) throw new BadRequestError({ message: "Org not found" }); if (!organization) throw new BadRequestError({ message: "Org not found" });
if (user) { if (user) {
await userDAL.transaction(async (tx) => { await userDAL.transaction(async (tx) => {
const [orgMembership] = await orgDAL.findMembership({ userId: user.id, orgId }, { tx }); const [orgMembership] = await orgDAL.findMembership(
{
userId: user.id,
[`${TableName.OrgMembership}.orgId` as "id"]: orgId
},
{ tx }
);
if (!orgMembership) { if (!orgMembership) {
await orgDAL.createMembership( await orgDAL.createMembership(
{ {
@@ -334,7 +349,7 @@ export const samlConfigServiceFactory = ({
user = await userDAL.transaction(async (tx) => { user = await userDAL.transaction(async (tx) => {
const newUser = await userDAL.create( const newUser = await userDAL.create(
{ {
username: email, username,
email, email,
firstName, firstName,
lastName, lastName,
@@ -36,7 +36,8 @@ export type TGetSamlCfgDTO =
}; };
export type TSamlLoginDTO = { export type TSamlLoginDTO = {
email: string; username: string;
email?: string;
firstName: string; firstName: string;
lastName?: string; lastName?: string;
authProvider: string; authProvider: string;
+13 -9
View File
@@ -28,12 +28,12 @@ export const buildScimUser = ({
}: { }: {
userId: string; userId: string;
username: string; username: string;
email: string; email?: string | null;
firstName: string; firstName: string;
lastName: string; lastName: string;
active: boolean; active: boolean;
}): TScimUser => { }): TScimUser => {
return { const scimUser = {
schemas: ["urn:ietf:params:scim:schemas:core:2.0:User"], schemas: ["urn:ietf:params:scim:schemas:core:2.0:User"],
id: userId, id: userId,
userName: username, userName: username,
@@ -43,13 +43,15 @@ export const buildScimUser = ({
middleName: null, middleName: null,
familyName: lastName familyName: lastName
}, },
emails: [ emails: email
{ ? [
primary: true, {
value: email, primary: true,
type: "work" value: email,
} type: "work"
], }
]
: [],
active, active,
groups: [], groups: [],
meta: { meta: {
@@ -57,4 +59,6 @@ export const buildScimUser = ({
location: null location: null
} }
}; };
return scimUser;
}; };
+20 -15
View File
@@ -1,7 +1,7 @@
import { ForbiddenError } from "@casl/ability"; import { ForbiddenError } from "@casl/ability";
import jwt from "jsonwebtoken"; import jwt from "jsonwebtoken";
import { OrgMembershipRole, OrgMembershipStatus } from "@app/db/schemas"; import { OrgMembershipRole, OrgMembershipStatus, TableName } from "@app/db/schemas";
import { TScimDALFactory } from "@app/ee/services/scim/scim-dal"; import { TScimDALFactory } from "@app/ee/services/scim/scim-dal";
import { getConfig } from "@app/lib/config/env"; import { getConfig } from "@app/lib/config/env";
import { BadRequestError, ScimRequestError, UnauthorizedError } from "@app/lib/errors"; import { BadRequestError, ScimRequestError, UnauthorizedError } from "@app/lib/errors";
@@ -146,7 +146,7 @@ export const scimServiceFactory = ({
const users = await orgDAL.findMembership( const users = await orgDAL.findMembership(
{ {
orgId, [`${TableName.OrgMembership}.orgId` as "id"]: orgId,
...parseFilter(filter) ...parseFilter(filter)
}, },
findOpts findOpts
@@ -155,10 +155,10 @@ export const scimServiceFactory = ({
const scimUsers = users.map(({ userId, username, firstName, lastName, email }) => const scimUsers = users.map(({ userId, username, firstName, lastName, email }) =>
buildScimUser({ buildScimUser({
userId: userId ?? "", userId: userId ?? "",
username: username ?? "", username,
firstName: firstName ?? "", firstName: firstName ?? "",
lastName: lastName ?? "", lastName: lastName ?? "",
email: email ?? "", email,
active: true active: true
}) })
); );
@@ -174,7 +174,7 @@ export const scimServiceFactory = ({
const [membership] = await orgDAL const [membership] = await orgDAL
.findMembership({ .findMembership({
userId, userId,
orgId [`${TableName.OrgMembership}.orgId` as "id"]: orgId
}) })
.catch(() => { .catch(() => {
throw new ScimRequestError({ throw new ScimRequestError({
@@ -205,8 +205,7 @@ export const scimServiceFactory = ({
}); });
}; };
// TODO: update SCIM endpoints to add username const createScimUser = async ({ username, email, firstName, lastName, orgId }: TCreateScimUserDTO) => {
const createScimUser = async ({ firstName, lastName, email, orgId }: TCreateScimUserDTO) => {
const org = await orgDAL.findById(orgId); const org = await orgDAL.findById(orgId);
if (!org) if (!org)
@@ -222,12 +221,18 @@ export const scimServiceFactory = ({
}); });
let user = await userDAL.findOne({ let user = await userDAL.findOne({
email username
}); });
if (user) { if (user) {
await userDAL.transaction(async (tx) => { await userDAL.transaction(async (tx) => {
const [orgMembership] = await orgDAL.findMembership({ userId: user.id, orgId }, { tx }); const [orgMembership] = await orgDAL.findMembership(
{
userId: user.id,
[`${TableName.OrgMembership}.orgId` as "id"]: orgId
},
{ tx }
);
if (orgMembership) if (orgMembership)
throw new ScimRequestError({ throw new ScimRequestError({
detail: "User already exists in the database", detail: "User already exists in the database",
@@ -251,7 +256,7 @@ export const scimServiceFactory = ({
user = await userDAL.transaction(async (tx) => { user = await userDAL.transaction(async (tx) => {
const newUser = await userDAL.create( const newUser = await userDAL.create(
{ {
username: email, username,
email, email,
firstName, firstName,
lastName, lastName,
@@ -279,7 +284,7 @@ export const scimServiceFactory = ({
await smtpService.sendMail({ await smtpService.sendMail({
template: SmtpTemplates.ScimUserProvisioned, template: SmtpTemplates.ScimUserProvisioned,
subjectLine: "Infisical organization invitation", subjectLine: "Infisical organization invitation",
recipients: [email], recipients: email ? [email] : [],
substitutions: { substitutions: {
organizationName: org.name, organizationName: org.name,
callback_url: `${appCfg.SITE_URL}/api/v1/sso/redirect/saml2/organizations/${org.slug}` callback_url: `${appCfg.SITE_URL}/api/v1/sso/redirect/saml2/organizations/${org.slug}`
@@ -300,7 +305,7 @@ export const scimServiceFactory = ({
const [membership] = await orgDAL const [membership] = await orgDAL
.findMembership({ .findMembership({
userId, userId,
orgId [`${TableName.OrgMembership}.orgId` as "id"]: orgId
}) })
.catch(() => { .catch(() => {
throw new ScimRequestError({ throw new ScimRequestError({
@@ -348,7 +353,7 @@ export const scimServiceFactory = ({
return buildScimUser({ return buildScimUser({
userId: membership.userId as string, userId: membership.userId as string,
username: membership.username, username: membership.username,
email: membership.email ?? "", email: membership.email,
firstName: membership.firstName as string, firstName: membership.firstName as string,
lastName: membership.lastName as string, lastName: membership.lastName as string,
active active
@@ -359,7 +364,7 @@ export const scimServiceFactory = ({
const [membership] = await orgDAL const [membership] = await orgDAL
.findMembership({ .findMembership({
userId, userId,
orgId [`${TableName.OrgMembership}.orgId` as "id"]: orgId
}) })
.catch(() => { .catch(() => {
throw new ScimRequestError({ throw new ScimRequestError({
@@ -394,7 +399,7 @@ export const scimServiceFactory = ({
return buildScimUser({ return buildScimUser({
userId: membership.userId as string, userId: membership.userId as string,
username: membership.username, username: membership.username,
email: membership.email ?? "", email: membership.email,
firstName: membership.firstName as string, firstName: membership.firstName as string,
lastName: membership.lastName as string, lastName: membership.lastName as string,
active active
+2 -1
View File
@@ -32,7 +32,8 @@ export type TGetScimUserDTO = {
}; };
export type TCreateScimUserDTO = { export type TCreateScimUserDTO = {
email: string; username: string;
email?: string;
firstName: string; firstName: string;
lastName: string; lastName: string;
orgId: string; orgId: string;
@@ -301,7 +301,6 @@ export const authLoginServiceFactory = ({ userDAL, tokenService, smtpService }:
{ {
authTokenType: AuthTokenType.PROVIDER_TOKEN, authTokenType: AuthTokenType.PROVIDER_TOKEN,
userId: user.id, userId: user.id,
// email: user.email,
username: user.username, username: user.username,
firstName: user.firstName, firstName: user.firstName,
lastName: user.lastName, lastName: user.lastName,
+3 -1
View File
@@ -250,7 +250,9 @@ export const orgDALFactory = (db: TDbClient) => {
db.ref("firstName").withSchema(TableName.Users), db.ref("firstName").withSchema(TableName.Users),
db.ref("lastName").withSchema(TableName.Users), db.ref("lastName").withSchema(TableName.Users),
db.ref("scimEnabled").withSchema(TableName.Organization) db.ref("scimEnabled").withSchema(TableName.Organization)
); )
.where({ isGhost: false });
if (limit) void query.limit(limit); if (limit) void query.limit(limit);
if (offset) void query.offset(offset); if (offset) void query.offset(offset);
if (sort) { if (sort) {
@@ -30,6 +30,15 @@ export const LDAPStep = ({
password password
}); });
if (!nextUrl) {
createNotification({
text: "Login unsuccessful. Double-check your credentials and try again.",
type: "error"
});
return;
}
createNotification({ createNotification({
text: "Successfully logged in", text: "Successfully logged in",
type: "success" type: "success"