Merge remote-tracking branch 'origin/main' into feat/gateway-v2

This commit is contained in:
Sheen Capadngan
2025-09-11 02:35:17 +08:00
64 changed files with 808 additions and 182 deletions
+9
View File
@@ -56,6 +56,15 @@ export const mockKeyStore = (): TKeyStoreFactory => {
incrementBy: async () => { incrementBy: async () => {
return 1; return 1;
}, },
pgGetIntItem: async (key) => {
const value = store[key];
if (typeof value === "number") {
return Number(value);
}
},
pgIncrementBy: async () => {
return 1;
},
getItems: async (keys) => { getItems: async (keys) => {
const values = keys.map((key) => { const values = keys.map((key) => {
const value = store[key]; const value = store[key];
+3 -1
View File
@@ -15,6 +15,7 @@ import { mockSmtpServer } from "./mocks/smtp";
import { initDbConnection } from "@app/db"; import { initDbConnection } from "@app/db";
import { queueServiceFactory } from "@app/queue"; import { queueServiceFactory } from "@app/queue";
import { keyStoreFactory } from "@app/keystore/keystore"; import { keyStoreFactory } from "@app/keystore/keystore";
import { keyValueStoreDALFactory } from "@app/keystore/key-value-store-dal";
import { initializeHsmModule } from "@app/ee/services/hsm/hsm-fns"; import { initializeHsmModule } from "@app/ee/services/hsm/hsm-fns";
import { buildRedisFromConfig } from "@app/lib/config/redis"; import { buildRedisFromConfig } from "@app/lib/config/redis";
import { superAdminDALFactory } from "@app/services/super-admin/super-admin-dal"; import { superAdminDALFactory } from "@app/services/super-admin/super-admin-dal";
@@ -62,7 +63,8 @@ export default {
const smtp = mockSmtpServer(); const smtp = mockSmtpServer();
const queue = queueServiceFactory(envCfg, { dbConnectionUrl: envCfg.DB_CONNECTION_URI }); const queue = queueServiceFactory(envCfg, { dbConnectionUrl: envCfg.DB_CONNECTION_URI });
const keyStore = keyStoreFactory(envCfg); const keyValueStoreDAL = keyValueStoreDALFactory(db);
const keyStore = keyStoreFactory(envCfg, keyValueStoreDAL);
await queue.initialize(); await queue.initialize();
+8
View File
@@ -197,6 +197,9 @@ import {
TInternalKms, TInternalKms,
TInternalKmsInsert, TInternalKmsInsert,
TInternalKmsUpdate, TInternalKmsUpdate,
TKeyValueStore,
TKeyValueStoreInsert,
TKeyValueStoreUpdate,
TKmipClientCertificates, TKmipClientCertificates,
TKmipClientCertificatesInsert, TKmipClientCertificatesInsert,
TKmipClientCertificatesUpdate, TKmipClientCertificatesUpdate,
@@ -1296,5 +1299,10 @@ declare module "knex/types/tables" {
TUserNotificationsInsert, TUserNotificationsInsert,
TUserNotificationsUpdate TUserNotificationsUpdate
>; >;
[TableName.KeyValueStore]: KnexOriginal.CompositeTableType<
TKeyValueStore,
TKeyValueStoreInsert,
TKeyValueStoreUpdate
>;
} }
} }
@@ -0,0 +1,18 @@
import { Knex } from "knex";
import { TableName } from "../schemas";
export async function up(knex: Knex): Promise<void> {
if (!(await knex.schema.hasTable(TableName.KeyValueStore))) {
await knex.schema.createTable(TableName.KeyValueStore, (t) => {
t.text("key").primary();
t.bigint("integerValue");
t.datetime("expiresAt");
t.timestamps(true, true, true);
});
}
}
export async function down(knex: Knex): Promise<void> {
await knex.schema.dropTableIfExists(TableName.KeyValueStore);
}
+1
View File
@@ -63,6 +63,7 @@ export * from "./integration-auths";
export * from "./integrations"; export * from "./integrations";
export * from "./internal-certificate-authorities"; export * from "./internal-certificate-authorities";
export * from "./internal-kms"; export * from "./internal-kms";
export * from "./key-value-store";
export * from "./kmip-client-certificates"; export * from "./kmip-client-certificates";
export * from "./kmip-clients"; export * from "./kmip-clients";
export * from "./kmip-org-configs"; export * from "./kmip-org-configs";
+20
View File
@@ -0,0 +1,20 @@
// Code generated by automation script, DO NOT EDIT.
// Automated by pulling database and generating zod schema
// To update. Just run npm run generate:schema
// Written by akhilmhdh.
import { z } from "zod";
import { TImmutableDBKeys } from "./models";
export const KeyValueStoreSchema = z.object({
key: z.string(),
integerValue: z.coerce.number().nullable().optional(),
expiresAt: z.date().nullable().optional(),
createdAt: z.date(),
updatedAt: z.date()
});
export type TKeyValueStore = z.infer<typeof KeyValueStoreSchema>;
export type TKeyValueStoreInsert = Omit<z.input<typeof KeyValueStoreSchema>, TImmutableDBKeys>;
export type TKeyValueStoreUpdate = Partial<Omit<z.input<typeof KeyValueStoreSchema>, TImmutableDBKeys>>;
+3 -1
View File
@@ -186,7 +186,9 @@ export enum TableName {
OrgRelayConfig = "org_relay_config", OrgRelayConfig = "org_relay_config",
OrgGatewayConfigV2 = "org_gateway_config_v2", OrgGatewayConfigV2 = "org_gateway_config_v2",
Relay = "relays", Relay = "relays",
GatewayV2 = "gateways_v2" GatewayV2 = "gateways_v2",
KeyValueStore = "key_value_store"
} }
export type TImmutableDBKeys = "id" | "createdAt" | "updatedAt" | "commitId"; export type TImmutableDBKeys = "id" | "createdAt" | "updatedAt" | "commitId";
@@ -84,7 +84,9 @@ export const registerDynamicSecretRouter = async (server: FastifyZodProvider) =>
}), }),
response: { response: {
200: z.object({ 200: z.object({
dynamicSecret: SanitizedDynamicSecretSchema dynamicSecret: SanitizedDynamicSecretSchema.extend({
inputs: z.unknown()
})
}) })
} }
}, },
@@ -151,7 +153,9 @@ export const registerDynamicSecretRouter = async (server: FastifyZodProvider) =>
}), }),
response: { response: {
200: z.object({ 200: z.object({
dynamicSecret: SanitizedDynamicSecretSchema dynamicSecret: SanitizedDynamicSecretSchema.extend({
inputs: z.unknown()
})
}) })
} }
}, },
+8 -1
View File
@@ -43,6 +43,12 @@ export const registerLicenseRouter = async (server: FastifyZodProvider) => {
}, },
schema: { schema: {
params: z.object({ organizationId: z.string().trim() }), params: z.object({ organizationId: z.string().trim() }),
querystring: z.object({
refreshCache: z
.enum(["true", "false"])
.default("false")
.transform((value) => value === "true")
}),
response: { response: {
200: z.object({ plan: z.any() }) 200: z.object({ plan: z.any() })
} }
@@ -54,7 +60,8 @@ export const registerLicenseRouter = async (server: FastifyZodProvider) => {
actor: req.permission.type, actor: req.permission.type,
actorOrgId: req.permission.orgId, actorOrgId: req.permission.orgId,
actorAuthMethod: req.permission.authMethod, actorAuthMethod: req.permission.authMethod,
orgId: req.params.organizationId orgId: req.params.organizationId,
refreshCache: req.query.refreshCache
}); });
return { plan }; return { plan };
} }
@@ -46,7 +46,10 @@ export const dynamicSecretLeaseDALFactory = (db: TDbClient) => {
const countLeasesForDynamicSecret = async (dynamicSecretId: string, tx?: Knex) => { const countLeasesForDynamicSecret = async (dynamicSecretId: string, tx?: Knex) => {
try { try {
const doc = await (tx || db)(TableName.DynamicSecretLease).count("*").where({ dynamicSecretId }).first(); const doc = await (tx || db.replicaNode())(TableName.DynamicSecretLease)
.count("*")
.where({ dynamicSecretId })
.first();
return parseInt(doc || "0", 10); return parseInt(doc || "0", 10);
} catch (error) { } catch (error) {
throw new DatabaseError({ error, name: "DynamicSecretCountLeases" }); throw new DatabaseError({ error, name: "DynamicSecretCountLeases" });
@@ -55,7 +58,7 @@ export const dynamicSecretLeaseDALFactory = (db: TDbClient) => {
const findById = async (id: string, tx?: Knex) => { const findById = async (id: string, tx?: Knex) => {
try { try {
const doc = await (tx || db)(TableName.DynamicSecretLease) const doc = await (tx || db.replicaNode())(TableName.DynamicSecretLease)
.where({ [`${TableName.DynamicSecretLease}.id` as "id"]: id }) .where({ [`${TableName.DynamicSecretLease}.id` as "id"]: id })
.first() .first()
.join( .join(
@@ -190,7 +190,7 @@ export const dynamicSecretServiceFactory = ({
return cfg; return cfg;
}); });
return dynamicSecretCfg; return { ...dynamicSecretCfg, inputs };
}; };
const updateByName: TDynamicSecretServiceFactory["updateByName"] = async ({ const updateByName: TDynamicSecretServiceFactory["updateByName"] = async ({
@@ -355,7 +355,7 @@ export const dynamicSecretServiceFactory = ({
return cfg; return cfg;
}); });
return updatedDynamicCfg; return { ...updatedDynamicCfg, inputs: updatedInput };
}; };
const deleteByName: TDynamicSecretServiceFactory["deleteByName"] = async ({ const deleteByName: TDynamicSecretServiceFactory["deleteByName"] = async ({
@@ -165,6 +165,7 @@ export const DynamicSecretSqlDBSchema = z.object({
revocationStatement: z.string().trim(), revocationStatement: z.string().trim(),
renewStatement: z.string().trim().optional(), renewStatement: z.string().trim().optional(),
ca: z.string().optional(), ca: z.string().optional(),
sslEnabled: z.boolean().optional(),
gatewayId: z.string().nullable().optional() gatewayId: z.string().nullable().optional()
}); });
@@ -1,4 +1,5 @@
import handlebars from "handlebars"; import handlebars from "handlebars";
import RE2 from "re2";
import knex from "knex"; import knex from "knex";
import { z } from "zod"; import { z } from "zod";
@@ -156,19 +157,40 @@ export const SqlDatabaseProvider = ({
return { ...providerInputs, hostIp }; return { ...providerInputs, hostIp };
}; };
const $getClient = async (providerInputs: z.infer<typeof DynamicSecretSqlDBSchema> & { hostIp: string }) => { const $getClient = async (
providerInputs: z.infer<typeof DynamicSecretSqlDBSchema> & { hostIp: string; originalHost: string }
) => {
const ssl = providerInputs.ca const ssl = providerInputs.ca
? { rejectUnauthorized: false, ca: providerInputs.ca, servername: providerInputs.host } ? { rejectUnauthorized: false, ca: providerInputs.ca, servername: providerInputs.host }
: undefined; : undefined;
const isMsSQLClient = providerInputs.client === SqlProviders.MsSQL; const isMsSQLClient = providerInputs.client === SqlProviders.MsSQL;
/*
We route through the gateway by setting connection.host = "localhost".
Azure SQL identifies the logical server from the TDS login name when the host
isn’t the Azure FQDN. Therefore, when using the gateway, ensure username is
"user@<azure-server-name>" so Azure opens the correct logical server.
Direct connections to the Azure FQDN usually don’t require this suffix.
*/
const isAzureSql = isMsSQLClient && new RE2(/\.database\.windows\.net$/i).test(providerInputs.originalHost);
const azureServerLabel =
isAzureSql && providerInputs.gatewayId ? providerInputs.originalHost?.split(".")[0] : undefined;
const effectiveUser =
isAzureSql && !providerInputs.username.includes("@") && azureServerLabel
? `${providerInputs.username}@${azureServerLabel}`
: providerInputs.username;
const db = knex({ const db = knex({
client: providerInputs.client, client: providerInputs.client,
connection: { connection: {
database: providerInputs.database, database: providerInputs.database,
port: providerInputs.port, port: providerInputs.port,
host: providerInputs.client === SqlProviders.Postgres ? providerInputs.hostIp : providerInputs.host, host:
user: providerInputs.username, providerInputs.client === SqlProviders.Postgres && !providerInputs.gatewayId
? providerInputs.hostIp
: providerInputs.host,
user: effectiveUser,
password: providerInputs.password, password: providerInputs.password,
ssl, ssl,
// @ts-expect-error this is because of knexjs type signature issue. This is directly passed to driver // @ts-expect-error this is because of knexjs type signature issue. This is directly passed to driver
@@ -176,6 +198,7 @@ export const SqlDatabaseProvider = ({
// https://github.com/tediousjs/tedious/blob/ebb023ed90969a7ec0e4b036533ad52739d921f7/test/config.ci.ts#L19 // https://github.com/tediousjs/tedious/blob/ebb023ed90969a7ec0e4b036533ad52739d921f7/test/config.ci.ts#L19
options: isMsSQLClient options: isMsSQLClient
? { ? {
...(providerInputs.sslEnabled !== undefined ? { encrypt: providerInputs.sslEnabled } : {}),
trustServerCertificate: !providerInputs.ca, trustServerCertificate: !providerInputs.ca,
cryptoCredentialsDetails: providerInputs.ca ? { ca: providerInputs.ca } : {} cryptoCredentialsDetails: providerInputs.ca ? { ca: providerInputs.ca } : {}
} }
@@ -238,7 +261,13 @@ export const SqlDatabaseProvider = ({
const providerInputs = await validateProviderInputs(inputs); const providerInputs = await validateProviderInputs(inputs);
let isConnected = false; let isConnected = false;
const gatewayCallback = async (host = providerInputs.host, port = providerInputs.port) => { const gatewayCallback = async (host = providerInputs.host, port = providerInputs.port) => {
const db = await $getClient({ ...providerInputs, port, host, hostIp: providerInputs.hostIp }); const db = await $getClient({
...providerInputs,
port,
host,
hostIp: providerInputs.hostIp,
originalHost: providerInputs.host
});
// oracle needs from keyword // oracle needs from keyword
const testStatement = providerInputs.client === SqlProviders.Oracle ? "SELECT 1 FROM DUAL" : "SELECT 1"; const testStatement = providerInputs.client === SqlProviders.Oracle ? "SELECT 1 FROM DUAL" : "SELECT 1";
@@ -279,7 +308,12 @@ export const SqlDatabaseProvider = ({
const password = generatePassword(providerInputs.client, providerInputs.passwordRequirements); const password = generatePassword(providerInputs.client, providerInputs.passwordRequirements);
const gatewayCallback = async (host = providerInputs.host, port = providerInputs.port) => { const gatewayCallback = async (host = providerInputs.host, port = providerInputs.port) => {
const db = await $getClient({ ...providerInputs, port, host }); const db = await $getClient({
...providerInputs,
port,
host,
originalHost: providerInputs.host
});
try { try {
const expiration = new Date(expireAt).toISOString(); const expiration = new Date(expireAt).toISOString();
@@ -322,7 +356,12 @@ export const SqlDatabaseProvider = ({
const username = entityId; const username = entityId;
const { database } = providerInputs; const { database } = providerInputs;
const gatewayCallback = async (host = providerInputs.host, port = providerInputs.port) => { const gatewayCallback = async (host = providerInputs.host, port = providerInputs.port) => {
const db = await $getClient({ ...providerInputs, port, host }); const db = await $getClient({
...providerInputs,
port,
host,
originalHost: providerInputs.host
});
try { try {
const revokeStatement = handlebars.compile(providerInputs.revocationStatement)({ username, database }); const revokeStatement = handlebars.compile(providerInputs.revocationStatement)({ username, database });
const queries = revokeStatement.toString().split(";").filter(Boolean); const queries = revokeStatement.toString().split(";").filter(Boolean);
@@ -357,7 +396,12 @@ export const SqlDatabaseProvider = ({
if (!providerInputs.renewStatement) return { entityId }; if (!providerInputs.renewStatement) return { entityId };
const gatewayCallback = async (host = providerInputs.host, port = providerInputs.port) => { const gatewayCallback = async (host = providerInputs.host, port = providerInputs.port) => {
const db = await $getClient({ ...providerInputs, port, host }); const db = await $getClient({
...providerInputs,
port,
host,
originalHost: providerInputs.host
});
const expiration = new Date(expireAt).toISOString(); const expiration = new Date(expireAt).toISOString();
const { database } = providerInputs; const { database } = providerInputs;
@@ -13,7 +13,7 @@ export const gatewayDALFactory = (db: TDbClient) => {
{ offset, limit, sort, tx }: TFindOpt<TGateways> = {} { offset, limit, sort, tx }: TFindOpt<TGateways> = {}
) => { ) => {
try { try {
const query = (tx || db)(TableName.Gateway) const query = (tx || db.replicaNode())(TableName.Gateway)
// eslint-disable-next-line @typescript-eslint/no-misused-promises // eslint-disable-next-line @typescript-eslint/no-misused-promises
.where(buildFindFilter(filter, TableName.Gateway, ["orgId"])) .where(buildFindFilter(filter, TableName.Gateway, ["orgId"]))
.join(TableName.Identity, `${TableName.Identity}.id`, `${TableName.Gateway}.identityId`) .join(TableName.Identity, `${TableName.Identity}.id`, `${TableName.Gateway}.identityId`)
@@ -23,7 +23,7 @@ export const userGroupMembershipDALFactory = (db: TDbClient) => {
.whereIn(`${TableName.ProjectMembership}.projectId`, projectIds) .whereIn(`${TableName.ProjectMembership}.projectId`, projectIds)
.pluck(`${TableName.ProjectMembership}.projectId`); .pluck(`${TableName.ProjectMembership}.projectId`);
const userGroupMemberships: string[] = await (tx || db)(TableName.UserGroupMembership) const userGroupMemberships: string[] = await (tx || db.replicaNode())(TableName.UserGroupMembership)
.where(`${TableName.UserGroupMembership}.userId`, userId) .where(`${TableName.UserGroupMembership}.userId`, userId)
.whereNot(`${TableName.UserGroupMembership}.groupId`, groupId) .whereNot(`${TableName.UserGroupMembership}.groupId`, groupId)
.join( .join(
@@ -79,7 +79,7 @@ export const userGroupMembershipDALFactory = (db: TDbClient) => {
.pluck(`${TableName.GroupProjectMembership}.groupId`); .pluck(`${TableName.GroupProjectMembership}.groupId`);
// main query // main query
const members = await (tx || db)(TableName.UserGroupMembership) const members = await (tx || db.replicaNode())(TableName.UserGroupMembership)
.where(`${TableName.UserGroupMembership}.groupId`, groupId) .where(`${TableName.UserGroupMembership}.groupId`, groupId)
.where(`${TableName.UserGroupMembership}.isPending`, false) .where(`${TableName.UserGroupMembership}.isPending`, false)
.join(TableName.Users, `${TableName.UserGroupMembership}.userId`, `${TableName.Users}.id`) .join(TableName.Users, `${TableName.UserGroupMembership}.userId`, `${TableName.Users}.id`)
@@ -127,6 +127,20 @@ export const ldapConfigServiceFactory = ({
message: message:
"Failed to create LDAP configuration due to plan restriction. Upgrade plan to create LDAP configuration." "Failed to create LDAP configuration due to plan restriction. Upgrade plan to create LDAP configuration."
}); });
const org = await orgDAL.findOrgById(orgId);
if (!org) {
throw new NotFoundError({ message: `Could not find organization with ID "${orgId}"` });
}
if (org.googleSsoAuthEnforced && isActive) {
throw new BadRequestError({
message:
"You cannot enable LDAP SSO while Google OAuth is enforced. Disable Google OAuth enforcement to enable LDAP SSO."
});
}
const { encryptor } = await kmsService.createCipherPairWithDataKey({ const { encryptor } = await kmsService.createCipherPairWithDataKey({
type: KmsDataKey.Organization, type: KmsDataKey.Organization,
orgId orgId
@@ -233,6 +247,19 @@ export const ldapConfigServiceFactory = ({
"Failed to update LDAP configuration due to plan restriction. Upgrade plan to update LDAP configuration." "Failed to update LDAP configuration due to plan restriction. Upgrade plan to update LDAP configuration."
}); });
const org = await orgDAL.findOrgById(orgId);
if (!org) {
throw new NotFoundError({ message: `Could not find organization with ID "${orgId}"` });
}
if (org.googleSsoAuthEnforced && isActive) {
throw new BadRequestError({
message:
"You cannot enable LDAP SSO while Google OAuth is enforced. Disable Google OAuth enforcement to enable LDAP SSO."
});
}
const updateQuery: TLdapConfigsUpdate = { const updateQuery: TLdapConfigsUpdate = {
isActive, isActive,
url, url,
@@ -28,7 +28,7 @@ export const licenseDALFactory = (db: TDbClient) => {
const countOrgUsersAndIdentities = async (orgId: string | null, tx?: Knex) => { const countOrgUsersAndIdentities = async (orgId: string | null, tx?: Knex) => {
try { try {
// count org users // count org users
const userDoc = await (tx || db)(TableName.OrgMembership) const userDoc = await (tx || db.replicaNode())(TableName.OrgMembership)
.where({ status: OrgMembershipStatus.Accepted }) .where({ status: OrgMembershipStatus.Accepted })
.andWhere((bd) => { .andWhere((bd) => {
if (orgId) { if (orgId) {
@@ -42,7 +42,7 @@ export const licenseDALFactory = (db: TDbClient) => {
const userCount = Number(userDoc?.[0].count); const userCount = Number(userDoc?.[0].count);
// count org identities // count org identities
const identityDoc = await (tx || db)(TableName.IdentityOrgMembership) const identityDoc = await (tx || db.replicaNode())(TableName.IdentityOrgMembership)
.where((bd) => { .where((bd) => {
if (orgId) { if (orgId) {
void bd.where({ orgId }); void bd.where({ orgId });
@@ -99,6 +99,17 @@ export const licenseServiceFactory = ({
const workspacesUsed = await projectDAL.countOfOrgProjects(null); const workspacesUsed = await projectDAL.countOfOrgProjects(null);
currentPlan.workspacesUsed = workspacesUsed; currentPlan.workspacesUsed = workspacesUsed;
const usedIdentitySeats = await licenseDAL.countOrgUsersAndIdentities(null);
if (usedIdentitySeats !== currentPlan.identitiesUsed) {
const usedSeats = await licenseDAL.countOfOrgMembers(null);
await licenseServerOnPremApi.request.patch(`/api/license/v1/license`, {
usedSeats,
usedIdentitySeats
});
currentPlan.identitiesUsed = usedIdentitySeats;
currentPlan.membersUsed = usedSeats;
}
onPremFeatures = currentPlan; onPremFeatures = currentPlan;
logger.info("Successfully synchronized license key features"); logger.info("Successfully synchronized license key features");
} catch (error) { } catch (error) {
@@ -226,10 +237,13 @@ export const licenseServiceFactory = ({
}; };
const refreshPlan = async (orgId: string) => { const refreshPlan = async (orgId: string) => {
await keyStore.deleteItem(FEATURE_CACHE_KEY(orgId));
if (instanceType === InstanceType.Cloud) { if (instanceType === InstanceType.Cloud) {
await keyStore.deleteItem(FEATURE_CACHE_KEY(orgId));
await getPlan(orgId); await getPlan(orgId);
} }
if (instanceType === InstanceType.EnterpriseOnPrem) {
await syncLicenseKeyOnPremFeatures(true);
}
}; };
const generateOrgCustomerId = async (orgName: string, email?: string | null) => { const generateOrgCustomerId = async (orgName: string, email?: string | null) => {
@@ -296,8 +310,19 @@ export const licenseServiceFactory = ({
return data; return data;
}; };
const getOrgPlan = async ({ orgId, actor, actorId, actorOrgId, actorAuthMethod, projectId }: TOrgPlanDTO) => { const getOrgPlan = async ({
orgId,
actor,
actorId,
actorOrgId,
actorAuthMethod,
projectId,
refreshCache
}: TOrgPlanDTO) => {
await permissionService.getOrgPermission(actor, actorId, orgId, actorAuthMethod, actorOrgId); await permissionService.getOrgPermission(actor, actorId, orgId, actorAuthMethod, actorOrgId);
if (refreshCache) {
await refreshPlan(orgId);
}
const plan = await getPlan(orgId, projectId); const plan = await getPlan(orgId, projectId);
return plan; return plan;
}; };
@@ -87,6 +87,7 @@ export type TOrgPlansTableDTO = {
export type TOrgPlanDTO = { export type TOrgPlanDTO = {
projectId?: string; projectId?: string;
refreshCache?: boolean;
} & TOrgPermission; } & TOrgPermission;
export type TStartOrgTrialDTO = { export type TStartOrgTrialDTO = {
@@ -499,6 +499,13 @@ export const oidcConfigServiceFactory = ({
); );
ForbiddenError.from(permission).throwUnlessCan(OrgPermissionActions.Edit, OrgPermissionSubjects.Sso); ForbiddenError.from(permission).throwUnlessCan(OrgPermissionActions.Edit, OrgPermissionSubjects.Sso);
if (org.googleSsoAuthEnforced && isActive) {
throw new BadRequestError({
message:
"You cannot enable OIDC SSO while Google OAuth is enforced. Disable Google OAuth enforcement to enable OIDC SSO."
});
}
const { encryptor } = await kmsService.createCipherPairWithDataKey({ const { encryptor } = await kmsService.createCipherPairWithDataKey({
type: KmsDataKey.Organization, type: KmsDataKey.Organization,
orgId: org.id orgId: org.id
@@ -586,6 +593,13 @@ export const oidcConfigServiceFactory = ({
); );
ForbiddenError.from(permission).throwUnlessCan(OrgPermissionActions.Create, OrgPermissionSubjects.Sso); ForbiddenError.from(permission).throwUnlessCan(OrgPermissionActions.Create, OrgPermissionSubjects.Sso);
if (org.googleSsoAuthEnforced && isActive) {
throw new BadRequestError({
message:
"You cannot enable OIDC SSO while Google OAuth is enforced. Disable Google OAuth enforcement to enable OIDC SSO."
});
}
const { encryptor } = await kmsService.createCipherPairWithDataKey({ const { encryptor } = await kmsService.createCipherPairWithDataKey({
type: KmsDataKey.Organization, type: KmsDataKey.Organization,
orgId: org.id orgId: org.id
@@ -82,6 +82,19 @@ export const samlConfigServiceFactory = ({
"Failed to create SAML SSO configuration due to plan restriction. Upgrade plan to create SSO configuration." "Failed to create SAML SSO configuration due to plan restriction. Upgrade plan to create SSO configuration."
}); });
const org = await orgDAL.findOrgById(orgId);
if (!org) {
throw new NotFoundError({ message: `Could not find organization with ID "${orgId}"` });
}
if (org.googleSsoAuthEnforced && isActive) {
throw new BadRequestError({
message:
"You cannot enable SAML SSO while Google OAuth is enforced. Disable Google OAuth enforcement to enable SAML SSO."
});
}
const { encryptor } = await kmsService.createCipherPairWithDataKey({ const { encryptor } = await kmsService.createCipherPairWithDataKey({
type: KmsDataKey.Organization, type: KmsDataKey.Organization,
orgId orgId
@@ -120,6 +133,19 @@ export const samlConfigServiceFactory = ({
"Failed to update SAML SSO configuration due to plan restriction. Upgrade plan to update SSO configuration." "Failed to update SAML SSO configuration due to plan restriction. Upgrade plan to update SSO configuration."
}); });
const org = await orgDAL.findOrgById(orgId);
if (!org) {
throw new NotFoundError({ message: `Could not find organization with ID "${orgId}"` });
}
if (org.googleSsoAuthEnforced && isActive) {
throw new BadRequestError({
message:
"Cannot enable SAML SSO while Google OAuth is enforced. Disable Google OAuth enforcement to enable SAML SSO."
});
}
const updateQuery: TSamlConfigsUpdate = { authProvider, isActive, lastUsed: null }; const updateQuery: TSamlConfigsUpdate = { authProvider, isActive, lastUsed: null };
const { encryptor } = await kmsService.createCipherPairWithDataKey({ const { encryptor } = await kmsService.createCipherPairWithDataKey({
type: KmsDataKey.Organization, type: KmsDataKey.Organization,
@@ -345,7 +345,7 @@ export const secretApprovalRequestDALFactory = (db: TDbClient) => {
const findProjectRequestCount = async (projectId: string, userId: string, policyId?: string, tx?: Knex) => { const findProjectRequestCount = async (projectId: string, userId: string, policyId?: string, tx?: Knex) => {
try { try {
const docs = await (tx || db) const docs = await (tx || db.replicaNode())
.with( .with(
"temp", "temp",
(tx || db.replicaNode())(TableName.SecretApprovalRequest) (tx || db.replicaNode())(TableName.SecretApprovalRequest)
@@ -494,7 +494,7 @@ export const secretApprovalRequestDALFactory = (db: TDbClient) => {
.distinctOn(`${TableName.SecretApprovalRequest}.id`) .distinctOn(`${TableName.SecretApprovalRequest}.id`)
.as("inner"); .as("inner");
const query = (tx || db) const query = (tx || db.replicaNode())
.select("*") .select("*")
.select(db.raw("count(*) OVER() as total_count")) .select(db.raw("count(*) OVER() as total_count"))
.from(innerQuery) .from(innerQuery)
@@ -377,7 +377,7 @@ export const secretApprovalRequestSecretDALFactory = (db: TDbClient) => {
// special query for migration to v2 secret // special query for migration to v2 secret
const findByProjectId = async (projectId: string, tx?: Knex) => { const findByProjectId = async (projectId: string, tx?: Knex) => {
try { try {
const docs = await (tx || db)(TableName.SecretApprovalRequestSecret) const docs = await (tx || db.replicaNode())(TableName.SecretApprovalRequestSecret)
.join( .join(
TableName.SecretApprovalRequest, TableName.SecretApprovalRequest,
`${TableName.SecretApprovalRequest}.id`, `${TableName.SecretApprovalRequest}.id`,
@@ -787,6 +787,7 @@ export const secretApprovalRequestServiceFactory = ({
}, },
tx tx
); );
await secretV2BridgeDAL.invalidateSecretCacheByProjectId(projectId, tx);
return { return {
secrets: { created: newSecrets, updated: updatedSecrets, deleted: deletedSecret }, secrets: { created: newSecrets, updated: updatedSecrets, deleted: deletedSecret },
approval: updatedSecretApproval approval: updatedSecretApproval
@@ -976,6 +977,7 @@ export const secretApprovalRequestServiceFactory = ({
}, },
tx tx
); );
await secretV2BridgeDAL.invalidateSecretCacheByProjectId(projectId, tx);
return { return {
secrets: { created: newSecrets, updated: updatedSecrets, deleted: deletedSecret }, secrets: { created: newSecrets, updated: updatedSecrets, deleted: deletedSecret },
approval: updatedSecretApproval approval: updatedSecretApproval
@@ -983,7 +985,6 @@ export const secretApprovalRequestServiceFactory = ({
}); });
} }
await secretV2BridgeDAL.invalidateSecretCacheByProjectId(projectId);
await snapshotService.performSnapshot(folderId); await snapshotService.performSnapshot(folderId);
const [folder] = await folderDAL.findSecretPathByFolderIds(projectId, [folderId]); const [folder] = await folderDAL.findSecretPathByFolderIds(projectId, [folderId]);
if (!folder) { if (!folder) {
@@ -509,9 +509,9 @@ export const secretReplicationServiceFactory = ({
tx tx
); );
} }
await secretV2BridgeDAL.invalidateSecretCacheByProjectId(projectId, tx);
}); });
await secretV2BridgeDAL.invalidateSecretCacheByProjectId(projectId);
await secretQueueService.syncSecrets({ await secretQueueService.syncSecrets({
projectId, projectId,
orgId, orgId,
@@ -361,9 +361,8 @@ export const secretRotationQueueFactory = ({
}, },
tx tx
); );
await secretV2BridgeDAL.invalidateSecretCacheByProjectId(secretRotation.projectId, tx);
}); });
await secretV2BridgeDAL.invalidateSecretCacheByProjectId(secretRotation.projectId);
} else { } else {
if (!botKey) if (!botKey)
throw new NotFoundError({ throw new NotFoundError({
@@ -265,7 +265,7 @@ export const snapshotDALFactory = (db: TDbClient) => {
// then joins with respective secrets and folder // then joins with respective secrets and folder
const findRecursivelySnapshots = async (snapshotId: string, tx?: Knex) => { const findRecursivelySnapshots = async (snapshotId: string, tx?: Knex) => {
try { try {
const data = await (tx || db) const data = await (tx || db.replicaNode())
.withRecursive("parent", (qb) => { .withRecursive("parent", (qb) => {
void qb void qb
.from(TableName.Snapshot) .from(TableName.Snapshot)
@@ -419,7 +419,7 @@ export const snapshotDALFactory = (db: TDbClient) => {
// then joins with respective secrets and folder // then joins with respective secrets and folder
const findRecursivelySnapshotsV2Bridge = async (snapshotId: string, tx?: Knex) => { const findRecursivelySnapshotsV2Bridge = async (snapshotId: string, tx?: Knex) => {
try { try {
const data = await (tx || db) const data = await (tx || db.replicaNode())
.withRecursive("parent", (qb) => { .withRecursive("parent", (qb) => {
void qb void qb
.from(TableName.Snapshot) .from(TableName.Snapshot)
@@ -581,7 +581,11 @@ export const snapshotDALFactory = (db: TDbClient) => {
const docs = await (tx || db.replicaNode())(TableName.Snapshot) const docs = await (tx || db.replicaNode())(TableName.Snapshot)
.where(`${TableName.Snapshot}.folderId`, folderId) .where(`${TableName.Snapshot}.folderId`, folderId)
.join<TSecretSnapshots>( .join<TSecretSnapshots>(
(tx || db)(TableName.Snapshot).groupBy("folderId").max("createdAt").select("folderId").as("latestVersion"), (tx || db.replicaNode())(TableName.Snapshot)
.groupBy("folderId")
.max("createdAt")
.select("folderId")
.as("latestVersion"),
(bd) => { (bd) => {
bd.on(`${TableName.Snapshot}.folderId`, "latestVersion.folderId").andOn( bd.on(`${TableName.Snapshot}.folderId`, "latestVersion.folderId").andOn(
`${TableName.Snapshot}.createdAt`, `${TableName.Snapshot}.createdAt`,
@@ -766,7 +770,7 @@ export const snapshotDALFactory = (db: TDbClient) => {
) )
.orderBy(`${TableName.Snapshot}.createdAt`, "desc") .orderBy(`${TableName.Snapshot}.createdAt`, "desc")
.where(`${TableName.Snapshot}.folderId`, folderId); .where(`${TableName.Snapshot}.folderId`, folderId);
const data = await (tx || db) const data = await (tx || db.replicaNode())
.with("w", query) .with("w", query)
.select("*") .select("*")
.from<Awaited<typeof query>[number]>("w") .from<Awaited<typeof query>[number]>("w")
@@ -0,0 +1,91 @@
import { Knex } from "knex";
import { TDbClient } from "@app/db";
import { TableName } from "@app/db/schemas";
import { ormify, TOrmify } from "@app/lib/knex";
import { logger } from "@app/lib/logger";
import { QueueName } from "@app/queue";
export interface TKeyValueStoreDALFactory extends TOrmify<TableName.KeyValueStore> {
incrementBy: (key: string, dto: { incr?: number; tx?: Knex; expiresAt?: Date }) => Promise<number>;
findOneInt: (key: string, tx?: Knex) => Promise<number | undefined>;
pruneExpiredKeys: () => Promise<void>;
}
const QUERY_TIMEOUT_MS = 10 * 60 * 1000; // 10 minutes
const CACHE_KEY_PRUNE_BATCH_SIZE = 10000;
const MAX_RETRY_ON_FAILURE = 3;
export const keyValueStoreDALFactory = (db: TDbClient): TKeyValueStoreDALFactory => {
const keyValueStoreOrm = ormify(db, TableName.KeyValueStore);
const incrementBy: TKeyValueStoreDALFactory["incrementBy"] = async (key, { incr = 1, tx, expiresAt }) => {
return (tx || db)(TableName.KeyValueStore)
.insert({ key, integerValue: 1, expiresAt })
.onConflict("key")
.merge({
integerValue: db.raw(`"${TableName.KeyValueStore}"."integerValue" + ?`, [incr]),
expiresAt
})
.returning("integerValue")
.then((result) => Number(result[0]?.integerValue || 0));
};
const findOneInt: TKeyValueStoreDALFactory["findOneInt"] = async (key, tx) => {
const doc = await (tx || db.replicaNode())(TableName.KeyValueStore)
.where({ key })
.andWhere(
(builder) =>
void builder
.whereNull("expiresAt") // no expiry
.orWhere("expiresAt", ">", db.fn.now()) // or not expired
)
.first()
.select("integerValue");
return Number(doc?.integerValue || 0);
};
const pruneExpiredKeys: TKeyValueStoreDALFactory["pruneExpiredKeys"] = async () => {
let deletedIds: { key: string }[] = [];
let numberOfRetryOnFailure = 0;
let isRetrying = false;
logger.info(`${QueueName.DailyResourceCleanUp}: db key value store clean up started`);
do {
try {
// eslint-disable-next-line no-await-in-loop
deletedIds = await db.transaction(async (trx) => {
await trx.raw(`SET statement_timeout = ${QUERY_TIMEOUT_MS}`);
const findExpiredKeysSubQuery = trx(TableName.KeyValueStore)
.where("expiresAt", "<", db.fn.now())
.select("key")
.limit(CACHE_KEY_PRUNE_BATCH_SIZE);
// eslint-disable-next-line no-await-in-loop
const results = await trx(TableName.KeyValueStore)
.whereIn("key", findExpiredKeysSubQuery)
.del()
.returning("key");
return results;
});
numberOfRetryOnFailure = 0; // reset
} catch (error) {
numberOfRetryOnFailure += 1;
deletedIds = [];
logger.error(error, "Failed to clean up db key value");
} finally {
// eslint-disable-next-line no-await-in-loop
await new Promise((resolve) => {
setTimeout(resolve, 10); // time to breathe for db
});
}
isRetrying = numberOfRetryOnFailure > 0;
} while (deletedIds.length > 0 || (isRetrying && numberOfRetryOnFailure < MAX_RETRY_ON_FAILURE));
logger.info(`${QueueName.DailyResourceCleanUp}: db key value store clean up completed`);
};
return { ...keyValueStoreOrm, incrementBy, findOneInt, pruneExpiredKeys };
};
+49 -27
View File
@@ -1,9 +1,14 @@
import { Cluster, Redis } from "ioredis";
import { Knex } from "knex";
import { buildRedisFromConfig, TRedisConfigKeys } from "@app/lib/config/redis"; import { buildRedisFromConfig, TRedisConfigKeys } from "@app/lib/config/redis";
import { pgAdvisoryLockHashText } from "@app/lib/crypto/hashtext"; import { pgAdvisoryLockHashText } from "@app/lib/crypto/hashtext";
import { applyJitter } from "@app/lib/dates"; import { applyJitter } from "@app/lib/dates";
import { delay as delayMs } from "@app/lib/delay"; import { delay as delayMs } from "@app/lib/delay";
import { ms } from "@app/lib/ms";
import { ExecutionResult, Redlock, Settings } from "@app/lib/red-lock"; import { ExecutionResult, Redlock, Settings } from "@app/lib/red-lock";
import { Redis, Cluster } from "ioredis";
import { TKeyValueStoreDALFactory } from "./key-value-store-dal";
export const PgSqlLock = { export const PgSqlLock = {
BootUpMigration: 2023, BootUpMigration: 2023,
@@ -97,13 +102,17 @@ export type TKeyStoreFactory = {
deleteItemsByKeyIn: (keys: string[]) => Promise<number>; deleteItemsByKeyIn: (keys: string[]) => Promise<number>;
deleteItems: (arg: TDeleteItems) => Promise<number>; deleteItems: (arg: TDeleteItems) => Promise<number>;
incrementBy: (key: string, value: number) => Promise<number>; incrementBy: (key: string, value: number) => Promise<number>;
getKeysByPattern: (pattern: string, limit?: number) => Promise<string[]>;
// pg
pgIncrementBy: (key: string, dto: { incr?: number; expiry?: string; tx?: Knex }) => Promise<number>;
pgGetIntItem: (key: string, prefix?: string) => Promise<number | undefined>;
// locks
acquireLock( acquireLock(
resources: string[], resources: string[],
duration: number, duration: number,
settings?: Partial<Settings> settings?: Partial<Settings>
): Promise<{ release: () => Promise<ExecutionResult> }>; ): Promise<{ release: () => Promise<ExecutionResult> }>;
waitTillReady: ({ key, waitingCb, keyCheckCb, waitIteration, delay, jitter }: TWaitTillReady) => Promise<void>; waitTillReady: ({ key, waitingCb, keyCheckCb, waitIteration, delay, jitter }: TWaitTillReady) => Promise<void>;
getKeysByPattern: (pattern: string, limit?: number) => Promise<string[]>;
}; };
const pickPrimaryOrSecondaryRedis = (primary: Redis | Cluster, secondaries?: Array<Redis | Cluster>) => { const pickPrimaryOrSecondaryRedis = (primary: Redis | Cluster, secondaries?: Array<Redis | Cluster>) => {
@@ -116,7 +125,10 @@ interface TKeyStoreFactoryDTO extends TRedisConfigKeys {
REDIS_READ_REPLICAS?: { host: string; port: number }[]; REDIS_READ_REPLICAS?: { host: string; port: number }[];
} }
export const keyStoreFactory = (redisConfigKeys: TKeyStoreFactoryDTO): TKeyStoreFactory => { export const keyStoreFactory = (
redisConfigKeys: TKeyStoreFactoryDTO,
keyValueStoreDAL: TKeyValueStoreDALFactory
): TKeyStoreFactory => {
const primaryRedis = buildRedisFromConfig(redisConfigKeys); const primaryRedis = buildRedisFromConfig(redisConfigKeys);
const redisReadReplicas = redisConfigKeys.REDIS_READ_REPLICAS?.map((el) => { const redisReadReplicas = redisConfigKeys.REDIS_READ_REPLICAS?.map((el) => {
if (redisConfigKeys.REDIS_URL) { if (redisConfigKeys.REDIS_URL) {
@@ -191,29 +203,6 @@ export const keyStoreFactory = (redisConfigKeys: TKeyStoreFactoryDTO): TKeyStore
const setExpiry = async (key: string, expiryInSeconds: number) => primaryRedis.expire(key, expiryInSeconds); const setExpiry = async (key: string, expiryInSeconds: number) => primaryRedis.expire(key, expiryInSeconds);
const waitTillReady = async ({
key,
waitingCb,
keyCheckCb,
waitIteration = 10,
delay = 1000,
jitter = 200
}: TWaitTillReady) => {
let attempts = 0;
let isReady = keyCheckCb(await getItem(key));
while (!isReady) {
if (attempts > waitIteration) return;
// eslint-disable-next-line
await new Promise((resolve) => {
waitingCb?.();
setTimeout(resolve, Math.max(0, applyJitter(delay, jitter)));
});
attempts += 1;
// eslint-disable-next-line
isReady = keyCheckCb(await getItem(key));
}
};
const getKeysByPattern = async (pattern: string, limit?: number) => { const getKeysByPattern = async (pattern: string, limit?: number) => {
let cursor = "0"; let cursor = "0";
const allKeys: string[] = []; const allKeys: string[] = [];
@@ -238,6 +227,37 @@ export const keyStoreFactory = (redisConfigKeys: TKeyStoreFactoryDTO): TKeyStore
return allKeys; return allKeys;
}; };
const pgIncrementBy: TKeyStoreFactory["pgIncrementBy"] = async (key, { incr = 1, tx, expiry }) => {
const expiresAt = expiry ? new Date(Date.now() + ms(expiry)) : undefined;
return keyValueStoreDAL.incrementBy(key, { incr, expiresAt, tx });
};
const pgGetIntItem = async (key: string, prefix?: string) =>
keyValueStoreDAL.findOneInt(prefix ? `${prefix}:${key}` : key);
const waitTillReady = async ({
key,
waitingCb,
keyCheckCb,
waitIteration = 10,
delay = 1000,
jitter = 200
}: TWaitTillReady) => {
let attempts = 0;
let isReady = keyCheckCb(await getItem(key));
while (!isReady) {
if (attempts > waitIteration) return;
// eslint-disable-next-line
await new Promise((resolve) => {
waitingCb?.();
setTimeout(resolve, Math.max(0, applyJitter(delay, jitter)));
});
attempts += 1;
// eslint-disable-next-line
isReady = keyCheckCb(await getItem(key));
}
};
return { return {
setItem, setItem,
getItem, getItem,
@@ -252,6 +272,8 @@ export const keyStoreFactory = (redisConfigKeys: TKeyStoreFactoryDTO): TKeyStore
waitTillReady, waitTillReady,
getKeysByPattern, getKeysByPattern,
deleteItemsByKeyIn, deleteItemsByKeyIn,
getItems getItems,
pgGetIntItem,
pgIncrementBy
}; };
}; };
+9
View File
@@ -53,6 +53,15 @@ export const inMemoryKeyStore = (): TKeyStoreFactory => {
} }
return null; return null;
}, },
pgGetIntItem: async (key) => {
const value = store[key];
if (typeof value === "number") {
return Number(value);
}
},
pgIncrementBy: async () => {
return 1;
},
incrementBy: async () => { incrementBy: async () => {
return 1; return 1;
}, },
+1
View File
@@ -412,6 +412,7 @@ const envSchema = z
Boolean(data.INF_APP_CONNECTION_GITHUB_RADAR_APP_CLIENT_ID) && Boolean(data.INF_APP_CONNECTION_GITHUB_RADAR_APP_CLIENT_ID) &&
Boolean(data.INF_APP_CONNECTION_GITHUB_RADAR_APP_CLIENT_SECRET) && Boolean(data.INF_APP_CONNECTION_GITHUB_RADAR_APP_CLIENT_SECRET) &&
Boolean(data.INF_APP_CONNECTION_GITHUB_RADAR_APP_WEBHOOK_SECRET), Boolean(data.INF_APP_CONNECTION_GITHUB_RADAR_APP_WEBHOOK_SECRET),
isSecondaryInstance: Boolean(data.INFISICAL_PRIMARY_INSTANCE_URL),
isHsmConfigured: isHsmConfigured:
Boolean(data.HSM_LIB_PATH) && Boolean(data.HSM_PIN) && Boolean(data.HSM_KEY_LABEL) && data.HSM_SLOT !== undefined, Boolean(data.HSM_LIB_PATH) && Boolean(data.HSM_PIN) && Boolean(data.HSM_KEY_LABEL) && data.HSM_SLOT !== undefined,
samlDefaultOrgSlug: data.DEFAULT_SAML_ORG_SLUG, samlDefaultOrgSlug: data.DEFAULT_SAML_ORG_SLUG,
+4 -4
View File
@@ -250,12 +250,12 @@ export const ormify = <DbOps extends object, Tname extends keyof Tables>(
.returning("*"); .returning("*");
if ($incr) { if ($incr) {
Object.entries($incr).forEach(([incrementField, incrementValue]) => { Object.entries($incr).forEach(([incrementField, incrementValue]) => {
void query.increment(incrementField, incrementValue); void query.increment(incrementField, incrementValue as number);
}); });
} }
if ($decr) { if ($decr) {
Object.entries($decr).forEach(([incrementField, incrementValue]) => { Object.entries($decr).forEach(([incrementField, incrementValue]) => {
void query.decrement(incrementField, incrementValue); void query.decrement(incrementField, incrementValue as number);
}); });
} }
const [docs] = await query; const [docs] = await query;
@@ -273,12 +273,12 @@ export const ormify = <DbOps extends object, Tname extends keyof Tables>(
// increment and decrement operation in update // increment and decrement operation in update
if ($incr) { if ($incr) {
Object.entries($incr).forEach(([incrementField, incrementValue]) => { Object.entries($incr).forEach(([incrementField, incrementValue]) => {
void query.increment(incrementField, incrementValue); void query.increment(incrementField, incrementValue as number);
}); });
} }
if ($decr) { if ($decr) {
Object.entries($decr).forEach(([incrementField, incrementValue]) => { Object.entries($decr).forEach(([incrementField, incrementValue]) => {
void query.increment(incrementField, incrementValue); void query.decrement(incrementField, incrementValue as number);
}); });
} }
return (await query) as Tables[Tname]["base"][]; return (await query) as Tables[Tname]["base"][];
+3 -1
View File
@@ -5,6 +5,7 @@ import "./lib/telemetry/instrumentation";
import dotenv from "dotenv"; import dotenv from "dotenv";
import { initializeHsmModule } from "@app/ee/services/hsm/hsm-fns"; import { initializeHsmModule } from "@app/ee/services/hsm/hsm-fns";
import { keyValueStoreDALFactory } from "@app/keystore/key-value-store-dal";
import { runMigrations } from "./auto-start-migrations"; import { runMigrations } from "./auto-start-migrations";
import { initAuditLogDbConnection, initDbConnection } from "./db"; import { initAuditLogDbConnection, initDbConnection } from "./db";
@@ -54,7 +55,8 @@ const run = async () => {
await queue.initialize(); await queue.initialize();
const keyStore = keyStoreFactory(envConfig); const keyValueStoreDAL = keyValueStoreDALFactory(db);
const keyStore = keyStoreFactory(envConfig, keyValueStoreDAL);
const redis = buildRedisFromConfig(envConfig); const redis = buildRedisFromConfig(envConfig);
const hsmModule = initializeHsmModule(envConfig); const hsmModule = initializeHsmModule(envConfig);
+7 -1
View File
@@ -130,6 +130,7 @@ import { sshHostGroupMembershipDALFactory } from "@app/ee/services/ssh-host-grou
import { sshHostGroupServiceFactory } from "@app/ee/services/ssh-host-group/ssh-host-group-service"; import { sshHostGroupServiceFactory } from "@app/ee/services/ssh-host-group/ssh-host-group-service";
import { trustedIpDALFactory } from "@app/ee/services/trusted-ip/trusted-ip-dal"; import { trustedIpDALFactory } from "@app/ee/services/trusted-ip/trusted-ip-dal";
import { trustedIpServiceFactory } from "@app/ee/services/trusted-ip/trusted-ip-service"; import { trustedIpServiceFactory } from "@app/ee/services/trusted-ip/trusted-ip-service";
import { keyValueStoreDALFactory } from "@app/keystore/key-value-store-dal";
import { TKeyStoreFactory } from "@app/keystore/keystore"; import { TKeyStoreFactory } from "@app/keystore/keystore";
import { getConfig, TEnvConfig } from "@app/lib/config/env"; import { getConfig, TEnvConfig } from "@app/lib/config/env";
import { crypto } from "@app/lib/crypto/cryptography"; import { crypto } from "@app/lib/crypto/cryptography";
@@ -514,6 +515,7 @@ export const registerRoutes = async (
const microsoftTeamsIntegrationDAL = microsoftTeamsIntegrationDALFactory(db); const microsoftTeamsIntegrationDAL = microsoftTeamsIntegrationDALFactory(db);
const projectMicrosoftTeamsConfigDAL = projectMicrosoftTeamsConfigDALFactory(db); const projectMicrosoftTeamsConfigDAL = projectMicrosoftTeamsConfigDALFactory(db);
const secretScanningV2DAL = secretScanningV2DALFactory(db); const secretScanningV2DAL = secretScanningV2DALFactory(db);
const keyValueStoreDAL = keyValueStoreDALFactory(db);
const eventBusService = eventBusFactory(server.redis); const eventBusService = eventBusFactory(server.redis);
const sseService = sseServiceFactory(eventBusService, server.redis); const sseService = sseServiceFactory(eventBusService, server.redis);
@@ -650,6 +652,7 @@ export const registerRoutes = async (
const folderTreeCheckpointDAL = folderTreeCheckpointDALFactory(db); const folderTreeCheckpointDAL = folderTreeCheckpointDALFactory(db);
const folderCommitDAL = folderCommitDALFactory(db); const folderCommitDAL = folderCommitDALFactory(db);
const folderTreeCheckpointResourcesDAL = folderTreeCheckpointResourcesDALFactory(db); const folderTreeCheckpointResourcesDAL = folderTreeCheckpointResourcesDALFactory(db);
const folderCommitQueueService = folderCommitQueueServiceFactory({ const folderCommitQueueService = folderCommitQueueServiceFactory({
queueService, queueService,
folderTreeCheckpointDAL, folderTreeCheckpointDAL,
@@ -814,6 +817,7 @@ export const registerRoutes = async (
groupDAL, groupDAL,
orgBotDAL, orgBotDAL,
oidcConfigDAL, oidcConfigDAL,
ldapConfigDAL,
loginService, loginService,
projectBotService, projectBotService,
reminderService reminderService
@@ -1719,6 +1723,7 @@ export const registerRoutes = async (
userDAL, userDAL,
identityDAL identityDAL
}); });
const dailyResourceCleanUp = dailyResourceCleanUpQueueServiceFactory({ const dailyResourceCleanUp = dailyResourceCleanUpQueueServiceFactory({
auditLogDAL, auditLogDAL,
queueService, queueService,
@@ -1731,7 +1736,8 @@ export const registerRoutes = async (
identityUniversalAuthClientSecretDAL: identityUaClientSecretDAL, identityUniversalAuthClientSecretDAL: identityUaClientSecretDAL,
serviceTokenService, serviceTokenService,
orgService, orgService,
userNotificationDAL userNotificationDAL,
keyValueStoreDAL
}); });
const dailyReminderQueueService = dailyReminderQueueServiceFactory({ const dailyReminderQueueService = dailyReminderQueueServiceFactory({
@@ -661,7 +661,7 @@ describe("folderCommitServiceFactory", () => {
// Assert // Assert
expect(mockFolderCommitDAL.create).toHaveBeenCalled(); expect(mockFolderCommitDAL.create).toHaveBeenCalled();
expect(mockSecretV2BridgeDAL.invalidateSecretCacheByProjectId).toHaveBeenCalledWith(projectId); expect(mockSecretV2BridgeDAL.invalidateSecretCacheByProjectId).toHaveBeenCalledWith(projectId, {});
// Check that we got the right counts // Check that we got the right counts
expect(result.totalChanges).toEqual(2); expect(result.totalChanges).toEqual(2);
@@ -1386,7 +1386,7 @@ export const folderCommitServiceFactory = ({
); );
// Invalidate cache to reflect the changes // Invalidate cache to reflect the changes
await secretV2BridgeDAL.invalidateSecretCacheByProjectId(projectId); await secretV2BridgeDAL.invalidateSecretCacheByProjectId(projectId, tx);
return { return {
secretChangesCount: secretChanges.length, secretChangesCount: secretChanges.length,
@@ -84,18 +84,20 @@ export const identityUaServiceFactory = ({
const LOCKOUT_KEY = `lockout:identity:${identityUa.identityId}:${IdentityAuthMethod.UNIVERSAL_AUTH}:${clientId}`; const LOCKOUT_KEY = `lockout:identity:${identityUa.identityId}:${IdentityAuthMethod.UNIVERSAL_AUTH}:${clientId}`;
let lock: Awaited<ReturnType<typeof keyStore.acquireLock>>; let lock: Awaited<ReturnType<typeof keyStore.acquireLock>> | undefined;
try { if (identityUa.lockoutEnabled) {
lock = await keyStore.acquireLock([KeyStorePrefixes.IdentityLockoutLock(LOCKOUT_KEY)], 500, { try {
retryCount: 3, lock = await keyStore.acquireLock([KeyStorePrefixes.IdentityLockoutLock(LOCKOUT_KEY)], 500, {
retryDelay: 300, retryCount: 3,
retryJitter: 100 retryDelay: 300,
}); retryJitter: 100
} catch (e) { });
logger.info( } catch (e) {
`identity login failed to acquire lock [identityId=${identityUa.identityId}] [authMethod=${IdentityAuthMethod.UNIVERSAL_AUTH}]` logger.info(
); `identity login failed to acquire lock [identityId=${identityUa.identityId}] [authMethod=${IdentityAuthMethod.UNIVERSAL_AUTH}]`
throw new RateLimitError({ message: "Rate limit exceeded" }); );
throw new RateLimitError({ message: "Failed to acquire lock: rate limit exceeded" });
}
} }
try { try {
@@ -257,7 +259,7 @@ export const identityUaServiceFactory = ({
...accessTokenTTLParams ...accessTokenTTLParams
}; };
} finally { } finally {
await lock.release(); if (lock) await lock.release();
} }
}; };
@@ -25,7 +25,7 @@ export const identityDALFactory = (db: TDbClient) => {
} as const; } as const;
const tableName = authMethodToTableName[authMethod]; const tableName = authMethodToTableName[authMethod];
if (!tableName) return; if (!tableName) return;
const data = await db(tableName).where({ identityId }).first(); const data = await db.replicaNode()(tableName).where({ identityId }).first();
if (!data) return; if (!data) return;
return data.accessTokenTrustedIps; return data.accessTokenTrustedIps;
}; };
@@ -30,7 +30,7 @@ export const integrationAuthDALFactory = (db: TDbClient) => {
const getByOrg = async (orgId: string, tx?: Knex) => { const getByOrg = async (orgId: string, tx?: Knex) => {
try { try {
const integrationAuths = await (tx || db)(TableName.IntegrationAuth) const integrationAuths = await (tx || db.replicaNode())(TableName.IntegrationAuth)
.join(TableName.Project, `${TableName.Project}.id`, `${TableName.IntegrationAuth}.projectId`) .join(TableName.Project, `${TableName.Project}.id`, `${TableName.IntegrationAuth}.projectId`)
.join(TableName.Organization, `${TableName.Organization}.id`, `${TableName.Project}.orgId`) .join(TableName.Organization, `${TableName.Organization}.id`, `${TableName.Project}.orgId`)
.where(`${TableName.Organization}.id`, "=", orgId) .where(`${TableName.Organization}.id`, "=", orgId)
@@ -12,7 +12,7 @@ export const kmsRootConfigDALFactory = (db: TDbClient) => {
const findById = async (id: string, tx?: Knex) => { const findById = async (id: string, tx?: Knex) => {
try { try {
const result = await (tx || db)(TableName.KmsServerRootConfig) const result = await (tx || db?.replicaNode?.() || db)(TableName.KmsServerRootConfig)
.where({ id } as never) .where({ id } as never)
.first("*"); .first("*");
return result; return result;
+38 -3
View File
@@ -9,11 +9,14 @@ import {
ProjectMembershipRole, ProjectMembershipRole,
ProjectVersion, ProjectVersion,
TableName, TableName,
TOidcConfigs,
TProjectMemberships, TProjectMemberships,
TProjectUserMembershipRolesInsert, TProjectUserMembershipRolesInsert,
TSamlConfigs,
TUsers TUsers
} from "@app/db/schemas"; } from "@app/db/schemas";
import { TGroupDALFactory } from "@app/ee/services/group/group-dal"; import { TGroupDALFactory } from "@app/ee/services/group/group-dal";
import { TLdapConfigDALFactory } from "@app/ee/services/ldap-config/ldap-config-dal";
import { TLicenseServiceFactory } from "@app/ee/services/license/license-service"; import { TLicenseServiceFactory } from "@app/ee/services/license/license-service";
import { TOidcConfigDALFactory } from "@app/ee/services/oidc/oidc-config-dal"; import { TOidcConfigDALFactory } from "@app/ee/services/oidc/oidc-config-dal";
import { import {
@@ -125,6 +128,7 @@ type TOrgServiceFactoryDep = {
incidentContactDAL: TIncidentContactsDALFactory; incidentContactDAL: TIncidentContactsDALFactory;
samlConfigDAL: Pick<TSamlConfigDALFactory, "findOne">; samlConfigDAL: Pick<TSamlConfigDALFactory, "findOne">;
oidcConfigDAL: Pick<TOidcConfigDALFactory, "findOne">; oidcConfigDAL: Pick<TOidcConfigDALFactory, "findOne">;
ldapConfigDAL: Pick<TLdapConfigDALFactory, "findOne">;
smtpService: TSmtpService; smtpService: TSmtpService;
tokenService: TAuthTokenServiceFactory; tokenService: TAuthTokenServiceFactory;
permissionService: TPermissionServiceFactory; permissionService: TPermissionServiceFactory;
@@ -165,6 +169,7 @@ export const orgServiceFactory = ({
projectRoleDAL, projectRoleDAL,
samlConfigDAL, samlConfigDAL,
oidcConfigDAL, oidcConfigDAL,
ldapConfigDAL,
projectUserMembershipRoleDAL, projectUserMembershipRoleDAL,
identityMetadataDAL, identityMetadataDAL,
projectBotService, projectBotService,
@@ -446,16 +451,20 @@ export const orgServiceFactory = ({
}); });
} }
if (authEnforced) { let samlCfg: TSamlConfigs | undefined;
const samlCfg = await samlConfigDAL.findOne({ let oidcCfg: TOidcConfigs | undefined;
if (authEnforced || googleSsoAuthEnforced) {
samlCfg = await samlConfigDAL.findOne({
orgId, orgId,
isActive: true isActive: true
}); });
const oidcCfg = await oidcConfigDAL.findOne({ oidcCfg = await oidcConfigDAL.findOne({
orgId, orgId,
isActive: true isActive: true
}); });
}
if (authEnforced) {
if (!samlCfg && !oidcCfg) if (!samlCfg && !oidcCfg)
throw new NotFoundError({ throw new NotFoundError({
message: `SAML or OIDC configuration for organization with ID '${orgId}' not found` message: `SAML or OIDC configuration for organization with ID '${orgId}' not found`
@@ -483,6 +492,32 @@ export const orgServiceFactory = ({
}); });
} }
if (samlCfg) {
throw new BadRequestError({
message:
"Cannot enable Google OAuth enforcement while SAML SSO is configured. Disable SAML SSO to enforce Google OAuth."
});
}
if (oidcCfg) {
throw new BadRequestError({
message:
"Cannot enable Google OAuth enforcement while OIDC SSO is configured. Disable OIDC SSO to enforce Google OAuth."
});
}
const ldapCfg = await ldapConfigDAL.findOne({
orgId,
isActive: true
});
if (ldapCfg) {
throw new BadRequestError({
message:
"Cannot enable Google OAuth enforcement while LDAP SSO is configured. Disable LDAP SSO to enforce Google OAuth."
});
}
if (!currentOrg.googleSsoAuthLastUsed) { if (!currentOrg.googleSsoAuthLastUsed) {
throw new BadRequestError({ throw new BadRequestError({
message: message:
@@ -39,7 +39,7 @@ export const reminderDALFactory = (db: TDbClient) => {
const findSecretDailyReminders = async (tx?: Knex) => { const findSecretDailyReminders = async (tx?: Knex) => {
const { startOfDay, endOfDay } = getTodayDateRange(); const { startOfDay, endOfDay } = getTodayDateRange();
const rawReminders = await (tx || db)(TableName.Reminder) const rawReminders = await (tx || db.replicaNode())(TableName.Reminder)
.whereBetween("nextReminderDate", [startOfDay, endOfDay]) .whereBetween("nextReminderDate", [startOfDay, endOfDay])
.leftJoin(TableName.ReminderRecipient, `${TableName.Reminder}.id`, `${TableName.ReminderRecipient}.reminderId`) .leftJoin(TableName.ReminderRecipient, `${TableName.Reminder}.id`, `${TableName.ReminderRecipient}.reminderId`)
.leftJoin<TUsers>(TableName.Users, `${TableName.ReminderRecipient}.userId`, `${TableName.Users}.id`) .leftJoin<TUsers>(TableName.Users, `${TableName.ReminderRecipient}.userId`, `${TableName.Users}.id`)
@@ -90,7 +90,7 @@ export const reminderDALFactory = (db: TDbClient) => {
const futureDate = new Date(startOfDay); const futureDate = new Date(startOfDay);
futureDate.setDate(futureDate.getDate() + daysAhead); futureDate.setDate(futureDate.getDate() + daysAhead);
const reminders = await (tx || db)(TableName.Reminder) const reminders = await (tx || db.replicaNode())(TableName.Reminder)
.where("nextReminderDate", ">=", startOfDay) .where("nextReminderDate", ">=", startOfDay)
.where("nextReminderDate", "<=", futureDate) .where("nextReminderDate", "<=", futureDate)
.orderBy("nextReminderDate", "asc") .orderBy("nextReminderDate", "asc")
@@ -101,7 +101,7 @@ export const reminderDALFactory = (db: TDbClient) => {
}; };
const findSecretReminder = async (secretId: string, tx?: Knex) => { const findSecretReminder = async (secretId: string, tx?: Knex) => {
const rawReminders = await (tx || db)(TableName.Reminder) const rawReminders = await (tx || db.replicaNode())(TableName.Reminder)
.where(`${TableName.Reminder}.secretId`, secretId) .where(`${TableName.Reminder}.secretId`, secretId)
.leftJoin(TableName.ReminderRecipient, `${TableName.Reminder}.id`, `${TableName.ReminderRecipient}.reminderId`) .leftJoin(TableName.ReminderRecipient, `${TableName.Reminder}.id`, `${TableName.ReminderRecipient}.reminderId`)
.select(selectAllTableCols(TableName.Reminder)) .select(selectAllTableCols(TableName.Reminder))
@@ -125,7 +125,7 @@ export const reminderDALFactory = (db: TDbClient) => {
}; };
const findSecretReminders = async (secretIds: string[], tx?: Knex) => { const findSecretReminders = async (secretIds: string[], tx?: Knex) => {
const rawReminders = await (tx || db)(TableName.Reminder) const rawReminders = await (tx || db.replicaNode())(TableName.Reminder)
.whereIn(`${TableName.Reminder}.secretId`, secretIds) .whereIn(`${TableName.Reminder}.secretId`, secretIds)
.leftJoin(TableName.ReminderRecipient, `${TableName.Reminder}.id`, `${TableName.ReminderRecipient}.reminderId`) .leftJoin(TableName.ReminderRecipient, `${TableName.Reminder}.id`, `${TableName.ReminderRecipient}.reminderId`)
.select(selectAllTableCols(TableName.Reminder)) .select(selectAllTableCols(TableName.Reminder))
@@ -1,5 +1,6 @@
import { TAuditLogDALFactory } from "@app/ee/services/audit-log/audit-log-dal"; import { TAuditLogDALFactory } from "@app/ee/services/audit-log/audit-log-dal";
import { TSnapshotDALFactory } from "@app/ee/services/secret-snapshot/snapshot-dal"; import { TSnapshotDALFactory } from "@app/ee/services/secret-snapshot/snapshot-dal";
import { TKeyValueStoreDALFactory } from "@app/keystore/key-value-store-dal";
import { getConfig } from "@app/lib/config/env"; import { getConfig } from "@app/lib/config/env";
import { logger } from "@app/lib/logger"; import { logger } from "@app/lib/logger";
import { QueueJobs, QueueName, TQueueServiceFactory } from "@app/queue"; import { QueueJobs, QueueName, TQueueServiceFactory } from "@app/queue";
@@ -27,6 +28,7 @@ type TDailyResourceCleanUpQueueServiceFactoryDep = {
queueService: TQueueServiceFactory; queueService: TQueueServiceFactory;
orgService: TOrgServiceFactory; orgService: TOrgServiceFactory;
userNotificationDAL: Pick<TUserNotificationDALFactory, "pruneNotifications">; userNotificationDAL: Pick<TUserNotificationDALFactory, "pruneNotifications">;
keyValueStoreDAL: Pick<TKeyValueStoreDALFactory, "pruneExpiredKeys">;
}; };
export type TDailyResourceCleanUpQueueServiceFactory = ReturnType<typeof dailyResourceCleanUpQueueServiceFactory>; export type TDailyResourceCleanUpQueueServiceFactory = ReturnType<typeof dailyResourceCleanUpQueueServiceFactory>;
@@ -43,7 +45,8 @@ export const dailyResourceCleanUpQueueServiceFactory = ({
identityUniversalAuthClientSecretDAL, identityUniversalAuthClientSecretDAL,
serviceTokenService, serviceTokenService,
orgService, orgService,
userNotificationDAL userNotificationDAL,
keyValueStoreDAL
}: TDailyResourceCleanUpQueueServiceFactoryDep) => { }: TDailyResourceCleanUpQueueServiceFactoryDep) => {
const appCfg = getConfig(); const appCfg = getConfig();
@@ -52,6 +55,10 @@ export const dailyResourceCleanUpQueueServiceFactory = ({
} }
const init = async () => { const init = async () => {
if (appCfg.isSecondaryInstance) {
return;
}
await queueService.stopRepeatableJob( await queueService.stopRepeatableJob(
QueueName.AuditLogPrune, QueueName.AuditLogPrune,
QueueJobs.AuditLogPrune, QueueJobs.AuditLogPrune,
@@ -82,6 +89,7 @@ export const dailyResourceCleanUpQueueServiceFactory = ({
await orgService.notifyInvitedUsers(); await orgService.notifyInvitedUsers();
await auditLogDAL.pruneAuditLog(); await auditLogDAL.pruneAuditLog();
await userNotificationDAL.pruneNotifications(); await userNotificationDAL.pruneNotifications();
await keyValueStoreDAL.pruneExpiredKeys();
logger.info(`${QueueName.DailyResourceCleanUp}: queue task completed`); logger.info(`${QueueName.DailyResourceCleanUp}: queue task completed`);
} catch (error) { } catch (error) {
logger.error(error, `${QueueName.DailyResourceCleanUp}: resource cleanup failed`); logger.error(error, `${QueueName.DailyResourceCleanUp}: resource cleanup failed`);
@@ -45,7 +45,7 @@ export const secretFolderVersionDALFactory = (db: TDbClient) => {
) )
.whereIn(`${TableName.SecretFolderVersion}.folderId`, folderIds) .whereIn(`${TableName.SecretFolderVersion}.folderId`, folderIds)
.join( .join(
(tx || db)(TableName.SecretFolderVersion) (tx || db.replicaNode())(TableName.SecretFolderVersion)
.groupBy("folderId") .groupBy("folderId")
.max("version") .max("version")
.select("folderId") .select("folderId")
@@ -15,7 +15,7 @@ export const secretImportDALFactory = (db: TDbClient) => {
// we are using postion based sorting as its a small list // we are using postion based sorting as its a small list
// this will return the last value of the position in a folder with secret imports // this will return the last value of the position in a folder with secret imports
const findLastImportPosition = async (folderId: string, tx?: Knex) => { const findLastImportPosition = async (folderId: string, tx?: Knex) => {
const lastPos = await (tx || db)(TableName.SecretImport) const lastPos = await (tx || db.replicaNode())(TableName.SecretImport)
.where({ folderId }) .where({ folderId })
.max("position", { as: "position" }) .max("position", { as: "position" })
.first(); .first();
@@ -119,7 +119,7 @@ export const secretSharingDALFactory = (db: TDbClient) => {
const findActiveSharedSecrets = async (filters: Partial<TSecretSharing>, tx?: Knex) => { const findActiveSharedSecrets = async (filters: Partial<TSecretSharing>, tx?: Knex) => {
try { try {
const now = new Date(); const now = new Date();
return await (tx || db)(TableName.SecretSharing) return await (tx || db.replicaNode())(TableName.SecretSharing)
.where(filters) .where(filters)
.andWhere("expiresAt", ">", now) .andWhere("expiresAt", ">", now)
.andWhere("encryptedValue", "<>", "") .andWhere("encryptedValue", "<>", "")
@@ -50,20 +50,19 @@ interface TSecretV2DalArg {
} }
export const SECRET_DAL_TTL = () => applyJitter(10 * 60, 2 * 60); export const SECRET_DAL_TTL = () => applyJitter(10 * 60, 2 * 60);
export const SECRET_DAL_VERSION_TTL = 15 * 60; export const SECRET_DAL_VERSION_TTL = "15m";
export const MAX_SECRET_CACHE_BYTES = 25 * 1024 * 1024; export const MAX_SECRET_CACHE_BYTES = 25 * 1024 * 1024;
export const secretV2BridgeDALFactory = ({ db, keyStore }: TSecretV2DalArg) => { export const secretV2BridgeDALFactory = ({ db, keyStore }: TSecretV2DalArg) => {
const secretOrm = ormify(db, TableName.SecretV2); const secretOrm = ormify(db, TableName.SecretV2);
const invalidateSecretCacheByProjectId = async (projectId: string) => { const invalidateSecretCacheByProjectId = async (projectId: string, tx?: Knex) => {
const secretDalVersionKey = SecretServiceCacheKeys.getSecretDalVersion(projectId); const secretDalVersionKey = SecretServiceCacheKeys.getSecretDalVersion(projectId);
await keyStore.incrementBy(secretDalVersionKey, 1); await keyStore.pgIncrementBy(secretDalVersionKey, { incr: 1, tx, expiry: SECRET_DAL_VERSION_TTL });
await keyStore.setExpiry(secretDalVersionKey, SECRET_DAL_VERSION_TTL);
}; };
const findOne = async (filter: Partial<TSecretsV2>, tx?: Knex) => { const findOne = async (filter: Partial<TSecretsV2>, tx?: Knex) => {
try { try {
const docs = await (tx || db)(TableName.SecretV2) const docs = await (tx || db.replicaNode())(TableName.SecretV2)
// eslint-disable-next-line @typescript-eslint/no-misused-promises // eslint-disable-next-line @typescript-eslint/no-misused-promises
.where(buildFindFilter(filter, TableName.SecretV2)) .where(buildFindFilter(filter, TableName.SecretV2))
.leftJoin( .leftJoin(
@@ -144,7 +143,7 @@ export const secretV2BridgeDALFactory = ({ db, keyStore }: TSecretV2DalArg) => {
const find = async (filter: TFindFilter<TSecretsV2>, opts: TFindOpt<TSecretsV2> = {}) => { const find = async (filter: TFindFilter<TSecretsV2>, opts: TFindOpt<TSecretsV2> = {}) => {
const { offset, limit, sort, tx } = opts; const { offset, limit, sort, tx } = opts;
try { try {
const query = (tx || db)(TableName.SecretV2) const query = (tx || db.replicaNode())(TableName.SecretV2)
// eslint-disable-next-line @typescript-eslint/no-misused-promises // eslint-disable-next-line @typescript-eslint/no-misused-promises
.where(buildFindFilter(filter)) .where(buildFindFilter(filter))
.leftJoin( .leftJoin(
@@ -888,13 +887,13 @@ export const secretV2BridgeDALFactory = ({ db, keyStore }: TSecretV2DalArg) => {
const findSecretsWithReminderRecipients = async (ids: string[], limit: number, tx?: Knex) => { const findSecretsWithReminderRecipients = async (ids: string[], limit: number, tx?: Knex) => {
try { try {
// Create a subquery to get limited secret IDs // Create a subquery to get limited secret IDs
const limitedSecretIds = (tx || db)(TableName.SecretV2) const limitedSecretIds = (tx || db.replicaNode())(TableName.SecretV2)
.whereIn(`${TableName.SecretV2}.id`, ids) .whereIn(`${TableName.SecretV2}.id`, ids)
.limit(limit) .limit(limit)
.select("id"); .select("id");
// Join with all recipients for the limited secrets // Join with all recipients for the limited secrets
const docs = await (tx || db)(TableName.SecretV2) const docs = await (tx || db.replicaNode())(TableName.SecretV2)
.whereIn(`${TableName.SecretV2}.id`, limitedSecretIds) .whereIn(`${TableName.SecretV2}.id`, limitedSecretIds)
.leftJoin(TableName.Reminder, `${TableName.SecretV2}.id`, `${TableName.Reminder}.secretId`) .leftJoin(TableName.Reminder, `${TableName.SecretV2}.id`, `${TableName.Reminder}.secretId`)
.leftJoin(TableName.ReminderRecipient, `${TableName.Reminder}.id`, `${TableName.ReminderRecipient}.reminderId`) .leftJoin(TableName.ReminderRecipient, `${TableName.Reminder}.id`, `${TableName.ReminderRecipient}.reminderId`)
@@ -926,13 +925,13 @@ export const secretV2BridgeDALFactory = ({ db, keyStore }: TSecretV2DalArg) => {
const findSecretsWithReminderRecipientsOld = async (ids: string[], limit: number, tx?: Knex) => { const findSecretsWithReminderRecipientsOld = async (ids: string[], limit: number, tx?: Knex) => {
try { try {
// Create a subquery to get limited secret IDs // Create a subquery to get limited secret IDs
const limitedSecretIds = (tx || db)(TableName.SecretV2) const limitedSecretIds = (tx || db.replicaNode())(TableName.SecretV2)
.whereIn(`${TableName.SecretV2}.id`, ids) .whereIn(`${TableName.SecretV2}.id`, ids)
.limit(limit) .limit(limit)
.select("id"); .select("id");
// Join with all recipients for the limited secrets // Join with all recipients for the limited secrets
const docs = await (tx || db)(TableName.SecretV2) const docs = await (tx || db.replicaNode())(TableName.SecretV2)
.whereIn(`${TableName.SecretV2}.id`, limitedSecretIds) .whereIn(`${TableName.SecretV2}.id`, limitedSecretIds)
.leftJoin(TableName.Reminder, `${TableName.SecretV2}.id`, `${TableName.Reminder}.secretId`) .leftJoin(TableName.Reminder, `${TableName.SecretV2}.id`, `${TableName.Reminder}.secretId`)
.leftJoin( .leftJoin(
@@ -118,7 +118,7 @@ type TSecretV2BridgeServiceFactoryDep = {
>; >;
snapshotService: Pick<TSecretSnapshotServiceFactory, "performSnapshot">; snapshotService: Pick<TSecretSnapshotServiceFactory, "performSnapshot">;
resourceMetadataDAL: Pick<TResourceMetadataDALFactory, "insertMany" | "delete">; resourceMetadataDAL: Pick<TResourceMetadataDALFactory, "insertMany" | "delete">;
keyStore: Pick<TKeyStoreFactory, "getItem" | "setExpiry" | "setItemWithExpiry" | "deleteItem">; keyStore: Pick<TKeyStoreFactory, "getItem" | "setExpiry" | "setItemWithExpiry" | "deleteItem" | "pgGetIntItem">;
reminderService: Pick<TReminderServiceFactory, "createReminder" | "getReminder">; reminderService: Pick<TReminderServiceFactory, "createReminder" | "getReminder">;
}; };
@@ -360,6 +360,7 @@ export const secretV2BridgeServiceFactory = ({
tx tx
}); });
await secretDAL.invalidateSecretCacheByProjectId(projectId, tx);
return createdSecret; return createdSecret;
}); });
@@ -377,7 +378,6 @@ export const secretV2BridgeServiceFactory = ({
}); });
} }
await secretDAL.invalidateSecretCacheByProjectId(projectId);
if (inputSecret.type === SecretType.Shared) { if (inputSecret.type === SecretType.Shared) {
await snapshotService.performSnapshot(folderId); await snapshotService.performSnapshot(folderId);
await secretQueueService.syncSecrets({ await secretQueueService.syncSecrets({
@@ -566,8 +566,8 @@ export const secretV2BridgeServiceFactory = ({
await $validateSecretReferences(projectId, permission, allSecretReferences); await $validateSecretReferences(projectId, permission, allSecretReferences);
} }
const updatedSecret = await secretDAL.transaction(async (tx) => const updatedSecret = await secretDAL.transaction(async (tx) => {
fnSecretBulkUpdate({ const modifiedSecretsInDB = await fnSecretBulkUpdate({
folderId, folderId,
orgId: actorOrgId, orgId: actorOrgId,
resourceMetadataDAL, resourceMetadataDAL,
@@ -598,8 +598,11 @@ export const secretV2BridgeServiceFactory = ({
actorId actorId
}, },
tx tx
}) });
);
await secretDAL.invalidateSecretCacheByProjectId(projectId, tx);
return modifiedSecretsInDB;
});
if (inputSecret.secretReminderRepeatDays) { if (inputSecret.secretReminderRepeatDays) {
await reminderService.createReminder({ await reminderService.createReminder({
actor, actor,
@@ -615,7 +618,6 @@ export const secretV2BridgeServiceFactory = ({
}); });
} }
await secretDAL.invalidateSecretCacheByProjectId(projectId);
if (inputSecret.type === SecretType.Shared) { if (inputSecret.type === SecretType.Shared) {
await snapshotService.performSnapshot(folderId); await snapshotService.performSnapshot(folderId);
await secretQueueService.syncSecrets({ await secretQueueService.syncSecrets({
@@ -715,8 +717,8 @@ export const secretV2BridgeServiceFactory = ({
); );
try { try {
const deletedSecret = await secretDAL.transaction(async (tx) => const deletedSecret = await secretDAL.transaction(async (tx) => {
fnSecretBulkDelete({ const modifiedSecretsInDB = await fnSecretBulkDelete({
projectId, projectId,
folderId, folderId,
actorId, actorId,
@@ -732,10 +734,11 @@ export const secretV2BridgeServiceFactory = ({
} }
], ],
tx tx
}) });
); await secretDAL.invalidateSecretCacheByProjectId(projectId, tx);
return modifiedSecretsInDB;
});
await secretDAL.invalidateSecretCacheByProjectId(projectId);
if (inputSecret.type === SecretType.Shared) { if (inputSecret.type === SecretType.Shared) {
await snapshotService.performSnapshot(folderId); await snapshotService.performSnapshot(folderId);
await secretQueueService.syncSecrets({ await secretQueueService.syncSecrets({
@@ -1027,7 +1030,7 @@ export const secretV2BridgeServiceFactory = ({
}); });
throwIfMissingSecretReadValueOrDescribePermission(permission, ProjectPermissionSecretActions.DescribeSecret); throwIfMissingSecretReadValueOrDescribePermission(permission, ProjectPermissionSecretActions.DescribeSecret);
const cachedSecretDalVersion = await keyStore.getItem(SecretServiceCacheKeys.getSecretDalVersion(projectId)); const cachedSecretDalVersion = await keyStore.pgGetIntItem(SecretServiceCacheKeys.getSecretDalVersion(projectId));
const secretDalVersion = Number(cachedSecretDalVersion || 0); const secretDalVersion = Number(cachedSecretDalVersion || 0);
const cacheKey = SecretServiceCacheKeys.getSecretsOfServiceLayer(projectId, secretDalVersion, { const cacheKey = SecretServiceCacheKeys.getSecretsOfServiceLayer(projectId, secretDalVersion, {
...dto, ...dto,
@@ -1692,7 +1695,7 @@ export const secretV2BridgeServiceFactory = ({
await kmsService.createCipherPairWithDataKey({ type: KmsDataKey.SecretManager, projectId }); await kmsService.createCipherPairWithDataKey({ type: KmsDataKey.SecretManager, projectId });
const executeBulkInsert = async (tx: Knex) => { const executeBulkInsert = async (tx: Knex) => {
return fnSecretBulkInsert({ const modifiedSecretsInDB = await fnSecretBulkInsert({
inputSecrets: inputSecrets.map((el) => { inputSecrets: inputSecrets.map((el) => {
const references = secretReferencesGroupByInputSecretKey[el.secretKey]?.nestedReferences; const references = secretReferencesGroupByInputSecretKey[el.secretKey]?.nestedReferences;
@@ -1728,13 +1731,14 @@ export const secretV2BridgeServiceFactory = ({
}, },
tx tx
}); });
await secretDAL.invalidateSecretCacheByProjectId(projectId, tx);
return modifiedSecretsInDB;
}; };
const newSecrets = providedTx const newSecrets = providedTx
? await executeBulkInsert(providedTx) ? await executeBulkInsert(providedTx)
: await secretDAL.transaction(executeBulkInsert); : await secretDAL.transaction(executeBulkInsert);
await secretDAL.invalidateSecretCacheByProjectId(projectId);
await snapshotService.performSnapshot(folderId); await snapshotService.performSnapshot(folderId);
await secretQueueService.syncSecrets({ await secretQueueService.syncSecrets({
actor, actor,
@@ -2099,6 +2103,7 @@ export const secretV2BridgeServiceFactory = ({
} }
} }
await secretDAL.invalidateSecretCacheByProjectId(projectId, tx);
return updatedSecrets; return updatedSecrets;
}; };
@@ -2106,7 +2111,6 @@ export const secretV2BridgeServiceFactory = ({
? await executeBulkUpdate(providedTx) ? await executeBulkUpdate(providedTx)
: await secretDAL.transaction(executeBulkUpdate); : await secretDAL.transaction(executeBulkUpdate);
await secretDAL.invalidateSecretCacheByProjectId(projectId);
await Promise.allSettled(folders.map((el) => (el?.id ? snapshotService.performSnapshot(el.id) : undefined))); await Promise.allSettled(folders.map((el) => (el?.id ? snapshotService.performSnapshot(el.id) : undefined)));
await Promise.allSettled( await Promise.allSettled(
folders.map((el) => folders.map((el) =>
@@ -2233,7 +2237,7 @@ export const secretV2BridgeServiceFactory = ({
}); });
const executeBulkDelete = async (tx: Knex) => { const executeBulkDelete = async (tx: Knex) => {
return fnSecretBulkDelete({ const modifiedSecretsInDB = await fnSecretBulkDelete({
secretDAL, secretDAL,
secretQueueService, secretQueueService,
folderCommitService, folderCommitService,
@@ -2249,6 +2253,8 @@ export const secretV2BridgeServiceFactory = ({
commitChanges, commitChanges,
tx tx
}); });
await secretDAL.invalidateSecretCacheByProjectId(projectId, tx);
return modifiedSecretsInDB;
}; };
try { try {
@@ -2256,7 +2262,6 @@ export const secretV2BridgeServiceFactory = ({
? await executeBulkDelete(providedTx) ? await executeBulkDelete(providedTx)
: await secretDAL.transaction(executeBulkDelete); : await secretDAL.transaction(executeBulkDelete);
await secretDAL.invalidateSecretCacheByProjectId(projectId);
await snapshotService.performSnapshot(folderId); await snapshotService.performSnapshot(folderId);
await secretQueueService.syncSecrets({ await secretQueueService.syncSecrets({
actor, actor,
@@ -72,7 +72,7 @@ export const secretVersionV2BridgeDALFactory = (db: TDbClient) => {
.where(`${TableName.SecretVersionV2}.folderId`, folderId) .where(`${TableName.SecretVersionV2}.folderId`, folderId)
.join(TableName.SecretV2, `${TableName.SecretV2}.id`, `${TableName.SecretVersionV2}.secretId`) .join(TableName.SecretV2, `${TableName.SecretV2}.id`, `${TableName.SecretVersionV2}.secretId`)
.join<TSecretVersionsV2, TSecretVersionsV2 & { secretId: string; max: number }>( .join<TSecretVersionsV2, TSecretVersionsV2 & { secretId: string; max: number }>(
(tx || db)(TableName.SecretVersionV2) (tx || db.replicaNode())(TableName.SecretVersionV2)
.where(`${TableName.SecretVersionV2}.folderId`, folderId) .where(`${TableName.SecretVersionV2}.folderId`, folderId)
.groupBy("secretId") .groupBy("secretId")
.max("version") .max("version")
@@ -121,7 +121,7 @@ export const secretVersionV2BridgeDALFactory = (db: TDbClient) => {
.where("folderId", folderId) .where("folderId", folderId)
.whereIn(`${TableName.SecretVersionV2}.secretId`, secretIds) .whereIn(`${TableName.SecretVersionV2}.secretId`, secretIds)
.join( .join(
(tx || db)(TableName.SecretVersionV2) (tx || db.replicaNode())(TableName.SecretVersionV2)
.groupBy("secretId") .groupBy("secretId")
.max("version") .max("version")
.select("secretId") .select("secretId")
@@ -189,7 +189,7 @@ export const secretVersionV2BridgeDALFactory = (db: TDbClient) => {
}) => { }) => {
try { try {
const { offset, limit, sort = [["createdAt", "desc"]] } = findOpt; const { offset, limit, sort = [["createdAt", "desc"]] } = findOpt;
const query = (tx || db)(TableName.SecretVersionV2) const query = (tx || db.replicaNode())(TableName.SecretVersionV2)
.leftJoin(TableName.Users, `${TableName.Users}.id`, `${TableName.SecretVersionV2}.userActorId`) .leftJoin(TableName.Users, `${TableName.Users}.id`, `${TableName.SecretVersionV2}.userActorId`)
.leftJoin( .leftJoin(
TableName.ProjectMembership, TableName.ProjectMembership,
@@ -11,7 +11,7 @@ export const superAdminDALFactory = (db: TDbClient) => {
const superAdminOrm = ormify(db, TableName.SuperAdmin); const superAdminOrm = ormify(db, TableName.SuperAdmin);
const findById = async (id: string, tx?: Knex) => { const findById = async (id: string, tx?: Knex) => {
const config = await (tx || db)(TableName.SuperAdmin) const config = await (tx || db.replicaNode())(TableName.SuperAdmin)
.where(`${TableName.SuperAdmin}.id`, id) .where(`${TableName.SuperAdmin}.id`, id)
.leftJoin(TableName.Organization, `${TableName.SuperAdmin}.defaultAuthOrgId`, `${TableName.Organization}.id`) .leftJoin(TableName.Organization, `${TableName.SuperAdmin}.defaultAuthOrgId`, `${TableName.Organization}.id`)
.leftJoin(TableName.SamlConfig, (qb) => { .leftJoin(TableName.SamlConfig, (qb) => {
+8 -4
View File
@@ -19,12 +19,16 @@ export type TUserDALFactory = ReturnType<typeof userDALFactory>;
export const userDALFactory = (db: TDbClient) => { export const userDALFactory = (db: TDbClient) => {
const userOrm = ormify(db, TableName.Users); const userOrm = ormify(db, TableName.Users);
const findUserByUsername = async (username: string, tx?: Knex) => const findUserByUsername = async (username: string, tx?: Knex) =>
(tx || db)(TableName.Users).whereRaw('lower("username") = :username', { username: username.toLowerCase() }); (tx || db.replicaNode())(TableName.Users).whereRaw('lower("username") = :username', {
username: username.toLowerCase()
});
const findUserByEmail = async (email: string, tx?: Knex) => const findUserByEmail = async (email: string, tx?: Knex) =>
(tx || db)(TableName.Users).whereRaw('lower("email") = :email', { email: email.toLowerCase() }).where({ (tx || db.replicaNode())(TableName.Users)
isEmailVerified: true .whereRaw('lower("email") = :email', { email: email.toLowerCase() })
}); .where({
isEmailVerified: true
});
const getUsersByFilter = async ({ const getUsersByFilter = async ({
limit, limit,
@@ -1,6 +1,7 @@
import { FunctionComponent, ReactNode } from "react"; import { FunctionComponent, ReactNode } from "react";
import { BoundCanProps, Can } from "@casl/react"; import { BoundCanProps, Can } from "@casl/react";
import { TooltipProps } from "@app/components/v2/Tooltip/Tooltip";
import { TOrgPermission, useOrgPermission } from "@app/context/OrgPermissionContext"; import { TOrgPermission, useOrgPermission } from "@app/context/OrgPermissionContext";
import { AccessRestrictedBanner, Tooltip } from "../v2"; import { AccessRestrictedBanner, Tooltip } from "../v2";
@@ -20,6 +21,7 @@ type Props = {
renderTooltip?: boolean; renderTooltip?: boolean;
allowedLabel?: string; allowedLabel?: string;
renderGuardBanner?: boolean; renderGuardBanner?: boolean;
tooltipProps?: Omit<TooltipProps, "children">;
} & BoundCanProps<TOrgPermission>; } & BoundCanProps<TOrgPermission>;
export const OrgPermissionCan: FunctionComponent<Props> = ({ export const OrgPermissionCan: FunctionComponent<Props> = ({
@@ -29,6 +31,7 @@ export const OrgPermissionCan: FunctionComponent<Props> = ({
renderTooltip, renderTooltip,
allowedLabel, allowedLabel,
renderGuardBanner, renderGuardBanner,
tooltipProps,
...props ...props
}) => { }) => {
const { permission } = useOrgPermission(); const { permission } = useOrgPermission();
@@ -43,11 +46,19 @@ export const OrgPermissionCan: FunctionComponent<Props> = ({
: children; : children;
if (!isAllowed && passThrough) { if (!isAllowed && passThrough) {
return <Tooltip content={label}>{finalChild}</Tooltip>; return (
<Tooltip content={label} {...tooltipProps}>
{finalChild}
</Tooltip>
);
} }
if (isAllowed && renderTooltip && allowedLabel) { if (isAllowed && renderTooltip && allowedLabel) {
return <Tooltip content={allowedLabel}>{finalChild}</Tooltip>; return (
<Tooltip content={allowedLabel} {...tooltipProps}>
{finalChild}
</Tooltip>
);
} }
if (!isAllowed && renderGuardBanner) { if (!isAllowed && renderGuardBanner) {
@@ -3,7 +3,7 @@ import { useRouteContext } from "@tanstack/react-router";
import { fetchOrgSubscription, subscriptionQueryKeys } from "@app/hooks/api/subscriptions/queries"; import { fetchOrgSubscription, subscriptionQueryKeys } from "@app/hooks/api/subscriptions/queries";
export const useSubscription = () => { export const useSubscription = (refreshCache?: boolean) => {
const organizationId = useRouteContext({ const organizationId = useRouteContext({
from: "/_authenticate/_inject-org-details", from: "/_authenticate/_inject-org-details",
select: (el) => el.organizationId select: (el) => el.organizationId
@@ -11,7 +11,7 @@ export const useSubscription = () => {
const { data: subscription } = useSuspenseQuery({ const { data: subscription } = useSuspenseQuery({
queryKey: subscriptionQueryKeys.getOrgSubsription(organizationId), queryKey: subscriptionQueryKeys.getOrgSubsription(organizationId),
queryFn: () => fetchOrgSubscription(organizationId), queryFn: () => fetchOrgSubscription(organizationId, refreshCache),
staleTime: Infinity staleTime: Infinity
}); });
@@ -10,9 +10,9 @@ export const subscriptionQueryKeys = {
getOrgSubsription: (orgID: string) => ["plan", { orgID }] as const getOrgSubsription: (orgID: string) => ["plan", { orgID }] as const
}; };
export const fetchOrgSubscription = async (orgID: string) => { export const fetchOrgSubscription = async (orgID: string, refreshCache: boolean = false) => {
const { data } = await apiRequest.get<{ plan: SubscriptionPlan }>( const { data } = await apiRequest.get<{ plan: SubscriptionPlan }>(
`/api/v1/organizations/${orgID}/plan` `/api/v1/organizations/${orgID}/plan${refreshCache ? "?refreshCache=true" : ""}`
); );
return data.plan; return data.plan;
@@ -54,5 +54,7 @@ export type SubscriptionPlan = {
secretScanning: boolean; secretScanning: boolean;
enterpriseSecretSyncs: boolean; enterpriseSecretSyncs: boolean;
enterpriseAppConnections: boolean; enterpriseAppConnections: boolean;
cardDeclined?: boolean;
cardDeclinedReason?: string;
machineIdentityAuthTemplates: boolean; machineIdentityAuthTemplates: boolean;
}; };
@@ -1,4 +1,4 @@
import { useState } from "react"; import { useEffect, useState } from "react";
import { faGithub, faSlack } from "@fortawesome/free-brands-svg-icons"; import { faGithub, faSlack } from "@fortawesome/free-brands-svg-icons";
import { faCircleQuestion, faUserCircle } from "@fortawesome/free-regular-svg-icons"; import { faCircleQuestion, faUserCircle } from "@fortawesome/free-regular-svg-icons";
import { import {
@@ -8,6 +8,7 @@ import {
faCaretDown, faCaretDown,
faCheck, faCheck,
faEnvelope, faEnvelope,
faExclamationTriangle,
faInfo, faInfo,
faInfoCircle, faInfoCircle,
faServer, faServer,
@@ -111,6 +112,14 @@ export const Navbar = () => {
const { subscription } = useSubscription(); const { subscription } = useSubscription();
const { currentOrg } = useOrganization(); const { currentOrg } = useOrganization();
const [showAdminsModal, setShowAdminsModal] = useState(false); const [showAdminsModal, setShowAdminsModal] = useState(false);
const [showCardDeclinedModal, setShowCardDeclinedModal] = useState(false);
useEffect(() => {
if (subscription?.cardDeclined && !sessionStorage.getItem("paymentFailed")) {
sessionStorage.setItem("paymentFailed", "true");
setShowCardDeclinedModal(true);
}
}, [subscription]);
const { data: orgs } = useGetOrganizations(); const { data: orgs } = useGetOrganizations();
const navigate = useNavigate(); const navigate = useNavigate();
@@ -222,6 +231,19 @@ export const Navbar = () => {
<div className="mr-1 rounded border border-mineshaft-500 px-1 text-xs text-bunker-300 !no-underline"> <div className="mr-1 rounded border border-mineshaft-500 px-1 text-xs text-bunker-300 !no-underline">
{getPlan(subscription)} {getPlan(subscription)}
</div> </div>
{subscription.cardDeclined && (
<Tooltip
content={`Your payment could not be processed${subscription.cardDeclinedReason ? `: ${subscription.cardDeclinedReason}` : ""}. Please update your payment method to continue enjoying premium features.`}
className="max-w-xs"
>
<div className="flex items-center">
<FontAwesomeIcon
icon={faExclamationTriangle}
className="animate-pulse cursor-help text-xs text-primary-400"
/>
</div>
</Tooltip>
)}
</div> </div>
</Link> </Link>
<DropdownMenuTrigger asChild> <DropdownMenuTrigger asChild>
@@ -428,6 +450,49 @@ export const Navbar = () => {
</DropdownMenuItem> </DropdownMenuItem>
</DropdownMenuContent> </DropdownMenuContent>
</DropdownMenu> </DropdownMenu>
<Modal isOpen={showCardDeclinedModal} onOpenChange={setShowCardDeclinedModal}>
<ModalContent
title={
<div className="flex items-center gap-2">
<FontAwesomeIcon icon={faExclamationTriangle} className="text-lg text-primary-400" />
Your payment could not be processed.
</div>
}
>
<div>
<div>
<div className="mb-1">
<p>
We were unable to process your last payment
{subscription.cardDeclinedReason ? `: ${subscription.cardDeclinedReason}` : ""}.
Please update your payment information to continue using premium features.
</p>
</div>
<div className="mt-4">
<div className="flex space-x-3">
<Link to="/organization/billing" className="inline-flex">
<Button
colorSchema="primary"
variant="solid"
onClick={() => setShowCardDeclinedModal(false)}
>
Update Payment Method
</Button>
<Button
colorSchema="secondary"
variant="outline"
className="ml-2"
onClick={() => setShowCardDeclinedModal(false)}
>
Dismiss
</Button>
</Link>
</div>
</div>
</div>
</div>
</ModalContent>
</Modal>
<Modal isOpen={showAdminsModal} onOpenChange={setShowAdminsModal}> <Modal isOpen={showAdminsModal} onOpenChange={setShowAdminsModal}>
<ModalContent title="Server Administrators" subTitle="View all server administrators"> <ModalContent title="Server Administrators" subTitle="View all server administrators">
<div className="mb-2"> <div className="mb-2">
@@ -1,5 +1,7 @@
import { useEffect } from "react";
import { faArrowUpRightFromSquare } from "@fortawesome/free-solid-svg-icons"; import { faArrowUpRightFromSquare } from "@fortawesome/free-solid-svg-icons";
import { FontAwesomeIcon } from "@fortawesome/react-fontawesome"; import { FontAwesomeIcon } from "@fortawesome/react-fontawesome";
import { useQueryClient } from "@tanstack/react-query";
import { OrgPermissionCan } from "@app/components/permissions"; import { OrgPermissionCan } from "@app/components/permissions";
import { Button } from "@app/components/v2"; import { Button } from "@app/components/v2";
@@ -15,13 +17,15 @@ import {
useGetOrgPlanBillingInfo, useGetOrgPlanBillingInfo,
useGetOrgTrialUrl useGetOrgTrialUrl
} from "@app/hooks/api"; } from "@app/hooks/api";
import { subscriptionQueryKeys } from "@app/hooks/api/subscriptions/queries";
import { usePopUp } from "@app/hooks/usePopUp"; import { usePopUp } from "@app/hooks/usePopUp";
import { ManagePlansModal } from "./ManagePlansModal"; import { ManagePlansModal } from "./ManagePlansModal";
export const PreviewSection = () => { export const PreviewSection = () => {
const { currentOrg } = useOrganization(); const { currentOrg } = useOrganization();
const { subscription } = useSubscription(); const { subscription } = useSubscription(true);
const queryClient = useQueryClient();
const { data, isPending } = useGetOrgPlanBillingInfo(currentOrg?.id ?? ""); const { data, isPending } = useGetOrgPlanBillingInfo(currentOrg?.id ?? "");
const getOrgTrialUrl = useGetOrgTrialUrl(); const getOrgTrialUrl = useGetOrgTrialUrl();
const createCustomerPortalSession = useCreateCustomerPortalSession(); const createCustomerPortalSession = useCreateCustomerPortalSession();
@@ -37,6 +41,12 @@ export const PreviewSection = () => {
return formattedTotal; return formattedTotal;
}; };
useEffect(() => {
queryClient.invalidateQueries({
queryKey: subscriptionQueryKeys.getOrgSubsription(currentOrg?.id ?? "")
});
}, []);
const formatDate = (date: number) => { const formatDate = (date: number) => {
const createdDate = new Date(date * 1000); const createdDate = new Date(date * 1000);
const day: number = createdDate.getDate(); const day: number = createdDate.getDate();
@@ -24,11 +24,17 @@ enum EnforceAuthType {
export const OrgGeneralAuthSection = ({ export const OrgGeneralAuthSection = ({
isSamlConfigured, isSamlConfigured,
isOidcConfigured, isOidcConfigured,
isGoogleConfigured isGoogleConfigured,
isSamlActive,
isOidcActive,
isLdapActive
}: { }: {
isSamlConfigured: boolean; isSamlConfigured: boolean;
isOidcConfigured: boolean; isOidcConfigured: boolean;
isGoogleConfigured: boolean; isGoogleConfigured: boolean;
isSamlActive: boolean;
isOidcActive: boolean;
isLdapActive: boolean;
}) => { }) => {
const { currentOrg } = useOrganization(); const { currentOrg } = useOrganization();
const { subscription } = useSubscription(); const { subscription } = useSubscription();
@@ -126,6 +132,15 @@ export const OrgGeneralAuthSection = ({
} }
}; };
const isGoogleOAuthEnforced = currentOrg.googleSsoAuthEnforced;
const getActiveSsoLabel = () => {
if (isSamlActive) return "SAML";
if (isOidcActive) return "OIDC";
if (isLdapActive) return "LDAP";
return "";
};
return ( return (
<div className="rounded-lg border border-mineshaft-600 bg-mineshaft-900 p-6"> <div className="rounded-lg border border-mineshaft-600 bg-mineshaft-900 p-6">
<div> <div>
@@ -135,7 +150,7 @@ export const OrgGeneralAuthSection = ({
</p> </p>
</div> </div>
<div className="flex flex-col gap-2 py-4"> <div className="flex flex-col gap-2 py-4">
<div className={twMerge("mt-4", !isSamlConfigured && "hidden")}> <div className={twMerge("mt-4", (!isSamlConfigured || isGoogleOAuthEnforced) && "hidden")}>
<div className="mb-2 flex justify-between"> <div className="mb-2 flex justify-between">
<div className="flex items-center gap-1"> <div className="flex items-center gap-1">
<span className="text-md text-mineshaft-100">Enforce SAML SSO</span> <span className="text-md text-mineshaft-100">Enforce SAML SSO</span>
@@ -160,7 +175,7 @@ export const OrgGeneralAuthSection = ({
</p> </p>
</div> </div>
<div className={twMerge("mt-4", !isOidcConfigured && "hidden")}> <div className={twMerge("mt-4", (!isOidcConfigured || isGoogleOAuthEnforced) && "hidden")}>
<div className="mb-2 flex justify-between"> <div className="mb-2 flex justify-between">
<div className="flex items-center gap-1"> <div className="flex items-center gap-1">
<span className="text-md text-mineshaft-100">Enforce OIDC SSO</span> <span className="text-md text-mineshaft-100">Enforce OIDC SSO</span>
@@ -188,26 +203,47 @@ export const OrgGeneralAuthSection = ({
<div className={twMerge("mt-2", !isGoogleConfigured && "hidden")}> <div className={twMerge("mt-2", !isGoogleConfigured && "hidden")}>
<div className="mb-2 flex justify-between"> <div className="mb-2 flex justify-between">
<div className="flex items-center gap-1"> <div className="flex items-center gap-1">
<span className="text-md text-mineshaft-100">Enforce Google SSO</span> <span className="text-md text-mineshaft-100">Enforce Google OAuth</span>
</div> </div>
<OrgPermissionCan I={OrgPermissionActions.Edit} a={OrgPermissionSubjects.Sso}> <OrgPermissionCan
I={OrgPermissionActions.Edit}
a={OrgPermissionSubjects.Sso}
tooltipProps={{
className: "max-w-sm",
side: "left"
}}
allowedLabel={
isOidcActive || isSamlActive || isLdapActive
? `You cannot enforce Google OAuth while ${getActiveSsoLabel()} SSO is enabled. Disable ${getActiveSsoLabel()} SSO to enforce Google OAuth.`
: undefined
}
renderTooltip={isOidcActive || isSamlActive || isLdapActive}
>
{(isAllowed) => ( {(isAllowed) => (
<Switch <div>
id="enforce-google-sso" <Switch
onCheckedChange={(value) => id="enforce-google-sso"
handleEnforceOrgAuthToggle(value, EnforceAuthType.GOOGLE) onCheckedChange={(value) =>
} handleEnforceOrgAuthToggle(value, EnforceAuthType.GOOGLE)
isChecked={currentOrg?.googleSsoAuthEnforced ?? false} }
isDisabled={!isAllowed || currentOrg?.authEnforced} isChecked={currentOrg?.googleSsoAuthEnforced ?? false}
/> isDisabled={
!isAllowed ||
currentOrg?.authEnforced ||
isOidcActive ||
isSamlActive ||
isLdapActive
}
/>
</div>
)} )}
</OrgPermissionCan> </OrgPermissionCan>
</div> </div>
<p className="text-sm text-mineshaft-300"> <p className="text-sm text-mineshaft-300">
Enforce users to authenticate via Google OAuth SSO to access this organization. Enforce users to authenticate via Google OAuth to access this organization.
<br /> <br />
When this is enabled your organization members will only be able to login with Google When this is enabled your organization members will only be able to login with Google
SSO (not Google SAML). OAuth (not Google SAML).
</p> </p>
</div> </div>
</div> </div>
@@ -267,8 +303,8 @@ export const OrgGeneralAuthSection = ({
</div> </div>
<p className="text-sm text-mineshaft-300"> <p className="text-sm text-mineshaft-300">
<span> <span>
Allow organization admins to bypass SAML enforcement when SSO is unavailable, Allow organization admins to bypass SSO login enforcement when your SSO provider is
misconfigured, or inaccessible. unavailable, misconfigured, or inaccessible.
</span> </span>
</p> </p>
</div> </div>
@@ -94,6 +94,8 @@ export const OrgLDAPSection = (): JSX.Element => {
handlePopUpOpen("ldapGroupMap"); handlePopUpOpen("ldapGroupMap");
}; };
const isGoogleOAuthEnabled = currentOrg.googleSsoAuthEnforced;
return ( return (
<div className="mb-4"> <div className="mb-4">
<div className="py-4"> <div className="py-4">
@@ -116,16 +118,31 @@ export const OrgLDAPSection = (): JSX.Element => {
<div className="pt-4"> <div className="pt-4">
<div className="mb-2 flex items-center justify-between"> <div className="mb-2 flex items-center justify-between">
<h2 className="text-md text-mineshaft-100">Enable LDAP</h2> <h2 className="text-md text-mineshaft-100">Enable LDAP</h2>
<OrgPermissionCan I={OrgPermissionActions.Edit} a={OrgPermissionSubjects.Ldap}> <OrgPermissionCan
I={OrgPermissionActions.Edit}
a={OrgPermissionSubjects.Ldap}
tooltipProps={{
className: "max-w-sm",
side: "left"
}}
allowedLabel={
isGoogleOAuthEnabled
? "You cannot enable LDAP SSO while Google OAuth is enforced. Disable Google OAuth enforcement to enable LDAP SSO."
: undefined
}
renderTooltip={isGoogleOAuthEnabled}
>
{(isAllowed) => ( {(isAllowed) => (
<Switch <div>
id="enable-saml-sso" <Switch
onCheckedChange={(value) => handleLDAPToggle(value)} id="enable-ldap-sso"
isChecked={data ? data.isActive : false} onCheckedChange={(value) => handleLDAPToggle(value)}
isDisabled={!isAllowed} isChecked={data ? data.isActive : false}
> isDisabled={!isAllowed || isGoogleOAuthEnabled}
Enable >
</Switch> Enable
</Switch>
</div>
)} )}
</OrgPermissionCan> </OrgPermissionCan>
</div> </div>
@@ -83,6 +83,8 @@ export const OrgOIDCSection = (): JSX.Element => {
} }
}; };
const isGoogleOAuthEnabled = currentOrg.googleSsoAuthEnforced;
return ( return (
<div className="mb-4 rounded-lg border-mineshaft-600 bg-mineshaft-900"> <div className="mb-4 rounded-lg border-mineshaft-600 bg-mineshaft-900">
<div className="mb-4 flex items-center justify-between"> <div className="mb-4 flex items-center justify-between">
@@ -106,14 +108,29 @@ export const OrgOIDCSection = (): JSX.Element => {
<div className="mb-2 flex items-center justify-between"> <div className="mb-2 flex items-center justify-between">
<h2 className="text-md text-mineshaft-100">Enable OIDC</h2> <h2 className="text-md text-mineshaft-100">Enable OIDC</h2>
{!isPending && ( {!isPending && (
<OrgPermissionCan I={OrgPermissionActions.Edit} a={OrgPermissionSubjects.Sso}> <OrgPermissionCan
I={OrgPermissionActions.Edit}
a={OrgPermissionSubjects.Sso}
tooltipProps={{
className: "max-w-sm",
side: "left"
}}
allowedLabel={
isGoogleOAuthEnabled
? "You cannot enable OIDC SSO while Google OAuth is enforced. Disable Google OAuth enforcement to enable OIDC SSO."
: undefined
}
renderTooltip={isGoogleOAuthEnabled}
>
{(isAllowed) => ( {(isAllowed) => (
<Switch <div>
id="enable-oidc-sso" <Switch
onCheckedChange={(value) => handleOIDCToggle(value)} id="enable-oidc-sso"
isChecked={data ? data.isActive : false} onCheckedChange={(value) => handleOIDCToggle(value)}
isDisabled={!isAllowed} isChecked={data ? data.isActive : false}
/> isDisabled={!isAllowed || isGoogleOAuthEnabled}
/>
</div>
)} )}
</OrgPermissionCan> </OrgPermissionCan>
)} )}
@@ -78,6 +78,8 @@ export const OrgSSOSection = (): JSX.Element => {
} }
}; };
const isGoogleOAuthEnabled = currentOrg.googleSsoAuthEnforced;
return ( return (
<div className="space-y-4"> <div className="space-y-4">
<div className="mb-4 flex items-center justify-between"> <div className="mb-4 flex items-center justify-between">
@@ -99,14 +101,29 @@ export const OrgSSOSection = (): JSX.Element => {
<div className="mb-2 flex items-center justify-between pt-4"> <div className="mb-2 flex items-center justify-between pt-4">
<h2 className="text-md text-mineshaft-100">Enable SAML</h2> <h2 className="text-md text-mineshaft-100">Enable SAML</h2>
{!isPending && ( {!isPending && (
<OrgPermissionCan I={OrgPermissionActions.Edit} a={OrgPermissionSubjects.Sso}> <OrgPermissionCan
I={OrgPermissionActions.Edit}
a={OrgPermissionSubjects.Sso}
tooltipProps={{
className: "max-w-sm",
side: "left"
}}
allowedLabel={
isGoogleOAuthEnabled
? "You cannot enable SAML SSO while Google OAuth is enforced. Disable Google OAuth enforcement to enable SAML SSO."
: undefined
}
renderTooltip={isGoogleOAuthEnabled}
>
{(isAllowed) => ( {(isAllowed) => (
<Switch <div>
id="enable-saml-sso" <Switch
onCheckedChange={(value) => handleSamlSSOToggle(value)} id="enable-saml-sso"
isChecked={data ? data.isActive : false} onCheckedChange={(value) => handleSamlSSOToggle(value)}
isDisabled={!isAllowed} isChecked={data ? data.isActive : false}
/> isDisabled={!isAllowed || isGoogleOAuthEnabled}
/>
</div>
)} )}
</OrgPermissionCan> </OrgPermissionCan>
)} )}
@@ -184,6 +184,9 @@ export const OrgSsoTab = withPermission(
isSamlConfigured={isSamlConfigured} isSamlConfigured={isSamlConfigured}
isOidcConfigured={isOidcConfigured} isOidcConfigured={isOidcConfigured}
isGoogleConfigured={isGoogleConfigured} isGoogleConfigured={isGoogleConfigured}
isSamlActive={Boolean(samlConfig?.isActive)}
isOidcActive={Boolean(oidcConfig?.isActive)}
isLdapActive={Boolean(ldapConfig?.isActive)}
/> />
)} )}
@@ -19,6 +19,7 @@ import {
SecretInput, SecretInput,
Select, Select,
SelectItem, SelectItem,
Switch,
TextArea, TextArea,
Tooltip Tooltip
} from "@app/components/v2"; } from "@app/components/v2";
@@ -66,6 +67,7 @@ const formSchema = z.object({
creationStatement: z.string().min(1), creationStatement: z.string().min(1),
revocationStatement: z.string().min(1), revocationStatement: z.string().min(1),
renewStatement: z.string().optional(), renewStatement: z.string().optional(),
sslEnabled: z.boolean().optional(),
ca: z.string().optional(), ca: z.string().optional(),
gatewayId: z.string().optional() gatewayId: z.string().optional()
}), }),
@@ -200,6 +202,7 @@ export const SqlDatabaseInputForm = ({
const createDynamicSecret = useCreateDynamicSecret(); const createDynamicSecret = useCreateDynamicSecret();
const { data: gateways, isPending: isGatewaysLoading } = useQuery(gatewaysQueryKeys.list()); const { data: gateways, isPending: isGatewaysLoading } = useQuery(gatewaysQueryKeys.list());
const selectedClient = watch("provider.client");
const handleCreateDynamicSecret = async ({ const handleCreateDynamicSecret = async ({
name, name,
@@ -458,13 +461,34 @@ export const SqlDatabaseInputForm = ({
/> />
</div> </div>
<div> <div>
{selectedClient === SqlProviders.MsSQL && (
<div className="mb-2 mt-2">
<Controller
control={control}
name="provider.sslEnabled"
render={({ field: { value, onChange }, fieldState: { error } }) => (
<FormControl isError={Boolean(error?.message)} errorText={error?.message}>
<Switch
className="bg-mineshaft-400/50 shadow-inner data-[state=checked]:bg-green/80"
id="sql-ds-ssl-enabled"
thumbClassName="bg-mineshaft-800"
isChecked={value}
onCheckedChange={onChange}
>
Encrypt Connection (SSL)
</Switch>
</FormControl>
)}
/>
</div>
)}
<Controller <Controller
control={control} control={control}
name="provider.ca" name="provider.ca"
render={({ field, fieldState: { error } }) => ( render={({ field, fieldState: { error } }) => (
<FormControl <FormControl
isOptional isOptional
label="CA(SSL)" label="CA (SSL)"
isError={Boolean(error?.message)} isError={Boolean(error?.message)}
errorText={error?.message} errorText={error?.message}
> >
@@ -18,6 +18,7 @@ import {
SecretInput, SecretInput,
Select, Select,
SelectItem, SelectItem,
Switch,
TextArea, TextArea,
Tooltip Tooltip
} from "@app/components/v2"; } from "@app/components/v2";
@@ -63,6 +64,7 @@ const formSchema = z.object({
creationStatement: z.string().min(1), creationStatement: z.string().min(1),
revocationStatement: z.string().min(1), revocationStatement: z.string().min(1),
renewStatement: z.string().optional(), renewStatement: z.string().optional(),
sslEnabled: z.boolean().optional(),
ca: z.string().optional(), ca: z.string().optional(),
gatewayId: z.string().optional().nullable() gatewayId: z.string().optional().nullable()
}) })
@@ -151,6 +153,7 @@ export const EditDynamicSecretSqlProviderForm = ({
}); });
const { data: gateways, isPending: isGatewaysLoading } = useQuery(gatewaysQueryKeys.list()); const { data: gateways, isPending: isGatewaysLoading } = useQuery(gatewaysQueryKeys.list());
const selectedClient = watch("inputs.client");
const updateDynamicSecret = useUpdateDynamicSecret(); const updateDynamicSecret = useUpdateDynamicSecret();
const selectedGatewayId = watch("inputs.gatewayId"); const selectedGatewayId = watch("inputs.gatewayId");
@@ -407,13 +410,34 @@ export const EditDynamicSecretSqlProviderForm = ({
/> />
</div> </div>
<div> <div>
{selectedClient === SqlProviders.MsSQL && (
<div className="mb-2 mt-2">
<Controller
control={control}
name="inputs.sslEnabled"
render={({ field: { value, onChange }, fieldState: { error } }) => (
<FormControl isError={Boolean(error?.message)} errorText={error?.message}>
<Switch
className="bg-mineshaft-400/50 shadow-inner data-[state=checked]:bg-green/80"
id="sql-ds-ssl-enabled"
thumbClassName="bg-mineshaft-800"
isChecked={Boolean(value)}
onCheckedChange={onChange}
>
Encrypt Connection (SSL)
</Switch>
</FormControl>
)}
/>
</div>
)}
<Controller <Controller
control={control} control={control}
name="inputs.ca" name="inputs.ca"
render={({ field, fieldState: { error } }) => ( render={({ field, fieldState: { error } }) => (
<FormControl <FormControl
isOptional isOptional
label="CA(SSL)" label="CA (SSL)"
isError={Boolean(error?.message)} isError={Boolean(error?.message)}
errorText={error?.message} errorText={error?.message}
> >