Make sql secret rotation use gateway

This commit is contained in:
x032205
2025-07-16 02:24:16 -04:00
parent aab204a68a
commit fce6738562
8 changed files with 85 additions and 58 deletions
@@ -107,6 +107,7 @@ export type TSecretRotationV2ServiceFactoryDep = {
queueService: Pick<TQueueServiceFactory, "queuePg">; queueService: Pick<TQueueServiceFactory, "queuePg">;
appConnectionDAL: Pick<TAppConnectionDALFactory, "findById" | "update" | "updateById">; appConnectionDAL: Pick<TAppConnectionDALFactory, "findById" | "update" | "updateById">;
folderCommitService: Pick<TFolderCommitServiceFactory, "createCommit">; folderCommitService: Pick<TFolderCommitServiceFactory, "createCommit">;
gatewayService: Pick<TGatewayServiceFactory, "fnGetGatewayClientTlsByGatewayId">;
}; };
export type TSecretRotationV2ServiceFactory = ReturnType<typeof secretRotationV2ServiceFactory>; export type TSecretRotationV2ServiceFactory = ReturnType<typeof secretRotationV2ServiceFactory>;
@@ -148,7 +149,8 @@ export const secretRotationV2ServiceFactory = ({
keyStore, keyStore,
queueService, queueService,
folderCommitService, folderCommitService,
appConnectionDAL appConnectionDAL,
gatewayService
}: TSecretRotationV2ServiceFactoryDep) => { }: TSecretRotationV2ServiceFactoryDep) => {
const $queueSendSecretRotationStatusNotification = async (secretRotation: TSecretRotationV2Raw) => { const $queueSendSecretRotationStatusNotification = async (secretRotation: TSecretRotationV2Raw) => {
const appCfg = getConfig(); const appCfg = getConfig();
@@ -461,7 +463,8 @@ export const secretRotationV2ServiceFactory = ({
rotationInterval: payload.rotationInterval rotationInterval: payload.rotationInterval
} as TSecretRotationV2WithConnection, } as TSecretRotationV2WithConnection,
appConnectionDAL, appConnectionDAL,
kmsService kmsService,
gatewayService
); );
// even though we have a db constraint we want to check before any rotation of credentials is attempted // even though we have a db constraint we want to check before any rotation of credentials is attempted
@@ -824,7 +827,8 @@ export const secretRotationV2ServiceFactory = ({
connection: appConnection connection: appConnection
} as TSecretRotationV2WithConnection, } as TSecretRotationV2WithConnection,
appConnectionDAL, appConnectionDAL,
kmsService kmsService,
gatewayService
); );
const generatedCredentials = await decryptSecretRotationCredentials({ const generatedCredentials = await decryptSecretRotationCredentials({
@@ -907,7 +911,8 @@ export const secretRotationV2ServiceFactory = ({
connection: appConnection connection: appConnection
} as TSecretRotationV2WithConnection, } as TSecretRotationV2WithConnection,
appConnectionDAL, appConnectionDAL,
kmsService kmsService,
gatewayService
); );
const updatedRotation = await rotationFactory.rotateCredentials( const updatedRotation = await rotationFactory.rotateCredentials(
@@ -1,4 +1,5 @@
import { AuditLogInfo } from "@app/ee/services/audit-log/audit-log-types"; import { AuditLogInfo } from "@app/ee/services/audit-log/audit-log-types";
import { TGatewayServiceFactory } from "@app/ee/services/gateway/gateway-service";
import { TSqlCredentialsRotationGeneratedCredentials } from "@app/ee/services/secret-rotation-v2/shared/sql-credentials/sql-credentials-rotation-types"; import { TSqlCredentialsRotationGeneratedCredentials } from "@app/ee/services/secret-rotation-v2/shared/sql-credentials/sql-credentials-rotation-types";
import { OrderByDirection } from "@app/lib/types"; import { OrderByDirection } from "@app/lib/types";
import { TAppConnectionDALFactory } from "@app/services/app-connection/app-connection-dal"; import { TAppConnectionDALFactory } from "@app/services/app-connection/app-connection-dal";
@@ -239,7 +240,8 @@ export type TRotationFactory<
> = ( > = (
secretRotation: T, secretRotation: T,
appConnectionDAL: Pick<TAppConnectionDALFactory, "findById" | "update" | "updateById">, appConnectionDAL: Pick<TAppConnectionDALFactory, "findById" | "update" | "updateById">,
kmsService: Pick<TKmsServiceFactory, "createCipherPairWithDataKey"> kmsService: Pick<TKmsServiceFactory, "createCipherPairWithDataKey">,
gatewayService: Pick<TGatewayServiceFactory, "fnGetGatewayClientTlsByGatewayId">
) => { ) => {
issueCredentials: TRotationFactoryIssueCredentials<C, P>; issueCredentials: TRotationFactoryIssueCredentials<C, P>;
revokeCredentials: TRotationFactoryRevokeCredentials<C>; revokeCredentials: TRotationFactoryRevokeCredentials<C>;
@@ -1,3 +1,5 @@
import { Knex } from "knex";
import { import {
TRotationFactory, TRotationFactory,
TRotationFactoryGetSecretsPayload, TRotationFactoryGetSecretsPayload,
@@ -5,7 +7,10 @@ import {
TRotationFactoryRevokeCredentials, TRotationFactoryRevokeCredentials,
TRotationFactoryRotateCredentials TRotationFactoryRotateCredentials
} from "@app/ee/services/secret-rotation-v2/secret-rotation-v2-types"; } from "@app/ee/services/secret-rotation-v2/secret-rotation-v2-types";
import { getSqlConnectionClient, SQL_CONNECTION_ALTER_LOGIN_STATEMENT } from "@app/services/app-connection/shared/sql"; import {
executeWithPotentialGateway,
SQL_CONNECTION_ALTER_LOGIN_STATEMENT
} from "@app/services/app-connection/shared/sql";
import { generatePassword } from "../utils"; import { generatePassword } from "../utils";
import { import {
@@ -30,7 +35,7 @@ const redactPasswords = (e: unknown, credentials: TSqlCredentialsRotationGenerat
export const sqlCredentialsRotationFactory: TRotationFactory< export const sqlCredentialsRotationFactory: TRotationFactory<
TSqlCredentialsRotationWithConnection, TSqlCredentialsRotationWithConnection,
TSqlCredentialsRotationGeneratedCredentials TSqlCredentialsRotationGeneratedCredentials
> = (secretRotation) => { > = (secretRotation, _appConnectionDAL, _kmsService, gatewayService) => {
const { const {
connection, connection,
parameters: { username1, username2 }, parameters: { username1, username2 },
@@ -38,29 +43,38 @@ export const sqlCredentialsRotationFactory: TRotationFactory<
secretsMapping secretsMapping
} = secretRotation; } = secretRotation;
const $validateCredentials = async (credentials: TSqlCredentialsRotationGeneratedCredentials[number]) => { const executeOperation = <T>(
const client = await getSqlConnectionClient({ operation: (client: Knex) => Promise<T>,
...connection, credentialsOverride?: TSqlCredentialsRotationGeneratedCredentials[number]
credentials: { ) => {
...connection.credentials, const finalCredentials = {
...credentials ...connection.credentials,
} ...credentialsOverride
}); };
return executeWithPotentialGateway(
{
...connection,
credentials: finalCredentials
},
gatewayService,
(client) => operation(client)
);
};
const $validateCredentials = async (credentials: TSqlCredentialsRotationGeneratedCredentials[number]) => {
try { try {
await client.raw("SELECT 1"); await executeOperation(async (client) => {
await client.raw("SELECT 1");
}, credentials);
} catch (error) { } catch (error) {
throw new Error(redactPasswords(error, [credentials])); throw new Error(redactPasswords(error, [credentials]));
} finally {
await client.destroy();
} }
}; };
const issueCredentials: TRotationFactoryIssueCredentials<TSqlCredentialsRotationGeneratedCredentials> = async ( const issueCredentials: TRotationFactoryIssueCredentials<TSqlCredentialsRotationGeneratedCredentials> = async (
callback callback
) => { ) => {
const client = await getSqlConnectionClient(connection);
// For SQL, since we get existing users, we change both their passwords // For SQL, since we get existing users, we change both their passwords
// on issue to invalidate their existing passwords // on issue to invalidate their existing passwords
const credentialsSet = [ const credentialsSet = [
@@ -69,15 +83,15 @@ export const sqlCredentialsRotationFactory: TRotationFactory<
]; ];
try { try {
await client.transaction(async (tx) => { await executeOperation(async (client) => {
for await (const credentials of credentialsSet) { await client.transaction(async (tx) => {
await tx.raw(...SQL_CONNECTION_ALTER_LOGIN_STATEMENT[connection.app](credentials)); for await (const credentials of credentialsSet) {
} await tx.raw(...SQL_CONNECTION_ALTER_LOGIN_STATEMENT[connection.app](credentials));
}
});
}); });
} catch (error) { } catch (error) {
throw new Error(redactPasswords(error, credentialsSet)); throw new Error(redactPasswords(error, credentialsSet));
} finally {
await client.destroy();
} }
for await (const credentials of credentialsSet) { for await (const credentials of credentialsSet) {
@@ -91,21 +105,19 @@ export const sqlCredentialsRotationFactory: TRotationFactory<
credentialsToRevoke, credentialsToRevoke,
callback callback
) => { ) => {
const client = await getSqlConnectionClient(connection);
const revokedCredentials = credentialsToRevoke.map(({ username }) => ({ username, password: generatePassword() })); const revokedCredentials = credentialsToRevoke.map(({ username }) => ({ username, password: generatePassword() }));
try { try {
await client.transaction(async (tx) => { await executeOperation(async (client) => {
for await (const credentials of revokedCredentials) { await client.transaction(async (tx) => {
// invalidate previous passwords for await (const credentials of revokedCredentials) {
await tx.raw(...SQL_CONNECTION_ALTER_LOGIN_STATEMENT[connection.app](credentials)); // invalidate previous passwords
} await tx.raw(...SQL_CONNECTION_ALTER_LOGIN_STATEMENT[connection.app](credentials));
}
});
}); });
} catch (error) { } catch (error) {
throw new Error(redactPasswords(error, revokedCredentials)); throw new Error(redactPasswords(error, revokedCredentials));
} finally {
await client.destroy();
} }
return callback(); return callback();
@@ -115,17 +127,15 @@ export const sqlCredentialsRotationFactory: TRotationFactory<
_, _,
callback callback
) => { ) => {
const client = await getSqlConnectionClient(connection);
// generate new password for the next active user // generate new password for the next active user
const credentials = { username: activeIndex === 0 ? username2 : username1, password: generatePassword() }; const credentials = { username: activeIndex === 0 ? username2 : username1, password: generatePassword() };
try { try {
await client.raw(...SQL_CONNECTION_ALTER_LOGIN_STATEMENT[connection.app](credentials)); await executeOperation(async (client) => {
await client.raw(...SQL_CONNECTION_ALTER_LOGIN_STATEMENT[connection.app](credentials));
});
} catch (error) { } catch (error) {
throw new Error(redactPasswords(error, [credentials])); throw new Error(redactPasswords(error, [credentials]));
} finally {
await client.destroy();
} }
await $validateCredentials(credentials); await $validateCredentials(credentials);
+2 -1
View File
@@ -1805,7 +1805,8 @@ export const registerRoutes = async (
snapshotService, snapshotService,
secretQueueService, secretQueueService,
queueService, queueService,
appConnectionDAL appConnectionDAL,
gatewayService
}); });
const certificateAuthorityService = certificateAuthorityServiceFactory({ const certificateAuthorityService = certificateAuthorityServiceFactory({
@@ -5,6 +5,7 @@ import {
validateOCIConnectionCredentials validateOCIConnectionCredentials
} from "@app/ee/services/app-connections/oci"; } from "@app/ee/services/app-connections/oci";
import { getOracleDBConnectionListItem, OracleDBConnectionMethod } from "@app/ee/services/app-connections/oracledb"; import { getOracleDBConnectionListItem, OracleDBConnectionMethod } from "@app/ee/services/app-connections/oracledb";
import { TGatewayServiceFactory } from "@app/ee/services/gateway/gateway-service";
import { TLicenseServiceFactory } from "@app/ee/services/license/license-service"; import { TLicenseServiceFactory } from "@app/ee/services/license/license-service";
import { crypto } from "@app/lib/crypto/cryptography"; import { crypto } from "@app/lib/crypto/cryptography";
import { BadRequestError } from "@app/lib/errors"; import { BadRequestError } from "@app/lib/errors";
@@ -193,7 +194,8 @@ export const decryptAppConnectionCredentials = async ({
}; };
export const validateAppConnectionCredentials = async ( export const validateAppConnectionCredentials = async (
appConnection: TAppConnectionConfig appConnection: TAppConnectionConfig,
gatewayService: Pick<TGatewayServiceFactory, "fnGetGatewayClientTlsByGatewayId">
): Promise<TAppConnection["credentials"]> => { ): Promise<TAppConnection["credentials"]> => {
const VALIDATE_APP_CONNECTION_CREDENTIALS_MAP: Record<AppConnection, TAppConnectionCredentialsValidator> = { const VALIDATE_APP_CONNECTION_CREDENTIALS_MAP: Record<AppConnection, TAppConnectionCredentialsValidator> = {
[AppConnection.AWS]: validateAwsConnectionCredentials as TAppConnectionCredentialsValidator, [AppConnection.AWS]: validateAwsConnectionCredentials as TAppConnectionCredentialsValidator,
@@ -232,7 +234,7 @@ export const validateAppConnectionCredentials = async (
[AppConnection.Bitbucket]: validateBitbucketConnectionCredentials as TAppConnectionCredentialsValidator [AppConnection.Bitbucket]: validateBitbucketConnectionCredentials as TAppConnectionCredentialsValidator
}; };
return VALIDATE_APP_CONNECTION_CREDENTIALS_MAP[appConnection.app](appConnection); return VALIDATE_APP_CONNECTION_CREDENTIALS_MAP[appConnection.app](appConnection, gatewayService);
}; };
export const getAppConnectionMethodName = (method: TAppConnection["method"]) => { export const getAppConnectionMethodName = (method: TAppConnection["method"]) => {
@@ -242,12 +242,15 @@ export const appConnectionServiceFactory = ({
"Failed to create app connection due to plan restriction. Upgrade plan to access enterprise app connections." "Failed to create app connection due to plan restriction. Upgrade plan to access enterprise app connections."
); );
const validatedCredentials = await validateAppConnectionCredentials({ const validatedCredentials = await validateAppConnectionCredentials(
app, {
credentials, app,
method, credentials,
orgId: actor.orgId method,
} as TAppConnectionConfig); orgId: actor.orgId
} as TAppConnectionConfig,
gatewayService
);
try { try {
const createConnection = async (connectionCredentials: TAppConnection["credentials"]) => { const createConnection = async (connectionCredentials: TAppConnection["credentials"]) => {
@@ -349,12 +352,15 @@ export const appConnectionServiceFactory = ({
} Connection with method ${getAppConnectionMethodName(method)}` } Connection with method ${getAppConnectionMethodName(method)}`
}); });
updatedCredentials = await validateAppConnectionCredentials({ updatedCredentials = await validateAppConnectionCredentials(
app, {
orgId: actor.orgId, app,
credentials, orgId: actor.orgId,
method credentials,
} as TAppConnectionConfig); method
} as TAppConnectionConfig,
gatewayService
);
if (!updatedCredentials) if (!updatedCredentials)
throw new BadRequestError({ message: "Unable to validate connection - check credentials" }); throw new BadRequestError({ message: "Unable to validate connection - check credentials" });
@@ -350,7 +350,8 @@ export type TListAwsConnectionIamUsers = {
}; };
export type TAppConnectionCredentialsValidator = ( export type TAppConnectionCredentialsValidator = (
appConnection: TAppConnectionConfig appConnection: TAppConnectionConfig,
gatewayService: Pick<TGatewayServiceFactory, "fnGetGatewayClientTlsByGatewayId">
) => Promise<TAppConnection["credentials"]>; ) => Promise<TAppConnection["credentials"]>;
export type TAppConnectionTransitionCredentialsToPlatform = ( export type TAppConnectionTransitionCredentialsToPlatform = (
@@ -100,7 +100,7 @@ export const getSqlConnectionClient = async (appConnection: Pick<TSqlConnection,
return client; return client;
}; };
const executeWithPotentialGateway = async <T>( export const executeWithPotentialGateway = async <T>(
config: TSqlConnectionConfig, config: TSqlConnectionConfig,
gatewayService: Pick<TGatewayServiceFactory, "fnGetGatewayClientTlsByGatewayId">, gatewayService: Pick<TGatewayServiceFactory, "fnGetGatewayClientTlsByGatewayId">,
operation: (client: Knex) => Promise<T> operation: (client: Knex) => Promise<T>
@@ -159,7 +159,7 @@ const executeWithPotentialGateway = async <T>(
}; };
export const validateSqlConnectionCredentials = async ( export const validateSqlConnectionCredentials = async (
config: TSqlConnectionConfig & { gatewayId?: string }, config: TSqlConnectionConfig,
gatewayService: Pick<TGatewayServiceFactory, "fnGetGatewayClientTlsByGatewayId"> gatewayService: Pick<TGatewayServiceFactory, "fnGetGatewayClientTlsByGatewayId">
) => { ) => {
try { try {