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 () => {
return 1;
},
pgGetIntItem: async (key) => {
const value = store[key];
if (typeof value === "number") {
return Number(value);
}
},
pgIncrementBy: async () => {
return 1;
},
getItems: async (keys) => {
const values = keys.map((key) => {
const value = store[key];
+3 -1
View File
@@ -15,6 +15,7 @@ import { mockSmtpServer } from "./mocks/smtp";
import { initDbConnection } from "@app/db";
import { queueServiceFactory } from "@app/queue";
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 { buildRedisFromConfig } from "@app/lib/config/redis";
import { superAdminDALFactory } from "@app/services/super-admin/super-admin-dal";
@@ -62,7 +63,8 @@ export default {
const smtp = mockSmtpServer();
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();
+8
View File
@@ -197,6 +197,9 @@ import {
TInternalKms,
TInternalKmsInsert,
TInternalKmsUpdate,
TKeyValueStore,
TKeyValueStoreInsert,
TKeyValueStoreUpdate,
TKmipClientCertificates,
TKmipClientCertificatesInsert,
TKmipClientCertificatesUpdate,
@@ -1296,5 +1299,10 @@ declare module "knex/types/tables" {
TUserNotificationsInsert,
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 "./internal-certificate-authorities";
export * from "./internal-kms";
export * from "./key-value-store";
export * from "./kmip-client-certificates";
export * from "./kmip-clients";
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",
OrgGatewayConfigV2 = "org_gateway_config_v2",
Relay = "relays",
GatewayV2 = "gateways_v2"
GatewayV2 = "gateways_v2",
KeyValueStore = "key_value_store"
}
export type TImmutableDBKeys = "id" | "createdAt" | "updatedAt" | "commitId";
@@ -84,7 +84,9 @@ export const registerDynamicSecretRouter = async (server: FastifyZodProvider) =>
}),
response: {
200: z.object({
dynamicSecret: SanitizedDynamicSecretSchema
dynamicSecret: SanitizedDynamicSecretSchema.extend({
inputs: z.unknown()
})
})
}
},
@@ -151,7 +153,9 @@ export const registerDynamicSecretRouter = async (server: FastifyZodProvider) =>
}),
response: {
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: {
params: z.object({ organizationId: z.string().trim() }),
querystring: z.object({
refreshCache: z
.enum(["true", "false"])
.default("false")
.transform((value) => value === "true")
}),
response: {
200: z.object({ plan: z.any() })
}
@@ -54,7 +60,8 @@ export const registerLicenseRouter = async (server: FastifyZodProvider) => {
actor: req.permission.type,
actorOrgId: req.permission.orgId,
actorAuthMethod: req.permission.authMethod,
orgId: req.params.organizationId
orgId: req.params.organizationId,
refreshCache: req.query.refreshCache
});
return { plan };
}
@@ -46,7 +46,10 @@ export const dynamicSecretLeaseDALFactory = (db: TDbClient) => {
const countLeasesForDynamicSecret = async (dynamicSecretId: string, tx?: Knex) => {
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);
} catch (error) {
throw new DatabaseError({ error, name: "DynamicSecretCountLeases" });
@@ -55,7 +58,7 @@ export const dynamicSecretLeaseDALFactory = (db: TDbClient) => {
const findById = async (id: string, tx?: Knex) => {
try {
const doc = await (tx || db)(TableName.DynamicSecretLease)
const doc = await (tx || db.replicaNode())(TableName.DynamicSecretLease)
.where({ [`${TableName.DynamicSecretLease}.id` as "id"]: id })
.first()
.join(
@@ -190,7 +190,7 @@ export const dynamicSecretServiceFactory = ({
return cfg;
});
return dynamicSecretCfg;
return { ...dynamicSecretCfg, inputs };
};
const updateByName: TDynamicSecretServiceFactory["updateByName"] = async ({
@@ -355,7 +355,7 @@ export const dynamicSecretServiceFactory = ({
return cfg;
});
return updatedDynamicCfg;
return { ...updatedDynamicCfg, inputs: updatedInput };
};
const deleteByName: TDynamicSecretServiceFactory["deleteByName"] = async ({
@@ -165,6 +165,7 @@ export const DynamicSecretSqlDBSchema = z.object({
revocationStatement: z.string().trim(),
renewStatement: z.string().trim().optional(),
ca: z.string().optional(),
sslEnabled: z.boolean().optional(),
gatewayId: z.string().nullable().optional()
});
@@ -1,4 +1,5 @@
import handlebars from "handlebars";
import RE2 from "re2";
import knex from "knex";
import { z } from "zod";
@@ -156,19 +157,40 @@ export const SqlDatabaseProvider = ({
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
? { rejectUnauthorized: false, ca: providerInputs.ca, servername: providerInputs.host }
: undefined;
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({
client: providerInputs.client,
connection: {
database: providerInputs.database,
port: providerInputs.port,
host: providerInputs.client === SqlProviders.Postgres ? providerInputs.hostIp : providerInputs.host,
user: providerInputs.username,
host:
providerInputs.client === SqlProviders.Postgres && !providerInputs.gatewayId
? providerInputs.hostIp
: providerInputs.host,
user: effectiveUser,
password: providerInputs.password,
ssl,
// @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
options: isMsSQLClient
? {
...(providerInputs.sslEnabled !== undefined ? { encrypt: providerInputs.sslEnabled } : {}),
trustServerCertificate: !providerInputs.ca,
cryptoCredentialsDetails: providerInputs.ca ? { ca: providerInputs.ca } : {}
}
@@ -238,7 +261,13 @@ export const SqlDatabaseProvider = ({
const providerInputs = await validateProviderInputs(inputs);
let isConnected = false;
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
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 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 {
const expiration = new Date(expireAt).toISOString();
@@ -322,7 +356,12 @@ export const SqlDatabaseProvider = ({
const username = entityId;
const { database } = providerInputs;
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 {
const revokeStatement = handlebars.compile(providerInputs.revocationStatement)({ username, database });
const queries = revokeStatement.toString().split(";").filter(Boolean);
@@ -357,7 +396,12 @@ export const SqlDatabaseProvider = ({
if (!providerInputs.renewStatement) return { entityId };
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 { database } = providerInputs;
@@ -13,7 +13,7 @@ export const gatewayDALFactory = (db: TDbClient) => {
{ offset, limit, sort, tx }: TFindOpt<TGateways> = {}
) => {
try {
const query = (tx || db)(TableName.Gateway)
const query = (tx || db.replicaNode())(TableName.Gateway)
// eslint-disable-next-line @typescript-eslint/no-misused-promises
.where(buildFindFilter(filter, TableName.Gateway, ["orgId"]))
.join(TableName.Identity, `${TableName.Identity}.id`, `${TableName.Gateway}.identityId`)
@@ -23,7 +23,7 @@ export const userGroupMembershipDALFactory = (db: TDbClient) => {
.whereIn(`${TableName.ProjectMembership}.projectId`, projectIds)
.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)
.whereNot(`${TableName.UserGroupMembership}.groupId`, groupId)
.join(
@@ -79,7 +79,7 @@ export const userGroupMembershipDALFactory = (db: TDbClient) => {
.pluck(`${TableName.GroupProjectMembership}.groupId`);
// 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}.isPending`, false)
.join(TableName.Users, `${TableName.UserGroupMembership}.userId`, `${TableName.Users}.id`)
@@ -127,6 +127,20 @@ export const ldapConfigServiceFactory = ({
message:
"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({
type: KmsDataKey.Organization,
orgId
@@ -233,6 +247,19 @@ export const ldapConfigServiceFactory = ({
"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 = {
isActive,
url,
@@ -28,7 +28,7 @@ export const licenseDALFactory = (db: TDbClient) => {
const countOrgUsersAndIdentities = async (orgId: string | null, tx?: Knex) => {
try {
// count org users
const userDoc = await (tx || db)(TableName.OrgMembership)
const userDoc = await (tx || db.replicaNode())(TableName.OrgMembership)
.where({ status: OrgMembershipStatus.Accepted })
.andWhere((bd) => {
if (orgId) {
@@ -42,7 +42,7 @@ export const licenseDALFactory = (db: TDbClient) => {
const userCount = Number(userDoc?.[0].count);
// count org identities
const identityDoc = await (tx || db)(TableName.IdentityOrgMembership)
const identityDoc = await (tx || db.replicaNode())(TableName.IdentityOrgMembership)
.where((bd) => {
if (orgId) {
void bd.where({ orgId });
@@ -99,6 +99,17 @@ export const licenseServiceFactory = ({
const workspacesUsed = await projectDAL.countOfOrgProjects(null);
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;
logger.info("Successfully synchronized license key features");
} catch (error) {
@@ -226,10 +237,13 @@ export const licenseServiceFactory = ({
};
const refreshPlan = async (orgId: string) => {
if (instanceType === InstanceType.Cloud) {
await keyStore.deleteItem(FEATURE_CACHE_KEY(orgId));
if (instanceType === InstanceType.Cloud) {
await getPlan(orgId);
}
if (instanceType === InstanceType.EnterpriseOnPrem) {
await syncLicenseKeyOnPremFeatures(true);
}
};
const generateOrgCustomerId = async (orgName: string, email?: string | null) => {
@@ -296,8 +310,19 @@ export const licenseServiceFactory = ({
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);
if (refreshCache) {
await refreshPlan(orgId);
}
const plan = await getPlan(orgId, projectId);
return plan;
};
@@ -87,6 +87,7 @@ export type TOrgPlansTableDTO = {
export type TOrgPlanDTO = {
projectId?: string;
refreshCache?: boolean;
} & TOrgPermission;
export type TStartOrgTrialDTO = {
@@ -499,6 +499,13 @@ export const oidcConfigServiceFactory = ({
);
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({
type: KmsDataKey.Organization,
orgId: org.id
@@ -586,6 +593,13 @@ export const oidcConfigServiceFactory = ({
);
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({
type: KmsDataKey.Organization,
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."
});
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({
type: KmsDataKey.Organization,
orgId
@@ -120,6 +133,19 @@ export const samlConfigServiceFactory = ({
"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 { encryptor } = await kmsService.createCipherPairWithDataKey({
type: KmsDataKey.Organization,
@@ -345,7 +345,7 @@ export const secretApprovalRequestDALFactory = (db: TDbClient) => {
const findProjectRequestCount = async (projectId: string, userId: string, policyId?: string, tx?: Knex) => {
try {
const docs = await (tx || db)
const docs = await (tx || db.replicaNode())
.with(
"temp",
(tx || db.replicaNode())(TableName.SecretApprovalRequest)
@@ -494,7 +494,7 @@ export const secretApprovalRequestDALFactory = (db: TDbClient) => {
.distinctOn(`${TableName.SecretApprovalRequest}.id`)
.as("inner");
const query = (tx || db)
const query = (tx || db.replicaNode())
.select("*")
.select(db.raw("count(*) OVER() as total_count"))
.from(innerQuery)
@@ -377,7 +377,7 @@ export const secretApprovalRequestSecretDALFactory = (db: TDbClient) => {
// special query for migration to v2 secret
const findByProjectId = async (projectId: string, tx?: Knex) => {
try {
const docs = await (tx || db)(TableName.SecretApprovalRequestSecret)
const docs = await (tx || db.replicaNode())(TableName.SecretApprovalRequestSecret)
.join(
TableName.SecretApprovalRequest,
`${TableName.SecretApprovalRequest}.id`,
@@ -787,6 +787,7 @@ export const secretApprovalRequestServiceFactory = ({
},
tx
);
await secretV2BridgeDAL.invalidateSecretCacheByProjectId(projectId, tx);
return {
secrets: { created: newSecrets, updated: updatedSecrets, deleted: deletedSecret },
approval: updatedSecretApproval
@@ -976,6 +977,7 @@ export const secretApprovalRequestServiceFactory = ({
},
tx
);
await secretV2BridgeDAL.invalidateSecretCacheByProjectId(projectId, tx);
return {
secrets: { created: newSecrets, updated: updatedSecrets, deleted: deletedSecret },
approval: updatedSecretApproval
@@ -983,7 +985,6 @@ export const secretApprovalRequestServiceFactory = ({
});
}
await secretV2BridgeDAL.invalidateSecretCacheByProjectId(projectId);
await snapshotService.performSnapshot(folderId);
const [folder] = await folderDAL.findSecretPathByFolderIds(projectId, [folderId]);
if (!folder) {
@@ -509,9 +509,9 @@ export const secretReplicationServiceFactory = ({
tx
);
}
await secretV2BridgeDAL.invalidateSecretCacheByProjectId(projectId, tx);
});
await secretV2BridgeDAL.invalidateSecretCacheByProjectId(projectId);
await secretQueueService.syncSecrets({
projectId,
orgId,
@@ -361,9 +361,8 @@ export const secretRotationQueueFactory = ({
},
tx
);
await secretV2BridgeDAL.invalidateSecretCacheByProjectId(secretRotation.projectId, tx);
});
await secretV2BridgeDAL.invalidateSecretCacheByProjectId(secretRotation.projectId);
} else {
if (!botKey)
throw new NotFoundError({
@@ -265,7 +265,7 @@ export const snapshotDALFactory = (db: TDbClient) => {
// then joins with respective secrets and folder
const findRecursivelySnapshots = async (snapshotId: string, tx?: Knex) => {
try {
const data = await (tx || db)
const data = await (tx || db.replicaNode())
.withRecursive("parent", (qb) => {
void qb
.from(TableName.Snapshot)
@@ -419,7 +419,7 @@ export const snapshotDALFactory = (db: TDbClient) => {
// then joins with respective secrets and folder
const findRecursivelySnapshotsV2Bridge = async (snapshotId: string, tx?: Knex) => {
try {
const data = await (tx || db)
const data = await (tx || db.replicaNode())
.withRecursive("parent", (qb) => {
void qb
.from(TableName.Snapshot)
@@ -581,7 +581,11 @@ export const snapshotDALFactory = (db: TDbClient) => {
const docs = await (tx || db.replicaNode())(TableName.Snapshot)
.where(`${TableName.Snapshot}.folderId`, folderId)
.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.on(`${TableName.Snapshot}.folderId`, "latestVersion.folderId").andOn(
`${TableName.Snapshot}.createdAt`,
@@ -766,7 +770,7 @@ export const snapshotDALFactory = (db: TDbClient) => {
)
.orderBy(`${TableName.Snapshot}.createdAt`, "desc")
.where(`${TableName.Snapshot}.folderId`, folderId);
const data = await (tx || db)
const data = await (tx || db.replicaNode())
.with("w", query)
.select("*")
.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 { pgAdvisoryLockHashText } from "@app/lib/crypto/hashtext";
import { applyJitter } from "@app/lib/dates";
import { delay as delayMs } from "@app/lib/delay";
import { ms } from "@app/lib/ms";
import { ExecutionResult, Redlock, Settings } from "@app/lib/red-lock";
import { Redis, Cluster } from "ioredis";
import { TKeyValueStoreDALFactory } from "./key-value-store-dal";
export const PgSqlLock = {
BootUpMigration: 2023,
@@ -97,13 +102,17 @@ export type TKeyStoreFactory = {
deleteItemsByKeyIn: (keys: string[]) => Promise<number>;
deleteItems: (arg: TDeleteItems) => 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(
resources: string[],
duration: number,
settings?: Partial<Settings>
): Promise<{ release: () => Promise<ExecutionResult> }>;
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>) => {
@@ -116,7 +125,10 @@ interface TKeyStoreFactoryDTO extends TRedisConfigKeys {
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 redisReadReplicas = redisConfigKeys.REDIS_READ_REPLICAS?.map((el) => {
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 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) => {
let cursor = "0";
const allKeys: string[] = [];
@@ -238,6 +227,37 @@ export const keyStoreFactory = (redisConfigKeys: TKeyStoreFactoryDTO): TKeyStore
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 {
setItem,
getItem,
@@ -252,6 +272,8 @@ export const keyStoreFactory = (redisConfigKeys: TKeyStoreFactoryDTO): TKeyStore
waitTillReady,
getKeysByPattern,
deleteItemsByKeyIn,
getItems
getItems,
pgGetIntItem,
pgIncrementBy
};
};
+9
View File
@@ -53,6 +53,15 @@ export const inMemoryKeyStore = (): TKeyStoreFactory => {
}
return null;
},
pgGetIntItem: async (key) => {
const value = store[key];
if (typeof value === "number") {
return Number(value);
}
},
pgIncrementBy: async () => {
return 1;
},
incrementBy: async () => {
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_SECRET) &&
Boolean(data.INF_APP_CONNECTION_GITHUB_RADAR_APP_WEBHOOK_SECRET),
isSecondaryInstance: Boolean(data.INFISICAL_PRIMARY_INSTANCE_URL),
isHsmConfigured:
Boolean(data.HSM_LIB_PATH) && Boolean(data.HSM_PIN) && Boolean(data.HSM_KEY_LABEL) && data.HSM_SLOT !== undefined,
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("*");
if ($incr) {
Object.entries($incr).forEach(([incrementField, incrementValue]) => {
void query.increment(incrementField, incrementValue);
void query.increment(incrementField, incrementValue as number);
});
}
if ($decr) {
Object.entries($decr).forEach(([incrementField, incrementValue]) => {
void query.decrement(incrementField, incrementValue);
void query.decrement(incrementField, incrementValue as number);
});
}
const [docs] = await query;
@@ -273,12 +273,12 @@ export const ormify = <DbOps extends object, Tname extends keyof Tables>(
// increment and decrement operation in update
if ($incr) {
Object.entries($incr).forEach(([incrementField, incrementValue]) => {
void query.increment(incrementField, incrementValue);
void query.increment(incrementField, incrementValue as number);
});
}
if ($decr) {
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"][];
+3 -1
View File
@@ -5,6 +5,7 @@ import "./lib/telemetry/instrumentation";
import dotenv from "dotenv";
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 { initAuditLogDbConnection, initDbConnection } from "./db";
@@ -54,7 +55,8 @@ const run = async () => {
await queue.initialize();
const keyStore = keyStoreFactory(envConfig);
const keyValueStoreDAL = keyValueStoreDALFactory(db);
const keyStore = keyStoreFactory(envConfig, keyValueStoreDAL);
const redis = buildRedisFromConfig(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 { trustedIpDALFactory } from "@app/ee/services/trusted-ip/trusted-ip-dal";
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 { getConfig, TEnvConfig } from "@app/lib/config/env";
import { crypto } from "@app/lib/crypto/cryptography";
@@ -514,6 +515,7 @@ export const registerRoutes = async (
const microsoftTeamsIntegrationDAL = microsoftTeamsIntegrationDALFactory(db);
const projectMicrosoftTeamsConfigDAL = projectMicrosoftTeamsConfigDALFactory(db);
const secretScanningV2DAL = secretScanningV2DALFactory(db);
const keyValueStoreDAL = keyValueStoreDALFactory(db);
const eventBusService = eventBusFactory(server.redis);
const sseService = sseServiceFactory(eventBusService, server.redis);
@@ -650,6 +652,7 @@ export const registerRoutes = async (
const folderTreeCheckpointDAL = folderTreeCheckpointDALFactory(db);
const folderCommitDAL = folderCommitDALFactory(db);
const folderTreeCheckpointResourcesDAL = folderTreeCheckpointResourcesDALFactory(db);
const folderCommitQueueService = folderCommitQueueServiceFactory({
queueService,
folderTreeCheckpointDAL,
@@ -814,6 +817,7 @@ export const registerRoutes = async (
groupDAL,
orgBotDAL,
oidcConfigDAL,
ldapConfigDAL,
loginService,
projectBotService,
reminderService
@@ -1719,6 +1723,7 @@ export const registerRoutes = async (
userDAL,
identityDAL
});
const dailyResourceCleanUp = dailyResourceCleanUpQueueServiceFactory({
auditLogDAL,
queueService,
@@ -1731,7 +1736,8 @@ export const registerRoutes = async (
identityUniversalAuthClientSecretDAL: identityUaClientSecretDAL,
serviceTokenService,
orgService,
userNotificationDAL
userNotificationDAL,
keyValueStoreDAL
});
const dailyReminderQueueService = dailyReminderQueueServiceFactory({
@@ -661,7 +661,7 @@ describe("folderCommitServiceFactory", () => {
// Assert
expect(mockFolderCommitDAL.create).toHaveBeenCalled();
expect(mockSecretV2BridgeDAL.invalidateSecretCacheByProjectId).toHaveBeenCalledWith(projectId);
expect(mockSecretV2BridgeDAL.invalidateSecretCacheByProjectId).toHaveBeenCalledWith(projectId, {});
// Check that we got the right counts
expect(result.totalChanges).toEqual(2);
@@ -1386,7 +1386,7 @@ export const folderCommitServiceFactory = ({
);
// Invalidate cache to reflect the changes
await secretV2BridgeDAL.invalidateSecretCacheByProjectId(projectId);
await secretV2BridgeDAL.invalidateSecretCacheByProjectId(projectId, tx);
return {
secretChangesCount: secretChanges.length,
@@ -84,7 +84,8 @@ export const identityUaServiceFactory = ({
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;
if (identityUa.lockoutEnabled) {
try {
lock = await keyStore.acquireLock([KeyStorePrefixes.IdentityLockoutLock(LOCKOUT_KEY)], 500, {
retryCount: 3,
@@ -95,7 +96,8 @@ export const identityUaServiceFactory = ({
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 {
@@ -257,7 +259,7 @@ export const identityUaServiceFactory = ({
...accessTokenTTLParams
};
} finally {
await lock.release();
if (lock) await lock.release();
}
};
@@ -25,7 +25,7 @@ export const identityDALFactory = (db: TDbClient) => {
} as const;
const tableName = authMethodToTableName[authMethod];
if (!tableName) return;
const data = await db(tableName).where({ identityId }).first();
const data = await db.replicaNode()(tableName).where({ identityId }).first();
if (!data) return;
return data.accessTokenTrustedIps;
};
@@ -30,7 +30,7 @@ export const integrationAuthDALFactory = (db: TDbClient) => {
const getByOrg = async (orgId: string, tx?: Knex) => {
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.Organization, `${TableName.Organization}.id`, `${TableName.Project}.orgId`)
.where(`${TableName.Organization}.id`, "=", orgId)
@@ -12,7 +12,7 @@ export const kmsRootConfigDALFactory = (db: TDbClient) => {
const findById = async (id: string, tx?: Knex) => {
try {
const result = await (tx || db)(TableName.KmsServerRootConfig)
const result = await (tx || db?.replicaNode?.() || db)(TableName.KmsServerRootConfig)
.where({ id } as never)
.first("*");
return result;
+38 -3
View File
@@ -9,11 +9,14 @@ import {
ProjectMembershipRole,
ProjectVersion,
TableName,
TOidcConfigs,
TProjectMemberships,
TProjectUserMembershipRolesInsert,
TSamlConfigs,
TUsers
} from "@app/db/schemas";
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 { TOidcConfigDALFactory } from "@app/ee/services/oidc/oidc-config-dal";
import {
@@ -125,6 +128,7 @@ type TOrgServiceFactoryDep = {
incidentContactDAL: TIncidentContactsDALFactory;
samlConfigDAL: Pick<TSamlConfigDALFactory, "findOne">;
oidcConfigDAL: Pick<TOidcConfigDALFactory, "findOne">;
ldapConfigDAL: Pick<TLdapConfigDALFactory, "findOne">;
smtpService: TSmtpService;
tokenService: TAuthTokenServiceFactory;
permissionService: TPermissionServiceFactory;
@@ -165,6 +169,7 @@ export const orgServiceFactory = ({
projectRoleDAL,
samlConfigDAL,
oidcConfigDAL,
ldapConfigDAL,
projectUserMembershipRoleDAL,
identityMetadataDAL,
projectBotService,
@@ -446,16 +451,20 @@ export const orgServiceFactory = ({
});
}
if (authEnforced) {
const samlCfg = await samlConfigDAL.findOne({
let samlCfg: TSamlConfigs | undefined;
let oidcCfg: TOidcConfigs | undefined;
if (authEnforced || googleSsoAuthEnforced) {
samlCfg = await samlConfigDAL.findOne({
orgId,
isActive: true
});
const oidcCfg = await oidcConfigDAL.findOne({
oidcCfg = await oidcConfigDAL.findOne({
orgId,
isActive: true
});
}
if (authEnforced) {
if (!samlCfg && !oidcCfg)
throw new NotFoundError({
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) {
throw new BadRequestError({
message:
@@ -39,7 +39,7 @@ export const reminderDALFactory = (db: TDbClient) => {
const findSecretDailyReminders = async (tx?: Knex) => {
const { startOfDay, endOfDay } = getTodayDateRange();
const rawReminders = await (tx || db)(TableName.Reminder)
const rawReminders = await (tx || db.replicaNode())(TableName.Reminder)
.whereBetween("nextReminderDate", [startOfDay, endOfDay])
.leftJoin(TableName.ReminderRecipient, `${TableName.Reminder}.id`, `${TableName.ReminderRecipient}.reminderId`)
.leftJoin<TUsers>(TableName.Users, `${TableName.ReminderRecipient}.userId`, `${TableName.Users}.id`)
@@ -90,7 +90,7 @@ export const reminderDALFactory = (db: TDbClient) => {
const futureDate = new Date(startOfDay);
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", "<=", futureDate)
.orderBy("nextReminderDate", "asc")
@@ -101,7 +101,7 @@ export const reminderDALFactory = (db: TDbClient) => {
};
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)
.leftJoin(TableName.ReminderRecipient, `${TableName.Reminder}.id`, `${TableName.ReminderRecipient}.reminderId`)
.select(selectAllTableCols(TableName.Reminder))
@@ -125,7 +125,7 @@ export const reminderDALFactory = (db: TDbClient) => {
};
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)
.leftJoin(TableName.ReminderRecipient, `${TableName.Reminder}.id`, `${TableName.ReminderRecipient}.reminderId`)
.select(selectAllTableCols(TableName.Reminder))
@@ -1,5 +1,6 @@
import { TAuditLogDALFactory } from "@app/ee/services/audit-log/audit-log-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 { logger } from "@app/lib/logger";
import { QueueJobs, QueueName, TQueueServiceFactory } from "@app/queue";
@@ -27,6 +28,7 @@ type TDailyResourceCleanUpQueueServiceFactoryDep = {
queueService: TQueueServiceFactory;
orgService: TOrgServiceFactory;
userNotificationDAL: Pick<TUserNotificationDALFactory, "pruneNotifications">;
keyValueStoreDAL: Pick<TKeyValueStoreDALFactory, "pruneExpiredKeys">;
};
export type TDailyResourceCleanUpQueueServiceFactory = ReturnType<typeof dailyResourceCleanUpQueueServiceFactory>;
@@ -43,7 +45,8 @@ export const dailyResourceCleanUpQueueServiceFactory = ({
identityUniversalAuthClientSecretDAL,
serviceTokenService,
orgService,
userNotificationDAL
userNotificationDAL,
keyValueStoreDAL
}: TDailyResourceCleanUpQueueServiceFactoryDep) => {
const appCfg = getConfig();
@@ -52,6 +55,10 @@ export const dailyResourceCleanUpQueueServiceFactory = ({
}
const init = async () => {
if (appCfg.isSecondaryInstance) {
return;
}
await queueService.stopRepeatableJob(
QueueName.AuditLogPrune,
QueueJobs.AuditLogPrune,
@@ -82,6 +89,7 @@ export const dailyResourceCleanUpQueueServiceFactory = ({
await orgService.notifyInvitedUsers();
await auditLogDAL.pruneAuditLog();
await userNotificationDAL.pruneNotifications();
await keyValueStoreDAL.pruneExpiredKeys();
logger.info(`${QueueName.DailyResourceCleanUp}: queue task completed`);
} catch (error) {
logger.error(error, `${QueueName.DailyResourceCleanUp}: resource cleanup failed`);
@@ -45,7 +45,7 @@ export const secretFolderVersionDALFactory = (db: TDbClient) => {
)
.whereIn(`${TableName.SecretFolderVersion}.folderId`, folderIds)
.join(
(tx || db)(TableName.SecretFolderVersion)
(tx || db.replicaNode())(TableName.SecretFolderVersion)
.groupBy("folderId")
.max("version")
.select("folderId")
@@ -15,7 +15,7 @@ export const secretImportDALFactory = (db: TDbClient) => {
// 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
const findLastImportPosition = async (folderId: string, tx?: Knex) => {
const lastPos = await (tx || db)(TableName.SecretImport)
const lastPos = await (tx || db.replicaNode())(TableName.SecretImport)
.where({ folderId })
.max("position", { as: "position" })
.first();
@@ -119,7 +119,7 @@ export const secretSharingDALFactory = (db: TDbClient) => {
const findActiveSharedSecrets = async (filters: Partial<TSecretSharing>, tx?: Knex) => {
try {
const now = new Date();
return await (tx || db)(TableName.SecretSharing)
return await (tx || db.replicaNode())(TableName.SecretSharing)
.where(filters)
.andWhere("expiresAt", ">", now)
.andWhere("encryptedValue", "<>", "")
@@ -50,20 +50,19 @@ interface TSecretV2DalArg {
}
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 secretV2BridgeDALFactory = ({ db, keyStore }: TSecretV2DalArg) => {
const secretOrm = ormify(db, TableName.SecretV2);
const invalidateSecretCacheByProjectId = async (projectId: string) => {
const invalidateSecretCacheByProjectId = async (projectId: string, tx?: Knex) => {
const secretDalVersionKey = SecretServiceCacheKeys.getSecretDalVersion(projectId);
await keyStore.incrementBy(secretDalVersionKey, 1);
await keyStore.setExpiry(secretDalVersionKey, SECRET_DAL_VERSION_TTL);
await keyStore.pgIncrementBy(secretDalVersionKey, { incr: 1, tx, expiry: SECRET_DAL_VERSION_TTL });
};
const findOne = async (filter: Partial<TSecretsV2>, tx?: Knex) => {
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
.where(buildFindFilter(filter, TableName.SecretV2))
.leftJoin(
@@ -144,7 +143,7 @@ export const secretV2BridgeDALFactory = ({ db, keyStore }: TSecretV2DalArg) => {
const find = async (filter: TFindFilter<TSecretsV2>, opts: TFindOpt<TSecretsV2> = {}) => {
const { offset, limit, sort, tx } = opts;
try {
const query = (tx || db)(TableName.SecretV2)
const query = (tx || db.replicaNode())(TableName.SecretV2)
// eslint-disable-next-line @typescript-eslint/no-misused-promises
.where(buildFindFilter(filter))
.leftJoin(
@@ -888,13 +887,13 @@ export const secretV2BridgeDALFactory = ({ db, keyStore }: TSecretV2DalArg) => {
const findSecretsWithReminderRecipients = async (ids: string[], limit: number, tx?: Knex) => {
try {
// 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)
.limit(limit)
.select("id");
// 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)
.leftJoin(TableName.Reminder, `${TableName.SecretV2}.id`, `${TableName.Reminder}.secretId`)
.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) => {
try {
// 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)
.limit(limit)
.select("id");
// 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)
.leftJoin(TableName.Reminder, `${TableName.SecretV2}.id`, `${TableName.Reminder}.secretId`)
.leftJoin(
@@ -118,7 +118,7 @@ type TSecretV2BridgeServiceFactoryDep = {
>;
snapshotService: Pick<TSecretSnapshotServiceFactory, "performSnapshot">;
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">;
};
@@ -360,6 +360,7 @@ export const secretV2BridgeServiceFactory = ({
tx
});
await secretDAL.invalidateSecretCacheByProjectId(projectId, tx);
return createdSecret;
});
@@ -377,7 +378,6 @@ export const secretV2BridgeServiceFactory = ({
});
}
await secretDAL.invalidateSecretCacheByProjectId(projectId);
if (inputSecret.type === SecretType.Shared) {
await snapshotService.performSnapshot(folderId);
await secretQueueService.syncSecrets({
@@ -566,8 +566,8 @@ export const secretV2BridgeServiceFactory = ({
await $validateSecretReferences(projectId, permission, allSecretReferences);
}
const updatedSecret = await secretDAL.transaction(async (tx) =>
fnSecretBulkUpdate({
const updatedSecret = await secretDAL.transaction(async (tx) => {
const modifiedSecretsInDB = await fnSecretBulkUpdate({
folderId,
orgId: actorOrgId,
resourceMetadataDAL,
@@ -598,8 +598,11 @@ export const secretV2BridgeServiceFactory = ({
actorId
},
tx
})
);
});
await secretDAL.invalidateSecretCacheByProjectId(projectId, tx);
return modifiedSecretsInDB;
});
if (inputSecret.secretReminderRepeatDays) {
await reminderService.createReminder({
actor,
@@ -615,7 +618,6 @@ export const secretV2BridgeServiceFactory = ({
});
}
await secretDAL.invalidateSecretCacheByProjectId(projectId);
if (inputSecret.type === SecretType.Shared) {
await snapshotService.performSnapshot(folderId);
await secretQueueService.syncSecrets({
@@ -715,8 +717,8 @@ export const secretV2BridgeServiceFactory = ({
);
try {
const deletedSecret = await secretDAL.transaction(async (tx) =>
fnSecretBulkDelete({
const deletedSecret = await secretDAL.transaction(async (tx) => {
const modifiedSecretsInDB = await fnSecretBulkDelete({
projectId,
folderId,
actorId,
@@ -732,10 +734,11 @@ export const secretV2BridgeServiceFactory = ({
}
],
tx
})
);
});
await secretDAL.invalidateSecretCacheByProjectId(projectId, tx);
return modifiedSecretsInDB;
});
await secretDAL.invalidateSecretCacheByProjectId(projectId);
if (inputSecret.type === SecretType.Shared) {
await snapshotService.performSnapshot(folderId);
await secretQueueService.syncSecrets({
@@ -1027,7 +1030,7 @@ export const secretV2BridgeServiceFactory = ({
});
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 cacheKey = SecretServiceCacheKeys.getSecretsOfServiceLayer(projectId, secretDalVersion, {
...dto,
@@ -1692,7 +1695,7 @@ export const secretV2BridgeServiceFactory = ({
await kmsService.createCipherPairWithDataKey({ type: KmsDataKey.SecretManager, projectId });
const executeBulkInsert = async (tx: Knex) => {
return fnSecretBulkInsert({
const modifiedSecretsInDB = await fnSecretBulkInsert({
inputSecrets: inputSecrets.map((el) => {
const references = secretReferencesGroupByInputSecretKey[el.secretKey]?.nestedReferences;
@@ -1728,13 +1731,14 @@ export const secretV2BridgeServiceFactory = ({
},
tx
});
await secretDAL.invalidateSecretCacheByProjectId(projectId, tx);
return modifiedSecretsInDB;
};
const newSecrets = providedTx
? await executeBulkInsert(providedTx)
: await secretDAL.transaction(executeBulkInsert);
await secretDAL.invalidateSecretCacheByProjectId(projectId);
await snapshotService.performSnapshot(folderId);
await secretQueueService.syncSecrets({
actor,
@@ -2099,6 +2103,7 @@ export const secretV2BridgeServiceFactory = ({
}
}
await secretDAL.invalidateSecretCacheByProjectId(projectId, tx);
return updatedSecrets;
};
@@ -2106,7 +2111,6 @@ export const secretV2BridgeServiceFactory = ({
? await executeBulkUpdate(providedTx)
: 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) =>
@@ -2233,7 +2237,7 @@ export const secretV2BridgeServiceFactory = ({
});
const executeBulkDelete = async (tx: Knex) => {
return fnSecretBulkDelete({
const modifiedSecretsInDB = await fnSecretBulkDelete({
secretDAL,
secretQueueService,
folderCommitService,
@@ -2249,6 +2253,8 @@ export const secretV2BridgeServiceFactory = ({
commitChanges,
tx
});
await secretDAL.invalidateSecretCacheByProjectId(projectId, tx);
return modifiedSecretsInDB;
};
try {
@@ -2256,7 +2262,6 @@ export const secretV2BridgeServiceFactory = ({
? await executeBulkDelete(providedTx)
: await secretDAL.transaction(executeBulkDelete);
await secretDAL.invalidateSecretCacheByProjectId(projectId);
await snapshotService.performSnapshot(folderId);
await secretQueueService.syncSecrets({
actor,
@@ -72,7 +72,7 @@ export const secretVersionV2BridgeDALFactory = (db: TDbClient) => {
.where(`${TableName.SecretVersionV2}.folderId`, folderId)
.join(TableName.SecretV2, `${TableName.SecretV2}.id`, `${TableName.SecretVersionV2}.secretId`)
.join<TSecretVersionsV2, TSecretVersionsV2 & { secretId: string; max: number }>(
(tx || db)(TableName.SecretVersionV2)
(tx || db.replicaNode())(TableName.SecretVersionV2)
.where(`${TableName.SecretVersionV2}.folderId`, folderId)
.groupBy("secretId")
.max("version")
@@ -121,7 +121,7 @@ export const secretVersionV2BridgeDALFactory = (db: TDbClient) => {
.where("folderId", folderId)
.whereIn(`${TableName.SecretVersionV2}.secretId`, secretIds)
.join(
(tx || db)(TableName.SecretVersionV2)
(tx || db.replicaNode())(TableName.SecretVersionV2)
.groupBy("secretId")
.max("version")
.select("secretId")
@@ -189,7 +189,7 @@ export const secretVersionV2BridgeDALFactory = (db: TDbClient) => {
}) => {
try {
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.ProjectMembership,
@@ -11,7 +11,7 @@ export const superAdminDALFactory = (db: TDbClient) => {
const superAdminOrm = ormify(db, TableName.SuperAdmin);
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)
.leftJoin(TableName.Organization, `${TableName.SuperAdmin}.defaultAuthOrgId`, `${TableName.Organization}.id`)
.leftJoin(TableName.SamlConfig, (qb) => {
+6 -2
View File
@@ -19,10 +19,14 @@ export type TUserDALFactory = ReturnType<typeof userDALFactory>;
export const userDALFactory = (db: TDbClient) => {
const userOrm = ormify(db, TableName.Users);
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) =>
(tx || db)(TableName.Users).whereRaw('lower("email") = :email', { email: email.toLowerCase() }).where({
(tx || db.replicaNode())(TableName.Users)
.whereRaw('lower("email") = :email', { email: email.toLowerCase() })
.where({
isEmailVerified: true
});
@@ -1,6 +1,7 @@
import { FunctionComponent, ReactNode } from "react";
import { BoundCanProps, Can } from "@casl/react";
import { TooltipProps } from "@app/components/v2/Tooltip/Tooltip";
import { TOrgPermission, useOrgPermission } from "@app/context/OrgPermissionContext";
import { AccessRestrictedBanner, Tooltip } from "../v2";
@@ -20,6 +21,7 @@ type Props = {
renderTooltip?: boolean;
allowedLabel?: string;
renderGuardBanner?: boolean;
tooltipProps?: Omit<TooltipProps, "children">;
} & BoundCanProps<TOrgPermission>;
export const OrgPermissionCan: FunctionComponent<Props> = ({
@@ -29,6 +31,7 @@ export const OrgPermissionCan: FunctionComponent<Props> = ({
renderTooltip,
allowedLabel,
renderGuardBanner,
tooltipProps,
...props
}) => {
const { permission } = useOrgPermission();
@@ -43,11 +46,19 @@ export const OrgPermissionCan: FunctionComponent<Props> = ({
: children;
if (!isAllowed && passThrough) {
return <Tooltip content={label}>{finalChild}</Tooltip>;
return (
<Tooltip content={label} {...tooltipProps}>
{finalChild}
</Tooltip>
);
}
if (isAllowed && renderTooltip && allowedLabel) {
return <Tooltip content={allowedLabel}>{finalChild}</Tooltip>;
return (
<Tooltip content={allowedLabel} {...tooltipProps}>
{finalChild}
</Tooltip>
);
}
if (!isAllowed && renderGuardBanner) {
@@ -3,7 +3,7 @@ import { useRouteContext } from "@tanstack/react-router";
import { fetchOrgSubscription, subscriptionQueryKeys } from "@app/hooks/api/subscriptions/queries";
export const useSubscription = () => {
export const useSubscription = (refreshCache?: boolean) => {
const organizationId = useRouteContext({
from: "/_authenticate/_inject-org-details",
select: (el) => el.organizationId
@@ -11,7 +11,7 @@ export const useSubscription = () => {
const { data: subscription } = useSuspenseQuery({
queryKey: subscriptionQueryKeys.getOrgSubsription(organizationId),
queryFn: () => fetchOrgSubscription(organizationId),
queryFn: () => fetchOrgSubscription(organizationId, refreshCache),
staleTime: Infinity
});
@@ -10,9 +10,9 @@ export const subscriptionQueryKeys = {
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 }>(
`/api/v1/organizations/${orgID}/plan`
`/api/v1/organizations/${orgID}/plan${refreshCache ? "?refreshCache=true" : ""}`
);
return data.plan;
@@ -54,5 +54,7 @@ export type SubscriptionPlan = {
secretScanning: boolean;
enterpriseSecretSyncs: boolean;
enterpriseAppConnections: boolean;
cardDeclined?: boolean;
cardDeclinedReason?: string;
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 { faCircleQuestion, faUserCircle } from "@fortawesome/free-regular-svg-icons";
import {
@@ -8,6 +8,7 @@ import {
faCaretDown,
faCheck,
faEnvelope,
faExclamationTriangle,
faInfo,
faInfoCircle,
faServer,
@@ -111,6 +112,14 @@ export const Navbar = () => {
const { subscription } = useSubscription();
const { currentOrg } = useOrganization();
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 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">
{getPlan(subscription)}
</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>
</Link>
<DropdownMenuTrigger asChild>
@@ -428,6 +450,49 @@ export const Navbar = () => {
</DropdownMenuItem>
</DropdownMenuContent>
</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}>
<ModalContent title="Server Administrators" subTitle="View all server administrators">
<div className="mb-2">
@@ -1,5 +1,7 @@
import { useEffect } from "react";
import { faArrowUpRightFromSquare } from "@fortawesome/free-solid-svg-icons";
import { FontAwesomeIcon } from "@fortawesome/react-fontawesome";
import { useQueryClient } from "@tanstack/react-query";
import { OrgPermissionCan } from "@app/components/permissions";
import { Button } from "@app/components/v2";
@@ -15,13 +17,15 @@ import {
useGetOrgPlanBillingInfo,
useGetOrgTrialUrl
} from "@app/hooks/api";
import { subscriptionQueryKeys } from "@app/hooks/api/subscriptions/queries";
import { usePopUp } from "@app/hooks/usePopUp";
import { ManagePlansModal } from "./ManagePlansModal";
export const PreviewSection = () => {
const { currentOrg } = useOrganization();
const { subscription } = useSubscription();
const { subscription } = useSubscription(true);
const queryClient = useQueryClient();
const { data, isPending } = useGetOrgPlanBillingInfo(currentOrg?.id ?? "");
const getOrgTrialUrl = useGetOrgTrialUrl();
const createCustomerPortalSession = useCreateCustomerPortalSession();
@@ -37,6 +41,12 @@ export const PreviewSection = () => {
return formattedTotal;
};
useEffect(() => {
queryClient.invalidateQueries({
queryKey: subscriptionQueryKeys.getOrgSubsription(currentOrg?.id ?? "")
});
}, []);
const formatDate = (date: number) => {
const createdDate = new Date(date * 1000);
const day: number = createdDate.getDate();
@@ -24,11 +24,17 @@ enum EnforceAuthType {
export const OrgGeneralAuthSection = ({
isSamlConfigured,
isOidcConfigured,
isGoogleConfigured
isGoogleConfigured,
isSamlActive,
isOidcActive,
isLdapActive
}: {
isSamlConfigured: boolean;
isOidcConfigured: boolean;
isGoogleConfigured: boolean;
isSamlActive: boolean;
isOidcActive: boolean;
isLdapActive: boolean;
}) => {
const { currentOrg } = useOrganization();
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 (
<div className="rounded-lg border border-mineshaft-600 bg-mineshaft-900 p-6">
<div>
@@ -135,7 +150,7 @@ export const OrgGeneralAuthSection = ({
</p>
</div>
<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="flex items-center gap-1">
<span className="text-md text-mineshaft-100">Enforce SAML SSO</span>
@@ -160,7 +175,7 @@ export const OrgGeneralAuthSection = ({
</p>
</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="flex items-center gap-1">
<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="mb-2 flex justify-between">
<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>
<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) => (
<div>
<Switch
id="enforce-google-sso"
onCheckedChange={(value) =>
handleEnforceOrgAuthToggle(value, EnforceAuthType.GOOGLE)
}
isChecked={currentOrg?.googleSsoAuthEnforced ?? false}
isDisabled={!isAllowed || currentOrg?.authEnforced}
isDisabled={
!isAllowed ||
currentOrg?.authEnforced ||
isOidcActive ||
isSamlActive ||
isLdapActive
}
/>
</div>
)}
</OrgPermissionCan>
</div>
<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 />
When this is enabled your organization members will only be able to login with Google
SSO (not Google SAML).
OAuth (not Google SAML).
</p>
</div>
</div>
@@ -267,8 +303,8 @@ export const OrgGeneralAuthSection = ({
</div>
<p className="text-sm text-mineshaft-300">
<span>
Allow organization admins to bypass SAML enforcement when SSO is unavailable,
misconfigured, or inaccessible.
Allow organization admins to bypass SSO login enforcement when your SSO provider is
unavailable, misconfigured, or inaccessible.
</span>
</p>
</div>
@@ -94,6 +94,8 @@ export const OrgLDAPSection = (): JSX.Element => {
handlePopUpOpen("ldapGroupMap");
};
const isGoogleOAuthEnabled = currentOrg.googleSsoAuthEnforced;
return (
<div className="mb-4">
<div className="py-4">
@@ -116,16 +118,31 @@ export const OrgLDAPSection = (): JSX.Element => {
<div className="pt-4">
<div className="mb-2 flex items-center justify-between">
<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) => (
<div>
<Switch
id="enable-saml-sso"
id="enable-ldap-sso"
onCheckedChange={(value) => handleLDAPToggle(value)}
isChecked={data ? data.isActive : false}
isDisabled={!isAllowed}
isDisabled={!isAllowed || isGoogleOAuthEnabled}
>
Enable
</Switch>
</div>
)}
</OrgPermissionCan>
</div>
@@ -83,6 +83,8 @@ export const OrgOIDCSection = (): JSX.Element => {
}
};
const isGoogleOAuthEnabled = currentOrg.googleSsoAuthEnforced;
return (
<div className="mb-4 rounded-lg border-mineshaft-600 bg-mineshaft-900">
<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">
<h2 className="text-md text-mineshaft-100">Enable OIDC</h2>
{!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) => (
<div>
<Switch
id="enable-oidc-sso"
onCheckedChange={(value) => handleOIDCToggle(value)}
isChecked={data ? data.isActive : false}
isDisabled={!isAllowed}
isDisabled={!isAllowed || isGoogleOAuthEnabled}
/>
</div>
)}
</OrgPermissionCan>
)}
@@ -78,6 +78,8 @@ export const OrgSSOSection = (): JSX.Element => {
}
};
const isGoogleOAuthEnabled = currentOrg.googleSsoAuthEnforced;
return (
<div className="space-y-4">
<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">
<h2 className="text-md text-mineshaft-100">Enable SAML</h2>
{!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) => (
<div>
<Switch
id="enable-saml-sso"
onCheckedChange={(value) => handleSamlSSOToggle(value)}
isChecked={data ? data.isActive : false}
isDisabled={!isAllowed}
isDisabled={!isAllowed || isGoogleOAuthEnabled}
/>
</div>
)}
</OrgPermissionCan>
)}
@@ -184,6 +184,9 @@ export const OrgSsoTab = withPermission(
isSamlConfigured={isSamlConfigured}
isOidcConfigured={isOidcConfigured}
isGoogleConfigured={isGoogleConfigured}
isSamlActive={Boolean(samlConfig?.isActive)}
isOidcActive={Boolean(oidcConfig?.isActive)}
isLdapActive={Boolean(ldapConfig?.isActive)}
/>
)}
@@ -19,6 +19,7 @@ import {
SecretInput,
Select,
SelectItem,
Switch,
TextArea,
Tooltip
} from "@app/components/v2";
@@ -66,6 +67,7 @@ const formSchema = z.object({
creationStatement: z.string().min(1),
revocationStatement: z.string().min(1),
renewStatement: z.string().optional(),
sslEnabled: z.boolean().optional(),
ca: z.string().optional(),
gatewayId: z.string().optional()
}),
@@ -200,6 +202,7 @@ export const SqlDatabaseInputForm = ({
const createDynamicSecret = useCreateDynamicSecret();
const { data: gateways, isPending: isGatewaysLoading } = useQuery(gatewaysQueryKeys.list());
const selectedClient = watch("provider.client");
const handleCreateDynamicSecret = async ({
name,
@@ -458,6 +461,27 @@ export const SqlDatabaseInputForm = ({
/>
</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
control={control}
name="provider.ca"
@@ -18,6 +18,7 @@ import {
SecretInput,
Select,
SelectItem,
Switch,
TextArea,
Tooltip
} from "@app/components/v2";
@@ -63,6 +64,7 @@ const formSchema = z.object({
creationStatement: z.string().min(1),
revocationStatement: z.string().min(1),
renewStatement: z.string().optional(),
sslEnabled: z.boolean().optional(),
ca: z.string().optional(),
gatewayId: z.string().optional().nullable()
})
@@ -151,6 +153,7 @@ export const EditDynamicSecretSqlProviderForm = ({
});
const { data: gateways, isPending: isGatewaysLoading } = useQuery(gatewaysQueryKeys.list());
const selectedClient = watch("inputs.client");
const updateDynamicSecret = useUpdateDynamicSecret();
const selectedGatewayId = watch("inputs.gatewayId");
@@ -407,6 +410,27 @@ export const EditDynamicSecretSqlProviderForm = ({
/>
</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
control={control}
name="inputs.ca"