Fix KMS memory leak

Adds a clean up method because KMS clients like GCP use a persistent connection snd if not closed, will continue to eat up the memory.
This commit is contained in:
Maidul Islam
2025-04-28 07:48:31 -04:00
parent 13b7729af8
commit 694ab35f53
6 changed files with 71 additions and 20 deletions
@@ -83,18 +83,26 @@ export const externalKmsServiceFactory = ({
throw error; throw error;
}); });
// if missing kms key this generate a new kms key id and returns new provider input try {
const newProviderInput = await externalKms.generateInputKmsKey(); // if missing kms key this generate a new kms key id and returns new provider input
sanitizedProviderInput = JSON.stringify(newProviderInput); const newProviderInput = await externalKms.generateInputKmsKey();
sanitizedProviderInput = JSON.stringify(newProviderInput);
await externalKms.validateConnection(); await externalKms.validateConnection();
} finally {
await externalKms.cleanup();
}
} }
break; break;
case KmsProviders.Gcp: case KmsProviders.Gcp:
{ {
const externalKms = await GcpKmsProviderFactory({ inputs: provider.inputs }); const externalKms = await GcpKmsProviderFactory({ inputs: provider.inputs });
await externalKms.validateConnection(); try {
sanitizedProviderInput = JSON.stringify(provider.inputs); await externalKms.validateConnection();
sanitizedProviderInput = JSON.stringify(provider.inputs);
} finally {
await externalKms.cleanup();
}
} }
break; break;
default: default:
@@ -186,8 +194,12 @@ export const externalKmsServiceFactory = ({
); );
const updatedProviderInput = { ...decryptedProviderInput, ...provider.inputs }; const updatedProviderInput = { ...decryptedProviderInput, ...provider.inputs };
const externalKms = await AwsKmsProviderFactory({ inputs: updatedProviderInput }); const externalKms = await AwsKmsProviderFactory({ inputs: updatedProviderInput });
await externalKms.validateConnection(); try {
sanitizedProviderInput = JSON.stringify(updatedProviderInput); await externalKms.validateConnection();
sanitizedProviderInput = JSON.stringify(updatedProviderInput);
} finally {
await externalKms.cleanup();
}
} }
break; break;
case KmsProviders.Gcp: case KmsProviders.Gcp:
@@ -197,8 +209,12 @@ export const externalKmsServiceFactory = ({
); );
const updatedProviderInput = { ...decryptedProviderInput, ...provider.inputs }; const updatedProviderInput = { ...decryptedProviderInput, ...provider.inputs };
const externalKms = await GcpKmsProviderFactory({ inputs: updatedProviderInput }); const externalKms = await GcpKmsProviderFactory({ inputs: updatedProviderInput });
await externalKms.validateConnection(); try {
sanitizedProviderInput = JSON.stringify(updatedProviderInput); await externalKms.validateConnection();
sanitizedProviderInput = JSON.stringify(updatedProviderInput);
} finally {
await externalKms.cleanup();
}
} }
break; break;
default: default:
@@ -368,7 +384,11 @@ export const externalKmsServiceFactory = ({
const fetchGcpKeys = async ({ credential, gcpRegion }: Pick<TExternalKmsGcpSchema, "credential" | "gcpRegion">) => { const fetchGcpKeys = async ({ credential, gcpRegion }: Pick<TExternalKmsGcpSchema, "credential" | "gcpRegion">) => {
const externalKms = await GcpKmsProviderFactory({ inputs: { credential, gcpRegion, keyName: "" } }); const externalKms = await GcpKmsProviderFactory({ inputs: { credential, gcpRegion, keyName: "" } });
return externalKms.getKeysList(); try {
return await externalKms.getKeysList();
} finally {
await externalKms.cleanup();
}
}; };
return { return {
@@ -2,6 +2,8 @@ import { CreateKeyCommand, DecryptCommand, DescribeKeyCommand, EncryptCommand, K
import { AssumeRoleCommand, STSClient } from "@aws-sdk/client-sts"; import { AssumeRoleCommand, STSClient } from "@aws-sdk/client-sts";
import { randomUUID } from "crypto"; import { randomUUID } from "crypto";
import { logger } from "@app/lib/logger";
import { ExternalKmsAwsSchema, KmsAwsCredentialType, TExternalKmsAwsSchema, TExternalKmsProviderFns } from "./model"; import { ExternalKmsAwsSchema, KmsAwsCredentialType, TExternalKmsAwsSchema, TExternalKmsProviderFns } from "./model";
const getAwsKmsClient = async (providerInputs: TExternalKmsAwsSchema) => { const getAwsKmsClient = async (providerInputs: TExternalKmsAwsSchema) => {
@@ -102,10 +104,21 @@ export const AwsKmsProviderFactory = async ({ inputs }: AwsKmsProviderArgs): Pro
return { data: Buffer.from(decryptionCommand.Plaintext) }; return { data: Buffer.from(decryptionCommand.Plaintext) };
}; };
const cleanup = async () => {
try {
awsClient.destroy();
return true;
} catch (error) {
logger.error(error, "cleanup: failed to destroy AWS KMS client");
return false;
}
};
return { return {
generateInputKmsKey, generateInputKmsKey,
validateConnection, validateConnection,
encrypt, encrypt,
decrypt decrypt,
cleanup
}; };
}; };
@@ -45,6 +45,16 @@ export const GcpKmsProviderFactory = async ({ inputs }: GcpKmsProviderArgs): Pro
} }
}; };
const cleanup = async () => {
try {
await gcpKmsClient.close();
return true;
} catch (error) {
logger.error(error, "cleanup: failed to close GCP KMS client");
return false;
}
};
// Used when adding the KMS to fetch the list of keys in specified region // Used when adding the KMS to fetch the list of keys in specified region
const getKeysList = async () => { const getKeysList = async () => {
try { try {
@@ -108,6 +118,7 @@ export const GcpKmsProviderFactory = async ({ inputs }: GcpKmsProviderArgs): Pro
validateConnection, validateConnection,
getKeysList, getKeysList,
encrypt, encrypt,
decrypt decrypt,
cleanup
}; };
}; };
@@ -98,4 +98,5 @@ export type TExternalKmsProviderFns = {
validateConnection: () => Promise<boolean>; validateConnection: () => Promise<boolean>;
encrypt: (data: Buffer) => Promise<{ encryptedBlob: Buffer }>; encrypt: (data: Buffer) => Promise<{ encryptedBlob: Buffer }>;
decrypt: (encryptedBlob: Buffer) => Promise<{ data: Buffer }>; decrypt: (encryptedBlob: Buffer) => Promise<{ data: Buffer }>;
cleanup: () => Promise<boolean>;
}; };
@@ -12,6 +12,7 @@ import { generateSrpServerKey, srpCheckClientProof } from "@app/lib/crypto";
import { infisicalSymmetricEncypt } from "@app/lib/crypto/encryption"; import { infisicalSymmetricEncypt } from "@app/lib/crypto/encryption";
import { getUserPrivateKey } from "@app/lib/crypto/srp"; import { getUserPrivateKey } from "@app/lib/crypto/srp";
import { BadRequestError, DatabaseError, ForbiddenRequestError, UnauthorizedError } from "@app/lib/errors"; import { BadRequestError, DatabaseError, ForbiddenRequestError, UnauthorizedError } from "@app/lib/errors";
import { removeTrailingSlash } from "@app/lib/fn";
import { logger } from "@app/lib/logger"; import { logger } from "@app/lib/logger";
import { getUserAgentType } from "@app/server/plugins/audit-log"; import { getUserAgentType } from "@app/server/plugins/audit-log";
import { getServerCfg } from "@app/services/super-admin/super-admin-service"; import { getServerCfg } from "@app/services/super-admin/super-admin-service";
@@ -39,7 +40,6 @@ import {
AuthTokenType, AuthTokenType,
MfaMethod MfaMethod
} from "./auth-type"; } from "./auth-type";
import { removeTrailingSlash } from "@app/lib/fn";
type TAuthLoginServiceFactoryDep = { type TAuthLoginServiceFactoryDep = {
userDAL: TUserDALFactory; userDAL: TUserDALFactory;
+12 -6
View File
@@ -342,9 +342,12 @@ export const kmsServiceFactory = ({
} }
return async ({ cipherTextBlob }: Pick<TDecryptWithKmsDTO, "cipherTextBlob">) => { return async ({ cipherTextBlob }: Pick<TDecryptWithKmsDTO, "cipherTextBlob">) => {
const { data } = await externalKms.decrypt(cipherTextBlob); try {
const { data } = await externalKms.decrypt(cipherTextBlob);
return data; return data;
} finally {
await externalKms.cleanup();
}
}; };
} }
@@ -557,9 +560,12 @@ export const kmsServiceFactory = ({
} }
return async ({ plainText }: Pick<TEncryptWithKmsDTO, "plainText">) => { return async ({ plainText }: Pick<TEncryptWithKmsDTO, "plainText">) => {
const { encryptedBlob } = await externalKms.encrypt(plainText); try {
const { encryptedBlob } = await externalKms.encrypt(plainText);
return { cipherTextBlob: encryptedBlob }; return { cipherTextBlob: encryptedBlob };
} finally {
await externalKms.cleanup();
}
}; };
} }