Merge branch 'main' into ENG-3506-LDAP

This commit is contained in:
x032205
2025-09-15 11:52:18 -04:00
190 changed files with 7343 additions and 645 deletions
+14 -3
View File
@@ -36,11 +36,23 @@ jobs:
echo "Latest production tag: $LATEST_STABLE_TAG" echo "Latest production tag: $LATEST_STABLE_TAG"
# Extract version numbers and increment minor version
VERSION_NUMBERS=$(echo "$LATEST_STABLE_TAG" | sed 's/^v//')
MAJOR=$(echo "$VERSION_NUMBERS" | cut -d'.' -f1)
MINOR=$(echo "$VERSION_NUMBERS" | cut -d'.' -f2)
PATCH=$(echo "$VERSION_NUMBERS" | cut -d'.' -f3)
# Increment minor version, reset patch to 0
NEXT_MINOR=$((MINOR + 1))
NEXT_VERSION="v${MAJOR}.${NEXT_MINOR}.0"
echo "Next version for nightly: $NEXT_VERSION"
# Get current date in YYYYMMDD format # Get current date in YYYYMMDD format
DATE=$(date +%Y%m%d) DATE=$(date +%Y%m%d)
# Base nightly tag name # Base nightly tag name using next version
BASE_TAG="${LATEST_STABLE_TAG}-nightly-${DATE}" BASE_TAG="${NEXT_VERSION}-nightly-${DATE}"
# Check if this exact tag already exists # Check if this exact tag already exists
if git tag --list | grep -q "^${BASE_TAG}$"; then if git tag --list | grep -q "^${BASE_TAG}$"; then
@@ -65,7 +77,6 @@ jobs:
echo "Generated nightly tag: $NIGHTLY_TAG" echo "Generated nightly tag: $NIGHTLY_TAG"
echo "NIGHTLY_TAG=$NIGHTLY_TAG" >> $GITHUB_ENV echo "NIGHTLY_TAG=$NIGHTLY_TAG" >> $GITHUB_ENV
echo "LATEST_PRODUCTION_TAG=$LATEST_STABLE_TAG" >> $GITHUB_ENV
git tag "$NIGHTLY_TAG" git tag "$NIGHTLY_TAG"
git push origin "$NIGHTLY_TAG" git push origin "$NIGHTLY_TAG"
+9
View File
@@ -56,6 +56,15 @@ export const mockKeyStore = (): TKeyStoreFactory => {
incrementBy: async () => { incrementBy: async () => {
return 1; return 1;
}, },
pgGetIntItem: async (key) => {
const value = store[key];
if (typeof value === "number") {
return Number(value);
}
},
pgIncrementBy: async () => {
return 1;
},
getItems: async (keys) => { getItems: async (keys) => {
const values = keys.map((key) => { const values = keys.map((key) => {
const value = store[key]; const value = store[key];
+3 -1
View File
@@ -15,6 +15,7 @@ import { mockSmtpServer } from "./mocks/smtp";
import { initDbConnection } from "@app/db"; import { initDbConnection } from "@app/db";
import { queueServiceFactory } from "@app/queue"; import { queueServiceFactory } from "@app/queue";
import { keyStoreFactory } from "@app/keystore/keystore"; import { keyStoreFactory } from "@app/keystore/keystore";
import { keyValueStoreDALFactory } from "@app/keystore/key-value-store-dal";
import { initializeHsmModule } from "@app/ee/services/hsm/hsm-fns"; import { initializeHsmModule } from "@app/ee/services/hsm/hsm-fns";
import { buildRedisFromConfig } from "@app/lib/config/redis"; import { buildRedisFromConfig } from "@app/lib/config/redis";
import { superAdminDALFactory } from "@app/services/super-admin/super-admin-dal"; import { superAdminDALFactory } from "@app/services/super-admin/super-admin-dal";
@@ -62,7 +63,8 @@ export default {
const smtp = mockSmtpServer(); const smtp = mockSmtpServer();
const queue = queueServiceFactory(envCfg, { dbConnectionUrl: envCfg.DB_CONNECTION_URI }); const queue = queueServiceFactory(envCfg, { dbConnectionUrl: envCfg.DB_CONNECTION_URI });
const keyStore = keyStoreFactory(envCfg); const keyValueStoreDAL = keyValueStoreDALFactory(db);
const keyStore = keyStoreFactory(envCfg, keyValueStoreDAL);
await queue.initialize(); await queue.initialize();
+4
View File
@@ -16,6 +16,7 @@ import { TEventBusService } from "@app/ee/services/event/event-bus-service";
import { TServerSentEventsService } from "@app/ee/services/event/event-sse-service"; import { TServerSentEventsService } from "@app/ee/services/event/event-sse-service";
import { TExternalKmsServiceFactory } from "@app/ee/services/external-kms/external-kms-service"; import { TExternalKmsServiceFactory } from "@app/ee/services/external-kms/external-kms-service";
import { TGatewayServiceFactory } from "@app/ee/services/gateway/gateway-service"; import { TGatewayServiceFactory } from "@app/ee/services/gateway/gateway-service";
import { TGatewayV2ServiceFactory } from "@app/ee/services/gateway-v2/gateway-v2-service";
import { TGithubOrgSyncServiceFactory } from "@app/ee/services/github-org-sync/github-org-sync-service"; import { TGithubOrgSyncServiceFactory } from "@app/ee/services/github-org-sync/github-org-sync-service";
import { TGroupServiceFactory } from "@app/ee/services/group/group-service"; import { TGroupServiceFactory } from "@app/ee/services/group/group-service";
import { TIdentityAuthTemplateServiceFactory } from "@app/ee/services/identity-auth-template"; import { TIdentityAuthTemplateServiceFactory } from "@app/ee/services/identity-auth-template";
@@ -32,6 +33,7 @@ import { TPitServiceFactory } from "@app/ee/services/pit/pit-service";
import { TProjectTemplateServiceFactory } from "@app/ee/services/project-template/project-template-types"; import { TProjectTemplateServiceFactory } from "@app/ee/services/project-template/project-template-types";
import { TProjectUserAdditionalPrivilegeServiceFactory } from "@app/ee/services/project-user-additional-privilege/project-user-additional-privilege-types"; import { TProjectUserAdditionalPrivilegeServiceFactory } from "@app/ee/services/project-user-additional-privilege/project-user-additional-privilege-types";
import { RateLimitConfiguration, TRateLimitServiceFactory } from "@app/ee/services/rate-limit/rate-limit-types"; import { RateLimitConfiguration, TRateLimitServiceFactory } from "@app/ee/services/rate-limit/rate-limit-types";
import { TRelayServiceFactory } from "@app/ee/services/relay/relay-service";
import { TSamlConfigServiceFactory } from "@app/ee/services/saml-config/saml-config-types"; import { TSamlConfigServiceFactory } from "@app/ee/services/saml-config/saml-config-types";
import { TScimServiceFactory } from "@app/ee/services/scim/scim-types"; import { TScimServiceFactory } from "@app/ee/services/scim/scim-types";
import { TSecretApprovalPolicyServiceFactory } from "@app/ee/services/secret-approval-policy/secret-approval-policy-service"; import { TSecretApprovalPolicyServiceFactory } from "@app/ee/services/secret-approval-policy/secret-approval-policy-service";
@@ -296,6 +298,8 @@ declare module "fastify" {
secretRotationV2: TSecretRotationV2ServiceFactory; secretRotationV2: TSecretRotationV2ServiceFactory;
microsoftTeams: TMicrosoftTeamsServiceFactory; microsoftTeams: TMicrosoftTeamsServiceFactory;
assumePrivileges: TAssumePrivilegeServiceFactory; assumePrivileges: TAssumePrivilegeServiceFactory;
relay: TRelayServiceFactory;
gatewayV2: TGatewayV2ServiceFactory;
githubOrgSync: TGithubOrgSyncServiceFactory; githubOrgSync: TGithubOrgSyncServiceFactory;
folderCommit: TFolderCommitServiceFactory; folderCommit: TFolderCommitServiceFactory;
pit: TPitServiceFactory; pit: TPitServiceFactory;
+40
View File
@@ -101,6 +101,9 @@ import {
TGateways, TGateways,
TGatewaysInsert, TGatewaysInsert,
TGatewaysUpdate, TGatewaysUpdate,
TGatewaysV2,
TGatewaysV2Insert,
TGatewaysV2Update,
TGitAppInstallSessions, TGitAppInstallSessions,
TGitAppInstallSessionsInsert, TGitAppInstallSessionsInsert,
TGitAppInstallSessionsUpdate, TGitAppInstallSessionsUpdate,
@@ -179,6 +182,9 @@ import {
TIncidentContacts, TIncidentContacts,
TIncidentContactsInsert, TIncidentContactsInsert,
TIncidentContactsUpdate, TIncidentContactsUpdate,
TInstanceRelayConfig,
TInstanceRelayConfigInsert,
TInstanceRelayConfigUpdate,
TIntegrationAuths, TIntegrationAuths,
TIntegrationAuthsInsert, TIntegrationAuthsInsert,
TIntegrationAuthsUpdate, TIntegrationAuthsUpdate,
@@ -191,6 +197,9 @@ import {
TInternalKms, TInternalKms,
TInternalKmsInsert, TInternalKmsInsert,
TInternalKmsUpdate, TInternalKmsUpdate,
TKeyValueStore,
TKeyValueStoreInsert,
TKeyValueStoreUpdate,
TKmipClientCertificates, TKmipClientCertificates,
TKmipClientCertificatesInsert, TKmipClientCertificatesInsert,
TKmipClientCertificatesUpdate, TKmipClientCertificatesUpdate,
@@ -230,9 +239,15 @@ import {
TOrgGatewayConfig, TOrgGatewayConfig,
TOrgGatewayConfigInsert, TOrgGatewayConfigInsert,
TOrgGatewayConfigUpdate, TOrgGatewayConfigUpdate,
TOrgGatewayConfigV2,
TOrgGatewayConfigV2Insert,
TOrgGatewayConfigV2Update,
TOrgMemberships, TOrgMemberships,
TOrgMembershipsInsert, TOrgMembershipsInsert,
TOrgMembershipsUpdate, TOrgMembershipsUpdate,
TOrgRelayConfig,
TOrgRelayConfigInsert,
TOrgRelayConfigUpdate,
TOrgRoles, TOrgRoles,
TOrgRolesInsert, TOrgRolesInsert,
TOrgRolesUpdate, TOrgRolesUpdate,
@@ -290,6 +305,9 @@ import {
TRateLimit, TRateLimit,
TRateLimitInsert, TRateLimitInsert,
TRateLimitUpdate, TRateLimitUpdate,
TRelays,
TRelaysInsert,
TRelaysUpdate,
TResourceMetadata, TResourceMetadata,
TResourceMetadataInsert, TResourceMetadataInsert,
TResourceMetadataUpdate, TResourceMetadataUpdate,
@@ -1238,6 +1256,17 @@ declare module "knex/types/tables" {
TSecretScanningResourcesInsert, TSecretScanningResourcesInsert,
TSecretScanningResourcesUpdate TSecretScanningResourcesUpdate
>; >;
[TableName.InstanceRelayConfig]: KnexOriginal.CompositeTableType<
TInstanceRelayConfig,
TInstanceRelayConfigInsert,
TInstanceRelayConfigUpdate
>;
[TableName.OrgRelayConfig]: KnexOriginal.CompositeTableType<
TOrgRelayConfig,
TOrgRelayConfigInsert,
TOrgRelayConfigUpdate
>;
[TableName.Relay]: KnexOriginal.CompositeTableType<TRelays, TRelaysInsert, TRelaysUpdate>;
[TableName.SecretScanningScan]: KnexOriginal.CompositeTableType< [TableName.SecretScanningScan]: KnexOriginal.CompositeTableType<
TSecretScanningScans, TSecretScanningScans,
TSecretScanningScansInsert, TSecretScanningScansInsert,
@@ -1259,10 +1288,21 @@ declare module "knex/types/tables" {
TRemindersRecipientsInsert, TRemindersRecipientsInsert,
TRemindersRecipientsUpdate TRemindersRecipientsUpdate
>; >;
[TableName.OrgGatewayConfigV2]: KnexOriginal.CompositeTableType<
TOrgGatewayConfigV2,
TOrgGatewayConfigV2Insert,
TOrgGatewayConfigV2Update
>;
[TableName.GatewayV2]: KnexOriginal.CompositeTableType<TGatewaysV2, TGatewaysV2Insert, TGatewaysV2Update>;
[TableName.UserNotifications]: KnexOriginal.CompositeTableType< [TableName.UserNotifications]: KnexOriginal.CompositeTableType<
TUserNotifications, TUserNotifications,
TUserNotificationsInsert, TUserNotificationsInsert,
TUserNotificationsUpdate TUserNotificationsUpdate
>; >;
[TableName.KeyValueStore]: KnexOriginal.CompositeTableType<
TKeyValueStore,
TKeyValueStoreInsert,
TKeyValueStoreUpdate
>;
} }
} }
@@ -0,0 +1,150 @@
import { Knex } from "knex";
import { TableName } from "../schemas";
import { createOnUpdateTrigger, dropOnUpdateTrigger } from "../utils";
export async function up(knex: Knex): Promise<void> {
if (!(await knex.schema.hasTable(TableName.InstanceRelayConfig))) {
await knex.schema.createTable(TableName.InstanceRelayConfig, (t) => {
t.uuid("id", { primaryKey: true }).defaultTo(knex.fn.uuid());
t.timestamps(true, true, true);
// Root CA for relay PKI
t.binary("encryptedRootRelayPkiCaPrivateKey").notNullable();
t.binary("encryptedRootRelayPkiCaCertificate").notNullable();
// Instance CA for relay PKI
t.binary("encryptedInstanceRelayPkiCaPrivateKey").notNullable();
t.binary("encryptedInstanceRelayPkiCaCertificate").notNullable();
t.binary("encryptedInstanceRelayPkiCaCertificateChain").notNullable();
// Instance client/server intermediates for relay PKI
t.binary("encryptedInstanceRelayPkiClientCaPrivateKey").notNullable();
t.binary("encryptedInstanceRelayPkiClientCaCertificate").notNullable();
t.binary("encryptedInstanceRelayPkiClientCaCertificateChain").notNullable();
t.binary("encryptedInstanceRelayPkiServerCaPrivateKey").notNullable();
t.binary("encryptedInstanceRelayPkiServerCaCertificate").notNullable();
t.binary("encryptedInstanceRelayPkiServerCaCertificateChain").notNullable();
// Org Parent CAs for relay
t.binary("encryptedOrgRelayPkiCaPrivateKey").notNullable();
t.binary("encryptedOrgRelayPkiCaCertificate").notNullable();
t.binary("encryptedOrgRelayPkiCaCertificateChain").notNullable();
// Instance SSH CAs for relay
t.binary("encryptedInstanceRelaySshClientCaPrivateKey").notNullable();
t.binary("encryptedInstanceRelaySshClientCaPublicKey").notNullable();
t.binary("encryptedInstanceRelaySshServerCaPrivateKey").notNullable();
t.binary("encryptedInstanceRelaySshServerCaPublicKey").notNullable();
});
await createOnUpdateTrigger(knex, TableName.InstanceRelayConfig);
}
// Org-level relay configuration (one-to-one with organization)
if (!(await knex.schema.hasTable(TableName.OrgRelayConfig))) {
await knex.schema.createTable(TableName.OrgRelayConfig, (t) => {
t.uuid("id", { primaryKey: true }).defaultTo(knex.fn.uuid());
t.timestamps(true, true, true);
t.uuid("orgId").notNullable().unique();
t.foreign("orgId").references("id").inTable(TableName.Organization).onDelete("CASCADE");
// Org-scoped relay PKI (client + server)
t.binary("encryptedRelayPkiClientCaPrivateKey").notNullable();
t.binary("encryptedRelayPkiClientCaCertificate").notNullable();
t.binary("encryptedRelayPkiClientCaCertificateChain").notNullable();
t.binary("encryptedRelayPkiServerCaPrivateKey").notNullable();
t.binary("encryptedRelayPkiServerCaCertificate").notNullable();
t.binary("encryptedRelayPkiServerCaCertificateChain").notNullable();
// Org-scoped relay SSH (client + server)
t.binary("encryptedRelaySshClientCaPrivateKey").notNullable();
t.binary("encryptedRelaySshClientCaPublicKey").notNullable();
t.binary("encryptedRelaySshServerCaPrivateKey").notNullable();
t.binary("encryptedRelaySshServerCaPublicKey").notNullable();
});
await createOnUpdateTrigger(knex, TableName.OrgRelayConfig);
}
if (!(await knex.schema.hasTable(TableName.OrgGatewayConfigV2))) {
await knex.schema.createTable(TableName.OrgGatewayConfigV2, (t) => {
t.uuid("id", { primaryKey: true }).defaultTo(knex.fn.uuid());
t.uuid("orgId").notNullable().unique();
t.foreign("orgId").references("id").inTable(TableName.Organization).onDelete("CASCADE");
t.timestamps(true, true, true);
t.binary("encryptedRootGatewayCaPrivateKey").notNullable();
t.binary("encryptedRootGatewayCaCertificate").notNullable();
t.binary("encryptedGatewayServerCaPrivateKey").notNullable();
t.binary("encryptedGatewayServerCaCertificate").notNullable();
t.binary("encryptedGatewayServerCaCertificateChain").notNullable();
t.binary("encryptedGatewayClientCaPrivateKey").notNullable();
t.binary("encryptedGatewayClientCaCertificate").notNullable();
t.binary("encryptedGatewayClientCaCertificateChain").notNullable();
});
await createOnUpdateTrigger(knex, TableName.OrgGatewayConfigV2);
}
if (!(await knex.schema.hasTable(TableName.Relay))) {
await knex.schema.createTable(TableName.Relay, (t) => {
t.uuid("id", { primaryKey: true }).defaultTo(knex.fn.uuid());
t.timestamps(true, true, true);
t.uuid("orgId");
t.foreign("orgId").references("id").inTable(TableName.Organization).onDelete("CASCADE");
t.uuid("identityId");
t.foreign("identityId").references("id").inTable(TableName.Identity).onDelete("CASCADE");
t.string("name").notNullable();
t.string("host").notNullable();
t.unique(["orgId", "name"]);
});
await createOnUpdateTrigger(knex, TableName.Relay);
}
if (!(await knex.schema.hasTable(TableName.GatewayV2))) {
await knex.schema.createTable(TableName.GatewayV2, (t) => {
t.uuid("id", { primaryKey: true }).defaultTo(knex.fn.uuid());
t.timestamps(true, true, true);
t.uuid("orgId").notNullable();
t.foreign("orgId").references("id").inTable(TableName.Organization).onDelete("CASCADE");
t.uuid("identityId").notNullable().unique();
t.foreign("identityId").references("id").inTable(TableName.Identity).onDelete("CASCADE");
t.uuid("relayId");
t.foreign("relayId").references("id").inTable(TableName.Relay).onDelete("SET NULL");
t.string("name").notNullable();
t.unique(["orgId", "name"]);
t.dateTime("heartbeat");
});
await createOnUpdateTrigger(knex, TableName.GatewayV2);
}
}
export async function down(knex: Knex): Promise<void> {
await dropOnUpdateTrigger(knex, TableName.OrgRelayConfig);
await knex.schema.dropTableIfExists(TableName.OrgRelayConfig);
await dropOnUpdateTrigger(knex, TableName.InstanceRelayConfig);
await knex.schema.dropTableIfExists(TableName.InstanceRelayConfig);
await dropOnUpdateTrigger(knex, TableName.OrgGatewayConfigV2);
await knex.schema.dropTableIfExists(TableName.OrgGatewayConfigV2);
await dropOnUpdateTrigger(knex, TableName.GatewayV2);
await knex.schema.dropTableIfExists(TableName.GatewayV2);
await dropOnUpdateTrigger(knex, TableName.Relay);
await knex.schema.dropTableIfExists(TableName.Relay);
}
@@ -0,0 +1,33 @@
import { Knex } from "knex";
import { TableName } from "../schemas";
export async function up(knex: Knex): Promise<void> {
if (!(await knex.schema.hasColumn(TableName.DynamicSecret, "gatewayV2Id"))) {
await knex.schema.alterTable(TableName.DynamicSecret, (table) => {
table.uuid("gatewayV2Id");
table.foreign("gatewayV2Id").references("id").inTable(TableName.GatewayV2).onDelete("SET NULL");
});
}
if (!(await knex.schema.hasColumn(TableName.IdentityKubernetesAuth, "gatewayV2Id"))) {
await knex.schema.alterTable(TableName.IdentityKubernetesAuth, (table) => {
table.uuid("gatewayV2Id");
table.foreign("gatewayV2Id").references("id").inTable(TableName.GatewayV2).onDelete("SET NULL");
});
}
}
export async function down(knex: Knex): Promise<void> {
if (await knex.schema.hasColumn(TableName.DynamicSecret, "gatewayV2Id")) {
await knex.schema.alterTable(TableName.DynamicSecret, (table) => {
table.dropColumn("gatewayV2Id");
});
}
if (await knex.schema.hasColumn(TableName.IdentityKubernetesAuth, "gatewayV2Id")) {
await knex.schema.alterTable(TableName.IdentityKubernetesAuth, (table) => {
table.dropColumn("gatewayV2Id");
});
}
}
@@ -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);
}
@@ -0,0 +1,23 @@
import { Knex } from "knex";
import { TableName } from "../schemas";
export async function up(knex: Knex): Promise<void> {
const hasPayloadCol = await knex.schema.hasColumn(TableName.AuthTokens, "payload");
if (!hasPayloadCol) {
await knex.schema.alterTable(TableName.AuthTokens, (t) => {
t.text("payload").nullable();
});
}
}
export async function down(knex: Knex): Promise<void> {
const hasPayloadCol = await knex.schema.hasColumn(TableName.AuthTokens, "payload");
if (hasPayloadCol) {
await knex.schema.alterTable(TableName.AuthTokens, (t) => {
t.dropColumn("payload");
});
}
}
+2 -1
View File
@@ -18,7 +18,8 @@ export const AuthTokensSchema = z.object({
updatedAt: z.date(), updatedAt: z.date(),
userId: z.string().uuid().nullable().optional(), userId: z.string().uuid().nullable().optional(),
orgId: z.string().uuid().nullable().optional(), orgId: z.string().uuid().nullable().optional(),
aliasId: z.string().nullable().optional() aliasId: z.string().nullable().optional(),
payload: z.string().nullable().optional()
}); });
export type TAuthTokens = z.infer<typeof AuthTokensSchema>; export type TAuthTokens = z.infer<typeof AuthTokensSchema>;
+2 -1
View File
@@ -29,7 +29,8 @@ export const DynamicSecretsSchema = z.object({
encryptedInput: zodBuffer, encryptedInput: zodBuffer,
projectGatewayId: z.string().uuid().nullable().optional(), projectGatewayId: z.string().uuid().nullable().optional(),
gatewayId: z.string().uuid().nullable().optional(), gatewayId: z.string().uuid().nullable().optional(),
usernameTemplate: z.string().nullable().optional() usernameTemplate: z.string().nullable().optional(),
gatewayV2Id: z.string().uuid().nullable().optional()
}); });
export type TDynamicSecrets = z.infer<typeof DynamicSecretsSchema>; export type TDynamicSecrets = z.infer<typeof DynamicSecretsSchema>;
+23
View File
@@ -0,0 +1,23 @@
// 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 GatewaysV2Schema = z.object({
id: z.string().uuid(),
createdAt: z.date(),
updatedAt: z.date(),
orgId: z.string().uuid(),
identityId: z.string().uuid(),
relayId: z.string().uuid().nullable().optional(),
name: z.string(),
heartbeat: z.date().nullable().optional()
});
export type TGatewaysV2 = z.infer<typeof GatewaysV2Schema>;
export type TGatewaysV2Insert = Omit<z.input<typeof GatewaysV2Schema>, TImmutableDBKeys>;
export type TGatewaysV2Update = Partial<Omit<z.input<typeof GatewaysV2Schema>, TImmutableDBKeys>>;
@@ -32,7 +32,8 @@ export const IdentityKubernetesAuthsSchema = z.object({
encryptedKubernetesCaCertificate: zodBuffer.nullable().optional(), encryptedKubernetesCaCertificate: zodBuffer.nullable().optional(),
gatewayId: z.string().uuid().nullable().optional(), gatewayId: z.string().uuid().nullable().optional(),
accessTokenPeriod: z.coerce.number().default(0), accessTokenPeriod: z.coerce.number().default(0),
tokenReviewMode: z.string().default("api") tokenReviewMode: z.string().default("api"),
gatewayV2Id: z.string().uuid().nullable().optional()
}); });
export type TIdentityKubernetesAuths = z.infer<typeof IdentityKubernetesAuthsSchema>; export type TIdentityKubernetesAuths = z.infer<typeof IdentityKubernetesAuthsSchema>;
+6
View File
@@ -31,6 +31,7 @@ export * from "./folder-commits";
export * from "./folder-tree-checkpoint-resources"; export * from "./folder-tree-checkpoint-resources";
export * from "./folder-tree-checkpoints"; export * from "./folder-tree-checkpoints";
export * from "./gateways"; export * from "./gateways";
export * from "./gateways-v2";
export * from "./git-app-install-sessions"; export * from "./git-app-install-sessions";
export * from "./git-app-org"; export * from "./git-app-org";
export * from "./github-org-sync-configs"; export * from "./github-org-sync-configs";
@@ -57,10 +58,12 @@ export * from "./identity-token-auths";
export * from "./identity-ua-client-secrets"; export * from "./identity-ua-client-secrets";
export * from "./identity-universal-auths"; export * from "./identity-universal-auths";
export * from "./incident-contacts"; export * from "./incident-contacts";
export * from "./instance-relay-config";
export * from "./integration-auths"; export * from "./integration-auths";
export * from "./integrations"; export * from "./integrations";
export * from "./internal-certificate-authorities"; export * from "./internal-certificate-authorities";
export * from "./internal-kms"; export * from "./internal-kms";
export * from "./key-value-store";
export * from "./kmip-client-certificates"; export * from "./kmip-client-certificates";
export * from "./kmip-clients"; export * from "./kmip-clients";
export * from "./kmip-org-configs"; export * from "./kmip-org-configs";
@@ -75,7 +78,9 @@ export * from "./models";
export * from "./oidc-configs"; export * from "./oidc-configs";
export * from "./org-bots"; export * from "./org-bots";
export * from "./org-gateway-config"; export * from "./org-gateway-config";
export * from "./org-gateway-config-v2";
export * from "./org-memberships"; export * from "./org-memberships";
export * from "./org-relay-config";
export * from "./org-roles"; export * from "./org-roles";
export * from "./organizations"; export * from "./organizations";
export * from "./pki-alerts"; export * from "./pki-alerts";
@@ -96,6 +101,7 @@ export * from "./project-user-additional-privilege";
export * from "./project-user-membership-roles"; export * from "./project-user-membership-roles";
export * from "./projects"; export * from "./projects";
export * from "./rate-limit"; export * from "./rate-limit";
export * from "./relays";
export * from "./resource-metadata"; export * from "./resource-metadata";
export * from "./saml-configs"; export * from "./saml-configs";
export * from "./scim-tokens"; export * from "./scim-tokens";
@@ -0,0 +1,38 @@
// 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 { zodBuffer } from "@app/lib/zod";
import { TImmutableDBKeys } from "./models";
export const InstanceRelayConfigSchema = z.object({
id: z.string().uuid(),
createdAt: z.date(),
updatedAt: z.date(),
encryptedRootRelayPkiCaPrivateKey: zodBuffer,
encryptedRootRelayPkiCaCertificate: zodBuffer,
encryptedInstanceRelayPkiCaPrivateKey: zodBuffer,
encryptedInstanceRelayPkiCaCertificate: zodBuffer,
encryptedInstanceRelayPkiCaCertificateChain: zodBuffer,
encryptedInstanceRelayPkiClientCaPrivateKey: zodBuffer,
encryptedInstanceRelayPkiClientCaCertificate: zodBuffer,
encryptedInstanceRelayPkiClientCaCertificateChain: zodBuffer,
encryptedInstanceRelayPkiServerCaPrivateKey: zodBuffer,
encryptedInstanceRelayPkiServerCaCertificate: zodBuffer,
encryptedInstanceRelayPkiServerCaCertificateChain: zodBuffer,
encryptedOrgRelayPkiCaPrivateKey: zodBuffer,
encryptedOrgRelayPkiCaCertificate: zodBuffer,
encryptedOrgRelayPkiCaCertificateChain: zodBuffer,
encryptedInstanceRelaySshClientCaPrivateKey: zodBuffer,
encryptedInstanceRelaySshClientCaPublicKey: zodBuffer,
encryptedInstanceRelaySshServerCaPrivateKey: zodBuffer,
encryptedInstanceRelaySshServerCaPublicKey: zodBuffer
});
export type TInstanceRelayConfig = z.infer<typeof InstanceRelayConfigSchema>;
export type TInstanceRelayConfigInsert = Omit<z.input<typeof InstanceRelayConfigSchema>, TImmutableDBKeys>;
export type TInstanceRelayConfigUpdate = Partial<Omit<z.input<typeof InstanceRelayConfigSchema>, TImmutableDBKeys>>;
+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>>;
+10 -1
View File
@@ -179,7 +179,16 @@ export enum TableName {
SecretScanningConfig = "secret_scanning_configs", SecretScanningConfig = "secret_scanning_configs",
// reminders // reminders
Reminder = "reminders", Reminder = "reminders",
ReminderRecipient = "reminders_recipients" ReminderRecipient = "reminders_recipients",
// gateway v2
InstanceRelayConfig = "instance_relay_config",
OrgRelayConfig = "org_relay_config",
OrgGatewayConfigV2 = "org_gateway_config_v2",
Relay = "relays",
GatewayV2 = "gateways_v2",
KeyValueStore = "key_value_store"
} }
export type TImmutableDBKeys = "id" | "createdAt" | "updatedAt" | "commitId"; export type TImmutableDBKeys = "id" | "createdAt" | "updatedAt" | "commitId";
@@ -0,0 +1,29 @@
// 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 { zodBuffer } from "@app/lib/zod";
import { TImmutableDBKeys } from "./models";
export const OrgGatewayConfigV2Schema = z.object({
id: z.string().uuid(),
orgId: z.string().uuid(),
createdAt: z.date(),
updatedAt: z.date(),
encryptedRootGatewayCaPrivateKey: zodBuffer,
encryptedRootGatewayCaCertificate: zodBuffer,
encryptedGatewayServerCaPrivateKey: zodBuffer,
encryptedGatewayServerCaCertificate: zodBuffer,
encryptedGatewayServerCaCertificateChain: zodBuffer,
encryptedGatewayClientCaPrivateKey: zodBuffer,
encryptedGatewayClientCaCertificate: zodBuffer,
encryptedGatewayClientCaCertificateChain: zodBuffer
});
export type TOrgGatewayConfigV2 = z.infer<typeof OrgGatewayConfigV2Schema>;
export type TOrgGatewayConfigV2Insert = Omit<z.input<typeof OrgGatewayConfigV2Schema>, TImmutableDBKeys>;
export type TOrgGatewayConfigV2Update = Partial<Omit<z.input<typeof OrgGatewayConfigV2Schema>, TImmutableDBKeys>>;
@@ -0,0 +1,31 @@
// 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 { zodBuffer } from "@app/lib/zod";
import { TImmutableDBKeys } from "./models";
export const OrgRelayConfigSchema = z.object({
id: z.string().uuid(),
createdAt: z.date(),
updatedAt: z.date(),
orgId: z.string().uuid(),
encryptedRelayPkiClientCaPrivateKey: zodBuffer,
encryptedRelayPkiClientCaCertificate: zodBuffer,
encryptedRelayPkiClientCaCertificateChain: zodBuffer,
encryptedRelayPkiServerCaPrivateKey: zodBuffer,
encryptedRelayPkiServerCaCertificate: zodBuffer,
encryptedRelayPkiServerCaCertificateChain: zodBuffer,
encryptedRelaySshClientCaPrivateKey: zodBuffer,
encryptedRelaySshClientCaPublicKey: zodBuffer,
encryptedRelaySshServerCaPrivateKey: zodBuffer,
encryptedRelaySshServerCaPublicKey: zodBuffer
});
export type TOrgRelayConfig = z.infer<typeof OrgRelayConfigSchema>;
export type TOrgRelayConfigInsert = Omit<z.input<typeof OrgRelayConfigSchema>, TImmutableDBKeys>;
export type TOrgRelayConfigUpdate = Partial<Omit<z.input<typeof OrgRelayConfigSchema>, TImmutableDBKeys>>;
+22
View File
@@ -0,0 +1,22 @@
// 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 RelaysSchema = z.object({
id: z.string().uuid(),
createdAt: z.date(),
updatedAt: z.date(),
orgId: z.string().uuid().nullable().optional(),
identityId: z.string().uuid().nullable().optional(),
name: z.string(),
host: z.string()
});
export type TRelays = z.infer<typeof RelaysSchema>;
export type TRelaysInsert = Omit<z.input<typeof RelaysSchema>, TImmutableDBKeys>;
export type TRelaysUpdate = Partial<Omit<z.input<typeof RelaysSchema>, TImmutableDBKeys>>;
@@ -1,5 +1,13 @@
import { z } from "zod"; import { z } from "zod";
import {
AzureProviderListItemSchema,
SanitizedAzureProviderSchema
} from "@app/ee/services/audit-log-stream/azure/azure-provider-schemas";
import {
CriblProviderListItemSchema,
SanitizedCriblProviderSchema
} from "@app/ee/services/audit-log-stream/cribl/cribl-provider-schemas";
import { import {
CustomProviderListItemSchema, CustomProviderListItemSchema,
SanitizedCustomProviderSchema SanitizedCustomProviderSchema
@@ -19,13 +27,17 @@ import { AuthMode } from "@app/services/auth/auth-type";
const SanitizedAuditLogStreamSchema = z.union([ const SanitizedAuditLogStreamSchema = z.union([
SanitizedCustomProviderSchema, SanitizedCustomProviderSchema,
SanitizedDatadogProviderSchema, SanitizedDatadogProviderSchema,
SanitizedSplunkProviderSchema SanitizedSplunkProviderSchema,
SanitizedAzureProviderSchema,
SanitizedCriblProviderSchema
]); ]);
const ProviderOptionsSchema = z.discriminatedUnion("provider", [ const ProviderOptionsSchema = z.discriminatedUnion("provider", [
CustomProviderListItemSchema, CustomProviderListItemSchema,
DatadogProviderListItemSchema, DatadogProviderListItemSchema,
SplunkProviderListItemSchema SplunkProviderListItemSchema,
AzureProviderListItemSchema,
CriblProviderListItemSchema
]); ]);
export const registerAuditLogStreamRouter = async (server: FastifyZodProvider) => { export const registerAuditLogStreamRouter = async (server: FastifyZodProvider) => {
@@ -1,4 +1,14 @@
import { LogProvider } from "@app/ee/services/audit-log-stream/audit-log-stream-enums"; import { LogProvider } from "@app/ee/services/audit-log-stream/audit-log-stream-enums";
import {
CreateAzureProviderLogStreamSchema,
SanitizedAzureProviderSchema,
UpdateAzureProviderLogStreamSchema
} from "@app/ee/services/audit-log-stream/azure/azure-provider-schemas";
import {
CreateCriblProviderLogStreamSchema,
SanitizedCriblProviderSchema,
UpdateCriblProviderLogStreamSchema
} from "@app/ee/services/audit-log-stream/cribl/cribl-provider-schemas";
import { import {
CreateCustomProviderLogStreamSchema, CreateCustomProviderLogStreamSchema,
SanitizedCustomProviderSchema, SanitizedCustomProviderSchema,
@@ -21,6 +31,15 @@ export * from "./audit-log-stream-router";
export const AUDIT_LOG_STREAM_REGISTER_ROUTER_MAP: Record<LogProvider, (server: FastifyZodProvider) => Promise<void>> = export const AUDIT_LOG_STREAM_REGISTER_ROUTER_MAP: Record<LogProvider, (server: FastifyZodProvider) => Promise<void>> =
{ {
[LogProvider.Azure]: async (server: FastifyZodProvider) => {
registerAuditLogStreamEndpoints({
server,
provider: LogProvider.Azure,
sanitizedResponseSchema: SanitizedAzureProviderSchema,
createSchema: CreateAzureProviderLogStreamSchema,
updateSchema: UpdateAzureProviderLogStreamSchema
});
},
[LogProvider.Custom]: async (server: FastifyZodProvider) => { [LogProvider.Custom]: async (server: FastifyZodProvider) => {
registerAuditLogStreamEndpoints({ registerAuditLogStreamEndpoints({
server, server,
@@ -47,5 +66,14 @@ export const AUDIT_LOG_STREAM_REGISTER_ROUTER_MAP: Record<LogProvider, (server:
createSchema: CreateSplunkProviderLogStreamSchema, createSchema: CreateSplunkProviderLogStreamSchema,
updateSchema: UpdateSplunkProviderLogStreamSchema updateSchema: UpdateSplunkProviderLogStreamSchema
}); });
},
[LogProvider.Cribl]: async (server: FastifyZodProvider) => {
registerAuditLogStreamEndpoints({
server,
provider: LogProvider.Cribl,
sanitizedResponseSchema: SanitizedCriblProviderSchema,
createSchema: CreateCriblProviderLogStreamSchema,
updateSchema: UpdateCriblProviderLogStreamSchema
});
} }
}; };
@@ -84,7 +84,9 @@ export const registerDynamicSecretRouter = async (server: FastifyZodProvider) =>
}), }),
response: { response: {
200: z.object({ 200: z.object({
dynamicSecret: SanitizedDynamicSecretSchema dynamicSecret: SanitizedDynamicSecretSchema.extend({
inputs: z.unknown()
})
}) })
} }
}, },
@@ -151,7 +153,9 @@ export const registerDynamicSecretRouter = async (server: FastifyZodProvider) =>
}), }),
response: { response: {
200: z.object({ 200: z.object({
dynamicSecret: SanitizedDynamicSecretSchema dynamicSecret: SanitizedDynamicSecretSchema.extend({
inputs: z.unknown()
})
}) })
} }
}, },
+2
View File
@@ -24,6 +24,7 @@ import { registerPITRouter } from "./pit-router";
import { registerProjectRoleRouter } from "./project-role-router"; import { registerProjectRoleRouter } from "./project-role-router";
import { registerProjectRouter } from "./project-router"; import { registerProjectRouter } from "./project-router";
import { registerRateLimitRouter } from "./rate-limit-router"; import { registerRateLimitRouter } from "./rate-limit-router";
import { registerRelayRouter } from "./relay-router";
import { registerSamlRouter } from "./saml-router"; import { registerSamlRouter } from "./saml-router";
import { registerScimRouter } from "./scim-router"; import { registerScimRouter } from "./scim-router";
import { registerSecretApprovalPolicyRouter } from "./secret-approval-policy-router"; import { registerSecretApprovalPolicyRouter } from "./secret-approval-policy-router";
@@ -79,6 +80,7 @@ export const registerV1EERoutes = async (server: FastifyZodProvider) => {
); );
await server.register(registerGatewayRouter, { prefix: "/gateways" }); await server.register(registerGatewayRouter, { prefix: "/gateways" });
await server.register(registerRelayRouter, { prefix: "/relays" });
await server.register(registerGithubOrgSyncRouter, { prefix: "/github-org-sync-config" }); await server.register(registerGithubOrgSyncRouter, { prefix: "/github-org-sync-config" });
await server.register( await server.register(
+8 -1
View File
@@ -43,6 +43,12 @@ export const registerLicenseRouter = async (server: FastifyZodProvider) => {
}, },
schema: { schema: {
params: z.object({ organizationId: z.string().trim() }), params: z.object({ organizationId: z.string().trim() }),
querystring: z.object({
refreshCache: z
.enum(["true", "false"])
.default("false")
.transform((value) => value === "true")
}),
response: { response: {
200: z.object({ plan: z.any() }) 200: z.object({ plan: z.any() })
} }
@@ -54,7 +60,8 @@ export const registerLicenseRouter = async (server: FastifyZodProvider) => {
actor: req.permission.type, actor: req.permission.type,
actorOrgId: req.permission.orgId, actorOrgId: req.permission.orgId,
actorAuthMethod: req.permission.authMethod, actorAuthMethod: req.permission.authMethod,
orgId: req.params.organizationId orgId: req.params.organizationId,
refreshCache: req.query.refreshCache
}); });
return { plan }; return { plan };
} }
+103
View File
@@ -0,0 +1,103 @@
import { z } from "zod";
import { getConfig } from "@app/lib/config/env";
import { crypto } from "@app/lib/crypto/cryptography";
import { BadRequestError, UnauthorizedError } from "@app/lib/errors";
import { writeLimit } from "@app/server/config/rateLimiter";
import { slugSchema } from "@app/server/lib/schemas";
import { verifyAuth } from "@app/server/plugins/auth/verify-auth";
import { AuthMode } from "@app/services/auth/auth-type";
export const registerRelayRouter = async (server: FastifyZodProvider) => {
const appCfg = getConfig();
server.route({
method: "POST",
url: "/register-instance-relay",
config: {
rateLimit: writeLimit
},
schema: {
body: z.object({
host: z.string(),
name: slugSchema({ min: 1, max: 32, field: "name" })
}),
response: {
200: z.object({
pki: z.object({
serverCertificate: z.string(),
serverPrivateKey: z.string(),
clientCertificateChain: z.string()
}),
ssh: z.object({
serverCertificate: z.string(),
serverPrivateKey: z.string(),
clientCAPublicKey: z.string()
})
})
}
},
onRequest: (req, _, next) => {
const authHeader = req.headers.authorization;
if (appCfg.RELAY_AUTH_SECRET && authHeader) {
const expectedHeader = `Bearer ${appCfg.RELAY_AUTH_SECRET}`;
if (
authHeader.length === expectedHeader.length &&
crypto.nativeCrypto.timingSafeEqual(Buffer.from(authHeader), Buffer.from(expectedHeader))
) {
return next();
}
}
throw new UnauthorizedError({
message: "Invalid relay auth secret"
});
},
handler: async (req) => {
return server.services.relay.registerRelay({
...req.body
});
}
});
server.route({
method: "POST",
url: "/register-org-relay",
config: {
rateLimit: writeLimit
},
schema: {
body: z.object({
host: z.string(),
name: slugSchema({ min: 1, max: 32, field: "name" })
}),
response: {
200: z.object({
pki: z.object({
serverCertificate: z.string(),
serverPrivateKey: z.string(),
clientCertificateChain: z.string()
}),
ssh: z.object({
serverCertificate: z.string(),
serverPrivateKey: z.string(),
clientCAPublicKey: z.string()
})
})
}
},
onRequest: verifyAuth([AuthMode.IDENTITY_ACCESS_TOKEN]),
handler: async (req) => {
throw new BadRequestError({
message: "Org relay registration is not yet supported"
});
return server.services.relay.registerRelay({
...req.body,
identityId: req.permission.id,
orgId: req.permission.orgId
});
}
});
};
+133
View File
@@ -0,0 +1,133 @@
import z from "zod";
import { GatewaysV2Schema } from "@app/db/schemas";
import { readLimit, writeLimit } from "@app/server/config/rateLimiter";
import { slugSchema } from "@app/server/lib/schemas";
import { verifyAuth } from "@app/server/plugins/auth/verify-auth";
import { AuthMode } from "@app/services/auth/auth-type";
const SanitizedGatewayV2Schema = GatewaysV2Schema.pick({
id: true,
identityId: true,
name: true,
createdAt: true,
updatedAt: true,
heartbeat: true
});
export const registerGatewayV2Router = async (server: FastifyZodProvider) => {
server.route({
method: "POST",
url: "/",
schema: {
body: z.object({
relayName: slugSchema({ min: 1, max: 32, field: "relayName" }),
name: slugSchema({ min: 1, max: 32, field: "name" })
}),
response: {
200: z.object({
gatewayId: z.string(),
relayHost: z.string(),
pki: z.object({
serverCertificate: z.string(),
serverPrivateKey: z.string(),
clientCertificateChain: z.string()
}),
ssh: z.object({
clientCertificate: z.string(),
clientPrivateKey: z.string(),
serverCAPublicKey: z.string()
})
})
}
},
config: {
rateLimit: writeLimit
},
onRequest: verifyAuth([AuthMode.IDENTITY_ACCESS_TOKEN]),
handler: async (req) => {
const gateway = await server.services.gatewayV2.registerGateway({
orgId: req.permission.orgId,
relayName: req.body.relayName,
actorId: req.permission.id,
actorAuthMethod: req.permission.authMethod,
name: req.body.name
});
return gateway;
}
});
server.route({
method: "POST",
url: "/heartbeat",
config: {
rateLimit: writeLimit
},
schema: {
response: {
200: z.object({
message: z.string()
})
}
},
onRequest: verifyAuth([AuthMode.IDENTITY_ACCESS_TOKEN]),
handler: async (req) => {
await server.services.gatewayV2.heartbeat({
orgPermission: req.permission
});
return { message: "Successfully triggered heartbeat" };
}
});
server.route({
method: "GET",
url: "/",
schema: {
response: {
200: SanitizedGatewayV2Schema.extend({
identity: z.object({
name: z.string(),
id: z.string()
})
}).array()
}
},
config: {
rateLimit: readLimit
},
onRequest: verifyAuth([AuthMode.JWT, AuthMode.IDENTITY_ACCESS_TOKEN]),
handler: async (req) => {
const gateways = await server.services.gatewayV2.listGateways({
orgPermission: req.permission
});
return gateways;
}
});
server.route({
method: "DELETE",
url: "/:id",
config: {
rateLimit: writeLimit
},
schema: {
params: z.object({
id: z.string()
}),
response: {
200: SanitizedGatewayV2Schema
}
},
onRequest: verifyAuth([AuthMode.IDENTITY_ACCESS_TOKEN, AuthMode.JWT]),
handler: async (req) => {
const gateway = await server.services.gatewayV2.deleteGatewayById({
orgPermission: req.permission,
id: req.params.id
});
return gateway;
}
});
};
+3
View File
@@ -7,6 +7,7 @@ import {
SECRET_SCANNING_REGISTER_ROUTER_MAP SECRET_SCANNING_REGISTER_ROUTER_MAP
} from "@app/ee/routes/v2/secret-scanning-v2-routers"; } from "@app/ee/routes/v2/secret-scanning-v2-routers";
import { registerGatewayV2Router } from "./gateway-router";
import { registerIdentityProjectAdditionalPrivilegeRouter } from "./identity-project-additional-privilege-router"; import { registerIdentityProjectAdditionalPrivilegeRouter } from "./identity-project-additional-privilege-router";
import { registerProjectRoleRouter } from "./project-role-router"; import { registerProjectRoleRouter } from "./project-role-router";
@@ -23,6 +24,8 @@ export const registerV2EERoutes = async (server: FastifyZodProvider) => {
prefix: "/identity-project-additional-privilege" prefix: "/identity-project-additional-privilege"
}); });
await server.register(registerGatewayV2Router, { prefix: "/gateways" });
await server.register( await server.register(
async (secretRotationV2Router) => { async (secretRotationV2Router) => {
// register generic secret rotation endpoints // register generic secret rotation endpoints
@@ -1,5 +1,7 @@
export enum LogProvider { export enum LogProvider {
Azure = "azure",
Cribl = "cribl",
Custom = "custom",
Datadog = "datadog", Datadog = "datadog",
Splunk = "splunk", Splunk = "splunk"
Custom = "custom"
} }
@@ -1,5 +1,7 @@
import { LogProvider } from "./audit-log-stream-enums"; import { LogProvider } from "./audit-log-stream-enums";
import { TAuditLogStreamCredentials, TLogStreamFactory } from "./audit-log-stream-types"; import { TAuditLogStreamCredentials, TLogStreamFactory } from "./audit-log-stream-types";
import { AzureProviderFactory } from "./azure/azure-provider-factory";
import { CriblProviderFactory } from "./cribl/cribl-provider-factory";
import { CustomProviderFactory } from "./custom/custom-provider-factory"; import { CustomProviderFactory } from "./custom/custom-provider-factory";
import { DatadogProviderFactory } from "./datadog/datadog-provider-factory"; import { DatadogProviderFactory } from "./datadog/datadog-provider-factory";
import { SplunkProviderFactory } from "./splunk/splunk-provider-factory"; import { SplunkProviderFactory } from "./splunk/splunk-provider-factory";
@@ -7,7 +9,9 @@ import { SplunkProviderFactory } from "./splunk/splunk-provider-factory";
type TLogStreamFactoryImplementation = TLogStreamFactory<TAuditLogStreamCredentials>; type TLogStreamFactoryImplementation = TLogStreamFactory<TAuditLogStreamCredentials>;
export const LOG_STREAM_FACTORY_MAP: Record<LogProvider, TLogStreamFactoryImplementation> = { export const LOG_STREAM_FACTORY_MAP: Record<LogProvider, TLogStreamFactoryImplementation> = {
[LogProvider.Azure]: AzureProviderFactory as TLogStreamFactoryImplementation,
[LogProvider.Datadog]: DatadogProviderFactory as TLogStreamFactoryImplementation, [LogProvider.Datadog]: DatadogProviderFactory as TLogStreamFactoryImplementation,
[LogProvider.Splunk]: SplunkProviderFactory as TLogStreamFactoryImplementation, [LogProvider.Splunk]: SplunkProviderFactory as TLogStreamFactoryImplementation,
[LogProvider.Custom]: CustomProviderFactory as TLogStreamFactoryImplementation [LogProvider.Custom]: CustomProviderFactory as TLogStreamFactoryImplementation,
[LogProvider.Cribl]: CriblProviderFactory as TLogStreamFactoryImplementation
}; };
@@ -3,14 +3,20 @@ import { TKmsServiceFactory } from "@app/services/kms/kms-service";
import { KmsDataKey } from "@app/services/kms/kms-types"; import { KmsDataKey } from "@app/services/kms/kms-types";
import { TAuditLogStream, TAuditLogStreamCredentials } from "./audit-log-stream-types"; import { TAuditLogStream, TAuditLogStreamCredentials } from "./audit-log-stream-types";
import { getAzureProviderListItem } from "./azure/azure-provider-fns";
import { getCriblProviderListItem } from "./cribl/cribl-provider-fns";
import { getCustomProviderListItem } from "./custom/custom-provider-fns"; import { getCustomProviderListItem } from "./custom/custom-provider-fns";
import { getDatadogProviderListItem } from "./datadog/datadog-provider-fns"; import { getDatadogProviderListItem } from "./datadog/datadog-provider-fns";
import { getSplunkProviderListItem } from "./splunk/splunk-provider-fns"; import { getSplunkProviderListItem } from "./splunk/splunk-provider-fns";
export const listProviderOptions = () => { export const listProviderOptions = () => {
return [getDatadogProviderListItem(), getSplunkProviderListItem(), getCustomProviderListItem()].sort((a, b) => return [
a.name.localeCompare(b.name) getDatadogProviderListItem(),
); getSplunkProviderListItem(),
getCustomProviderListItem(),
getAzureProviderListItem(),
getCriblProviderListItem()
].sort((a, b) => a.name.localeCompare(b.name));
}; };
export const encryptLogStreamCredentials = async ({ export const encryptLogStreamCredentials = async ({
@@ -1,16 +1,20 @@
import { TAuditLogs } from "@app/db/schemas"; import { TAuditLogs } from "@app/db/schemas";
import { LogProvider } from "./audit-log-stream-enums"; import { LogProvider } from "./audit-log-stream-enums";
import { TAzureProvider, TAzureProviderCredentials } from "./azure/azure-provider-types";
import { TCriblProvider, TCriblProviderCredentials } from "./cribl/cribl-provider-types";
import { TCustomProvider, TCustomProviderCredentials } from "./custom/custom-provider-types"; import { TCustomProvider, TCustomProviderCredentials } from "./custom/custom-provider-types";
import { TDatadogProvider, TDatadogProviderCredentials } from "./datadog/datadog-provider-types"; import { TDatadogProvider, TDatadogProviderCredentials } from "./datadog/datadog-provider-types";
import { TSplunkProvider, TSplunkProviderCredentials } from "./splunk/splunk-provider-types"; import { TSplunkProvider, TSplunkProviderCredentials } from "./splunk/splunk-provider-types";
export type TAuditLogStream = TDatadogProvider | TSplunkProvider | TCustomProvider; export type TAuditLogStream = TDatadogProvider | TSplunkProvider | TCustomProvider | TAzureProvider | TCriblProvider;
export type TAuditLogStreamCredentials = export type TAuditLogStreamCredentials =
| TDatadogProviderCredentials | TDatadogProviderCredentials
| TSplunkProviderCredentials | TSplunkProviderCredentials
| TCustomProviderCredentials; | TCustomProviderCredentials
| TAzureProviderCredentials
| TCriblProviderCredentials;
export type TCreateAuditLogStreamDTO = { export type TCreateAuditLogStreamDTO = {
provider: LogProvider; provider: LogProvider;
@@ -0,0 +1,98 @@
import { RawAxiosRequestHeaders } from "axios";
import { request } from "@app/lib/config/request";
import { BadRequestError } from "@app/lib/errors";
import { blockLocalAndPrivateIpAddresses } from "@app/lib/validator";
import { AUDIT_LOG_STREAM_TIMEOUT } from "../../audit-log/audit-log-queue";
import { TLogStreamFactoryStreamLog, TLogStreamFactoryValidateCredentials } from "../audit-log-stream-types";
import { TAzureProviderCredentials } from "./azure-provider-types";
function createPayload(event: { createdAt?: Date | string } & Record<string, unknown>) {
return [
{
...event,
TimeGenerated: (event.createdAt ? new Date(event.createdAt) : new Date()).toISOString()
}
];
}
async function getAzureToken(tenantId: string, clientId: string, clientSecret: string) {
const { data } = await request.post<{ access_token: string }>(
`https://login.microsoftonline.com/${tenantId}/oauth2/v2.0/token`,
new URLSearchParams({
grant_type: "client_credentials",
client_id: clientId,
client_secret: clientSecret,
scope: "https://monitor.azure.com/.default"
}),
{
headers: {
"Content-Type": "application/x-www-form-urlencoded"
}
}
);
return data.access_token;
}
export const AzureProviderFactory = () => {
const validateCredentials: TLogStreamFactoryValidateCredentials<TAzureProviderCredentials> = async ({
credentials
}) => {
const { tenantId, clientId, clientSecret, dceUrl, dcrId, cltName } = credentials;
await blockLocalAndPrivateIpAddresses(dceUrl);
const token = await getAzureToken(tenantId, clientId, clientSecret);
const streamHeaders: RawAxiosRequestHeaders = {
"Content-Type": "application/json",
Authorization: `Bearer ${token}`
};
await request
.post(
`${dceUrl}/dataCollectionRules/${dcrId}/streams/Custom-${cltName}_CL?api-version=2023-01-01`,
createPayload({ ping: "ok" }),
{
headers: streamHeaders,
timeout: AUDIT_LOG_STREAM_TIMEOUT,
signal: AbortSignal.timeout(AUDIT_LOG_STREAM_TIMEOUT)
}
)
.catch((err) => {
throw new BadRequestError({ message: `Failed to connect with Azure: ${(err as Error)?.message}` });
});
return credentials;
};
const streamLog: TLogStreamFactoryStreamLog<TAzureProviderCredentials> = async ({ credentials, auditLog }) => {
const { tenantId, clientId, clientSecret, dceUrl, dcrId, cltName } = credentials;
await blockLocalAndPrivateIpAddresses(dceUrl);
const token = await getAzureToken(tenantId, clientId, clientSecret);
const streamHeaders: RawAxiosRequestHeaders = {
"Content-Type": "application/json",
Authorization: `Bearer ${token}`
};
await request.post(
`${dceUrl}/dataCollectionRules/${dcrId}/streams/Custom-${cltName}_CL?api-version=2023-01-01`,
createPayload(auditLog),
{
headers: streamHeaders,
timeout: AUDIT_LOG_STREAM_TIMEOUT,
signal: AbortSignal.timeout(AUDIT_LOG_STREAM_TIMEOUT)
}
);
};
return {
validateCredentials,
streamLog
};
};
@@ -0,0 +1,8 @@
import { LogProvider } from "../audit-log-stream-enums";
export const getAzureProviderListItem = () => {
return {
name: "Azure" as const,
provider: LogProvider.Azure as const
};
};
@@ -0,0 +1,52 @@
import RE2 from "re2";
import { z } from "zod";
import { LogProvider } from "../audit-log-stream-enums";
import { BaseProviderSchema } from "../audit-log-stream-schemas";
export const AzureProviderCredentialsSchema = z.object({
tenantId: z.string().trim().uuid(),
clientId: z.string().trim().uuid(),
clientSecret: z.string().trim().length(40),
// Data Collection Endpoint URL
dceUrl: z.string().trim().url().min(1).max(255),
// Data Collection Rule Immutable ID
dcrId: z
.string()
.trim()
.refine((val) => new RE2(/^dcr-[0-9a-f]{32}$/).test(val), "DCR ID must be in dcr-*** format"),
// Custom Log Table Name
cltName: z.string().trim().min(1).max(255)
});
const BaseAzureProviderSchema = BaseProviderSchema.extend({ provider: z.literal(LogProvider.Azure) });
export const AzureProviderSchema = BaseAzureProviderSchema.extend({
credentials: AzureProviderCredentialsSchema
});
export const SanitizedAzureProviderSchema = BaseAzureProviderSchema.extend({
credentials: AzureProviderCredentialsSchema.pick({
tenantId: true,
clientId: true,
dceUrl: true,
dcrId: true,
cltName: true
})
});
export const AzureProviderListItemSchema = z.object({
name: z.literal("Azure"),
provider: z.literal(LogProvider.Azure)
});
export const CreateAzureProviderLogStreamSchema = z.object({
credentials: AzureProviderCredentialsSchema
});
export const UpdateAzureProviderLogStreamSchema = z.object({
credentials: AzureProviderCredentialsSchema
});
@@ -0,0 +1,7 @@
import { z } from "zod";
import { AzureProviderCredentialsSchema, AzureProviderSchema } from "./azure-provider-schemas";
export type TAzureProvider = z.infer<typeof AzureProviderSchema>;
export type TAzureProviderCredentials = z.infer<typeof AzureProviderCredentialsSchema>;
@@ -0,0 +1,58 @@
import { RawAxiosRequestHeaders } from "axios";
import { request } from "@app/lib/config/request";
import { BadRequestError } from "@app/lib/errors";
import { blockLocalAndPrivateIpAddresses } from "@app/lib/validator";
import { AUDIT_LOG_STREAM_TIMEOUT } from "../../audit-log/audit-log-queue";
import { TLogStreamFactoryStreamLog, TLogStreamFactoryValidateCredentials } from "../audit-log-stream-types";
import { TCriblProviderCredentials } from "./cribl-provider-types";
export const CriblProviderFactory = () => {
const validateCredentials: TLogStreamFactoryValidateCredentials<TCriblProviderCredentials> = async ({
credentials
}) => {
const { url, token } = credentials;
await blockLocalAndPrivateIpAddresses(url);
const streamHeaders: RawAxiosRequestHeaders = {
"Content-Type": "application/json",
Authorization: `Bearer ${token}`
};
await request
.post(url, JSON.stringify({ ping: "ok" }), {
headers: streamHeaders,
timeout: AUDIT_LOG_STREAM_TIMEOUT,
signal: AbortSignal.timeout(AUDIT_LOG_STREAM_TIMEOUT)
})
.catch((err) => {
throw new BadRequestError({ message: `Failed to connect with Cribl: ${(err as Error)?.message}` });
});
return credentials;
};
const streamLog: TLogStreamFactoryStreamLog<TCriblProviderCredentials> = async ({ credentials, auditLog }) => {
const { url, token } = credentials;
await blockLocalAndPrivateIpAddresses(url);
const streamHeaders: RawAxiosRequestHeaders = {
"Content-Type": "application/json",
Authorization: `Bearer ${token}`
};
await request.post(url, JSON.stringify(auditLog), {
headers: streamHeaders,
timeout: AUDIT_LOG_STREAM_TIMEOUT,
signal: AbortSignal.timeout(AUDIT_LOG_STREAM_TIMEOUT)
});
};
return {
validateCredentials,
streamLog
};
};
@@ -0,0 +1,8 @@
import { LogProvider } from "../audit-log-stream-enums";
export const getCriblProviderListItem = () => {
return {
name: "Cribl" as const,
provider: LogProvider.Cribl as const
};
};
@@ -0,0 +1,34 @@
import { z } from "zod";
import { LogProvider } from "../audit-log-stream-enums";
import { BaseProviderSchema } from "../audit-log-stream-schemas";
export const CriblProviderCredentialsSchema = z.object({
url: z.string().url().trim().min(1).max(255),
token: z.string().trim().min(21).max(255)
});
const BaseCriblProviderSchema = BaseProviderSchema.extend({ provider: z.literal(LogProvider.Cribl) });
export const CriblProviderSchema = BaseCriblProviderSchema.extend({
credentials: CriblProviderCredentialsSchema
});
export const SanitizedCriblProviderSchema = BaseCriblProviderSchema.extend({
credentials: CriblProviderCredentialsSchema.pick({
url: true
})
});
export const CriblProviderListItemSchema = z.object({
name: z.literal("Cribl"),
provider: z.literal(LogProvider.Cribl)
});
export const CreateCriblProviderLogStreamSchema = z.object({
credentials: CriblProviderCredentialsSchema
});
export const UpdateCriblProviderLogStreamSchema = z.object({
credentials: CriblProviderCredentialsSchema
});
@@ -0,0 +1,7 @@
import { z } from "zod";
import { CriblProviderCredentialsSchema, CriblProviderSchema } from "./cribl-provider-schemas";
export type TCriblProvider = z.infer<typeof CriblProviderSchema>;
export type TCriblProviderCredentials = z.infer<typeof CriblProviderCredentialsSchema>;
@@ -46,7 +46,10 @@ export const dynamicSecretLeaseDALFactory = (db: TDbClient) => {
const countLeasesForDynamicSecret = async (dynamicSecretId: string, tx?: Knex) => { const countLeasesForDynamicSecret = async (dynamicSecretId: string, tx?: Knex) => {
try { try {
const doc = await (tx || db)(TableName.DynamicSecretLease).count("*").where({ dynamicSecretId }).first(); const doc = await (tx || db.replicaNode())(TableName.DynamicSecretLease)
.count("*")
.where({ dynamicSecretId })
.first();
return parseInt(doc || "0", 10); return parseInt(doc || "0", 10);
} catch (error) { } catch (error) {
throw new DatabaseError({ error, name: "DynamicSecretCountLeases" }); throw new DatabaseError({ error, name: "DynamicSecretCountLeases" });
@@ -55,7 +58,7 @@ export const dynamicSecretLeaseDALFactory = (db: TDbClient) => {
const findById = async (id: string, tx?: Knex) => { const findById = async (id: string, tx?: Knex) => {
try { try {
const doc = await (tx || db)(TableName.DynamicSecretLease) const doc = await (tx || db.replicaNode())(TableName.DynamicSecretLease)
.where({ [`${TableName.DynamicSecretLease}.id` as "id"]: id }) .where({ [`${TableName.DynamicSecretLease}.id` as "id"]: id })
.first() .first()
.join( .join(
@@ -19,6 +19,7 @@ import { TSecretFolderDALFactory } from "@app/services/secret-folder/secret-fold
import { TDynamicSecretLeaseDALFactory } from "../dynamic-secret-lease/dynamic-secret-lease-dal"; import { TDynamicSecretLeaseDALFactory } from "../dynamic-secret-lease/dynamic-secret-lease-dal";
import { TDynamicSecretLeaseQueueServiceFactory } from "../dynamic-secret-lease/dynamic-secret-lease-queue"; import { TDynamicSecretLeaseQueueServiceFactory } from "../dynamic-secret-lease/dynamic-secret-lease-queue";
import { TGatewayDALFactory } from "../gateway/gateway-dal"; import { TGatewayDALFactory } from "../gateway/gateway-dal";
import { TGatewayV2DALFactory } from "../gateway-v2/gateway-v2-dal";
import { OrgPermissionGatewayActions, OrgPermissionSubjects } from "../permission/org-permission"; import { OrgPermissionGatewayActions, OrgPermissionSubjects } from "../permission/org-permission";
import { TDynamicSecretDALFactory } from "./dynamic-secret-dal"; import { TDynamicSecretDALFactory } from "./dynamic-secret-dal";
import { DynamicSecretStatus, TDynamicSecretServiceFactory } from "./dynamic-secret-types"; import { DynamicSecretStatus, TDynamicSecretServiceFactory } from "./dynamic-secret-types";
@@ -39,6 +40,7 @@ type TDynamicSecretServiceFactoryDep = {
permissionService: Pick<TPermissionServiceFactory, "getProjectPermission" | "getOrgPermission">; permissionService: Pick<TPermissionServiceFactory, "getProjectPermission" | "getOrgPermission">;
kmsService: Pick<TKmsServiceFactory, "createCipherPairWithDataKey">; kmsService: Pick<TKmsServiceFactory, "createCipherPairWithDataKey">;
gatewayDAL: Pick<TGatewayDALFactory, "findOne" | "find">; gatewayDAL: Pick<TGatewayDALFactory, "findOne" | "find">;
gatewayV2DAL: Pick<TGatewayV2DALFactory, "findOne" | "find">;
resourceMetadataDAL: Pick<TResourceMetadataDALFactory, "insertMany" | "delete">; resourceMetadataDAL: Pick<TResourceMetadataDALFactory, "insertMany" | "delete">;
}; };
@@ -53,6 +55,7 @@ export const dynamicSecretServiceFactory = ({
projectDAL, projectDAL,
kmsService, kmsService,
gatewayDAL, gatewayDAL,
gatewayV2DAL,
resourceMetadataDAL resourceMetadataDAL
}: TDynamicSecretServiceFactoryDep): TDynamicSecretServiceFactory => { }: TDynamicSecretServiceFactoryDep): TDynamicSecretServiceFactory => {
const create: TDynamicSecretServiceFactory["create"] = async ({ const create: TDynamicSecretServiceFactory["create"] = async ({
@@ -70,6 +73,7 @@ export const dynamicSecretServiceFactory = ({
metadata, metadata,
usernameTemplate usernameTemplate
}) => { }) => {
let isGatewayV1 = true;
const project = await projectDAL.findProjectBySlug(projectSlug, actorOrgId); const project = await projectDAL.findProjectBySlug(projectSlug, actorOrgId);
if (!project) throw new NotFoundError({ message: `Project with slug '${projectSlug}' not found` }); if (!project) throw new NotFoundError({ message: `Project with slug '${projectSlug}' not found` });
@@ -118,17 +122,22 @@ export const dynamicSecretServiceFactory = ({
const gatewayId = inputs.gatewayId as string; const gatewayId = inputs.gatewayId as string;
const [gateway] = await gatewayDAL.find({ id: gatewayId, orgId: actorOrgId }); const [gateway] = await gatewayDAL.find({ id: gatewayId, orgId: actorOrgId });
const [gatewayv2] = await gatewayV2DAL.find({ id: gatewayId, orgId: actorOrgId });
if (!gateway) { if (!gateway && !gatewayv2) {
throw new NotFoundError({ throw new NotFoundError({
message: `Gateway with ID ${gatewayId} not found` message: `Gateway with ID ${gatewayId} not found`
}); });
} }
if (!gateway) {
isGatewayV1 = false;
}
const { permission: orgPermission } = await permissionService.getOrgPermission( const { permission: orgPermission } = await permissionService.getOrgPermission(
actor, actor,
actorId, actorId,
gateway.orgId, gateway?.orgId ?? gatewayv2?.orgId,
actorAuthMethod, actorAuthMethod,
actorOrgId actorOrgId
); );
@@ -138,7 +147,7 @@ export const dynamicSecretServiceFactory = ({
OrgPermissionSubjects.Gateway OrgPermissionSubjects.Gateway
); );
selectedGatewayId = gateway.id; selectedGatewayId = gateway?.id ?? gatewayv2?.id;
} }
const isConnected = await selectedProvider.validateConnection(provider.inputs, { projectId }); const isConnected = await selectedProvider.validateConnection(provider.inputs, { projectId });
@@ -159,7 +168,8 @@ export const dynamicSecretServiceFactory = ({
defaultTTL, defaultTTL,
folderId: folder.id, folderId: folder.id,
name, name,
gatewayId: selectedGatewayId, gatewayId: isGatewayV1 ? selectedGatewayId : undefined,
gatewayV2Id: isGatewayV1 ? undefined : selectedGatewayId,
usernameTemplate usernameTemplate
}, },
tx tx
@@ -180,7 +190,7 @@ export const dynamicSecretServiceFactory = ({
return cfg; return cfg;
}); });
return dynamicSecretCfg; return { ...dynamicSecretCfg, inputs };
}; };
const updateByName: TDynamicSecretServiceFactory["updateByName"] = async ({ const updateByName: TDynamicSecretServiceFactory["updateByName"] = async ({
@@ -270,20 +280,27 @@ export const dynamicSecretServiceFactory = ({
const updatedInput = await selectedProvider.validateProviderInputs(newInput, { projectId }); const updatedInput = await selectedProvider.validateProviderInputs(newInput, { projectId });
let selectedGatewayId: string | null = null; let selectedGatewayId: string | null = null;
let isGatewayV1 = true;
if (updatedInput && typeof updatedInput === "object" && "gatewayId" in updatedInput && updatedInput?.gatewayId) { if (updatedInput && typeof updatedInput === "object" && "gatewayId" in updatedInput && updatedInput?.gatewayId) {
const gatewayId = updatedInput.gatewayId as string; const gatewayId = updatedInput.gatewayId as string;
const [gateway] = await gatewayDAL.find({ id: gatewayId, orgId: actorOrgId }); const [gateway] = await gatewayDAL.find({ id: gatewayId, orgId: actorOrgId });
if (!gateway) { const [gatewayv2] = await gatewayV2DAL.find({ id: gatewayId, orgId: actorOrgId });
if (!gateway && !gatewayv2) {
throw new NotFoundError({ throw new NotFoundError({
message: `Gateway with ID ${gatewayId} not found` message: `Gateway with ID ${gatewayId} not found`
}); });
} }
if (!gateway) {
isGatewayV1 = false;
}
const { permission: orgPermission } = await permissionService.getOrgPermission( const { permission: orgPermission } = await permissionService.getOrgPermission(
actor, actor,
actorId, actorId,
gateway.orgId, actorOrgId,
actorAuthMethod, actorAuthMethod,
actorOrgId actorOrgId
); );
@@ -293,7 +310,7 @@ export const dynamicSecretServiceFactory = ({
OrgPermissionSubjects.Gateway OrgPermissionSubjects.Gateway
); );
selectedGatewayId = gateway.id; selectedGatewayId = gateway?.id ?? gatewayv2?.id;
} }
const isConnected = await selectedProvider.validateConnection(newInput, { projectId }); const isConnected = await selectedProvider.validateConnection(newInput, { projectId });
@@ -309,7 +326,8 @@ export const dynamicSecretServiceFactory = ({
defaultTTL, defaultTTL,
name: newName ?? name, name: newName ?? name,
status: null, status: null,
gatewayId: selectedGatewayId, gatewayId: isGatewayV1 ? selectedGatewayId : null,
gatewayV2Id: isGatewayV1 ? null : selectedGatewayId,
usernameTemplate usernameTemplate
}, },
tx tx
@@ -337,7 +355,7 @@ export const dynamicSecretServiceFactory = ({
return cfg; return cfg;
}); });
return updatedDynamicCfg; return { ...updatedDynamicCfg, inputs: updatedInput };
}; };
const deleteByName: TDynamicSecretServiceFactory["deleteByName"] = async ({ const deleteByName: TDynamicSecretServiceFactory["deleteByName"] = async ({
@@ -1,6 +1,7 @@
import { SnowflakeProvider } from "@app/ee/services/dynamic-secret/providers/snowflake"; import { SnowflakeProvider } from "@app/ee/services/dynamic-secret/providers/snowflake";
import { TGatewayServiceFactory } from "../../gateway/gateway-service"; import { TGatewayServiceFactory } from "../../gateway/gateway-service";
import { TGatewayV2ServiceFactory } from "../../gateway-v2/gateway-v2-service";
import { AwsElastiCacheDatabaseProvider } from "./aws-elasticache"; import { AwsElastiCacheDatabaseProvider } from "./aws-elasticache";
import { AwsIamProvider } from "./aws-iam"; import { AwsIamProvider } from "./aws-iam";
import { AzureEntraIDProvider } from "./azure-entra-id"; import { AzureEntraIDProvider } from "./azure-entra-id";
@@ -24,12 +25,14 @@ import { VerticaProvider } from "./vertica";
type TBuildDynamicSecretProviderDTO = { type TBuildDynamicSecretProviderDTO = {
gatewayService: Pick<TGatewayServiceFactory, "fnGetGatewayClientTlsByGatewayId">; gatewayService: Pick<TGatewayServiceFactory, "fnGetGatewayClientTlsByGatewayId">;
gatewayV2Service: Pick<TGatewayV2ServiceFactory, "getPlatformConnectionDetailsByGatewayId">;
}; };
export const buildDynamicSecretProviders = ({ export const buildDynamicSecretProviders = ({
gatewayService gatewayService,
gatewayV2Service
}: TBuildDynamicSecretProviderDTO): Record<DynamicSecretProviders, TDynamicProviderFns> => ({ }: TBuildDynamicSecretProviderDTO): Record<DynamicSecretProviders, TDynamicProviderFns> => ({
[DynamicSecretProviders.SqlDatabase]: SqlDatabaseProvider({ gatewayService }), [DynamicSecretProviders.SqlDatabase]: SqlDatabaseProvider({ gatewayService, gatewayV2Service }),
[DynamicSecretProviders.Cassandra]: CassandraProvider(), [DynamicSecretProviders.Cassandra]: CassandraProvider(),
[DynamicSecretProviders.AwsIam]: AwsIamProvider(), [DynamicSecretProviders.AwsIam]: AwsIamProvider(),
[DynamicSecretProviders.Redis]: RedisDatabaseProvider(), [DynamicSecretProviders.Redis]: RedisDatabaseProvider(),
@@ -44,7 +47,7 @@ export const buildDynamicSecretProviders = ({
[DynamicSecretProviders.Snowflake]: SnowflakeProvider(), [DynamicSecretProviders.Snowflake]: SnowflakeProvider(),
[DynamicSecretProviders.Totp]: TotpProvider(), [DynamicSecretProviders.Totp]: TotpProvider(),
[DynamicSecretProviders.SapAse]: SapAseProvider(), [DynamicSecretProviders.SapAse]: SapAseProvider(),
[DynamicSecretProviders.Kubernetes]: KubernetesProvider({ gatewayService }), [DynamicSecretProviders.Kubernetes]: KubernetesProvider({ gatewayService, gatewayV2Service }),
[DynamicSecretProviders.Vertica]: VerticaProvider({ gatewayService }), [DynamicSecretProviders.Vertica]: VerticaProvider({ gatewayService }),
[DynamicSecretProviders.GcpIam]: GcpIamProvider(), [DynamicSecretProviders.GcpIam]: GcpIamProvider(),
[DynamicSecretProviders.Github]: GithubProvider(), [DynamicSecretProviders.Github]: GithubProvider(),
@@ -5,12 +5,14 @@ import https from "https";
import { BadRequestError } from "@app/lib/errors"; import { BadRequestError } from "@app/lib/errors";
import { sanitizeString } from "@app/lib/fn"; import { sanitizeString } from "@app/lib/fn";
import { GatewayHttpProxyActions, GatewayProxyProtocol, withGatewayProxy } from "@app/lib/gateway"; import { GatewayHttpProxyActions, GatewayProxyProtocol, withGatewayProxy } from "@app/lib/gateway";
import { withGatewayV2Proxy } from "@app/lib/gateway-v2/gateway-v2";
import { alphaNumericNanoId } from "@app/lib/nanoid"; import { alphaNumericNanoId } from "@app/lib/nanoid";
import { blockLocalAndPrivateIpAddresses } from "@app/lib/validator"; import { blockLocalAndPrivateIpAddresses } from "@app/lib/validator";
import { TKubernetesTokenRequest } from "@app/services/identity-kubernetes-auth/identity-kubernetes-auth-types"; import { TKubernetesTokenRequest } from "@app/services/identity-kubernetes-auth/identity-kubernetes-auth-types";
import { TDynamicSecretKubernetesLeaseConfig } from "../../dynamic-secret-lease/dynamic-secret-lease-types"; import { TDynamicSecretKubernetesLeaseConfig } from "../../dynamic-secret-lease/dynamic-secret-lease-types";
import { TGatewayServiceFactory } from "../../gateway/gateway-service"; import { TGatewayServiceFactory } from "../../gateway/gateway-service";
import { TGatewayV2ServiceFactory } from "../../gateway-v2/gateway-v2-service";
import { import {
DynamicSecretKubernetesSchema, DynamicSecretKubernetesSchema,
KubernetesAuthMethod, KubernetesAuthMethod,
@@ -26,6 +28,7 @@ const GATEWAY_AUTH_DEFAULT_URL = "https://kubernetes.default.svc.cluster.local";
type TKubernetesProviderDTO = { type TKubernetesProviderDTO = {
gatewayService: Pick<TGatewayServiceFactory, "fnGetGatewayClientTlsByGatewayId">; gatewayService: Pick<TGatewayServiceFactory, "fnGetGatewayClientTlsByGatewayId">;
gatewayV2Service: Pick<TGatewayV2ServiceFactory, "getPlatformConnectionDetailsByGatewayId">;
}; };
const generateUsername = (usernameTemplate?: string | null) => { const generateUsername = (usernameTemplate?: string | null) => {
@@ -38,7 +41,10 @@ const generateUsername = (usernameTemplate?: string | null) => {
}); });
}; };
export const KubernetesProvider = ({ gatewayService }: TKubernetesProviderDTO): TDynamicProviderFns => { export const KubernetesProvider = ({
gatewayService,
gatewayV2Service
}: TKubernetesProviderDTO): TDynamicProviderFns => {
const validateProviderInputs = async (inputs: unknown) => { const validateProviderInputs = async (inputs: unknown) => {
const providerInputs = await DynamicSecretKubernetesSchema.parseAsync(inputs); const providerInputs = await DynamicSecretKubernetesSchema.parseAsync(inputs);
if (!providerInputs.gatewayId && providerInputs.url) { if (!providerInputs.gatewayId && providerInputs.url) {
@@ -58,6 +64,32 @@ export const KubernetesProvider = ({ gatewayService }: TKubernetesProviderDTO):
}, },
gatewayCallback: (host: string, port: number, httpsAgent?: https.Agent) => Promise<T> gatewayCallback: (host: string, port: number, httpsAgent?: https.Agent) => Promise<T>
): Promise<T> => { ): Promise<T> => {
const gatewayV2ConnectionDetails = await gatewayV2Service.getPlatformConnectionDetailsByGatewayId({
gatewayId: inputs.gatewayId,
targetHost: inputs.targetHost,
targetPort: inputs.targetPort
});
if (gatewayV2ConnectionDetails) {
const callbackResult = await withGatewayV2Proxy(
async (port) => {
return gatewayCallback(
inputs.reviewTokenThroughGateway ? "http://localhost" : "https://localhost",
port,
inputs.httpsAgent
);
},
{
relayHost: gatewayV2ConnectionDetails.relayHost,
gateway: gatewayV2ConnectionDetails.gateway,
relay: gatewayV2ConnectionDetails.relay,
protocol: inputs.reviewTokenThroughGateway ? GatewayProxyProtocol.Http : GatewayProxyProtocol.Tcp,
httpsAgent: inputs.httpsAgent
}
);
return callbackResult;
}
const relayDetails = await gatewayService.fnGetGatewayClientTlsByGatewayId(inputs.gatewayId); const relayDetails = await gatewayService.fnGetGatewayClientTlsByGatewayId(inputs.gatewayId);
const [relayHost, relayPort] = relayDetails.relayAddress.split(":"); const [relayHost, relayPort] = relayDetails.relayAddress.split(":");
@@ -353,8 +385,18 @@ export const KubernetesProvider = ({ gatewayService }: TKubernetesProviderDTO):
return true; return true;
} catch (error) { } catch (error) {
let errorMessage = error instanceof Error ? error.message : "Unknown error"; let errorMessage = error instanceof Error ? error.message : "Unknown error";
if (axios.isAxiosError(error) && (error.response?.data as { message: string })?.message) { if (axios.isAxiosError(error)) {
errorMessage = (error.response?.data as { message: string }).message; if (error.response) {
let { message } = error?.response?.data as unknown as { message?: string };
if (!message && typeof error.response.data === "string") {
message = error.response.data;
}
if (message) {
errorMessage = message;
}
}
} }
const sanitizedErrorMessage = sanitizeString({ const sanitizedErrorMessage = sanitizeString({
@@ -603,8 +645,18 @@ export const KubernetesProvider = ({ gatewayService }: TKubernetesProviderDTO):
}; };
} catch (error) { } catch (error) {
let errorMessage = error instanceof Error ? error.message : "Unknown error"; let errorMessage = error instanceof Error ? error.message : "Unknown error";
if (axios.isAxiosError(error) && (error.response?.data as { message: string })?.message) { if (axios.isAxiosError(error)) {
errorMessage = (error.response?.data as { message: string }).message; if (error.response) {
let { message } = error?.response?.data as unknown as { message?: string };
if (!message && typeof error.response.data === "string") {
message = error.response.data;
}
if (message) {
errorMessage = message;
}
}
} }
const sanitizedErrorMessage = sanitizeString({ const sanitizedErrorMessage = sanitizeString({
@@ -740,8 +792,18 @@ export const KubernetesProvider = ({ gatewayService }: TKubernetesProviderDTO):
} }
} catch (error) { } catch (error) {
let errorMessage = error instanceof Error ? error.message : "Unknown error"; let errorMessage = error instanceof Error ? error.message : "Unknown error";
if (axios.isAxiosError(error) && (error.response?.data as { message: string })?.message) { if (axios.isAxiosError(error)) {
errorMessage = (error.response?.data as { message: string }).message; if (error.response) {
let { message } = error?.response?.data as unknown as { message?: string };
if (!message && typeof error.response.data === "string") {
message = error.response.data;
}
if (message) {
errorMessage = message;
}
}
} }
const sanitizedErrorMessage = sanitizeString({ const sanitizedErrorMessage = sanitizeString({
@@ -165,6 +165,7 @@ export const DynamicSecretSqlDBSchema = z.object({
revocationStatement: z.string().trim(), revocationStatement: z.string().trim(),
renewStatement: z.string().trim().optional(), renewStatement: z.string().trim().optional(),
ca: z.string().optional(), ca: z.string().optional(),
sslEnabled: z.boolean().optional(),
gatewayId: z.string().nullable().optional() gatewayId: z.string().nullable().optional()
}); });
@@ -275,11 +276,11 @@ export const DynamicSecretMongoAtlasSchema = z.object({
export const DynamicSecretMongoDBSchema = z.object({ export const DynamicSecretMongoDBSchema = z.object({
host: z.string().min(1).trim().toLowerCase(), host: z.string().min(1).trim().toLowerCase(),
port: z.number().optional(), port: z.number().optional().nullable(),
username: z.string().min(1).trim(), username: z.string().min(1).trim(),
password: z.string().min(1).trim(), password: z.string().min(1).trim(),
database: z.string().min(1).trim(), database: z.string().min(1).trim(),
ca: z.string().min(1).optional(), ca: z.string().trim().optional().nullable(),
roles: z roles: z
.string() .string()
.array() .array()
@@ -44,7 +44,7 @@ export const MongoDBProvider = (): TDynamicProviderFns => {
password: providerInputs.password password: providerInputs.password
}, },
directConnection: !isSrv, directConnection: !isSrv,
ca: providerInputs.ca ca: providerInputs.ca || undefined
}); });
return client; return client;
}; };
@@ -1,15 +1,18 @@
import handlebars from "handlebars"; import handlebars from "handlebars";
import knex from "knex"; import knex from "knex";
import RE2 from "re2";
import { z } from "zod"; import { z } from "zod";
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";
import { sanitizeString } from "@app/lib/fn"; import { sanitizeString } from "@app/lib/fn";
import { GatewayProxyProtocol, withGatewayProxy } from "@app/lib/gateway"; import { GatewayProxyProtocol, withGatewayProxy } from "@app/lib/gateway";
import { withGatewayV2Proxy } from "@app/lib/gateway-v2/gateway-v2";
import { alphaNumericNanoId } from "@app/lib/nanoid"; import { alphaNumericNanoId } from "@app/lib/nanoid";
import { validateHandlebarTemplate } from "@app/lib/template/validate-handlebars"; import { validateHandlebarTemplate } from "@app/lib/template/validate-handlebars";
import { TGatewayServiceFactory } from "../../gateway/gateway-service"; import { TGatewayServiceFactory } from "../../gateway/gateway-service";
import { TGatewayV2ServiceFactory } from "../../gateway-v2/gateway-v2-service";
import { verifyHostInputValidity } from "../dynamic-secret-fns"; import { verifyHostInputValidity } from "../dynamic-secret-fns";
import { DynamicSecretSqlDBSchema, PasswordRequirements, SqlProviders, TDynamicProviderFns } from "./models"; import { DynamicSecretSqlDBSchema, PasswordRequirements, SqlProviders, TDynamicProviderFns } from "./models";
import { compileUsernameTemplate } from "./templateUtils"; import { compileUsernameTemplate } from "./templateUtils";
@@ -128,9 +131,13 @@ const generateUsername = (provider: SqlProviders, usernameTemplate?: string | nu
type TSqlDatabaseProviderDTO = { type TSqlDatabaseProviderDTO = {
gatewayService: Pick<TGatewayServiceFactory, "fnGetGatewayClientTlsByGatewayId">; gatewayService: Pick<TGatewayServiceFactory, "fnGetGatewayClientTlsByGatewayId">;
gatewayV2Service: Pick<TGatewayV2ServiceFactory, "getPlatformConnectionDetailsByGatewayId">;
}; };
export const SqlDatabaseProvider = ({ gatewayService }: TSqlDatabaseProviderDTO): TDynamicProviderFns => { export const SqlDatabaseProvider = ({
gatewayService,
gatewayV2Service
}: TSqlDatabaseProviderDTO): TDynamicProviderFns => {
const validateProviderInputs = async (inputs: unknown) => { const validateProviderInputs = async (inputs: unknown) => {
const providerInputs = await DynamicSecretSqlDBSchema.parseAsync(inputs); const providerInputs = await DynamicSecretSqlDBSchema.parseAsync(inputs);
@@ -150,19 +157,40 @@ export const SqlDatabaseProvider = ({ gatewayService }: TSqlDatabaseProviderDTO)
return { ...providerInputs, hostIp }; return { ...providerInputs, hostIp };
}; };
const $getClient = async (providerInputs: z.infer<typeof DynamicSecretSqlDBSchema> & { hostIp: string }) => { const $getClient = async (
providerInputs: z.infer<typeof DynamicSecretSqlDBSchema> & { hostIp: string; originalHost: string }
) => {
const ssl = providerInputs.ca const ssl = providerInputs.ca
? { rejectUnauthorized: false, ca: providerInputs.ca, servername: providerInputs.host } ? { rejectUnauthorized: false, ca: providerInputs.ca, servername: providerInputs.host }
: undefined; : undefined;
const isMsSQLClient = providerInputs.client === SqlProviders.MsSQL; const isMsSQLClient = providerInputs.client === SqlProviders.MsSQL;
/*
We route through the gateway by setting connection.host = "localhost".
Azure SQL identifies the logical server from the TDS login name when the host
isn’t the Azure FQDN. Therefore, when using the gateway, ensure username is
"user@<azure-server-name>" so Azure opens the correct logical server.
Direct connections to the Azure FQDN usually don’t require this suffix.
*/
const isAzureSql = isMsSQLClient && new RE2(/\.database\.windows\.net$/i).test(providerInputs.originalHost);
const azureServerLabel =
isAzureSql && providerInputs.gatewayId ? providerInputs.originalHost?.split(".")[0] : undefined;
const effectiveUser =
isAzureSql && !providerInputs.username.includes("@") && azureServerLabel
? `${providerInputs.username}@${azureServerLabel}`
: providerInputs.username;
const db = knex({ const db = knex({
client: providerInputs.client, client: providerInputs.client,
connection: { connection: {
database: providerInputs.database, database: providerInputs.database,
port: providerInputs.port, port: providerInputs.port,
host: providerInputs.client === SqlProviders.Postgres ? providerInputs.hostIp : providerInputs.host, host:
user: providerInputs.username, providerInputs.client === SqlProviders.Postgres && !providerInputs.gatewayId
? providerInputs.hostIp
: providerInputs.host,
user: effectiveUser,
password: providerInputs.password, password: providerInputs.password,
ssl, ssl,
// @ts-expect-error this is because of knexjs type signature issue. This is directly passed to driver // @ts-expect-error this is because of knexjs type signature issue. This is directly passed to driver
@@ -170,6 +198,7 @@ export const SqlDatabaseProvider = ({ gatewayService }: TSqlDatabaseProviderDTO)
// https://github.com/tediousjs/tedious/blob/ebb023ed90969a7ec0e4b036533ad52739d921f7/test/config.ci.ts#L19 // https://github.com/tediousjs/tedious/blob/ebb023ed90969a7ec0e4b036533ad52739d921f7/test/config.ci.ts#L19
options: isMsSQLClient options: isMsSQLClient
? { ? {
...(providerInputs.sslEnabled !== undefined ? { encrypt: providerInputs.sslEnabled } : {}),
trustServerCertificate: !providerInputs.ca, trustServerCertificate: !providerInputs.ca,
cryptoCredentialsDetails: providerInputs.ca ? { ca: providerInputs.ca } : {} cryptoCredentialsDetails: providerInputs.ca ? { ca: providerInputs.ca } : {}
} }
@@ -185,6 +214,26 @@ export const SqlDatabaseProvider = ({ gatewayService }: TSqlDatabaseProviderDTO)
providerInputs: z.infer<typeof DynamicSecretSqlDBSchema>, providerInputs: z.infer<typeof DynamicSecretSqlDBSchema>,
gatewayCallback: (host: string, port: number) => Promise<void> gatewayCallback: (host: string, port: number) => Promise<void>
) => { ) => {
const gatewayV2ConnectionDetails = await gatewayV2Service.getPlatformConnectionDetailsByGatewayId({
gatewayId: providerInputs.gatewayId as string,
targetHost: providerInputs.host,
targetPort: providerInputs.port
});
if (gatewayV2ConnectionDetails) {
return withGatewayV2Proxy(
async (port) => {
await gatewayCallback("localhost", port);
},
{
relayHost: gatewayV2ConnectionDetails.relayHost,
gateway: gatewayV2ConnectionDetails.gateway,
relay: gatewayV2ConnectionDetails.relay,
protocol: GatewayProxyProtocol.Tcp
}
);
}
const relayDetails = await gatewayService.fnGetGatewayClientTlsByGatewayId(providerInputs.gatewayId as string); const relayDetails = await gatewayService.fnGetGatewayClientTlsByGatewayId(providerInputs.gatewayId as string);
const [relayHost, relayPort] = relayDetails.relayAddress.split(":"); const [relayHost, relayPort] = relayDetails.relayAddress.split(":");
await withGatewayProxy( await withGatewayProxy(
@@ -212,7 +261,13 @@ export const SqlDatabaseProvider = ({ gatewayService }: TSqlDatabaseProviderDTO)
const providerInputs = await validateProviderInputs(inputs); const providerInputs = await validateProviderInputs(inputs);
let isConnected = false; let isConnected = false;
const gatewayCallback = async (host = providerInputs.host, port = providerInputs.port) => { const gatewayCallback = async (host = providerInputs.host, port = providerInputs.port) => {
const db = await $getClient({ ...providerInputs, port, host, hostIp: providerInputs.hostIp }); const db = await $getClient({
...providerInputs,
port,
host,
hostIp: providerInputs.hostIp,
originalHost: providerInputs.host
});
// oracle needs from keyword // oracle needs from keyword
const testStatement = providerInputs.client === SqlProviders.Oracle ? "SELECT 1 FROM DUAL" : "SELECT 1"; const testStatement = providerInputs.client === SqlProviders.Oracle ? "SELECT 1 FROM DUAL" : "SELECT 1";
@@ -253,7 +308,12 @@ export const SqlDatabaseProvider = ({ gatewayService }: TSqlDatabaseProviderDTO)
const password = generatePassword(providerInputs.client, providerInputs.passwordRequirements); const password = generatePassword(providerInputs.client, providerInputs.passwordRequirements);
const gatewayCallback = async (host = providerInputs.host, port = providerInputs.port) => { const gatewayCallback = async (host = providerInputs.host, port = providerInputs.port) => {
const db = await $getClient({ ...providerInputs, port, host }); const db = await $getClient({
...providerInputs,
port,
host,
originalHost: providerInputs.host
});
try { try {
const expiration = new Date(expireAt).toISOString(); const expiration = new Date(expireAt).toISOString();
@@ -296,7 +356,12 @@ export const SqlDatabaseProvider = ({ gatewayService }: TSqlDatabaseProviderDTO)
const username = entityId; const username = entityId;
const { database } = providerInputs; const { database } = providerInputs;
const gatewayCallback = async (host = providerInputs.host, port = providerInputs.port) => { const gatewayCallback = async (host = providerInputs.host, port = providerInputs.port) => {
const db = await $getClient({ ...providerInputs, port, host }); const db = await $getClient({
...providerInputs,
port,
host,
originalHost: providerInputs.host
});
try { try {
const revokeStatement = handlebars.compile(providerInputs.revocationStatement)({ username, database }); const revokeStatement = handlebars.compile(providerInputs.revocationStatement)({ username, database });
const queries = revokeStatement.toString().split(";").filter(Boolean); const queries = revokeStatement.toString().split(";").filter(Boolean);
@@ -331,7 +396,12 @@ export const SqlDatabaseProvider = ({ gatewayService }: TSqlDatabaseProviderDTO)
if (!providerInputs.renewStatement) return { entityId }; if (!providerInputs.renewStatement) return { entityId };
const gatewayCallback = async (host = providerInputs.host, port = providerInputs.port) => { const gatewayCallback = async (host = providerInputs.host, port = providerInputs.port) => {
const db = await $getClient({ ...providerInputs, port, host }); const db = await $getClient({
...providerInputs,
port,
host,
originalHost: providerInputs.host
});
const expiration = new Date(expireAt).toISOString(); const expiration = new Date(expireAt).toISOString();
const { database } = providerInputs; const { database } = providerInputs;
@@ -0,0 +1,2 @@
export const GATEWAY_ROUTING_INFO_OID = "1.3.6.1.4.1.12345.100.1";
export const GATEWAY_ACTOR_OID = "1.3.6.1.4.1.12345.100.2";
@@ -0,0 +1,60 @@
import { Knex } from "knex";
import { TDbClient } from "@app/db";
import { GatewaysV2Schema, TableName, TGatewaysV2 } from "@app/db/schemas";
import { DatabaseError } from "@app/lib/errors";
import { buildFindFilter, ormify, selectAllTableCols, TFindFilter, TFindOpt } from "@app/lib/knex";
export type TGatewayV2DALFactory = ReturnType<typeof gatewayV2DalFactory>;
export const gatewayV2DalFactory = (db: TDbClient) => {
const orm = ormify(db, TableName.GatewayV2);
const find = async (filter: TFindFilter<TGatewaysV2>, { offset, limit, sort, tx }: TFindOpt<TGatewaysV2> = {}) => {
try {
const query = (tx || db.replicaNode())(TableName.GatewayV2)
// eslint-disable-next-line @typescript-eslint/no-misused-promises
.where(buildFindFilter(filter, TableName.GatewayV2))
.join(TableName.Identity, `${TableName.Identity}.id`, `${TableName.GatewayV2}.identityId`)
.join(
TableName.IdentityOrgMembership,
`${TableName.IdentityOrgMembership}.identityId`,
`${TableName.GatewayV2}.identityId`
)
.select(selectAllTableCols(TableName.GatewayV2))
.select(db.ref("name").withSchema(TableName.Identity).as("identityName"));
if (limit) void query.limit(limit);
if (offset) void query.offset(offset);
if (sort) {
void query.orderBy(sort.map(([column, order, nulls]) => ({ column: column as string, order, nulls })));
}
const docs = await query;
return docs.map((el) => ({
...GatewaysV2Schema.parse(el),
identity: { id: el.identityId, name: el.identityName }
}));
} catch (error) {
throw new DatabaseError({ error, name: `${TableName.GatewayV2}: Find` });
}
};
const findById = async (id: string, tx?: Knex) => {
try {
const doc = await (tx || db.replicaNode())(TableName.GatewayV2)
.join(TableName.Organization, `${TableName.GatewayV2}.orgId`, `${TableName.Organization}.id`)
.where(`${TableName.GatewayV2}.id`, id)
.select(selectAllTableCols(TableName.GatewayV2))
.select(db.ref("name").withSchema(TableName.Organization).as("orgName"))
.first();
return doc;
} catch (error) {
throw new DatabaseError({ error, name: `${TableName.GatewayV2}: Find by id` });
}
};
return { ...orm, find, findById };
};
@@ -0,0 +1,656 @@
import net from "node:net";
import { ForbiddenError } from "@casl/ability";
import * as x509 from "@peculiar/x509";
import { TRelays } from "@app/db/schemas";
import { PgSqlLock } from "@app/keystore/keystore";
import { crypto } from "@app/lib/crypto";
import { DatabaseErrorCode } from "@app/lib/error-codes";
import { BadRequestError, DatabaseError, NotFoundError } from "@app/lib/errors";
import { GatewayProxyProtocol } from "@app/lib/gateway/types";
import { withGatewayV2Proxy } from "@app/lib/gateway-v2/gateway-v2";
import { OrgServiceActor } from "@app/lib/types";
import { ActorAuthMethod, ActorType } from "@app/services/auth/auth-type";
import { constructPemChainFromCerts } from "@app/services/certificate/certificate-fns";
import { CertExtendedKeyUsage, CertKeyAlgorithm, CertKeyUsage } from "@app/services/certificate/certificate-types";
import {
createSerialNumber,
keyAlgorithmToAlgCfg
} from "@app/services/certificate-authority/certificate-authority-fns";
import { TKmsServiceFactory } from "@app/services/kms/kms-service";
import { KmsDataKey } from "@app/services/kms/kms-types";
import { TLicenseServiceFactory } from "../license/license-service";
import { OrgPermissionGatewayActions, OrgPermissionSubjects } from "../permission/org-permission";
import { TPermissionServiceFactory } from "../permission/permission-service-types";
import { TRelayDALFactory } from "../relay/relay-dal";
import { TRelayServiceFactory } from "../relay/relay-service";
import { GATEWAY_ACTOR_OID, GATEWAY_ROUTING_INFO_OID } from "./gateway-v2-constants";
import { TGatewayV2DALFactory } from "./gateway-v2-dal";
import { TOrgGatewayConfigV2DALFactory } from "./org-gateway-config-v2-dal";
type TGatewayV2ServiceFactoryDep = {
orgGatewayConfigV2DAL: Pick<TOrgGatewayConfigV2DALFactory, "findOne" | "create" | "transaction" | "findById">;
licenseService: Pick<TLicenseServiceFactory, "onPremFeatures" | "getPlan">;
kmsService: TKmsServiceFactory;
relayService: TRelayServiceFactory;
gatewayV2DAL: TGatewayV2DALFactory;
relayDAL: TRelayDALFactory;
permissionService: TPermissionServiceFactory;
};
export type TGatewayV2ServiceFactory = ReturnType<typeof gatewayV2ServiceFactory>;
export const gatewayV2ServiceFactory = ({
orgGatewayConfigV2DAL,
licenseService,
kmsService,
relayService,
gatewayV2DAL,
relayDAL,
permissionService
}: TGatewayV2ServiceFactoryDep) => {
const $validateIdentityAccessToGateway = async (orgId: string, actorId: string, actorAuthMethod: ActorAuthMethod) => {
const orgLicensePlan = await licenseService.getPlan(orgId);
if (!orgLicensePlan.gateway) {
throw new BadRequestError({
message:
"Gateway operation failed due to organization plan restrictions. Please upgrade your instance to Infisical's Enterprise plan."
});
}
const { permission } = await permissionService.getOrgPermission(
ActorType.IDENTITY,
actorId,
orgId,
actorAuthMethod,
orgId
);
ForbiddenError.from(permission).throwUnlessCan(
OrgPermissionGatewayActions.CreateGateways,
OrgPermissionSubjects.Gateway
);
};
const $getOrgCAs = async (orgId: string) => {
const { encryptor: orgKmsEncryptor, decryptor: orgKmsDecryptor } = await kmsService.createCipherPairWithDataKey({
type: KmsDataKey.Organization,
orgId
});
const orgCAs = await orgGatewayConfigV2DAL.transaction(async (tx) => {
const orgGatewayConfigV2 = await orgGatewayConfigV2DAL.findOne({ orgId });
if (orgGatewayConfigV2) return orgGatewayConfigV2;
await tx.raw("SELECT pg_advisory_xact_lock(?)", [PgSqlLock.OrgGatewayV2Init(orgId)]);
// generate root CA
const rootCaKeyAlgorithm = CertKeyAlgorithm.RSA_2048;
const alg = keyAlgorithmToAlgCfg(rootCaKeyAlgorithm);
const rootCaKeys = await crypto.nativeCrypto.subtle.generateKey(alg, true, ["sign", "verify"]);
const rootCaSerialNumber = createSerialNumber();
const rootCaSkObj = crypto.nativeCrypto.KeyObject.from(rootCaKeys.privateKey);
const rootCaIssuedAt = new Date();
const rootCaExpiration = new Date(new Date().setFullYear(2045));
const rootCaCert = await x509.X509CertificateGenerator.createSelfSigned({
name: `O=${orgId},CN=Infisical Gateway Root CA`,
serialNumber: rootCaSerialNumber,
notBefore: rootCaIssuedAt,
notAfter: rootCaExpiration,
signingAlgorithm: alg,
keys: rootCaKeys,
extensions: [
// eslint-disable-next-line no-bitwise
new x509.KeyUsagesExtension(x509.KeyUsageFlags.keyCertSign | x509.KeyUsageFlags.cRLSign, true),
await x509.SubjectKeyIdentifierExtension.create(rootCaKeys.publicKey)
]
});
// generate server CA
const serverCaSerialNumber = createSerialNumber();
const serverCaIssuedAt = new Date();
const serverCaExpiration = new Date(new Date().setFullYear(2045));
const serverCaKeys = await crypto.nativeCrypto.subtle.generateKey(alg, true, ["sign", "verify"]);
const serverCaSkObj = crypto.nativeCrypto.KeyObject.from(serverCaKeys.privateKey);
const serverCaCert = await x509.X509CertificateGenerator.create({
serialNumber: serverCaSerialNumber,
subject: `O=${orgId},CN=Infisical Gateway Server CA`,
issuer: rootCaCert.subject,
notBefore: serverCaIssuedAt,
notAfter: serverCaExpiration,
signingKey: rootCaKeys.privateKey,
publicKey: serverCaKeys.publicKey,
signingAlgorithm: alg,
extensions: [
new x509.KeyUsagesExtension(
// eslint-disable-next-line no-bitwise
x509.KeyUsageFlags.keyCertSign |
x509.KeyUsageFlags.cRLSign |
x509.KeyUsageFlags.digitalSignature |
x509.KeyUsageFlags.keyEncipherment,
true
),
new x509.BasicConstraintsExtension(true, 0, true),
await x509.AuthorityKeyIdentifierExtension.create(rootCaCert, false),
await x509.SubjectKeyIdentifierExtension.create(serverCaKeys.publicKey)
]
});
// generate client CA
const clientCaSerialNumber = createSerialNumber();
const clientCaIssuedAt = new Date();
const clientCaExpiration = new Date(new Date().setFullYear(2045));
const clientCaKeys = await crypto.nativeCrypto.subtle.generateKey(alg, true, ["sign", "verify"]);
const clientCaSkObj = crypto.nativeCrypto.KeyObject.from(clientCaKeys.privateKey);
const clientCaCert = await x509.X509CertificateGenerator.create({
serialNumber: clientCaSerialNumber,
subject: `O=${orgId},CN=Infisical Gateway Client CA`,
issuer: rootCaCert.subject,
notBefore: clientCaIssuedAt,
notAfter: clientCaExpiration,
signingKey: rootCaKeys.privateKey,
publicKey: clientCaKeys.publicKey,
signingAlgorithm: alg,
extensions: [
new x509.KeyUsagesExtension(
// eslint-disable-next-line no-bitwise
x509.KeyUsageFlags.keyCertSign |
x509.KeyUsageFlags.cRLSign |
x509.KeyUsageFlags.digitalSignature |
x509.KeyUsageFlags.keyEncipherment,
true
),
new x509.BasicConstraintsExtension(true, 0, true),
await x509.AuthorityKeyIdentifierExtension.create(rootCaCert, false),
await x509.SubjectKeyIdentifierExtension.create(clientCaKeys.publicKey)
]
});
const encryptedRootGatewayCaPrivateKey = orgKmsEncryptor({
plainText: Buffer.from(
rootCaSkObj.export({
type: "pkcs8",
format: "der"
})
)
}).cipherTextBlob;
const encryptedRootGatewayCaCertificate = orgKmsEncryptor({
plainText: Buffer.from(rootCaCert.rawData)
}).cipherTextBlob;
const encryptedGatewayServerCaPrivateKey = orgKmsEncryptor({
plainText: Buffer.from(serverCaSkObj.export({ type: "pkcs8", format: "der" }))
}).cipherTextBlob;
const encryptedGatewayServerCaCertificate = orgKmsEncryptor({
plainText: Buffer.from(serverCaCert.rawData)
}).cipherTextBlob;
const encryptedGatewayServerCaCertificateChain = orgKmsEncryptor({
plainText: Buffer.from(constructPemChainFromCerts([rootCaCert]))
}).cipherTextBlob;
const encryptedGatewayClientCaPrivateKey = orgKmsEncryptor({
plainText: Buffer.from(clientCaSkObj.export({ type: "pkcs8", format: "der" }))
}).cipherTextBlob;
const encryptedGatewayClientCaCertificate = orgKmsEncryptor({
plainText: Buffer.from(clientCaCert.rawData)
}).cipherTextBlob;
const encryptedGatewayClientCaCertificateChain = orgKmsEncryptor({
plainText: Buffer.from(constructPemChainFromCerts([rootCaCert]))
}).cipherTextBlob;
return orgGatewayConfigV2DAL.create({
orgId,
encryptedRootGatewayCaPrivateKey,
encryptedRootGatewayCaCertificate,
encryptedGatewayServerCaPrivateKey,
encryptedGatewayServerCaCertificate,
encryptedGatewayServerCaCertificateChain,
encryptedGatewayClientCaPrivateKey,
encryptedGatewayClientCaCertificate,
encryptedGatewayClientCaCertificateChain
});
});
const rootGatewayCaPrivateKey = orgKmsDecryptor({ cipherTextBlob: orgCAs.encryptedRootGatewayCaPrivateKey });
const rootGatewayCaCertificate = orgKmsDecryptor({ cipherTextBlob: orgCAs.encryptedRootGatewayCaCertificate });
const gatewayServerCaPrivateKey = orgKmsDecryptor({ cipherTextBlob: orgCAs.encryptedGatewayServerCaPrivateKey });
const gatewayServerCaCertificate = orgKmsDecryptor({ cipherTextBlob: orgCAs.encryptedGatewayServerCaCertificate });
const gatewayServerCaCertificateChain = orgKmsDecryptor({
cipherTextBlob: orgCAs.encryptedGatewayServerCaCertificateChain
});
const gatewayClientCaPrivateKey = orgKmsDecryptor({ cipherTextBlob: orgCAs.encryptedGatewayClientCaPrivateKey });
const gatewayClientCaCertificate = orgKmsDecryptor({
cipherTextBlob: orgCAs.encryptedGatewayClientCaCertificate
});
const gatewayClientCaCertificateChain = orgKmsDecryptor({
cipherTextBlob: orgCAs.encryptedGatewayClientCaCertificateChain
});
return {
rootGatewayCaPrivateKey,
rootGatewayCaCertificate,
gatewayServerCaPrivateKey,
gatewayServerCaCertificate,
gatewayServerCaCertificateChain,
gatewayClientCaPrivateKey,
gatewayClientCaCertificate,
gatewayClientCaCertificateChain
};
};
const listGateways = async ({ orgPermission }: { orgPermission: OrgServiceActor }) => {
const { permission } = await permissionService.getOrgPermission(
orgPermission.type,
orgPermission.id,
orgPermission.orgId,
orgPermission.authMethod,
orgPermission.orgId
);
ForbiddenError.from(permission).throwUnlessCan(
OrgPermissionGatewayActions.ListGateways,
OrgPermissionSubjects.Gateway
);
const gateways = await gatewayV2DAL.find({
orgId: orgPermission.orgId
});
return gateways;
};
const getPlatformConnectionDetailsByGatewayId = async ({
gatewayId,
targetHost,
targetPort
}: {
gatewayId: string;
targetHost: string;
targetPort: number;
}) => {
const gateway = await gatewayV2DAL.findById(gatewayId);
if (!gateway) {
return;
}
const orgGatewayConfig = await orgGatewayConfigV2DAL.findOne({ orgId: gateway.orgId });
if (!orgGatewayConfig) {
throw new NotFoundError({ message: `Gateway Config for org ${gateway.orgId} not found.` });
}
if (!gateway.relayId) {
throw new BadRequestError({
message: "Gateway is not associated with a relay"
});
}
const orgLicensePlan = await licenseService.getPlan(orgGatewayConfig.orgId);
if (!orgLicensePlan.gateway) {
throw new BadRequestError({
message: "Please upgrade your instance to Infisical's Enterprise plan to use gateways."
});
}
const { decryptor: orgKmsDecryptor } = await kmsService.createCipherPairWithDataKey({
type: KmsDataKey.Organization,
orgId: orgGatewayConfig.orgId
});
const alg = keyAlgorithmToAlgCfg(CertKeyAlgorithm.RSA_2048);
const rootGatewayCaCert = new x509.X509Certificate(
orgKmsDecryptor({
cipherTextBlob: orgGatewayConfig.encryptedRootGatewayCaCertificate
})
);
const gatewayClientCaCert = new x509.X509Certificate(
orgKmsDecryptor({
cipherTextBlob: orgGatewayConfig.encryptedGatewayClientCaCertificate
})
);
const gatewayServerCaCert = new x509.X509Certificate(
orgKmsDecryptor({
cipherTextBlob: orgGatewayConfig.encryptedGatewayServerCaCertificate
})
);
const gatewayClientCaPrivateKey = orgKmsDecryptor({
cipherTextBlob: orgGatewayConfig.encryptedGatewayClientCaPrivateKey
});
const gatewayClientCaSkObj = crypto.nativeCrypto.createPrivateKey({
key: gatewayClientCaPrivateKey,
format: "der",
type: "pkcs8"
});
const importedGatewayClientCaPrivateKey = await crypto.nativeCrypto.subtle.importKey(
"pkcs8",
gatewayClientCaSkObj.export({ format: "der", type: "pkcs8" }),
alg,
true,
["sign"]
);
const clientCertIssuedAt = new Date();
const clientCertExpiration = new Date(new Date().getTime() + 5 * 60 * 1000);
const clientKeys = await crypto.nativeCrypto.subtle.generateKey(alg, true, ["sign", "verify"]);
const clientCertSerialNumber = createSerialNumber();
const routingInfo = {
targetHost,
targetPort
};
const routingExtension = new x509.Extension(
GATEWAY_ROUTING_INFO_OID,
false,
Buffer.from(JSON.stringify(routingInfo))
);
const actorExtension = new x509.Extension(
GATEWAY_ACTOR_OID,
false,
Buffer.from(JSON.stringify({ type: ActorType.PLATFORM }))
);
const clientCert = await x509.X509CertificateGenerator.create({
serialNumber: clientCertSerialNumber,
subject: `O=${orgGatewayConfig.orgId},OU=gateway-client,CN=${ActorType.PLATFORM}:${gatewayId}`,
issuer: gatewayClientCaCert.subject,
notAfter: clientCertExpiration,
notBefore: clientCertIssuedAt,
signingKey: importedGatewayClientCaPrivateKey,
publicKey: clientKeys.publicKey,
signingAlgorithm: alg,
extensions: [
new x509.BasicConstraintsExtension(false),
await x509.AuthorityKeyIdentifierExtension.create(gatewayClientCaCert, false),
await x509.SubjectKeyIdentifierExtension.create(clientKeys.publicKey),
new x509.CertificatePolicyExtension(["2.5.29.32.0"]), // anyPolicy
new x509.KeyUsagesExtension(
// eslint-disable-next-line no-bitwise
x509.KeyUsageFlags[CertKeyUsage.DIGITAL_SIGNATURE] |
x509.KeyUsageFlags[CertKeyUsage.KEY_ENCIPHERMENT] |
x509.KeyUsageFlags[CertKeyUsage.KEY_AGREEMENT],
true
),
new x509.ExtendedKeyUsageExtension([x509.ExtendedKeyUsage[CertExtendedKeyUsage.CLIENT_AUTH]], true),
routingExtension,
actorExtension
]
});
const gatewayClientCertPrivateKey = crypto.nativeCrypto.KeyObject.from(clientKeys.privateKey);
const relayCredentials = await relayService.getCredentialsForClient({
relayId: gateway.relayId,
orgId: gateway.orgId,
orgName: gateway.orgName,
gatewayId
});
return {
relayHost: relayCredentials.relayHost,
gateway: {
clientCertificate: clientCert.toString("pem"),
clientPrivateKey: gatewayClientCertPrivateKey.export({ format: "pem", type: "pkcs8" }).toString(),
serverCertificateChain: constructPemChainFromCerts([gatewayServerCaCert, rootGatewayCaCert])
},
relay: {
clientCertificate: relayCredentials.clientCertificate,
clientPrivateKey: relayCredentials.clientPrivateKey,
serverCertificateChain: relayCredentials.serverCertificateChain
}
};
};
const registerGateway = async ({
orgId,
actorId,
actorAuthMethod,
relayName,
name
}: {
orgId: string;
actorId: string;
actorAuthMethod: ActorAuthMethod;
relayName: string;
name: string;
}) => {
await $validateIdentityAccessToGateway(orgId, actorId, actorAuthMethod);
const orgCAs = await $getOrgCAs(orgId);
let relay: TRelays = await relayDAL.findOne({ orgId, name: relayName });
if (!relay) {
relay = await relayDAL.findOne({ name: relayName, orgId: null });
}
if (!relay) {
throw new NotFoundError({ message: `Relay ${relayName} not found` });
}
try {
const [gateway] = await gatewayV2DAL.upsert(
[
{
orgId,
name,
identityId: actorId,
relayId: relay.id
}
],
["identityId"]
);
const alg = keyAlgorithmToAlgCfg(CertKeyAlgorithm.RSA_2048);
const gatewayServerCaCert = new x509.X509Certificate(orgCAs.gatewayServerCaCertificate);
const rootGatewayCaCert = new x509.X509Certificate(orgCAs.rootGatewayCaCertificate);
const gatewayClientCaCert = new x509.X509Certificate(orgCAs.gatewayClientCaCertificate);
const gatewayServerCaSkObj = crypto.nativeCrypto.createPrivateKey({
key: orgCAs.gatewayServerCaPrivateKey,
format: "der",
type: "pkcs8"
});
const gatewayServerCaPrivateKey = await crypto.nativeCrypto.subtle.importKey(
"pkcs8",
gatewayServerCaSkObj.export({ format: "der", type: "pkcs8" }),
alg,
true,
["sign"]
);
const gatewayServerKeys = await crypto.nativeCrypto.subtle.generateKey(alg, true, ["sign", "verify"]);
const gatewayServerCertIssuedAt = new Date();
const gatewayServerCertExpireAt = new Date(new Date().setDate(new Date().getDate() + 1));
const gatewayServerCertPrivateKey = crypto.nativeCrypto.KeyObject.from(gatewayServerKeys.privateKey);
const gatewayServerCertExtensions: x509.Extension[] = [
new x509.BasicConstraintsExtension(false),
await x509.AuthorityKeyIdentifierExtension.create(gatewayServerCaCert, false),
await x509.SubjectKeyIdentifierExtension.create(gatewayServerKeys.publicKey),
new x509.CertificatePolicyExtension(["2.5.29.32.0"]), // anyPolicy
new x509.KeyUsagesExtension(
// eslint-disable-next-line no-bitwise
x509.KeyUsageFlags[CertKeyUsage.DIGITAL_SIGNATURE] | x509.KeyUsageFlags[CertKeyUsage.KEY_ENCIPHERMENT],
true
),
new x509.ExtendedKeyUsageExtension([x509.ExtendedKeyUsage[CertExtendedKeyUsage.SERVER_AUTH]], true),
new x509.SubjectAlternativeNameExtension([
{ type: "dns", value: "localhost" },
{ type: "ip", value: "127.0.0.1" },
{ type: "ip", value: "::1" }
])
];
const gatewayServerSerialNumber = createSerialNumber();
const gatewayServerCertificate = await x509.X509CertificateGenerator.create({
serialNumber: gatewayServerSerialNumber,
subject: `O=${orgId},CN=Gateway`,
issuer: gatewayServerCaCert.subject,
notBefore: gatewayServerCertIssuedAt,
notAfter: gatewayServerCertExpireAt,
signingKey: gatewayServerCaPrivateKey,
publicKey: gatewayServerKeys.publicKey,
signingAlgorithm: alg,
extensions: gatewayServerCertExtensions
});
const relayCredentials = await relayService.getCredentialsForGateway({
relayName,
orgId,
gatewayId: gateway.id
});
return {
gatewayId: gateway.id,
relayHost: relayCredentials.relayHost,
pki: {
serverCertificate: gatewayServerCertificate.toString("pem"),
serverPrivateKey: gatewayServerCertPrivateKey.export({ format: "pem", type: "pkcs8" }).toString(),
clientCertificateChain: constructPemChainFromCerts([gatewayClientCaCert, rootGatewayCaCert])
},
ssh: {
clientCertificate: relayCredentials.clientSshCert,
clientPrivateKey: relayCredentials.clientSshPrivateKey,
serverCAPublicKey: relayCredentials.serverCAPublicKey
}
};
} catch (err) {
if (err instanceof DatabaseError && (err.error as { code: string })?.code === DatabaseErrorCode.UniqueViolation) {
throw new BadRequestError({ message: `Gateway with name "${name}" already exists` });
}
throw err;
}
};
const heartbeat = async ({ orgPermission }: { orgPermission: OrgServiceActor }) => {
await $validateIdentityAccessToGateway(orgPermission.orgId, orgPermission.id, orgPermission.authMethod);
const gateway = await gatewayV2DAL.findOne({
orgId: orgPermission.orgId,
identityId: orgPermission.id
});
if (!gateway) {
throw new NotFoundError({ message: `Gateway for identity ${orgPermission.id} not found.` });
}
const gatewayV2ConnectionDetails = await getPlatformConnectionDetailsByGatewayId({
gatewayId: gateway.id,
targetHost: "health-check",
targetPort: 443
});
if (!gatewayV2ConnectionDetails) {
throw new NotFoundError({ message: `Gateway connection details for gateway ${gateway.id} not found.` });
}
const isGatewayReachable = await withGatewayV2Proxy(
async (port) => {
return new Promise<boolean>((resolve, reject) => {
const socket = new net.Socket();
let responseReceived = false;
let isResolved = false;
// Set socket timeout
socket.setTimeout(10000);
const cleanup = () => {
if (!socket.destroyed) {
socket.destroy();
}
};
socket.on("data", (data: Buffer) => {
const response = data.toString().trim();
if (response === "PONG" && !isResolved) {
isResolved = true;
responseReceived = true;
cleanup();
resolve(true);
}
});
socket.on("error", (err: Error) => {
if (!isResolved) {
isResolved = true;
cleanup();
reject(new Error(`TCP connection error: ${err.message}`));
}
});
socket.on("timeout", () => {
if (!isResolved) {
isResolved = true;
cleanup();
reject(new Error("TCP connection timeout"));
}
});
socket.on("close", () => {
if (!isResolved && !responseReceived) {
isResolved = true;
cleanup();
reject(new Error("Connection closed without receiving PONG"));
}
});
socket.connect(port, "localhost");
});
},
{
protocol: GatewayProxyProtocol.Ping,
relayHost: gatewayV2ConnectionDetails.relayHost,
gateway: gatewayV2ConnectionDetails.gateway,
relay: gatewayV2ConnectionDetails.relay
}
);
if (!isGatewayReachable) {
throw new BadRequestError({ message: `Gateway ${gateway.id} is not reachable` });
}
await gatewayV2DAL.updateById(gateway.id, { heartbeat: new Date() });
};
const deleteGatewayById = async ({ orgPermission, id }: { orgPermission: OrgServiceActor; id: string }) => {
const gateway = await gatewayV2DAL.findOne({ id, orgId: orgPermission.orgId });
if (!gateway) {
throw new NotFoundError({ message: `Gateway ${id} not found` });
}
const { permission } = await permissionService.getOrgPermission(
orgPermission.type,
orgPermission.id,
gateway.orgId,
orgPermission.authMethod,
orgPermission.orgId
);
ForbiddenError.from(permission).throwUnlessCan(
OrgPermissionGatewayActions.DeleteGateways,
OrgPermissionSubjects.Gateway
);
return gatewayV2DAL.deleteById(gateway.id);
};
return {
listGateways,
registerGateway,
getPlatformConnectionDetailsByGatewayId,
deleteGatewayById,
heartbeat
};
};
@@ -0,0 +1,11 @@
import { TDbClient } from "@app/db";
import { TableName } from "@app/db/schemas";
import { ormify } from "@app/lib/knex";
export type TOrgGatewayConfigV2DALFactory = ReturnType<typeof orgGatewayConfigV2DalFactory>;
export const orgGatewayConfigV2DalFactory = (db: TDbClient) => {
const orm = ormify(db, TableName.OrgGatewayConfigV2);
return orm;
};
@@ -13,7 +13,7 @@ export const gatewayDALFactory = (db: TDbClient) => {
{ offset, limit, sort, tx }: TFindOpt<TGateways> = {} { offset, limit, sort, tx }: TFindOpt<TGateways> = {}
) => { ) => {
try { try {
const query = (tx || db)(TableName.Gateway) const query = (tx || db.replicaNode())(TableName.Gateway)
// eslint-disable-next-line @typescript-eslint/no-misused-promises // eslint-disable-next-line @typescript-eslint/no-misused-promises
.where(buildFindFilter(filter, TableName.Gateway, ["orgId"])) .where(buildFindFilter(filter, TableName.Gateway, ["orgId"]))
.join(TableName.Identity, `${TableName.Identity}.id`, `${TableName.Gateway}.identityId`) .join(TableName.Identity, `${TableName.Identity}.id`, `${TableName.Gateway}.identityId`)
@@ -23,7 +23,7 @@ export const userGroupMembershipDALFactory = (db: TDbClient) => {
.whereIn(`${TableName.ProjectMembership}.projectId`, projectIds) .whereIn(`${TableName.ProjectMembership}.projectId`, projectIds)
.pluck(`${TableName.ProjectMembership}.projectId`); .pluck(`${TableName.ProjectMembership}.projectId`);
const userGroupMemberships: string[] = await (tx || db)(TableName.UserGroupMembership) const userGroupMemberships: string[] = await (tx || db.replicaNode())(TableName.UserGroupMembership)
.where(`${TableName.UserGroupMembership}.userId`, userId) .where(`${TableName.UserGroupMembership}.userId`, userId)
.whereNot(`${TableName.UserGroupMembership}.groupId`, groupId) .whereNot(`${TableName.UserGroupMembership}.groupId`, groupId)
.join( .join(
@@ -79,7 +79,7 @@ export const userGroupMembershipDALFactory = (db: TDbClient) => {
.pluck(`${TableName.GroupProjectMembership}.groupId`); .pluck(`${TableName.GroupProjectMembership}.groupId`);
// main query // main query
const members = await (tx || db)(TableName.UserGroupMembership) const members = await (tx || db.replicaNode())(TableName.UserGroupMembership)
.where(`${TableName.UserGroupMembership}.groupId`, groupId) .where(`${TableName.UserGroupMembership}.groupId`, groupId)
.where(`${TableName.UserGroupMembership}.isPending`, false) .where(`${TableName.UserGroupMembership}.isPending`, false)
.join(TableName.Users, `${TableName.UserGroupMembership}.userId`, `${TableName.Users}.id`) .join(TableName.Users, `${TableName.UserGroupMembership}.userId`, `${TableName.Users}.id`)
@@ -5,6 +5,7 @@ import { OrgMembershipStatus, TableName, TLdapConfigsUpdate, TUsers } from "@app
import { TGroupDALFactory } from "@app/ee/services/group/group-dal"; import { TGroupDALFactory } from "@app/ee/services/group/group-dal";
import { addUsersToGroupByUserIds, removeUsersFromGroupByUserIds } from "@app/ee/services/group/group-fns"; import { addUsersToGroupByUserIds, removeUsersFromGroupByUserIds } from "@app/ee/services/group/group-fns";
import { TUserGroupMembershipDALFactory } from "@app/ee/services/group/user-group-membership-dal"; import { TUserGroupMembershipDALFactory } from "@app/ee/services/group/user-group-membership-dal";
import { throwOnPlanSeatLimitReached } from "@app/ee/services/license/license-fns";
import { getConfig } from "@app/lib/config/env"; import { getConfig } from "@app/lib/config/env";
import { crypto } from "@app/lib/crypto"; import { crypto } from "@app/lib/crypto";
import { BadRequestError, ForbiddenRequestError, NotFoundError } from "@app/lib/errors"; import { BadRequestError, ForbiddenRequestError, NotFoundError } from "@app/lib/errors";
@@ -127,6 +128,20 @@ export const ldapConfigServiceFactory = ({
message: message:
"Failed to create LDAP configuration due to plan restriction. Upgrade plan to create LDAP configuration." "Failed to create LDAP configuration due to plan restriction. Upgrade plan to create LDAP configuration."
}); });
const org = await orgDAL.findOrgById(orgId);
if (!org) {
throw new NotFoundError({ message: `Could not find organization with ID "${orgId}"` });
}
if (org.googleSsoAuthEnforced && isActive) {
throw new BadRequestError({
message:
"You cannot enable LDAP SSO while Google OAuth is enforced. Disable Google OAuth enforcement to enable LDAP SSO."
});
}
const { encryptor } = await kmsService.createCipherPairWithDataKey({ const { encryptor } = await kmsService.createCipherPairWithDataKey({
type: KmsDataKey.Organization, type: KmsDataKey.Organization,
orgId orgId
@@ -233,6 +248,19 @@ export const ldapConfigServiceFactory = ({
"Failed to update LDAP configuration due to plan restriction. Upgrade plan to update LDAP configuration." "Failed to update LDAP configuration due to plan restriction. Upgrade plan to update LDAP configuration."
}); });
const org = await orgDAL.findOrgById(orgId);
if (!org) {
throw new NotFoundError({ message: `Could not find organization with ID "${orgId}"` });
}
if (org.googleSsoAuthEnforced && isActive) {
throw new BadRequestError({
message:
"You cannot enable LDAP SSO while Google OAuth is enforced. Disable Google OAuth enforcement to enable LDAP SSO."
});
}
const updateQuery: TLdapConfigsUpdate = { const updateQuery: TLdapConfigsUpdate = {
isActive, isActive,
url, url,
@@ -390,14 +418,6 @@ export const ldapConfigServiceFactory = ({
} }
}); });
} else { } else {
const plan = await licenseService.getPlan(orgId);
if (plan?.slug !== "enterprise" && plan?.identityLimit && plan.identitiesUsed >= plan.identityLimit) {
// limit imposed on number of identities allowed / number of identities used exceeds the number of identities allowed
throw new BadRequestError({
message: "Failed to create new member via LDAP due to member limit reached. Upgrade plan to add more members."
});
}
userAlias = await userDAL.transaction(async (tx) => { userAlias = await userDAL.transaction(async (tx) => {
let newUser: TUsers | undefined; let newUser: TUsers | undefined;
newUser = await userDAL.findOne( newUser = await userDAL.findOne(
@@ -446,6 +466,8 @@ export const ldapConfigServiceFactory = ({
); );
if (!orgMembership) { if (!orgMembership) {
await throwOnPlanSeatLimitReached(licenseService, orgId, UserAliasType.LDAP);
const { role, roleId } = await getDefaultOrgMembershipRole(organization.defaultMembershipRole); const { role, roleId } = await getDefaultOrgMembershipRole(organization.defaultMembershipRole);
await orgMembershipDAL.create( await orgMembershipDAL.create(
@@ -28,7 +28,7 @@ export const licenseDALFactory = (db: TDbClient) => {
const countOrgUsersAndIdentities = async (orgId: string | null, tx?: Knex) => { const countOrgUsersAndIdentities = async (orgId: string | null, tx?: Knex) => {
try { try {
// count org users // count org users
const userDoc = await (tx || db)(TableName.OrgMembership) const userDoc = await (tx || db.replicaNode())(TableName.OrgMembership)
.where({ status: OrgMembershipStatus.Accepted }) .where({ status: OrgMembershipStatus.Accepted })
.andWhere((bd) => { .andWhere((bd) => {
if (orgId) { if (orgId) {
@@ -42,7 +42,7 @@ export const licenseDALFactory = (db: TDbClient) => {
const userCount = Number(userDoc?.[0].count); const userCount = Number(userDoc?.[0].count);
// count org identities // count org identities
const identityDoc = await (tx || db)(TableName.IdentityOrgMembership) const identityDoc = await (tx || db.replicaNode())(TableName.IdentityOrgMembership)
.where((bd) => { .where((bd) => {
if (orgId) { if (orgId) {
void bd.where({ orgId }); void bd.where({ orgId });
@@ -1,8 +1,11 @@
import axios, { AxiosError } from "axios"; import axios, { AxiosError } from "axios";
import { TLicenseServiceFactory } from "@app/ee/services/license/license-service";
import { getConfig } from "@app/lib/config/env"; import { getConfig } from "@app/lib/config/env";
import { request } from "@app/lib/config/request"; import { request } from "@app/lib/config/request";
import { BadRequestError } from "@app/lib/errors";
import { logger } from "@app/lib/logger"; import { logger } from "@app/lib/logger";
import { UserAliasType } from "@app/services/user-alias/user-alias-types";
import { TFeatureSet } from "./license-types"; import { TFeatureSet } from "./license-types";
@@ -133,3 +136,18 @@ export const setupLicenseRequestWithStore = (
return { request: licenseReq, refreshLicense }; return { request: licenseReq, refreshLicense };
}; };
export const throwOnPlanSeatLimitReached = async (
licenseService: Pick<TLicenseServiceFactory, "getPlan">,
orgId: string,
type?: UserAliasType
) => {
const plan = await licenseService.getPlan(orgId);
if (plan?.slug !== "enterprise" && plan?.identityLimit && plan.identitiesUsed >= plan.identityLimit) {
// limit imposed on number of identities allowed / number of identities used exceeds the number of identities allowed
throw new BadRequestError({
message: `Failed to create new member${type ? ` via ${type.toUpperCase()}` : ""} due to member limit reached. Upgrade plan to add more members.`
});
}
};
@@ -99,6 +99,17 @@ export const licenseServiceFactory = ({
const workspacesUsed = await projectDAL.countOfOrgProjects(null); const workspacesUsed = await projectDAL.countOfOrgProjects(null);
currentPlan.workspacesUsed = workspacesUsed; currentPlan.workspacesUsed = workspacesUsed;
const usedIdentitySeats = await licenseDAL.countOrgUsersAndIdentities(null);
if (usedIdentitySeats !== currentPlan.identitiesUsed) {
const usedSeats = await licenseDAL.countOfOrgMembers(null);
await licenseServerOnPremApi.request.patch(`/api/license/v1/license`, {
usedSeats,
usedIdentitySeats
});
currentPlan.identitiesUsed = usedIdentitySeats;
currentPlan.membersUsed = usedSeats;
}
onPremFeatures = currentPlan; onPremFeatures = currentPlan;
logger.info("Successfully synchronized license key features"); logger.info("Successfully synchronized license key features");
} catch (error) { } catch (error) {
@@ -226,10 +237,13 @@ export const licenseServiceFactory = ({
}; };
const refreshPlan = async (orgId: string) => { const refreshPlan = async (orgId: string) => {
await keyStore.deleteItem(FEATURE_CACHE_KEY(orgId));
if (instanceType === InstanceType.Cloud) { if (instanceType === InstanceType.Cloud) {
await keyStore.deleteItem(FEATURE_CACHE_KEY(orgId));
await getPlan(orgId); await getPlan(orgId);
} }
if (instanceType === InstanceType.EnterpriseOnPrem) {
await syncLicenseKeyOnPremFeatures(true);
}
}; };
const generateOrgCustomerId = async (orgName: string, email?: string | null) => { const generateOrgCustomerId = async (orgName: string, email?: string | null) => {
@@ -296,8 +310,19 @@ export const licenseServiceFactory = ({
return data; return data;
}; };
const getOrgPlan = async ({ orgId, actor, actorId, actorOrgId, actorAuthMethod, projectId }: TOrgPlanDTO) => { const getOrgPlan = async ({
orgId,
actor,
actorId,
actorOrgId,
actorAuthMethod,
projectId,
refreshCache
}: TOrgPlanDTO) => {
await permissionService.getOrgPermission(actor, actorId, orgId, actorAuthMethod, actorOrgId); await permissionService.getOrgPermission(actor, actorId, orgId, actorAuthMethod, actorOrgId);
if (refreshCache) {
await refreshPlan(orgId);
}
const plan = await getPlan(orgId, projectId); const plan = await getPlan(orgId, projectId);
return plan; return plan;
}; };
@@ -417,6 +442,62 @@ export const licenseServiceFactory = ({
}; };
}; };
const calculateUsageValue = (
rowName: string,
field: string,
projectCount: number,
totalIdentities: number
): string => {
if (rowName === BillingPlanRows.WorkspaceLimit.name || field === BillingPlanRows.WorkspaceLimit.field) {
return projectCount.toString();
}
if (rowName === BillingPlanRows.IdentityLimit.name || field === BillingPlanRows.IdentityLimit.field) {
return totalIdentities.toString();
}
return "-";
};
const fetchPlanTableFromServer = async (customerId: string | null | undefined) => {
if (!customerId) {
throw new NotFoundError({ message: "Organization customer ID is required for plan table retrieval" });
}
const baseUrl = `/api/license-server/v1/customers/${customerId}`;
if (instanceType === InstanceType.Cloud) {
const { data } = await licenseServerCloudApi.request.get<{
head: { name: string }[];
rows: { name: string; allowed: boolean }[];
}>(`${baseUrl}/cloud-plan/table`);
return data;
}
if (instanceType === InstanceType.EnterpriseOnPrem) {
const { data } = await licenseServerOnPremApi.request.get<{
head: { name: string }[];
rows: { name: string; allowed: boolean }[];
}>(`${baseUrl}/on-prem-plan/table`);
return data;
}
throw new Error(`Unsupported instance type for server-based plan table: ${instanceType}`);
};
const getUsageMetrics = async (orgId: string) => {
const [orgMembersUsed, identityUsed, projectCount] = await Promise.all([
orgDAL.countAllOrgMembers(orgId),
identityOrgMembershipDAL.countAllOrgIdentities({ orgId }),
projectDAL.countOfOrgProjects(orgId)
]);
return {
orgMembersUsed,
identityUsed,
projectCount,
totalIdentities: identityUsed + orgMembersUsed
};
};
// returns org current plan feature table // returns org current plan feature table
const getOrgPlanTable = async ({ orgId, actor, actorId, actorAuthMethod, actorOrgId }: TGetOrgBillInfoDTO) => { const getOrgPlanTable = async ({ orgId, actor, actorId, actorAuthMethod, actorOrgId }: TGetOrgBillInfoDTO) => {
const { permission } = await permissionService.getOrgPermission(actor, actorId, orgId, actorAuthMethod, actorOrgId); const { permission } = await permissionService.getOrgPermission(actor, actorId, orgId, actorAuthMethod, actorOrgId);
@@ -429,55 +510,25 @@ export const licenseServiceFactory = ({
}); });
} }
const orgMembersUsed = await orgDAL.countAllOrgMembers(orgId); const { projectCount, totalIdentities } = await getUsageMetrics(orgId);
const identityUsed = await identityOrgMembershipDAL.countAllOrgIdentities({ orgId });
const projects = await projectDAL.find({ orgId });
const projectCount = projects.length;
if (instanceType === InstanceType.Cloud) { if (instanceType === InstanceType.Cloud || instanceType === InstanceType.EnterpriseOnPrem) {
const { data } = await licenseServerCloudApi.request.get<{ const tableResponse = await fetchPlanTableFromServer(organization.customerId);
head: { name: string }[];
rows: { name: string; allowed: boolean }[];
}>(`/api/license-server/v1/customers/${organization.customerId}/cloud-plan/table`);
const formattedData = { return {
head: data.head, head: tableResponse.head,
rows: data.rows.map((el) => { rows: tableResponse.rows.map((row) => ({
let used = "-"; ...row,
used: calculateUsageValue(row.name, "", projectCount, totalIdentities)
if (el.name === BillingPlanRows.WorkspaceLimit.name) { }))
used = projectCount.toString();
} else if (el.name === BillingPlanRows.IdentityLimit.name) {
used = (identityUsed + orgMembersUsed).toString();
}
return {
...el,
used
};
})
}; };
return formattedData;
} }
const mappedRows = await Promise.all( const mappedRows = Object.values(BillingPlanRows).map(({ name, field }) => ({
Object.values(BillingPlanRows).map(async ({ name, field }: { name: string; field: string }) => { name,
const allowed = onPremFeatures[field as keyof TFeatureSet]; allowed: onPremFeatures[field as keyof TFeatureSet] || false,
let used = "-"; used: calculateUsageValue(name, field, projectCount, totalIdentities)
}));
if (field === BillingPlanRows.WorkspaceLimit.field) {
used = projectCount.toString();
} else if (field === BillingPlanRows.IdentityLimit.field) {
used = (identityUsed + orgMembersUsed).toString();
}
return {
name,
allowed,
used
};
})
);
return { return {
head: Object.values(BillingPlanTableHead), head: Object.values(BillingPlanTableHead),
@@ -87,6 +87,7 @@ export type TOrgPlansTableDTO = {
export type TOrgPlanDTO = { export type TOrgPlanDTO = {
projectId?: string; projectId?: string;
refreshCache?: boolean;
} & TOrgPermission; } & TOrgPermission;
export type TStartOrgTrialDTO = { export type TStartOrgTrialDTO = {
@@ -8,6 +8,7 @@ import { EventType, TAuditLogServiceFactory } from "@app/ee/services/audit-log/a
import { TGroupDALFactory } from "@app/ee/services/group/group-dal"; import { TGroupDALFactory } from "@app/ee/services/group/group-dal";
import { addUsersToGroupByUserIds, removeUsersFromGroupByUserIds } from "@app/ee/services/group/group-fns"; import { addUsersToGroupByUserIds, removeUsersFromGroupByUserIds } from "@app/ee/services/group/group-fns";
import { TUserGroupMembershipDALFactory } from "@app/ee/services/group/user-group-membership-dal"; import { TUserGroupMembershipDALFactory } from "@app/ee/services/group/user-group-membership-dal";
import { throwOnPlanSeatLimitReached } from "@app/ee/services/license/license-fns";
import { TLicenseServiceFactory } from "@app/ee/services/license/license-service"; import { TLicenseServiceFactory } from "@app/ee/services/license/license-service";
import { OrgPermissionActions, OrgPermissionSubjects } from "@app/ee/services/permission/org-permission"; import { OrgPermissionActions, OrgPermissionSubjects } from "@app/ee/services/permission/org-permission";
import { TPermissionServiceFactory } from "@app/ee/services/permission/permission-service-types"; import { TPermissionServiceFactory } from "@app/ee/services/permission/permission-service-types";
@@ -294,6 +295,8 @@ export const oidcConfigServiceFactory = ({
); );
if (!orgMembership) { if (!orgMembership) {
await throwOnPlanSeatLimitReached(licenseService, orgId, UserAliasType.OIDC);
const { role, roleId } = await getDefaultOrgMembershipRole(organization.defaultMembershipRole); const { role, roleId } = await getDefaultOrgMembershipRole(organization.defaultMembershipRole);
await orgMembershipDAL.create( await orgMembershipDAL.create(
@@ -499,6 +502,13 @@ export const oidcConfigServiceFactory = ({
); );
ForbiddenError.from(permission).throwUnlessCan(OrgPermissionActions.Edit, OrgPermissionSubjects.Sso); ForbiddenError.from(permission).throwUnlessCan(OrgPermissionActions.Edit, OrgPermissionSubjects.Sso);
if (org.googleSsoAuthEnforced && isActive) {
throw new BadRequestError({
message:
"You cannot enable OIDC SSO while Google OAuth is enforced. Disable Google OAuth enforcement to enable OIDC SSO."
});
}
const { encryptor } = await kmsService.createCipherPairWithDataKey({ const { encryptor } = await kmsService.createCipherPairWithDataKey({
type: KmsDataKey.Organization, type: KmsDataKey.Organization,
orgId: org.id orgId: org.id
@@ -586,6 +596,13 @@ export const oidcConfigServiceFactory = ({
); );
ForbiddenError.from(permission).throwUnlessCan(OrgPermissionActions.Create, OrgPermissionSubjects.Sso); ForbiddenError.from(permission).throwUnlessCan(OrgPermissionActions.Create, OrgPermissionSubjects.Sso);
if (org.googleSsoAuthEnforced && isActive) {
throw new BadRequestError({
message:
"You cannot enable OIDC SSO while Google OAuth is enforced. Disable Google OAuth enforcement to enable OIDC SSO."
});
}
const { encryptor } = await kmsService.createCipherPairWithDataKey({ const { encryptor } = await kmsService.createCipherPairWithDataKey({
type: KmsDataKey.Organization, type: KmsDataKey.Organization,
orgId: org.id orgId: org.id
@@ -0,0 +1,11 @@
import { TDbClient } from "@app/db";
import { TableName } from "@app/db/schemas";
import { ormify } from "@app/lib/knex";
export type TInstanceRelayConfigDALFactory = ReturnType<typeof instanceRelayConfigDalFactory>;
export const instanceRelayConfigDalFactory = (db: TDbClient) => {
const orm = ormify(db, TableName.InstanceRelayConfig);
return orm;
};
@@ -0,0 +1,11 @@
import { TDbClient } from "@app/db";
import { TableName } from "@app/db/schemas";
import { ormify } from "@app/lib/knex";
export type TOrgRelayConfigDALFactory = ReturnType<typeof orgRelayConfigDalFactory>;
export const orgRelayConfigDalFactory = (db: TDbClient) => {
const orm = ormify(db, TableName.OrgRelayConfig);
return orm;
};
@@ -0,0 +1,11 @@
import { TDbClient } from "@app/db";
import { TableName } from "@app/db/schemas";
import { ormify } from "@app/lib/knex";
export type TRelayDALFactory = ReturnType<typeof relayDalFactory>;
export const relayDalFactory = (db: TDbClient) => {
const orm = ormify(db, TableName.Relay);
return orm;
};
File diff suppressed because it is too large Load Diff
@@ -1,6 +1,7 @@
import { ForbiddenError } from "@casl/ability"; import { ForbiddenError } from "@casl/ability";
import { OrgMembershipStatus, TableName, TSamlConfigs, TSamlConfigsUpdate, TUsers } from "@app/db/schemas"; import { OrgMembershipStatus, TableName, TSamlConfigs, TSamlConfigsUpdate, TUsers } from "@app/db/schemas";
import { throwOnPlanSeatLimitReached } from "@app/ee/services/license/license-fns";
import { getConfig } from "@app/lib/config/env"; import { getConfig } from "@app/lib/config/env";
import { crypto } from "@app/lib/crypto"; import { crypto } from "@app/lib/crypto";
import { BadRequestError, ForbiddenRequestError, NotFoundError } from "@app/lib/errors"; import { BadRequestError, ForbiddenRequestError, NotFoundError } from "@app/lib/errors";
@@ -82,6 +83,19 @@ export const samlConfigServiceFactory = ({
"Failed to create SAML SSO configuration due to plan restriction. Upgrade plan to create SSO configuration." "Failed to create SAML SSO configuration due to plan restriction. Upgrade plan to create SSO configuration."
}); });
const org = await orgDAL.findOrgById(orgId);
if (!org) {
throw new NotFoundError({ message: `Could not find organization with ID "${orgId}"` });
}
if (org.googleSsoAuthEnforced && isActive) {
throw new BadRequestError({
message:
"You cannot enable SAML SSO while Google OAuth is enforced. Disable Google OAuth enforcement to enable SAML SSO."
});
}
const { encryptor } = await kmsService.createCipherPairWithDataKey({ const { encryptor } = await kmsService.createCipherPairWithDataKey({
type: KmsDataKey.Organization, type: KmsDataKey.Organization,
orgId orgId
@@ -120,6 +134,19 @@ export const samlConfigServiceFactory = ({
"Failed to update SAML SSO configuration due to plan restriction. Upgrade plan to update SSO configuration." "Failed to update SAML SSO configuration due to plan restriction. Upgrade plan to update SSO configuration."
}); });
const org = await orgDAL.findOrgById(orgId);
if (!org) {
throw new NotFoundError({ message: `Could not find organization with ID "${orgId}"` });
}
if (org.googleSsoAuthEnforced && isActive) {
throw new BadRequestError({
message:
"Cannot enable SAML SSO while Google OAuth is enforced. Disable Google OAuth enforcement to enable SAML SSO."
});
}
const updateQuery: TSamlConfigsUpdate = { authProvider, isActive, lastUsed: null }; const updateQuery: TSamlConfigsUpdate = { authProvider, isActive, lastUsed: null };
const { encryptor } = await kmsService.createCipherPairWithDataKey({ const { encryptor } = await kmsService.createCipherPairWithDataKey({
type: KmsDataKey.Organization, type: KmsDataKey.Organization,
@@ -310,14 +337,6 @@ export const samlConfigServiceFactory = ({
return foundUser; return foundUser;
}); });
} else { } else {
const plan = await licenseService.getPlan(orgId);
if (plan?.slug !== "enterprise" && plan?.identityLimit && plan.identitiesUsed >= plan.identityLimit) {
// limit imposed on number of identities allowed / number of identities used exceeds the number of identities allowed
throw new BadRequestError({
message: "Failed to create new member via SAML due to member limit reached. Upgrade plan to add more members."
});
}
user = await userDAL.transaction(async (tx) => { user = await userDAL.transaction(async (tx) => {
let newUser: TUsers | undefined; let newUser: TUsers | undefined;
newUser = await userDAL.findOne( newUser = await userDAL.findOne(
@@ -365,6 +384,8 @@ export const samlConfigServiceFactory = ({
); );
if (!orgMembership) { if (!orgMembership) {
await throwOnPlanSeatLimitReached(licenseService, orgId, UserAliasType.SAML);
const { role, roleId } = await getDefaultOrgMembershipRole(organization.defaultMembershipRole); const { role, roleId } = await getDefaultOrgMembershipRole(organization.defaultMembershipRole);
await orgMembershipDAL.create( await orgMembershipDAL.create(
@@ -345,7 +345,7 @@ export const secretApprovalRequestDALFactory = (db: TDbClient) => {
const findProjectRequestCount = async (projectId: string, userId: string, policyId?: string, tx?: Knex) => { const findProjectRequestCount = async (projectId: string, userId: string, policyId?: string, tx?: Knex) => {
try { try {
const docs = await (tx || db) const docs = await (tx || db.replicaNode())
.with( .with(
"temp", "temp",
(tx || db.replicaNode())(TableName.SecretApprovalRequest) (tx || db.replicaNode())(TableName.SecretApprovalRequest)
@@ -494,7 +494,7 @@ export const secretApprovalRequestDALFactory = (db: TDbClient) => {
.distinctOn(`${TableName.SecretApprovalRequest}.id`) .distinctOn(`${TableName.SecretApprovalRequest}.id`)
.as("inner"); .as("inner");
const query = (tx || db) const query = (tx || db.replicaNode())
.select("*") .select("*")
.select(db.raw("count(*) OVER() as total_count")) .select(db.raw("count(*) OVER() as total_count"))
.from(innerQuery) .from(innerQuery)
@@ -377,7 +377,7 @@ export const secretApprovalRequestSecretDALFactory = (db: TDbClient) => {
// special query for migration to v2 secret // special query for migration to v2 secret
const findByProjectId = async (projectId: string, tx?: Knex) => { const findByProjectId = async (projectId: string, tx?: Knex) => {
try { try {
const docs = await (tx || db)(TableName.SecretApprovalRequestSecret) const docs = await (tx || db.replicaNode())(TableName.SecretApprovalRequestSecret)
.join( .join(
TableName.SecretApprovalRequest, TableName.SecretApprovalRequest,
`${TableName.SecretApprovalRequest}.id`, `${TableName.SecretApprovalRequest}.id`,
@@ -787,6 +787,7 @@ export const secretApprovalRequestServiceFactory = ({
}, },
tx tx
); );
await secretV2BridgeDAL.invalidateSecretCacheByProjectId(projectId, tx);
return { return {
secrets: { created: newSecrets, updated: updatedSecrets, deleted: deletedSecret }, secrets: { created: newSecrets, updated: updatedSecrets, deleted: deletedSecret },
approval: updatedSecretApproval approval: updatedSecretApproval
@@ -976,6 +977,7 @@ export const secretApprovalRequestServiceFactory = ({
}, },
tx tx
); );
await secretV2BridgeDAL.invalidateSecretCacheByProjectId(projectId, tx);
return { return {
secrets: { created: newSecrets, updated: updatedSecrets, deleted: deletedSecret }, secrets: { created: newSecrets, updated: updatedSecrets, deleted: deletedSecret },
approval: updatedSecretApproval approval: updatedSecretApproval
@@ -983,7 +985,6 @@ export const secretApprovalRequestServiceFactory = ({
}); });
} }
await secretV2BridgeDAL.invalidateSecretCacheByProjectId(projectId);
await snapshotService.performSnapshot(folderId); await snapshotService.performSnapshot(folderId);
const [folder] = await folderDAL.findSecretPathByFolderIds(projectId, [folderId]); const [folder] = await folderDAL.findSecretPathByFolderIds(projectId, [folderId]);
if (!folder) { if (!folder) {
@@ -509,9 +509,9 @@ export const secretReplicationServiceFactory = ({
tx tx
); );
} }
await secretV2BridgeDAL.invalidateSecretCacheByProjectId(projectId, tx);
}); });
await secretV2BridgeDAL.invalidateSecretCacheByProjectId(projectId);
await secretQueueService.syncSecrets({ await secretQueueService.syncSecrets({
projectId, projectId,
orgId, orgId,
@@ -82,6 +82,7 @@ import {
import { TSecretVersionV2DALFactory } from "@app/services/secret-v2-bridge/secret-version-dal"; import { TSecretVersionV2DALFactory } from "@app/services/secret-v2-bridge/secret-version-dal";
import { TSecretVersionV2TagDALFactory } from "@app/services/secret-v2-bridge/secret-version-tag-dal"; import { TSecretVersionV2TagDALFactory } from "@app/services/secret-v2-bridge/secret-version-tag-dal";
import { TGatewayV2ServiceFactory } from "../gateway-v2/gateway-v2-service";
import { awsIamUserSecretRotationFactory } from "./aws-iam-user-secret/aws-iam-user-secret-rotation-fns"; import { awsIamUserSecretRotationFactory } from "./aws-iam-user-secret/aws-iam-user-secret-rotation-fns";
import { oktaClientSecretRotationFactory } from "./okta-client-secret/okta-client-secret-rotation-fns"; import { oktaClientSecretRotationFactory } from "./okta-client-secret/okta-client-secret-rotation-fns";
import { TSecretRotationV2DALFactory } from "./secret-rotation-v2-dal"; import { TSecretRotationV2DALFactory } from "./secret-rotation-v2-dal";
@@ -110,6 +111,7 @@ export type TSecretRotationV2ServiceFactoryDep = {
appConnectionDAL: Pick<TAppConnectionDALFactory, "findById" | "update" | "updateById">; appConnectionDAL: Pick<TAppConnectionDALFactory, "findById" | "update" | "updateById">;
folderCommitService: Pick<TFolderCommitServiceFactory, "createCommit">; folderCommitService: Pick<TFolderCommitServiceFactory, "createCommit">;
gatewayService: Pick<TGatewayServiceFactory, "fnGetGatewayClientTlsByGatewayId">; gatewayService: Pick<TGatewayServiceFactory, "fnGetGatewayClientTlsByGatewayId">;
gatewayV2Service: Pick<TGatewayV2ServiceFactory, "getPlatformConnectionDetailsByGatewayId">;
}; };
export type TSecretRotationV2ServiceFactory = ReturnType<typeof secretRotationV2ServiceFactory>; export type TSecretRotationV2ServiceFactory = ReturnType<typeof secretRotationV2ServiceFactory>;
@@ -153,7 +155,8 @@ export const secretRotationV2ServiceFactory = ({
queueService, queueService,
folderCommitService, folderCommitService,
appConnectionDAL, appConnectionDAL,
gatewayService gatewayService,
gatewayV2Service
}: TSecretRotationV2ServiceFactoryDep) => { }: TSecretRotationV2ServiceFactoryDep) => {
const $queueSendSecretRotationStatusNotification = async (secretRotation: TSecretRotationV2Raw) => { const $queueSendSecretRotationStatusNotification = async (secretRotation: TSecretRotationV2Raw) => {
const appCfg = getConfig(); const appCfg = getConfig();
@@ -467,7 +470,8 @@ export const secretRotationV2ServiceFactory = ({
} as TSecretRotationV2WithConnection, } as TSecretRotationV2WithConnection,
appConnectionDAL, appConnectionDAL,
kmsService, kmsService,
gatewayService gatewayService,
gatewayV2Service
); );
// 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
@@ -831,7 +835,8 @@ export const secretRotationV2ServiceFactory = ({
} as TSecretRotationV2WithConnection, } as TSecretRotationV2WithConnection,
appConnectionDAL, appConnectionDAL,
kmsService, kmsService,
gatewayService gatewayService,
gatewayV2Service
); );
const generatedCredentials = await decryptSecretRotationCredentials({ const generatedCredentials = await decryptSecretRotationCredentials({
@@ -915,7 +920,8 @@ export const secretRotationV2ServiceFactory = ({
} as TSecretRotationV2WithConnection, } as TSecretRotationV2WithConnection,
appConnectionDAL, appConnectionDAL,
kmsService, kmsService,
gatewayService gatewayService,
gatewayV2Service
); );
const updatedRotation = await rotationFactory.rotateCredentials( const updatedRotation = await rotationFactory.rotateCredentials(
@@ -6,6 +6,7 @@ import { TAppConnectionDALFactory } from "@app/services/app-connection/app-conne
import { TKmsServiceFactory } from "@app/services/kms/kms-service"; import { TKmsServiceFactory } from "@app/services/kms/kms-service";
import { SecretsOrderBy } from "@app/services/secret/secret-types"; import { SecretsOrderBy } from "@app/services/secret/secret-types";
import { TGatewayV2ServiceFactory } from "../gateway-v2/gateway-v2-service";
import { import {
TAuth0ClientSecretRotation, TAuth0ClientSecretRotation,
TAuth0ClientSecretRotationGeneratedCredentials, TAuth0ClientSecretRotationGeneratedCredentials,
@@ -253,7 +254,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"> gatewayService: Pick<TGatewayServiceFactory, "fnGetGatewayClientTlsByGatewayId">,
gatewayV2Service: Pick<TGatewayV2ServiceFactory, "getPlatformConnectionDetailsByGatewayId">
) => { ) => {
issueCredentials: TRotationFactoryIssueCredentials<C, P>; issueCredentials: TRotationFactoryIssueCredentials<C, P>;
revokeCredentials: TRotationFactoryRevokeCredentials<C>; revokeCredentials: TRotationFactoryRevokeCredentials<C>;
@@ -41,7 +41,7 @@ const ORACLE_PASSWORD_REQUIREMENTS = {
export const sqlCredentialsRotationFactory: TRotationFactory< export const sqlCredentialsRotationFactory: TRotationFactory<
TSqlCredentialsRotationWithConnection, TSqlCredentialsRotationWithConnection,
TSqlCredentialsRotationGeneratedCredentials TSqlCredentialsRotationGeneratedCredentials
> = (secretRotation, _appConnectionDAL, _kmsService, gatewayService) => { > = (secretRotation, _appConnectionDAL, _kmsService, gatewayService, gatewayV2Service) => {
const { const {
connection, connection,
parameters: { username1, username2 }, parameters: { username1, username2 },
@@ -67,6 +67,7 @@ export const sqlCredentialsRotationFactory: TRotationFactory<
credentials: finalCredentials credentials: finalCredentials
}, },
gatewayService, gatewayService,
gatewayV2Service,
(client) => operation(client) (client) => operation(client)
); );
}; };
@@ -361,9 +361,8 @@ export const secretRotationQueueFactory = ({
}, },
tx tx
); );
await secretV2BridgeDAL.invalidateSecretCacheByProjectId(secretRotation.projectId, tx);
}); });
await secretV2BridgeDAL.invalidateSecretCacheByProjectId(secretRotation.projectId);
} else { } else {
if (!botKey) if (!botKey)
throw new NotFoundError({ throw new NotFoundError({
@@ -265,7 +265,7 @@ export const snapshotDALFactory = (db: TDbClient) => {
// then joins with respective secrets and folder // then joins with respective secrets and folder
const findRecursivelySnapshots = async (snapshotId: string, tx?: Knex) => { const findRecursivelySnapshots = async (snapshotId: string, tx?: Knex) => {
try { try {
const data = await (tx || db) const data = await (tx || db.replicaNode())
.withRecursive("parent", (qb) => { .withRecursive("parent", (qb) => {
void qb void qb
.from(TableName.Snapshot) .from(TableName.Snapshot)
@@ -419,7 +419,7 @@ export const snapshotDALFactory = (db: TDbClient) => {
// then joins with respective secrets and folder // then joins with respective secrets and folder
const findRecursivelySnapshotsV2Bridge = async (snapshotId: string, tx?: Knex) => { const findRecursivelySnapshotsV2Bridge = async (snapshotId: string, tx?: Knex) => {
try { try {
const data = await (tx || db) const data = await (tx || db.replicaNode())
.withRecursive("parent", (qb) => { .withRecursive("parent", (qb) => {
void qb void qb
.from(TableName.Snapshot) .from(TableName.Snapshot)
@@ -581,7 +581,11 @@ export const snapshotDALFactory = (db: TDbClient) => {
const docs = await (tx || db.replicaNode())(TableName.Snapshot) const docs = await (tx || db.replicaNode())(TableName.Snapshot)
.where(`${TableName.Snapshot}.folderId`, folderId) .where(`${TableName.Snapshot}.folderId`, folderId)
.join<TSecretSnapshots>( .join<TSecretSnapshots>(
(tx || db)(TableName.Snapshot).groupBy("folderId").max("createdAt").select("folderId").as("latestVersion"), (tx || db.replicaNode())(TableName.Snapshot)
.groupBy("folderId")
.max("createdAt")
.select("folderId")
.as("latestVersion"),
(bd) => { (bd) => {
bd.on(`${TableName.Snapshot}.folderId`, "latestVersion.folderId").andOn( bd.on(`${TableName.Snapshot}.folderId`, "latestVersion.folderId").andOn(
`${TableName.Snapshot}.createdAt`, `${TableName.Snapshot}.createdAt`,
@@ -766,7 +770,7 @@ export const snapshotDALFactory = (db: TDbClient) => {
) )
.orderBy(`${TableName.Snapshot}.createdAt`, "desc") .orderBy(`${TableName.Snapshot}.createdAt`, "desc")
.where(`${TableName.Snapshot}.folderId`, folderId); .where(`${TableName.Snapshot}.folderId`, folderId);
const data = await (tx || db) const data = await (tx || db.replicaNode())
.with("w", query) .with("w", query)
.select("*") .select("*")
.from<Awaited<typeof query>[number]>("w") .from<Awaited<typeof query>[number]>("w")
@@ -0,0 +1,91 @@
import { Knex } from "knex";
import { TDbClient } from "@app/db";
import { TableName } from "@app/db/schemas";
import { ormify, TOrmify } from "@app/lib/knex";
import { logger } from "@app/lib/logger";
import { QueueName } from "@app/queue";
export interface TKeyValueStoreDALFactory extends TOrmify<TableName.KeyValueStore> {
incrementBy: (key: string, dto: { incr?: number; tx?: Knex; expiresAt?: Date }) => Promise<number>;
findOneInt: (key: string, tx?: Knex) => Promise<number | undefined>;
pruneExpiredKeys: () => Promise<void>;
}
const QUERY_TIMEOUT_MS = 10 * 60 * 1000; // 10 minutes
const CACHE_KEY_PRUNE_BATCH_SIZE = 10000;
const MAX_RETRY_ON_FAILURE = 3;
export const keyValueStoreDALFactory = (db: TDbClient): TKeyValueStoreDALFactory => {
const keyValueStoreOrm = ormify(db, TableName.KeyValueStore);
const incrementBy: TKeyValueStoreDALFactory["incrementBy"] = async (key, { incr = 1, tx, expiresAt }) => {
return (tx || db)(TableName.KeyValueStore)
.insert({ key, integerValue: 1, expiresAt })
.onConflict("key")
.merge({
integerValue: db.raw(`"${TableName.KeyValueStore}"."integerValue" + ?`, [incr]),
expiresAt
})
.returning("integerValue")
.then((result) => Number(result[0]?.integerValue || 0));
};
const findOneInt: TKeyValueStoreDALFactory["findOneInt"] = async (key, tx) => {
const doc = await (tx || db.replicaNode())(TableName.KeyValueStore)
.where({ key })
.andWhere(
(builder) =>
void builder
.whereNull("expiresAt") // no expiry
.orWhere("expiresAt", ">", db.fn.now()) // or not expired
)
.first()
.select("integerValue");
return Number(doc?.integerValue || 0);
};
const pruneExpiredKeys: TKeyValueStoreDALFactory["pruneExpiredKeys"] = async () => {
let deletedIds: { key: string }[] = [];
let numberOfRetryOnFailure = 0;
let isRetrying = false;
logger.info(`${QueueName.DailyResourceCleanUp}: db key value store clean up started`);
do {
try {
// eslint-disable-next-line no-await-in-loop
deletedIds = await db.transaction(async (trx) => {
await trx.raw(`SET statement_timeout = ${QUERY_TIMEOUT_MS}`);
const findExpiredKeysSubQuery = trx(TableName.KeyValueStore)
.where("expiresAt", "<", db.fn.now())
.select("key")
.limit(CACHE_KEY_PRUNE_BATCH_SIZE);
// eslint-disable-next-line no-await-in-loop
const results = await trx(TableName.KeyValueStore)
.whereIn("key", findExpiredKeysSubQuery)
.del()
.returning("key");
return results;
});
numberOfRetryOnFailure = 0; // reset
} catch (error) {
numberOfRetryOnFailure += 1;
deletedIds = [];
logger.error(error, "Failed to clean up db key value");
} finally {
// eslint-disable-next-line no-await-in-loop
await new Promise((resolve) => {
setTimeout(resolve, 10); // time to breathe for db
});
}
isRetrying = numberOfRetryOnFailure > 0;
} while (deletedIds.length > 0 || (isRetrying && numberOfRetryOnFailure < MAX_RETRY_ON_FAILURE));
logger.info(`${QueueName.DailyResourceCleanUp}: db key value store clean up completed`);
};
return { ...keyValueStoreOrm, incrementBy, findOneInt, pruneExpiredKeys };
};
+50 -26
View File
@@ -1,11 +1,15 @@
import { Cluster, Redis } from "ioredis"; import { Cluster, Redis } from "ioredis";
import { Knex } from "knex";
import { buildRedisFromConfig, TRedisConfigKeys } from "@app/lib/config/redis"; import { buildRedisFromConfig, TRedisConfigKeys } from "@app/lib/config/redis";
import { pgAdvisoryLockHashText } from "@app/lib/crypto/hashtext"; import { pgAdvisoryLockHashText } from "@app/lib/crypto/hashtext";
import { applyJitter } from "@app/lib/dates"; import { applyJitter } from "@app/lib/dates";
import { delay as delayMs } from "@app/lib/delay"; import { delay as delayMs } from "@app/lib/delay";
import { ms } from "@app/lib/ms";
import { ExecutionResult, Redlock, Settings } from "@app/lib/red-lock"; import { ExecutionResult, Redlock, Settings } from "@app/lib/red-lock";
import { TKeyValueStoreDALFactory } from "./key-value-store-dal";
export const PgSqlLock = { export const PgSqlLock = {
BootUpMigration: 2023, BootUpMigration: 2023,
SuperAdminInit: 2024, SuperAdminInit: 2024,
@@ -16,6 +20,9 @@ export const PgSqlLock = {
CreateProject: (orgId: string) => pgAdvisoryLockHashText(`create-project:${orgId}`), CreateProject: (orgId: string) => pgAdvisoryLockHashText(`create-project:${orgId}`),
CreateFolder: (envId: string, projectId: string) => pgAdvisoryLockHashText(`create-folder:${envId}-${projectId}`), CreateFolder: (envId: string, projectId: string) => pgAdvisoryLockHashText(`create-folder:${envId}-${projectId}`),
SshInit: (projectId: string) => pgAdvisoryLockHashText(`ssh-bootstrap:${projectId}`), SshInit: (projectId: string) => pgAdvisoryLockHashText(`ssh-bootstrap:${projectId}`),
InstanceRelayConfigInit: () => pgAdvisoryLockHashText("instance-relay-config-init"),
OrgGatewayV2Init: (orgId: string) => pgAdvisoryLockHashText(`org-gateway-v2-init:${orgId}`),
OrgRelayConfigInit: (orgId: string) => pgAdvisoryLockHashText(`org-relay-config-init:${orgId}`),
IdentityLogin: (identityId: string, nonce: string) => pgAdvisoryLockHashText(`identity-login:${identityId}:${nonce}`) IdentityLogin: (identityId: string, nonce: string) => pgAdvisoryLockHashText(`identity-login:${identityId}:${nonce}`)
} as const; } as const;
@@ -95,13 +102,17 @@ export type TKeyStoreFactory = {
deleteItemsByKeyIn: (keys: string[]) => Promise<number>; deleteItemsByKeyIn: (keys: string[]) => Promise<number>;
deleteItems: (arg: TDeleteItems) => Promise<number>; deleteItems: (arg: TDeleteItems) => Promise<number>;
incrementBy: (key: string, value: number) => Promise<number>; incrementBy: (key: string, value: number) => Promise<number>;
getKeysByPattern: (pattern: string, limit?: number) => Promise<string[]>;
// pg
pgIncrementBy: (key: string, dto: { incr?: number; expiry?: string; tx?: Knex }) => Promise<number>;
pgGetIntItem: (key: string, prefix?: string) => Promise<number | undefined>;
// locks
acquireLock( acquireLock(
resources: string[], resources: string[],
duration: number, duration: number,
settings?: Partial<Settings> settings?: Partial<Settings>
): Promise<{ release: () => Promise<ExecutionResult> }>; ): Promise<{ release: () => Promise<ExecutionResult> }>;
waitTillReady: ({ key, waitingCb, keyCheckCb, waitIteration, delay, jitter }: TWaitTillReady) => Promise<void>; waitTillReady: ({ key, waitingCb, keyCheckCb, waitIteration, delay, jitter }: TWaitTillReady) => Promise<void>;
getKeysByPattern: (pattern: string, limit?: number) => Promise<string[]>;
}; };
const pickPrimaryOrSecondaryRedis = (primary: Redis | Cluster, secondaries?: Array<Redis | Cluster>) => { const pickPrimaryOrSecondaryRedis = (primary: Redis | Cluster, secondaries?: Array<Redis | Cluster>) => {
@@ -114,7 +125,10 @@ interface TKeyStoreFactoryDTO extends TRedisConfigKeys {
REDIS_READ_REPLICAS?: { host: string; port: number }[]; REDIS_READ_REPLICAS?: { host: string; port: number }[];
} }
export const keyStoreFactory = (redisConfigKeys: TKeyStoreFactoryDTO): TKeyStoreFactory => { export const keyStoreFactory = (
redisConfigKeys: TKeyStoreFactoryDTO,
keyValueStoreDAL: TKeyValueStoreDALFactory
): TKeyStoreFactory => {
const primaryRedis = buildRedisFromConfig(redisConfigKeys); const primaryRedis = buildRedisFromConfig(redisConfigKeys);
const redisReadReplicas = redisConfigKeys.REDIS_READ_REPLICAS?.map((el) => { const redisReadReplicas = redisConfigKeys.REDIS_READ_REPLICAS?.map((el) => {
if (redisConfigKeys.REDIS_URL) { if (redisConfigKeys.REDIS_URL) {
@@ -189,29 +203,6 @@ export const keyStoreFactory = (redisConfigKeys: TKeyStoreFactoryDTO): TKeyStore
const setExpiry = async (key: string, expiryInSeconds: number) => primaryRedis.expire(key, expiryInSeconds); const setExpiry = async (key: string, expiryInSeconds: number) => primaryRedis.expire(key, expiryInSeconds);
const waitTillReady = async ({
key,
waitingCb,
keyCheckCb,
waitIteration = 10,
delay = 1000,
jitter = 200
}: TWaitTillReady) => {
let attempts = 0;
let isReady = keyCheckCb(await getItem(key));
while (!isReady) {
if (attempts > waitIteration) return;
// eslint-disable-next-line
await new Promise((resolve) => {
waitingCb?.();
setTimeout(resolve, Math.max(0, applyJitter(delay, jitter)));
});
attempts += 1;
// eslint-disable-next-line
isReady = keyCheckCb(await getItem(key));
}
};
const getKeysByPattern = async (pattern: string, limit?: number) => { const getKeysByPattern = async (pattern: string, limit?: number) => {
let cursor = "0"; let cursor = "0";
const allKeys: string[] = []; const allKeys: string[] = [];
@@ -236,6 +227,37 @@ export const keyStoreFactory = (redisConfigKeys: TKeyStoreFactoryDTO): TKeyStore
return allKeys; return allKeys;
}; };
const pgIncrementBy: TKeyStoreFactory["pgIncrementBy"] = async (key, { incr = 1, tx, expiry }) => {
const expiresAt = expiry ? new Date(Date.now() + ms(expiry)) : undefined;
return keyValueStoreDAL.incrementBy(key, { incr, expiresAt, tx });
};
const pgGetIntItem = async (key: string, prefix?: string) =>
keyValueStoreDAL.findOneInt(prefix ? `${prefix}:${key}` : key);
const waitTillReady = async ({
key,
waitingCb,
keyCheckCb,
waitIteration = 10,
delay = 1000,
jitter = 200
}: TWaitTillReady) => {
let attempts = 0;
let isReady = keyCheckCb(await getItem(key));
while (!isReady) {
if (attempts > waitIteration) return;
// eslint-disable-next-line
await new Promise((resolve) => {
waitingCb?.();
setTimeout(resolve, Math.max(0, applyJitter(delay, jitter)));
});
attempts += 1;
// eslint-disable-next-line
isReady = keyCheckCb(await getItem(key));
}
};
return { return {
setItem, setItem,
getItem, getItem,
@@ -250,6 +272,8 @@ export const keyStoreFactory = (redisConfigKeys: TKeyStoreFactoryDTO): TKeyStore
waitTillReady, waitTillReady,
getKeysByPattern, getKeysByPattern,
deleteItemsByKeyIn, deleteItemsByKeyIn,
getItems getItems,
pgGetIntItem,
pgIncrementBy
}; };
}; };
+9
View File
@@ -53,6 +53,15 @@ export const inMemoryKeyStore = (): TKeyStoreFactory => {
} }
return null; return null;
}, },
pgGetIntItem: async (key) => {
const value = store[key];
if (typeof value === "number") {
return Number(value);
}
},
pgIncrementBy: async () => {
return 1;
},
incrementBy: async () => { incrementBy: async () => {
return 1; return 1;
}, },
+3
View File
@@ -259,6 +259,8 @@ const envSchema = z
GATEWAY_RELAY_REALM: zpStr(z.string().optional()), GATEWAY_RELAY_REALM: zpStr(z.string().optional()),
GATEWAY_RELAY_AUTH_SECRET: zpStr(z.string().optional()), GATEWAY_RELAY_AUTH_SECRET: zpStr(z.string().optional()),
RELAY_AUTH_SECRET: zpStr(z.string().optional()),
DYNAMIC_SECRET_ALLOW_INTERNAL_IP: zodStrBool.default("false"), DYNAMIC_SECRET_ALLOW_INTERNAL_IP: zodStrBool.default("false"),
DYNAMIC_SECRET_AWS_ACCESS_KEY_ID: zpStr(z.string().optional()).default( DYNAMIC_SECRET_AWS_ACCESS_KEY_ID: zpStr(z.string().optional()).default(
process.env.INF_APP_CONNECTION_AWS_ACCESS_KEY_ID process.env.INF_APP_CONNECTION_AWS_ACCESS_KEY_ID
@@ -410,6 +412,7 @@ const envSchema = z
Boolean(data.INF_APP_CONNECTION_GITHUB_RADAR_APP_CLIENT_ID) && Boolean(data.INF_APP_CONNECTION_GITHUB_RADAR_APP_CLIENT_ID) &&
Boolean(data.INF_APP_CONNECTION_GITHUB_RADAR_APP_CLIENT_SECRET) && Boolean(data.INF_APP_CONNECTION_GITHUB_RADAR_APP_CLIENT_SECRET) &&
Boolean(data.INF_APP_CONNECTION_GITHUB_RADAR_APP_WEBHOOK_SECRET), Boolean(data.INF_APP_CONNECTION_GITHUB_RADAR_APP_WEBHOOK_SECRET),
isSecondaryInstance: Boolean(data.INFISICAL_PRIMARY_INSTANCE_URL),
isHsmConfigured: isHsmConfigured:
Boolean(data.HSM_LIB_PATH) && Boolean(data.HSM_PIN) && Boolean(data.HSM_KEY_LABEL) && data.HSM_SLOT !== undefined, Boolean(data.HSM_LIB_PATH) && Boolean(data.HSM_PIN) && Boolean(data.HSM_KEY_LABEL) && data.HSM_SLOT !== undefined,
samlDefaultOrgSlug: data.DEFAULT_SAML_ORG_SLUG, samlDefaultOrgSlug: data.DEFAULT_SAML_ORG_SLUG,
@@ -424,7 +424,8 @@ const cryptographyFactory = () => {
constants: crypto.constants, constants: crypto.constants,
X509Certificate: crypto.X509Certificate, X509Certificate: crypto.X509Certificate,
KeyObject: crypto.KeyObject, KeyObject: crypto.KeyObject,
Hash: crypto.Hash Hash: crypto.Hash,
timingSafeEqual: crypto.timingSafeEqual
} }
}; };
}; };
+281
View File
@@ -0,0 +1,281 @@
import net from "node:net";
import tls from "node:tls";
import axios from "axios";
import https from "https";
import { verifyHostInputValidity } from "@app/ee/services/dynamic-secret/dynamic-secret-fns";
import { splitPemChain } from "@app/services/certificate/certificate-fns";
import { BadRequestError } from "../errors";
import { GatewayProxyProtocol } from "../gateway/types";
import { logger } from "../logger";
interface IGatewayRelayServer {
server: net.Server;
port: number;
cleanup: () => Promise<void>;
getRelayError: () => string;
}
const createRelayConnection = async ({
relayHost,
clientCertificate,
clientPrivateKey,
serverCertificateChain
}: {
relayHost: string;
clientCertificate: string;
clientPrivateKey: string;
serverCertificateChain: string;
}): Promise<net.Socket> => {
const [targetHost] = await verifyHostInputValidity(relayHost);
const [, portStr] = relayHost.split(":");
const port = parseInt(portStr, 10) || 8443;
const serverCAs = splitPemChain(serverCertificateChain);
const tlsOptions: tls.ConnectionOptions = {
host: targetHost,
servername: relayHost,
port,
cert: clientCertificate,
key: clientPrivateKey,
ca: serverCAs,
minVersion: "TLSv1.2",
rejectUnauthorized: true
};
return new Promise((resolve, reject) => {
try {
const socket = tls.connect(tlsOptions, () => {
logger.info("Relay TLS connection established successfully");
resolve(socket);
});
socket.on("error", (err: Error) => {
reject(new Error(`TLS connection error: ${err.message}`));
});
socket.on("close", (hadError: boolean) => {
if (hadError) {
logger.error("TLS connection closed with error");
}
});
socket.on("timeout", () => {
logger.error(`TLS connection timeout after 30 seconds`);
socket.destroy();
reject(new Error("TLS connection timeout"));
});
socket.setTimeout(30000);
} catch (error: unknown) {
reject(new Error(`Failed to create TLS connection: ${error instanceof Error ? error.message : String(error)}`));
}
});
};
const createGatewayConnection = async (
relayConn: net.Socket,
gateway: { clientCertificate: string; clientPrivateKey: string; serverCertificateChain: string },
protocol: GatewayProxyProtocol
): Promise<net.Socket> => {
const protocolToAlpn = {
[GatewayProxyProtocol.Http]: "infisical-http-proxy",
[GatewayProxyProtocol.Tcp]: "infisical-tcp-proxy",
[GatewayProxyProtocol.Ping]: "infisical-ping"
};
const tlsOptions: tls.ConnectionOptions = {
socket: relayConn,
cert: gateway.clientCertificate,
key: gateway.clientPrivateKey,
ca: splitPemChain(gateway.serverCertificateChain),
minVersion: "TLSv1.2",
maxVersion: "TLSv1.3",
rejectUnauthorized: true,
ALPNProtocols: [protocolToAlpn[protocol]]
};
return new Promise((resolve, reject) => {
try {
const gatewaySocket = tls.connect(tlsOptions, () => {
if (!gatewaySocket.authorized) {
const error = gatewaySocket.authorizationError;
gatewaySocket.destroy();
reject(new Error(`Gateway TLS authorization failed: ${error?.message}`));
return;
}
logger.info("Gateway mTLS connection established successfully");
resolve(gatewaySocket);
});
gatewaySocket.on("error", (err: Error) => {
reject(new Error(`Failed to establish gateway mTLS: ${err.message}`));
});
gatewaySocket.setTimeout(30000);
gatewaySocket.on("timeout", () => {
gatewaySocket.destroy();
reject(new Error("Gateway connection timeout"));
});
} catch (error: unknown) {
reject(
new Error(`Failed to create gateway TLS connection: ${error instanceof Error ? error.message : String(error)}`)
);
}
});
};
const setupRelayServer = async ({
protocol,
relayHost,
gateway,
relay,
httpsAgent
}: {
protocol: GatewayProxyProtocol;
relayHost: string;
gateway: { clientCertificate: string; clientPrivateKey: string; serverCertificateChain: string };
relay: { clientCertificate: string; clientPrivateKey: string; serverCertificateChain: string };
httpsAgent?: https.Agent;
}): Promise<IGatewayRelayServer> => {
const relayErrorMsg: string[] = [];
return new Promise((resolve, reject) => {
const server = net.createServer();
server.on("connection", (clientConn) => {
void (async () => {
try {
clientConn.setKeepAlive(true, 30000);
clientConn.setNoDelay(true);
// Stage 1: Connect to relay with TLS
const relayConn = await createRelayConnection({
relayHost,
clientCertificate: relay.clientCertificate,
clientPrivateKey: relay.clientPrivateKey,
serverCertificateChain: relay.serverCertificateChain
});
// Stage 2: Establish mTLS connection to gateway through the relay
const gatewayConn = await createGatewayConnection(relayConn, gateway, protocol);
// Send protocol-specific configuration for HTTP requests
if (protocol === GatewayProxyProtocol.Http) {
if (httpsAgent) {
const agentOptions = httpsAgent.options;
if (agentOptions && agentOptions.ca) {
const caCert = Array.isArray(agentOptions.ca) ? agentOptions.ca.join("\n") : agentOptions.ca;
const caB64 = Buffer.from(caCert as string).toString("base64");
const rejectUnauthorized = agentOptions.rejectUnauthorized !== false;
const configCommand = `CONFIG ca=${caB64} verify=${rejectUnauthorized}\n`;
gatewayConn.write(Buffer.from(configCommand));
} else {
// Send empty config to signal end of configuration
gatewayConn.write(Buffer.from("CONFIG\n"));
}
} else {
// Send empty config to signal end of configuration
gatewayConn.write(Buffer.from("CONFIG\n"));
}
}
// Bidirectional data forwarding
clientConn.pipe(gatewayConn);
gatewayConn.pipe(clientConn);
// Handle connection closure
clientConn.on("close", () => {
relayConn.destroy();
gatewayConn.destroy();
});
relayConn.on("close", () => {
clientConn.destroy();
gatewayConn.destroy();
});
gatewayConn.on("close", () => {
clientConn.destroy();
relayConn.destroy();
});
} catch (err) {
const errorMsg = err instanceof Error ? err.message : String(err);
relayErrorMsg.push(errorMsg);
clientConn.destroy();
}
})();
});
server.on("error", (err) => {
reject(err);
});
server.listen(0, () => {
const address = server.address();
if (!address || typeof address === "string") {
server.close();
reject(new Error("Failed to get server port"));
return;
}
resolve({
server,
port: address.port,
cleanup: async () => {
try {
server.close();
} catch (err) {
logger.debug("Error closing server:", err instanceof Error ? err.message : String(err));
}
},
getRelayError: () => relayErrorMsg.join(",")
});
});
});
};
export const withGatewayV2Proxy = async <T>(
callback: (port: number) => Promise<T>,
options: {
protocol: GatewayProxyProtocol;
relayHost: string;
gateway: { clientCertificate: string; clientPrivateKey: string; serverCertificateChain: string };
relay: { clientCertificate: string; clientPrivateKey: string; serverCertificateChain: string };
httpsAgent?: https.Agent;
}
): Promise<T> => {
const { protocol, relayHost, gateway, relay, httpsAgent } = options;
const { port, cleanup, getRelayError } = await setupRelayServer({
protocol,
relayHost,
gateway,
relay,
httpsAgent
});
try {
// Execute the callback with the allocated port
return await callback(port);
} catch (err) {
const relayErrorMessage = getRelayError();
if (relayErrorMessage) {
logger.error("Relay error:", relayErrorMessage);
}
logger.error("Gateway error:", err instanceof Error ? err.message : String(err));
let errorMessage = relayErrorMessage || (err instanceof Error ? err.message : String(err));
if (axios.isAxiosError(err) && (err.response?.data as { message?: string })?.message) {
errorMessage = (err.response?.data as { message: string }).message;
}
throw new BadRequestError({ message: errorMessage });
} finally {
// Ensure cleanup happens regardless of success or failure
await cleanup();
}
};
+2 -1
View File
@@ -6,7 +6,8 @@ export type TGatewayTlsOptions = { ca: string; cert: string; key: string };
export enum GatewayProxyProtocol { export enum GatewayProxyProtocol {
Http = "http", Http = "http",
Tcp = "tcp" Tcp = "tcp",
Ping = "ping"
} }
export enum GatewayHttpProxyActions { export enum GatewayHttpProxyActions {
+4 -4
View File
@@ -250,12 +250,12 @@ export const ormify = <DbOps extends object, Tname extends keyof Tables>(
.returning("*"); .returning("*");
if ($incr) { if ($incr) {
Object.entries($incr).forEach(([incrementField, incrementValue]) => { Object.entries($incr).forEach(([incrementField, incrementValue]) => {
void query.increment(incrementField, incrementValue); void query.increment(incrementField, incrementValue as number);
}); });
} }
if ($decr) { if ($decr) {
Object.entries($decr).forEach(([incrementField, incrementValue]) => { Object.entries($decr).forEach(([incrementField, incrementValue]) => {
void query.decrement(incrementField, incrementValue); void query.decrement(incrementField, incrementValue as number);
}); });
} }
const [docs] = await query; const [docs] = await query;
@@ -273,12 +273,12 @@ export const ormify = <DbOps extends object, Tname extends keyof Tables>(
// increment and decrement operation in update // increment and decrement operation in update
if ($incr) { if ($incr) {
Object.entries($incr).forEach(([incrementField, incrementValue]) => { Object.entries($incr).forEach(([incrementField, incrementValue]) => {
void query.increment(incrementField, incrementValue); void query.increment(incrementField, incrementValue as number);
}); });
} }
if ($decr) { if ($decr) {
Object.entries($decr).forEach(([incrementField, incrementValue]) => { Object.entries($decr).forEach(([incrementField, incrementValue]) => {
void query.increment(incrementField, incrementValue); void query.decrement(incrementField, incrementValue as number);
}); });
} }
return (await query) as Tables[Tname]["base"][]; return (await query) as Tables[Tname]["base"][];
+3 -1
View File
@@ -5,6 +5,7 @@ import "./lib/telemetry/instrumentation";
import dotenv from "dotenv"; import dotenv from "dotenv";
import { initializeHsmModule } from "@app/ee/services/hsm/hsm-fns"; import { initializeHsmModule } from "@app/ee/services/hsm/hsm-fns";
import { keyValueStoreDALFactory } from "@app/keystore/key-value-store-dal";
import { runMigrations } from "./auto-start-migrations"; import { runMigrations } from "./auto-start-migrations";
import { initAuditLogDbConnection, initDbConnection } from "./db"; import { initAuditLogDbConnection, initDbConnection } from "./db";
@@ -54,7 +55,8 @@ const run = async () => {
await queue.initialize(); await queue.initialize();
const keyStore = keyStoreFactory(envConfig); const keyValueStoreDAL = keyValueStoreDALFactory(db);
const keyStore = keyStoreFactory(envConfig, keyValueStoreDAL);
const redis = buildRedisFromConfig(envConfig); const redis = buildRedisFromConfig(envConfig);
const hsmModule = initializeHsmModule(envConfig); const hsmModule = initializeHsmModule(envConfig);
@@ -122,6 +122,11 @@ export const injectIdentity = fp(
return; return;
} }
// Authentication is handled on a route-level
if (req.url === "/api/v1/relays/register-instance-relay") {
return;
}
// Authentication is handled on a route-level here. // Authentication is handled on a route-level here.
if (req.url.includes("/api/v1/workflow-integrations/microsoft-teams/message-endpoint")) { if (req.url.includes("/api/v1/workflow-integrations/microsoft-teams/message-endpoint")) {
return; return;
+55 -5
View File
@@ -38,6 +38,9 @@ import { externalKmsServiceFactory } from "@app/ee/services/external-kms/externa
import { gatewayDALFactory } from "@app/ee/services/gateway/gateway-dal"; import { gatewayDALFactory } from "@app/ee/services/gateway/gateway-dal";
import { gatewayServiceFactory } from "@app/ee/services/gateway/gateway-service"; import { gatewayServiceFactory } from "@app/ee/services/gateway/gateway-service";
import { orgGatewayConfigDALFactory } from "@app/ee/services/gateway/org-gateway-config-dal"; import { orgGatewayConfigDALFactory } from "@app/ee/services/gateway/org-gateway-config-dal";
import { gatewayV2DalFactory } from "@app/ee/services/gateway-v2/gateway-v2-dal";
import { gatewayV2ServiceFactory } from "@app/ee/services/gateway-v2/gateway-v2-service";
import { orgGatewayConfigV2DalFactory } from "@app/ee/services/gateway-v2/org-gateway-config-v2-dal";
import { githubOrgSyncDALFactory } from "@app/ee/services/github-org-sync/github-org-sync-dal"; import { githubOrgSyncDALFactory } from "@app/ee/services/github-org-sync/github-org-sync-dal";
import { githubOrgSyncServiceFactory } from "@app/ee/services/github-org-sync/github-org-sync-service"; import { githubOrgSyncServiceFactory } from "@app/ee/services/github-org-sync/github-org-sync-service";
import { groupDALFactory } from "@app/ee/services/group/group-dal"; import { groupDALFactory } from "@app/ee/services/group/group-dal";
@@ -72,6 +75,10 @@ import { projectUserAdditionalPrivilegeDALFactory } from "@app/ee/services/proje
import { projectUserAdditionalPrivilegeServiceFactory } from "@app/ee/services/project-user-additional-privilege/project-user-additional-privilege-service"; import { projectUserAdditionalPrivilegeServiceFactory } from "@app/ee/services/project-user-additional-privilege/project-user-additional-privilege-service";
import { rateLimitDALFactory } from "@app/ee/services/rate-limit/rate-limit-dal"; import { rateLimitDALFactory } from "@app/ee/services/rate-limit/rate-limit-dal";
import { rateLimitServiceFactory } from "@app/ee/services/rate-limit/rate-limit-service"; import { rateLimitServiceFactory } from "@app/ee/services/rate-limit/rate-limit-service";
import { instanceRelayConfigDalFactory } from "@app/ee/services/relay/instance-relay-config-dal";
import { orgRelayConfigDalFactory } from "@app/ee/services/relay/org-relay-config-dal";
import { relayDalFactory } from "@app/ee/services/relay/relay-dal";
import { relayServiceFactory } from "@app/ee/services/relay/relay-service";
import { samlConfigDALFactory } from "@app/ee/services/saml-config/saml-config-dal"; import { samlConfigDALFactory } from "@app/ee/services/saml-config/saml-config-dal";
import { samlConfigServiceFactory } from "@app/ee/services/saml-config/saml-config-service"; import { samlConfigServiceFactory } from "@app/ee/services/saml-config/saml-config-service";
import { scimDALFactory } from "@app/ee/services/scim/scim-dal"; import { scimDALFactory } from "@app/ee/services/scim/scim-dal";
@@ -123,6 +130,7 @@ import { sshHostGroupMembershipDALFactory } from "@app/ee/services/ssh-host-grou
import { sshHostGroupServiceFactory } from "@app/ee/services/ssh-host-group/ssh-host-group-service"; import { sshHostGroupServiceFactory } from "@app/ee/services/ssh-host-group/ssh-host-group-service";
import { trustedIpDALFactory } from "@app/ee/services/trusted-ip/trusted-ip-dal"; import { trustedIpDALFactory } from "@app/ee/services/trusted-ip/trusted-ip-dal";
import { trustedIpServiceFactory } from "@app/ee/services/trusted-ip/trusted-ip-service"; import { trustedIpServiceFactory } from "@app/ee/services/trusted-ip/trusted-ip-service";
import { keyValueStoreDALFactory } from "@app/keystore/key-value-store-dal";
import { TKeyStoreFactory } from "@app/keystore/keystore"; import { TKeyStoreFactory } from "@app/keystore/keystore";
import { getConfig, TEnvConfig } from "@app/lib/config/env"; import { getConfig, TEnvConfig } from "@app/lib/config/env";
import { crypto } from "@app/lib/crypto/cryptography"; import { crypto } from "@app/lib/crypto/cryptography";
@@ -507,6 +515,7 @@ export const registerRoutes = async (
const microsoftTeamsIntegrationDAL = microsoftTeamsIntegrationDALFactory(db); const microsoftTeamsIntegrationDAL = microsoftTeamsIntegrationDALFactory(db);
const projectMicrosoftTeamsConfigDAL = projectMicrosoftTeamsConfigDALFactory(db); const projectMicrosoftTeamsConfigDAL = projectMicrosoftTeamsConfigDALFactory(db);
const secretScanningV2DAL = secretScanningV2DALFactory(db); const secretScanningV2DAL = secretScanningV2DALFactory(db);
const keyValueStoreDAL = keyValueStoreDALFactory(db);
const eventBusService = eventBusFactory(server.redis); const eventBusService = eventBusFactory(server.redis);
const sseService = sseServiceFactory(eventBusService, server.redis); const sseService = sseServiceFactory(eventBusService, server.redis);
@@ -643,6 +652,7 @@ export const registerRoutes = async (
const folderTreeCheckpointDAL = folderTreeCheckpointDALFactory(db); const folderTreeCheckpointDAL = folderTreeCheckpointDALFactory(db);
const folderCommitDAL = folderCommitDALFactory(db); const folderCommitDAL = folderCommitDALFactory(db);
const folderTreeCheckpointResourcesDAL = folderTreeCheckpointResourcesDALFactory(db); const folderTreeCheckpointResourcesDAL = folderTreeCheckpointResourcesDALFactory(db);
const folderCommitQueueService = folderCommitQueueServiceFactory({ const folderCommitQueueService = folderCommitQueueServiceFactory({
queueService, queueService,
folderTreeCheckpointDAL, folderTreeCheckpointDAL,
@@ -739,6 +749,7 @@ export const registerRoutes = async (
const userService = userServiceFactory({ const userService = userServiceFactory({
userDAL, userDAL,
orgDAL,
orgMembershipDAL, orgMembershipDAL,
tokenService, tokenService,
permissionService, permissionService,
@@ -807,6 +818,7 @@ export const registerRoutes = async (
groupDAL, groupDAL,
orgBotDAL, orgBotDAL,
oidcConfigDAL, oidcConfigDAL,
ldapConfigDAL,
loginService, loginService,
projectBotService, projectBotService,
reminderService reminderService
@@ -965,6 +977,13 @@ export const registerRoutes = async (
const pkiSubscriberDAL = pkiSubscriberDALFactory(db); const pkiSubscriberDAL = pkiSubscriberDALFactory(db);
const pkiTemplatesDAL = pkiTemplatesDALFactory(db); const pkiTemplatesDAL = pkiTemplatesDALFactory(db);
const instanceRelayConfigDAL = instanceRelayConfigDalFactory(db);
const orgRelayConfigDAL = orgRelayConfigDalFactory(db);
const relayDAL = relayDalFactory(db);
const gatewayV2DAL = gatewayV2DalFactory(db);
const orgGatewayConfigV2DAL = orgGatewayConfigV2DalFactory(db);
const certificateService = certificateServiceFactory({ const certificateService = certificateServiceFactory({
certificateDAL, certificateDAL,
certificateBodyDAL, certificateBodyDAL,
@@ -1083,6 +1102,23 @@ export const registerRoutes = async (
keyStore keyStore
}); });
const relayService = relayServiceFactory({
instanceRelayConfigDAL,
orgRelayConfigDAL,
relayDAL,
kmsService
});
const gatewayV2Service = gatewayV2ServiceFactory({
kmsService,
licenseService,
relayService,
orgGatewayConfigV2DAL,
gatewayV2DAL,
relayDAL,
permissionService
});
const secretSyncQueue = secretSyncQueueFactory({ const secretSyncQueue = secretSyncQueueFactory({
queueService, queueService,
secretSyncDAL, secretSyncDAL,
@@ -1107,7 +1143,8 @@ export const registerRoutes = async (
resourceMetadataDAL, resourceMetadataDAL,
appConnectionDAL, appConnectionDAL,
licenseService, licenseService,
gatewayService gatewayService,
gatewayV2Service
}); });
const secretQueueService = secretQueueFactory({ const secretQueueService = secretQueueFactory({
@@ -1531,6 +1568,7 @@ export const registerRoutes = async (
permissionService, permissionService,
licenseService licenseService
}); });
const identityUaService = identityUaServiceFactory({ const identityUaService = identityUaServiceFactory({
identityOrgMembershipDAL, identityOrgMembershipDAL,
permissionService, permissionService,
@@ -1548,6 +1586,8 @@ export const registerRoutes = async (
permissionService, permissionService,
licenseService, licenseService,
gatewayService, gatewayService,
gatewayV2Service,
gatewayV2DAL,
gatewayDAL, gatewayDAL,
kmsService kmsService
}); });
@@ -1645,8 +1685,10 @@ export const registerRoutes = async (
}); });
const dynamicSecretProviders = buildDynamicSecretProviders({ const dynamicSecretProviders = buildDynamicSecretProviders({
gatewayService gatewayService,
gatewayV2Service
}); });
const dynamicSecretQueueService = dynamicSecretLeaseQueueServiceFactory({ const dynamicSecretQueueService = dynamicSecretLeaseQueueServiceFactory({
queueService, queueService,
dynamicSecretLeaseDAL, dynamicSecretLeaseDAL,
@@ -1666,6 +1708,7 @@ export const registerRoutes = async (
licenseService, licenseService,
kmsService, kmsService,
gatewayDAL, gatewayDAL,
gatewayV2DAL,
resourceMetadataDAL resourceMetadataDAL
}); });
@@ -1682,6 +1725,7 @@ export const registerRoutes = async (
userDAL, userDAL,
identityDAL identityDAL
}); });
const dailyResourceCleanUp = dailyResourceCleanUpQueueServiceFactory({ const dailyResourceCleanUp = dailyResourceCleanUpQueueServiceFactory({
auditLogDAL, auditLogDAL,
queueService, queueService,
@@ -1694,7 +1738,8 @@ export const registerRoutes = async (
identityUniversalAuthClientSecretDAL: identityUaClientSecretDAL, identityUniversalAuthClientSecretDAL: identityUaClientSecretDAL,
serviceTokenService, serviceTokenService,
orgService, orgService,
userNotificationDAL userNotificationDAL,
keyValueStoreDAL
}); });
const dailyReminderQueueService = dailyReminderQueueServiceFactory({ const dailyReminderQueueService = dailyReminderQueueServiceFactory({
@@ -1791,7 +1836,9 @@ export const registerRoutes = async (
kmsService, kmsService,
licenseService, licenseService,
gatewayService, gatewayService,
gatewayDAL gatewayV2Service,
gatewayDAL,
gatewayV2DAL
}); });
const secretSyncService = secretSyncServiceFactory({ const secretSyncService = secretSyncServiceFactory({
@@ -1890,7 +1937,8 @@ export const registerRoutes = async (
secretQueueService, secretQueueService,
queueService, queueService,
appConnectionDAL, appConnectionDAL,
gatewayService gatewayService,
gatewayV2Service
}); });
const certificateAuthorityService = certificateAuthorityServiceFactory({ const certificateAuthorityService = certificateAuthorityServiceFactory({
@@ -2116,6 +2164,8 @@ export const registerRoutes = async (
kmip: kmipService, kmip: kmipService,
kmipOperation: kmipOperationService, kmipOperation: kmipOperationService,
gateway: gatewayService, gateway: gatewayService,
relay: relayService,
gatewayV2: gatewayV2Service,
secretRotationV2: secretRotationV2Service, secretRotationV2: secretRotationV2Service,
microsoftTeams: microsoftTeamsService, microsoftTeams: microsoftTeamsService,
assumePrivileges: assumePrivilegeService, assumePrivileges: assumePrivilegeService,
@@ -129,6 +129,63 @@ export const registerUserRouter = async (server: FastifyZodProvider) => {
} }
}); });
server.route({
method: "POST",
url: "/me/email-change/otp",
config: {
rateLimit: smtpRateLimit({
keyGenerator: (req) => req.permission.id
})
},
schema: {
body: z.object({
newEmail: z.string().email().trim()
}),
response: {
200: z.object({
success: z.boolean(),
message: z.string()
})
}
},
preHandler: verifyAuth([AuthMode.JWT], { requireOrg: false }),
handler: async (req) => {
const result = await server.services.user.requestEmailChangeOTP({
userId: req.permission.id,
newEmail: req.body.newEmail
});
return result;
}
});
server.route({
method: "PATCH",
url: "/me/email",
config: {
rateLimit: writeLimit
},
schema: {
body: z.object({
newEmail: z.string().email().trim(),
otpCode: z.string().trim().length(6)
}),
response: {
200: z.object({
user: UsersSchema
})
}
},
preHandler: verifyAuth([AuthMode.JWT], { requireOrg: false }),
handler: async (req) => {
const user = await server.services.user.updateUserEmail({
userId: req.permission.id,
newEmail: req.body.newEmail,
otpCode: req.body.otpCode
});
return { user };
}
});
server.route({ server.route({
method: "GET", method: "GET",
url: "/me/organizations", url: "/me/organizations",
@@ -6,6 +6,7 @@ import {
} 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 { TGatewayServiceFactory } from "@app/ee/services/gateway/gateway-service";
import { TGatewayV2ServiceFactory } from "@app/ee/services/gateway-v2/gateway-v2-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";
@@ -219,7 +220,8 @@ export const decryptAppConnectionCredentials = async ({
export const validateAppConnectionCredentials = async ( export const validateAppConnectionCredentials = async (
appConnection: TAppConnectionConfig, appConnection: TAppConnectionConfig,
gatewayService: Pick<TGatewayServiceFactory, "fnGetGatewayClientTlsByGatewayId"> gatewayService: Pick<TGatewayServiceFactory, "fnGetGatewayClientTlsByGatewayId">,
gatewayV2Service: Pick<TGatewayV2ServiceFactory, "getPlatformConnectionDetailsByGatewayId">
): 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,
@@ -264,7 +266,7 @@ export const validateAppConnectionCredentials = async (
[AppConnection.Netlify]: validateNetlifyConnectionCredentials as TAppConnectionCredentialsValidator [AppConnection.Netlify]: validateNetlifyConnectionCredentials as TAppConnectionCredentialsValidator
}; };
return VALIDATE_APP_CONNECTION_CREDENTIALS_MAP[appConnection.app](appConnection, gatewayService); return VALIDATE_APP_CONNECTION_CREDENTIALS_MAP[appConnection.app](appConnection, gatewayService, gatewayV2Service);
}; };
export const getAppConnectionMethodName = (method: TAppConnection["method"]) => { export const getAppConnectionMethodName = (method: TAppConnection["method"]) => {
@@ -5,6 +5,8 @@ import { ociConnectionService } from "@app/ee/services/app-connections/oci/oci-c
import { ValidateOracleDBConnectionCredentialsSchema } from "@app/ee/services/app-connections/oracledb"; import { ValidateOracleDBConnectionCredentialsSchema } from "@app/ee/services/app-connections/oracledb";
import { TGatewayDALFactory } from "@app/ee/services/gateway/gateway-dal"; import { TGatewayDALFactory } from "@app/ee/services/gateway/gateway-dal";
import { TGatewayServiceFactory } from "@app/ee/services/gateway/gateway-service"; import { TGatewayServiceFactory } from "@app/ee/services/gateway/gateway-service";
import { TGatewayV2DALFactory } from "@app/ee/services/gateway-v2/gateway-v2-dal";
import { TGatewayV2ServiceFactory } from "@app/ee/services/gateway-v2/gateway-v2-service";
import { TLicenseServiceFactory } from "@app/ee/services/license/license-service"; import { TLicenseServiceFactory } from "@app/ee/services/license/license-service";
import { import {
OrgPermissionAppConnectionActions, OrgPermissionAppConnectionActions,
@@ -110,7 +112,9 @@ export type TAppConnectionServiceFactoryDep = {
kmsService: Pick<TKmsServiceFactory, "createCipherPairWithDataKey">; kmsService: Pick<TKmsServiceFactory, "createCipherPairWithDataKey">;
licenseService: Pick<TLicenseServiceFactory, "getPlan">; licenseService: Pick<TLicenseServiceFactory, "getPlan">;
gatewayService: Pick<TGatewayServiceFactory, "fnGetGatewayClientTlsByGatewayId">; gatewayService: Pick<TGatewayServiceFactory, "fnGetGatewayClientTlsByGatewayId">;
gatewayV2Service: Pick<TGatewayV2ServiceFactory, "getPlatformConnectionDetailsByGatewayId">;
gatewayDAL: Pick<TGatewayDALFactory, "find">; gatewayDAL: Pick<TGatewayDALFactory, "find">;
gatewayV2DAL: Pick<TGatewayV2DALFactory, "find">;
}; };
export type TAppConnectionServiceFactory = ReturnType<typeof appConnectionServiceFactory>; export type TAppConnectionServiceFactory = ReturnType<typeof appConnectionServiceFactory>;
@@ -162,7 +166,9 @@ export const appConnectionServiceFactory = ({
kmsService, kmsService,
licenseService, licenseService,
gatewayService, gatewayService,
gatewayDAL gatewayV2Service,
gatewayDAL,
gatewayV2DAL
}: TAppConnectionServiceFactoryDep) => { }: TAppConnectionServiceFactoryDep) => {
const listAppConnectionsByOrg = async (actor: OrgServiceActor, app?: AppConnection) => { const listAppConnectionsByOrg = async (actor: OrgServiceActor, app?: AppConnection) => {
const { permission } = await permissionService.getOrgPermission( const { permission } = await permissionService.getOrgPermission(
@@ -266,7 +272,8 @@ export const appConnectionServiceFactory = ({
); );
const [gateway] = await gatewayDAL.find({ id: gatewayId, orgId: actor.orgId }); const [gateway] = await gatewayDAL.find({ id: gatewayId, orgId: actor.orgId });
if (!gateway) { const [gatewayV2] = await gatewayV2DAL.find({ id: gatewayId, orgId: actor.orgId });
if (!gateway && !gatewayV2) {
throw new NotFoundError({ throw new NotFoundError({
message: `Gateway with ID ${gatewayId} not found for org` message: `Gateway with ID ${gatewayId} not found for org`
}); });
@@ -288,7 +295,8 @@ export const appConnectionServiceFactory = ({
orgId: actor.orgId, orgId: actor.orgId,
gatewayId gatewayId
} as TAppConnectionConfig, } as TAppConnectionConfig,
gatewayService gatewayService,
gatewayV2Service
); );
try { try {
@@ -321,7 +329,8 @@ export const appConnectionServiceFactory = ({
gatewayId gatewayId
} as TAppConnectionConfig, } as TAppConnectionConfig,
(platformCredentials) => createConnection(platformCredentials), (platformCredentials) => createConnection(platformCredentials),
gatewayService gatewayService,
gatewayV2Service
); );
} else { } else {
connection = await createConnection(validatedCredentials); connection = await createConnection(validatedCredentials);
@@ -377,7 +386,8 @@ export const appConnectionServiceFactory = ({
if (gatewayId) { if (gatewayId) {
const [gateway] = await gatewayDAL.find({ id: gatewayId, orgId: actor.orgId }); const [gateway] = await gatewayDAL.find({ id: gatewayId, orgId: actor.orgId });
if (!gateway) { const [gatewayV2] = await gatewayV2DAL.find({ id: gatewayId, orgId: actor.orgId });
if (!gateway && !gatewayV2) {
throw new NotFoundError({ throw new NotFoundError({
message: `Gateway with ID ${gatewayId} not found for org` message: `Gateway with ID ${gatewayId} not found for org`
}); });
@@ -417,7 +427,8 @@ export const appConnectionServiceFactory = ({
method, method,
gatewayId gatewayId
} as TAppConnectionConfig, } as TAppConnectionConfig,
gatewayService gatewayService,
gatewayV2Service
); );
if (!updatedCredentials) if (!updatedCredentials)
@@ -458,7 +469,8 @@ export const appConnectionServiceFactory = ({
gatewayId gatewayId
} as TAppConnectionConfig, } as TAppConnectionConfig,
(platformCredentials) => updateConnection(platformCredentials), (platformCredentials) => updateConnection(platformCredentials),
gatewayService gatewayService,
gatewayV2Service
); );
} else { } else {
updatedConnection = await updateConnection(updatedCredentials); updatedConnection = await updateConnection(updatedCredentials);
@@ -588,7 +600,7 @@ export const appConnectionServiceFactory = ({
deleteAppConnection, deleteAppConnection,
connectAppConnectionById, connectAppConnectionById,
listAvailableAppConnectionsForUser, listAvailableAppConnectionsForUser,
github: githubConnectionService(connectAppConnectionById, gatewayService), github: githubConnectionService(connectAppConnectionById, gatewayService, gatewayV2Service),
githubRadar: githubRadarConnectionService(connectAppConnectionById), githubRadar: githubRadarConnectionService(connectAppConnectionById),
gcp: gcpConnectionService(connectAppConnectionById), gcp: gcpConnectionService(connectAppConnectionById),
databricks: databricksConnectionService(connectAppConnectionById, appConnectionDAL, kmsService), databricks: databricksConnectionService(connectAppConnectionById, appConnectionDAL, kmsService),
@@ -10,6 +10,7 @@ import {
TValidateOracleDBConnectionCredentialsSchema TValidateOracleDBConnectionCredentialsSchema
} from "@app/ee/services/app-connections/oracledb"; } from "@app/ee/services/app-connections/oracledb";
import { TGatewayServiceFactory } from "@app/ee/services/gateway/gateway-service"; import { TGatewayServiceFactory } from "@app/ee/services/gateway/gateway-service";
import { TGatewayV2ServiceFactory } from "@app/ee/services/gateway-v2/gateway-v2-service";
import { TAppConnectionDALFactory } from "@app/services/app-connection/app-connection-dal"; import { TAppConnectionDALFactory } from "@app/services/app-connection/app-connection-dal";
import { TSqlConnectionConfig } from "@app/services/app-connection/shared/sql/sql-connection-types"; import { TSqlConnectionConfig } from "@app/services/app-connection/shared/sql/sql-connection-types";
import { SecretSync } from "@app/services/secret-sync/secret-sync-enums"; import { SecretSync } from "@app/services/secret-sync/secret-sync-enums";
@@ -411,13 +412,15 @@ export type TListAwsConnectionIamUsers = {
export type TAppConnectionCredentialsValidator = ( export type TAppConnectionCredentialsValidator = (
appConnection: TAppConnectionConfig, appConnection: TAppConnectionConfig,
gatewayService: Pick<TGatewayServiceFactory, "fnGetGatewayClientTlsByGatewayId"> gatewayService: Pick<TGatewayServiceFactory, "fnGetGatewayClientTlsByGatewayId">,
gatewayV2Service: Pick<TGatewayV2ServiceFactory, "getPlatformConnectionDetailsByGatewayId">
) => Promise<TAppConnection["credentials"]>; ) => Promise<TAppConnection["credentials"]>;
export type TAppConnectionTransitionCredentialsToPlatform = ( export type TAppConnectionTransitionCredentialsToPlatform = (
appConnection: TAppConnectionConfig, appConnection: TAppConnectionConfig,
callback: (credentials: TAppConnection["credentials"]) => Promise<TAppConnectionRaw>, callback: (credentials: TAppConnection["credentials"]) => Promise<TAppConnectionRaw>,
gatewayService: Pick<TGatewayServiceFactory, "fnGetGatewayClientTlsByGatewayId"> gatewayService: Pick<TGatewayServiceFactory, "fnGetGatewayClientTlsByGatewayId">,
gatewayV2Service: Pick<TGatewayV2ServiceFactory, "getPlatformConnectionDetailsByGatewayId">
) => Promise<TAppConnectionRaw>; ) => Promise<TAppConnectionRaw>;
export type TAppConnectionBaseConfig = { export type TAppConnectionBaseConfig = {
@@ -4,11 +4,13 @@ import RE2 from "re2";
import { verifyHostInputValidity } from "@app/ee/services/dynamic-secret/dynamic-secret-fns"; import { verifyHostInputValidity } from "@app/ee/services/dynamic-secret/dynamic-secret-fns";
import { TGatewayServiceFactory } from "@app/ee/services/gateway/gateway-service"; import { TGatewayServiceFactory } from "@app/ee/services/gateway/gateway-service";
import { TGatewayV2ServiceFactory } from "@app/ee/services/gateway-v2/gateway-v2-service";
import { getConfig } from "@app/lib/config/env"; import { getConfig } from "@app/lib/config/env";
import { request as httpRequest } from "@app/lib/config/request"; import { request as httpRequest } from "@app/lib/config/request";
import { crypto } from "@app/lib/crypto"; import { crypto } from "@app/lib/crypto";
import { BadRequestError, ForbiddenRequestError, InternalServerError } from "@app/lib/errors"; import { BadRequestError, ForbiddenRequestError, InternalServerError } from "@app/lib/errors";
import { GatewayProxyProtocol, withGatewayProxy } from "@app/lib/gateway"; import { GatewayProxyProtocol, withGatewayProxy } from "@app/lib/gateway";
import { withGatewayV2Proxy } from "@app/lib/gateway-v2/gateway-v2";
import { logger } from "@app/lib/logger"; import { logger } from "@app/lib/logger";
import { blockLocalAndPrivateIpAddresses } from "@app/lib/validator"; import { blockLocalAndPrivateIpAddresses } from "@app/lib/validator";
import { getAppConnectionMethodName } from "@app/services/app-connection/app-connection-fns"; import { getAppConnectionMethodName } from "@app/services/app-connection/app-connection-fns";
@@ -49,6 +51,7 @@ export const getGitHubInstanceApiUrl = async (config: {
export const requestWithGitHubGateway = async <T>( export const requestWithGitHubGateway = async <T>(
appConnection: { gatewayId?: string | null }, appConnection: { gatewayId?: string | null },
gatewayService: Pick<TGatewayServiceFactory, "fnGetGatewayClientTlsByGatewayId">, gatewayService: Pick<TGatewayServiceFactory, "fnGetGatewayClientTlsByGatewayId">,
gatewayV2Service: Pick<TGatewayV2ServiceFactory, "getPlatformConnectionDetailsByGatewayId">,
requestConfig: AxiosRequestConfig requestConfig: AxiosRequestConfig
): Promise<AxiosResponse<T>> => { ): Promise<AxiosResponse<T>> => {
const { gatewayId } = appConnection; const { gatewayId } = appConnection;
@@ -63,6 +66,52 @@ export const requestWithGitHubGateway = async <T>(
await blockLocalAndPrivateIpAddresses(url.toString()); await blockLocalAndPrivateIpAddresses(url.toString());
const [targetHost] = await verifyHostInputValidity(url.host, true); const [targetHost] = await verifyHostInputValidity(url.host, true);
const gatewayConnectionDetails = await gatewayV2Service.getPlatformConnectionDetailsByGatewayId({
gatewayId,
targetHost,
targetPort: 443
});
if (gatewayConnectionDetails) {
return withGatewayV2Proxy(
async (proxyPort) => {
const httpsAgent = new https.Agent({
servername: targetHost
});
url.protocol = "https:";
url.host = `localhost:${proxyPort}`;
const finalRequestConfig: AxiosRequestConfig = {
...requestConfig,
url: url.toString(),
httpsAgent,
headers: {
...requestConfig.headers,
Host: targetHost
}
};
try {
return await httpRequest.request(finalRequestConfig);
} catch (error) {
const axiosError = error as AxiosError;
logger.error(
{ message: axiosError.message, data: axiosError.response?.data },
"Error during GitHub gateway request:"
);
throw error;
}
},
{
protocol: GatewayProxyProtocol.Tcp,
relayHost: gatewayConnectionDetails.relayHost,
gateway: gatewayConnectionDetails.gateway,
relay: gatewayConnectionDetails.relay
}
);
}
const relayDetails = await gatewayService.fnGetGatewayClientTlsByGatewayId(gatewayId); const relayDetails = await gatewayService.fnGetGatewayClientTlsByGatewayId(gatewayId);
const [relayHost, relayPort] = relayDetails.relayAddress.split(":"); const [relayHost, relayPort] = relayDetails.relayAddress.split(":");
@@ -115,7 +164,8 @@ export const requestWithGitHubGateway = async <T>(
export const getGitHubAppAuthToken = async ( export const getGitHubAppAuthToken = async (
appConnection: TGitHubConnection, appConnection: TGitHubConnection,
gatewayService: Pick<TGatewayServiceFactory, "fnGetGatewayClientTlsByGatewayId"> gatewayService: Pick<TGatewayServiceFactory, "fnGetGatewayClientTlsByGatewayId">,
gatewayV2Service: Pick<TGatewayV2ServiceFactory, "getPlatformConnectionDetailsByGatewayId">
) => { ) => {
const appCfg = getConfig(); const appCfg = getConfig();
const appId = appCfg.INF_APP_CONNECTION_GITHUB_APP_ID; const appId = appCfg.INF_APP_CONNECTION_GITHUB_APP_ID;
@@ -151,6 +201,7 @@ export const getGitHubAppAuthToken = async (
const response = await requestWithGitHubGateway<{ token: string; expires_at: string }>( const response = await requestWithGitHubGateway<{ token: string; expires_at: string }>(
appConnection, appConnection,
gatewayService, gatewayService,
gatewayV2Service,
{ {
url: `https://${apiBaseUrl}/app/installations/${installationId}/access_tokens`, url: `https://${apiBaseUrl}/app/installations/${installationId}/access_tokens`,
method: "POST", method: "POST",
@@ -191,6 +242,7 @@ function extractNextPageUrl(linkHeader: string | undefined): string | null {
export const makePaginatedGitHubRequest = async <T, R = T[]>( export const makePaginatedGitHubRequest = async <T, R = T[]>(
appConnection: TGitHubConnection, appConnection: TGitHubConnection,
gatewayService: Pick<TGatewayServiceFactory, "fnGetGatewayClientTlsByGatewayId">, gatewayService: Pick<TGatewayServiceFactory, "fnGetGatewayClientTlsByGatewayId">,
gatewayV2Service: Pick<TGatewayV2ServiceFactory, "getPlatformConnectionDetailsByGatewayId">,
path: string, path: string,
dataMapper?: (data: R) => T[] dataMapper?: (data: R) => T[]
): Promise<T[]> => { ): Promise<T[]> => {
@@ -199,7 +251,7 @@ export const makePaginatedGitHubRequest = async <T, R = T[]>(
const token = const token =
method === GitHubConnectionMethod.OAuth method === GitHubConnectionMethod.OAuth
? credentials.accessToken ? credentials.accessToken
: await getGitHubAppAuthToken(appConnection, gatewayService); : await getGitHubAppAuthToken(appConnection, gatewayService, gatewayV2Service);
const baseUrl = `https://${await getGitHubInstanceApiUrl(appConnection)}${path}`; const baseUrl = `https://${await getGitHubInstanceApiUrl(appConnection)}${path}`;
const initialUrlObj = new URL(baseUrl); const initialUrlObj = new URL(baseUrl);
@@ -209,15 +261,20 @@ export const makePaginatedGitHubRequest = async <T, R = T[]>(
const maxIterations = 1000; const maxIterations = 1000;
// Make initial request to get link header // Make initial request to get link header
const firstResponse: AxiosResponse<R> = await requestWithGitHubGateway<R>(appConnection, gatewayService, { const firstResponse: AxiosResponse<R> = await requestWithGitHubGateway<R>(
url: initialUrlObj.toString(), appConnection,
method: "GET", gatewayService,
headers: { gatewayV2Service,
Accept: "application/vnd.github+json", {
Authorization: `Bearer ${token}`, url: initialUrlObj.toString(),
"X-GitHub-Api-Version": "2022-11-28" method: "GET",
headers: {
Accept: "application/vnd.github+json",
Authorization: `Bearer ${token}`,
"X-GitHub-Api-Version": "2022-11-28"
}
} }
}); );
const firstPageItems = dataMapper ? dataMapper(firstResponse.data) : (firstResponse.data as unknown as T[]); const firstPageItems = dataMapper ? dataMapper(firstResponse.data) : (firstResponse.data as unknown as T[]);
results = results.concat(firstPageItems); results = results.concat(firstPageItems);
@@ -237,7 +294,7 @@ export const makePaginatedGitHubRequest = async <T, R = T[]>(
pageUrlObj.searchParams.set("page", pageNum.toString()); pageUrlObj.searchParams.set("page", pageNum.toString());
pageRequests.push( pageRequests.push(
requestWithGitHubGateway<R>(appConnection, gatewayService, { requestWithGitHubGateway<R>(appConnection, gatewayService, gatewayV2Service, {
url: pageUrlObj.toString(), url: pageUrlObj.toString(),
method: "GET", method: "GET",
headers: { headers: {
@@ -261,15 +318,20 @@ export const makePaginatedGitHubRequest = async <T, R = T[]>(
while (url && i < maxIterations) { while (url && i < maxIterations) {
// eslint-disable-next-line no-await-in-loop // eslint-disable-next-line no-await-in-loop
const response: AxiosResponse<R> = await requestWithGitHubGateway<R>(appConnection, gatewayService, { const response: AxiosResponse<R> = await requestWithGitHubGateway<R>(
url, appConnection,
method: "GET", gatewayService,
headers: { gatewayV2Service,
Accept: "application/vnd.github+json", {
Authorization: `Bearer ${token}`, url,
"X-GitHub-Api-Version": "2022-11-28" method: "GET",
headers: {
Accept: "application/vnd.github+json",
Authorization: `Bearer ${token}`,
"X-GitHub-Api-Version": "2022-11-28"
}
} }
}); );
const items = dataMapper ? dataMapper(response.data) : (response.data as unknown as T[]); const items = dataMapper ? dataMapper(response.data) : (response.data as unknown as T[]);
results = results.concat(items); results = results.concat(items);
@@ -308,30 +370,39 @@ type GitHubEnvironment = {
export const getGitHubRepositories = async ( export const getGitHubRepositories = async (
appConnection: TGitHubConnection, appConnection: TGitHubConnection,
gatewayService: Pick<TGatewayServiceFactory, "fnGetGatewayClientTlsByGatewayId"> gatewayService: Pick<TGatewayServiceFactory, "fnGetGatewayClientTlsByGatewayId">,
gatewayV2Service: Pick<TGatewayV2ServiceFactory, "getPlatformConnectionDetailsByGatewayId">
) => { ) => {
if (appConnection.method === GitHubConnectionMethod.App) { if (appConnection.method === GitHubConnectionMethod.App) {
return makePaginatedGitHubRequest<GitHubRepository, { repositories: GitHubRepository[] }>( return makePaginatedGitHubRequest<GitHubRepository, { repositories: GitHubRepository[] }>(
appConnection, appConnection,
gatewayService, gatewayService,
gatewayV2Service,
"/installation/repositories", "/installation/repositories",
(data) => data.repositories (data) => data.repositories
); );
} }
const repos = await makePaginatedGitHubRequest<GitHubRepository>(appConnection, gatewayService, "/user/repos"); const repos = await makePaginatedGitHubRequest<GitHubRepository>(
appConnection,
gatewayService,
gatewayV2Service,
"/user/repos"
);
return repos.filter((repo) => repo.permissions?.admin); return repos.filter((repo) => repo.permissions?.admin);
}; };
export const getGitHubOrganizations = async ( export const getGitHubOrganizations = async (
appConnection: TGitHubConnection, appConnection: TGitHubConnection,
gatewayService: Pick<TGatewayServiceFactory, "fnGetGatewayClientTlsByGatewayId"> gatewayService: Pick<TGatewayServiceFactory, "fnGetGatewayClientTlsByGatewayId">,
gatewayV2Service: Pick<TGatewayV2ServiceFactory, "getPlatformConnectionDetailsByGatewayId">
) => { ) => {
if (appConnection.method === GitHubConnectionMethod.App) { if (appConnection.method === GitHubConnectionMethod.App) {
const installationRepositories = await makePaginatedGitHubRequest< const installationRepositories = await makePaginatedGitHubRequest<
GitHubRepository, GitHubRepository,
{ repositories: GitHubRepository[] } { repositories: GitHubRepository[] }
>(appConnection, gatewayService, "/installation/repositories", (data) => data.repositories); >(appConnection, gatewayService, gatewayV2Service, "/installation/repositories", (data) => data.repositories);
const organizationMap: Record<string, GitHubOrganization> = {}; const organizationMap: Record<string, GitHubOrganization> = {};
installationRepositories.forEach((repo) => { installationRepositories.forEach((repo) => {
@@ -343,12 +414,13 @@ export const getGitHubOrganizations = async (
return Object.values(organizationMap); return Object.values(organizationMap);
} }
return makePaginatedGitHubRequest<GitHubOrganization>(appConnection, gatewayService, "/user/orgs"); return makePaginatedGitHubRequest<GitHubOrganization>(appConnection, gatewayService, gatewayV2Service, "/user/orgs");
}; };
export const getGitHubEnvironments = async ( export const getGitHubEnvironments = async (
appConnection: TGitHubConnection, appConnection: TGitHubConnection,
gatewayService: Pick<TGatewayServiceFactory, "fnGetGatewayClientTlsByGatewayId">, gatewayService: Pick<TGatewayServiceFactory, "fnGetGatewayClientTlsByGatewayId">,
gatewayV2Service: Pick<TGatewayV2ServiceFactory, "getPlatformConnectionDetailsByGatewayId">,
owner: string, owner: string,
repo: string repo: string
) => { ) => {
@@ -356,6 +428,7 @@ export const getGitHubEnvironments = async (
return await makePaginatedGitHubRequest<GitHubEnvironment, { environments: GitHubEnvironment[] }>( return await makePaginatedGitHubRequest<GitHubEnvironment, { environments: GitHubEnvironment[] }>(
appConnection, appConnection,
gatewayService, gatewayService,
gatewayV2Service,
`/repos/${encodeURIComponent(owner)}/${encodeURIComponent(repo)}/environments`, `/repos/${encodeURIComponent(owner)}/${encodeURIComponent(repo)}/environments`,
(data) => data.environments (data) => data.environments
); );
@@ -383,7 +456,8 @@ export function isGithubErrorResponse(data: GithubTokenRespData): data is Github
export const validateGitHubConnectionCredentials = async ( export const validateGitHubConnectionCredentials = async (
config: TGitHubConnectionConfig, config: TGitHubConnectionConfig,
gatewayService: Pick<TGatewayServiceFactory, "fnGetGatewayClientTlsByGatewayId"> gatewayService: Pick<TGatewayServiceFactory, "fnGetGatewayClientTlsByGatewayId">,
gatewayV2Service: Pick<TGatewayV2ServiceFactory, "getPlatformConnectionDetailsByGatewayId">
) => { ) => {
const { credentials, method } = config; const { credentials, method } = config;
const { const {
@@ -419,7 +493,7 @@ export const validateGitHubConnectionCredentials = async (
const host = credentials.host || "github.com"; const host = credentials.host || "github.com";
try { try {
tokenResp = await requestWithGitHubGateway<GithubTokenRespData>(config, gatewayService, { tokenResp = await requestWithGitHubGateway<GithubTokenRespData>(config, gatewayService, gatewayV2Service, {
url: `https://${host}/login/oauth/access_token`, url: `https://${host}/login/oauth/access_token`,
method: "POST", method: "POST",
data: { data: {
@@ -471,7 +545,7 @@ export const validateGitHubConnectionCredentials = async (
id: number; id: number;
}; };
}[]; }[];
}>(config, gatewayService, { }>(config, gatewayService, gatewayV2Service, {
url: `https://${await getGitHubInstanceApiUrl(config)}/user/installations`, url: `https://${await getGitHubInstanceApiUrl(config)}/user/installations`,
headers: { headers: {
Accept: "application/json", Accept: "application/json",
@@ -1,4 +1,5 @@
import { TGatewayServiceFactory } from "@app/ee/services/gateway/gateway-service"; import { TGatewayServiceFactory } from "@app/ee/services/gateway/gateway-service";
import { TGatewayV2ServiceFactory } from "@app/ee/services/gateway-v2/gateway-v2-service";
import { OrgServiceActor } from "@app/lib/types"; import { OrgServiceActor } from "@app/lib/types";
import { AppConnection } from "@app/services/app-connection/app-connection-enums"; import { AppConnection } from "@app/services/app-connection/app-connection-enums";
import { import {
@@ -22,12 +23,13 @@ type TListGitHubEnvironmentsDTO = {
export const githubConnectionService = ( export const githubConnectionService = (
getAppConnection: TGetAppConnectionFunc, getAppConnection: TGetAppConnectionFunc,
gatewayService: Pick<TGatewayServiceFactory, "fnGetGatewayClientTlsByGatewayId"> gatewayService: Pick<TGatewayServiceFactory, "fnGetGatewayClientTlsByGatewayId">,
gatewayV2Service: Pick<TGatewayV2ServiceFactory, "getPlatformConnectionDetailsByGatewayId">
) => { ) => {
const listRepositories = async (connectionId: string, actor: OrgServiceActor) => { const listRepositories = async (connectionId: string, actor: OrgServiceActor) => {
const appConnection = await getAppConnection(AppConnection.GitHub, connectionId, actor); const appConnection = await getAppConnection(AppConnection.GitHub, connectionId, actor);
const repositories = await getGitHubRepositories(appConnection, gatewayService); const repositories = await getGitHubRepositories(appConnection, gatewayService, gatewayV2Service);
return repositories; return repositories;
}; };
@@ -35,7 +37,7 @@ export const githubConnectionService = (
const listOrganizations = async (connectionId: string, actor: OrgServiceActor) => { const listOrganizations = async (connectionId: string, actor: OrgServiceActor) => {
const appConnection = await getAppConnection(AppConnection.GitHub, connectionId, actor); const appConnection = await getAppConnection(AppConnection.GitHub, connectionId, actor);
const organizations = await getGitHubOrganizations(appConnection, gatewayService); const organizations = await getGitHubOrganizations(appConnection, gatewayService, gatewayV2Service);
return organizations; return organizations;
}; };
@@ -46,7 +48,7 @@ export const githubConnectionService = (
) => { ) => {
const appConnection = await getAppConnection(AppConnection.GitHub, connectionId, actor); const appConnection = await getAppConnection(AppConnection.GitHub, connectionId, actor);
const environments = await getGitHubEnvironments(appConnection, gatewayService, owner, repo); const environments = await getGitHubEnvironments(appConnection, gatewayService, gatewayV2Service, owner, repo);
return environments; return environments;
}; };
@@ -3,6 +3,7 @@ import https from "https";
import { verifyHostInputValidity } from "@app/ee/services/dynamic-secret/dynamic-secret-fns"; import { verifyHostInputValidity } from "@app/ee/services/dynamic-secret/dynamic-secret-fns";
import { TGatewayServiceFactory } from "@app/ee/services/gateway/gateway-service"; import { TGatewayServiceFactory } from "@app/ee/services/gateway/gateway-service";
import { TGatewayV2ServiceFactory } from "@app/ee/services/gateway-v2/gateway-v2-service";
import { request } from "@app/lib/config/request"; import { request } from "@app/lib/config/request";
import { BadRequestError } from "@app/lib/errors"; import { BadRequestError } from "@app/lib/errors";
import { removeTrailingSlash } from "@app/lib/fn"; import { removeTrailingSlash } from "@app/lib/fn";
@@ -144,7 +145,9 @@ export const getHCVaultAccessToken = async (
export const validateHCVaultConnectionCredentials = async ( export const validateHCVaultConnectionCredentials = async (
connection: THCVaultConnection, connection: THCVaultConnection,
gatewayService: Pick<TGatewayServiceFactory, "fnGetGatewayClientTlsByGatewayId"> gatewayService: Pick<TGatewayServiceFactory, "fnGetGatewayClientTlsByGatewayId">,
// eslint-disable-next-line @typescript-eslint/no-unused-vars
_gatewayV2Service: Pick<TGatewayV2ServiceFactory, "getPlatformConnectionDetailsByGatewayId">
) => { ) => {
const instanceUrl = await getHCVaultInstanceUrl(connection); const instanceUrl = await getHCVaultInstanceUrl(connection);
@@ -2,12 +2,14 @@ import knex, { Knex } from "knex";
import { verifyHostInputValidity } from "@app/ee/services/dynamic-secret/dynamic-secret-fns"; import { verifyHostInputValidity } from "@app/ee/services/dynamic-secret/dynamic-secret-fns";
import { TGatewayServiceFactory } from "@app/ee/services/gateway/gateway-service"; import { TGatewayServiceFactory } from "@app/ee/services/gateway/gateway-service";
import { TGatewayV2ServiceFactory } from "@app/ee/services/gateway-v2/gateway-v2-service";
import { import {
TSqlCredentialsRotationGeneratedCredentials, TSqlCredentialsRotationGeneratedCredentials,
TSqlCredentialsRotationWithConnection TSqlCredentialsRotationWithConnection
} from "@app/ee/services/secret-rotation-v2/shared/sql-credentials/sql-credentials-rotation-types"; } from "@app/ee/services/secret-rotation-v2/shared/sql-credentials/sql-credentials-rotation-types";
import { BadRequestError, DatabaseError } from "@app/lib/errors"; import { BadRequestError, DatabaseError } from "@app/lib/errors";
import { GatewayProxyProtocol, withGatewayProxy } from "@app/lib/gateway"; import { GatewayProxyProtocol, withGatewayProxy } from "@app/lib/gateway";
import { withGatewayV2Proxy } from "@app/lib/gateway-v2/gateway-v2";
import { alphaNumericNanoId } from "@app/lib/nanoid"; import { alphaNumericNanoId } from "@app/lib/nanoid";
import { AppConnection } from "@app/services/app-connection/app-connection-enums"; import { AppConnection } from "@app/services/app-connection/app-connection-enums";
import { TAppConnectionRaw, TSqlConnection } from "@app/services/app-connection/app-connection-types"; import { TAppConnectionRaw, TSqlConnection } from "@app/services/app-connection/app-connection-types";
@@ -104,12 +106,49 @@ export const getSqlConnectionClient = async (appConnection: Pick<TSqlConnection,
export const executeWithPotentialGateway = async <T>( export const executeWithPotentialGateway = async <T>(
config: TSqlConnectionConfig, config: TSqlConnectionConfig,
gatewayService: Pick<TGatewayServiceFactory, "fnGetGatewayClientTlsByGatewayId">, gatewayService: Pick<TGatewayServiceFactory, "fnGetGatewayClientTlsByGatewayId">,
gatewayV2Service: Pick<TGatewayV2ServiceFactory, "getPlatformConnectionDetailsByGatewayId">,
operation: (client: Knex) => Promise<T> operation: (client: Knex) => Promise<T>
): Promise<T> => { ): Promise<T> => {
const { credentials, app, gatewayId } = config; const { credentials, app, gatewayId } = config;
if (gatewayId && gatewayService) { if (gatewayId && gatewayService && gatewayV2Service) {
const [targetHost] = await verifyHostInputValidity(credentials.host, true); const [targetHost] = await verifyHostInputValidity(credentials.host, true);
const platformConnectionDetails = await gatewayV2Service.getPlatformConnectionDetailsByGatewayId({
gatewayId,
targetHost,
targetPort: credentials.port
});
if (platformConnectionDetails) {
return withGatewayV2Proxy(
async (proxyPort) => {
const client = knex({
client: SQL_CONNECTION_CLIENT_MAP[app],
connection: {
database: credentials.database,
port: proxyPort,
host: "localhost",
user: credentials.username,
password: credentials.password,
connectionTimeoutMillis: EXTERNAL_REQUEST_TIMEOUT,
...getConnectionConfig({ app, credentials })
}
});
try {
return await operation(client);
} finally {
await client.destroy();
}
},
{
protocol: GatewayProxyProtocol.Tcp,
relayHost: platformConnectionDetails.relayHost,
gateway: platformConnectionDetails.gateway,
relay: platformConnectionDetails.relay
}
);
}
const relayDetails = await gatewayService.fnGetGatewayClientTlsByGatewayId(gatewayId); const relayDetails = await gatewayService.fnGetGatewayClientTlsByGatewayId(gatewayId);
const [relayHost, relayPort] = relayDetails.relayAddress.split(":"); const [relayHost, relayPort] = relayDetails.relayAddress.split(":");
@@ -161,10 +200,11 @@ export const executeWithPotentialGateway = async <T>(
export const validateSqlConnectionCredentials = async ( export const validateSqlConnectionCredentials = async (
config: TSqlConnectionConfig, config: TSqlConnectionConfig,
gatewayService: Pick<TGatewayServiceFactory, "fnGetGatewayClientTlsByGatewayId"> gatewayService: Pick<TGatewayServiceFactory, "fnGetGatewayClientTlsByGatewayId">,
gatewayV2Service: Pick<TGatewayV2ServiceFactory, "getPlatformConnectionDetailsByGatewayId">
) => { ) => {
try { try {
await executeWithPotentialGateway(config, gatewayService, async (client) => { await executeWithPotentialGateway(config, gatewayService, gatewayV2Service, async (client) => {
await client.raw(config.app === AppConnection.OracleDB ? `SELECT 1 FROM DUAL` : `Select 1`); await client.raw(config.app === AppConnection.OracleDB ? `SELECT 1 FROM DUAL` : `Select 1`);
}); });
return config.credentials; return config.credentials;
@@ -191,14 +231,15 @@ export const SQL_CONNECTION_ALTER_LOGIN_STATEMENT: Record<
export const transferSqlConnectionCredentialsToPlatform = async ( export const transferSqlConnectionCredentialsToPlatform = async (
config: TSqlConnectionConfig, config: TSqlConnectionConfig,
callback: (credentials: TSqlConnectionConfig["credentials"]) => Promise<TAppConnectionRaw>, callback: (credentials: TSqlConnectionConfig["credentials"]) => Promise<TAppConnectionRaw>,
gatewayService: Pick<TGatewayServiceFactory, "fnGetGatewayClientTlsByGatewayId"> gatewayService: Pick<TGatewayServiceFactory, "fnGetGatewayClientTlsByGatewayId">,
gatewayV2Service: Pick<TGatewayV2ServiceFactory, "getPlatformConnectionDetailsByGatewayId">
) => { ) => {
const { credentials, app } = config; const { credentials, app } = config;
const newPassword = alphaNumericNanoId(32); const newPassword = alphaNumericNanoId(32);
try { try {
return await executeWithPotentialGateway(config, gatewayService, (client) => { return await executeWithPotentialGateway(config, gatewayService, gatewayV2Service, (client) => {
return client.transaction(async (tx) => { return client.transaction(async (tx) => {
await tx.raw( await tx.raw(
...SQL_CONNECTION_ALTER_LOGIN_STATEMENT[app]({ username: credentials.username, password: newPassword }) ...SQL_CONNECTION_ALTER_LOGIN_STATEMENT[app]({ username: credentials.username, password: newPassword })
@@ -36,6 +36,12 @@ export const getTokenConfig = (tokenType: TokenType) => {
const expiresAt = new Date(new Date().getTime() + 86400000); const expiresAt = new Date(new Date().getTime() + 86400000);
return { token, triesLeft, expiresAt }; return { token, triesLeft, expiresAt };
} }
case TokenType.TOKEN_EMAIL_CHANGE_OTP: {
const token = String(crypto.randomInt(10 ** 5, 10 ** 6 - 1));
const triesLeft = 1;
const expiresAt = new Date(new Date().getTime() + 600000);
return { token, triesLeft, expiresAt };
}
case TokenType.TOKEN_EMAIL_MFA: { case TokenType.TOKEN_EMAIL_MFA: {
// generate random 6-digit code // generate random 6-digit code
const token = String(crypto.randomInt(10 ** 5, 10 ** 6 - 1)); const token = String(crypto.randomInt(10 ** 5, 10 ** 6 - 1));
@@ -75,7 +81,7 @@ export const getTokenConfig = (tokenType: TokenType) => {
}; };
export const tokenServiceFactory = ({ tokenDAL, userDAL, orgMembershipDAL }: TAuthTokenServiceFactoryDep) => { export const tokenServiceFactory = ({ tokenDAL, userDAL, orgMembershipDAL }: TAuthTokenServiceFactoryDep) => {
const createTokenForUser = async ({ type, userId, orgId, aliasId }: TCreateTokenForUserDTO) => { const createTokenForUser = async ({ type, userId, orgId, aliasId, payload }: TCreateTokenForUserDTO) => {
const { token, ...tkCfg } = getTokenConfig(type); const { token, ...tkCfg } = getTokenConfig(type);
const appCfg = getConfig(); const appCfg = getConfig();
const tokenHash = await crypto.hashing().createHash(token, appCfg.SALT_ROUNDS); const tokenHash = await crypto.hashing().createHash(token, appCfg.SALT_ROUNDS);
@@ -89,7 +95,8 @@ export const tokenServiceFactory = ({ tokenDAL, userDAL, orgMembershipDAL }: TAu
userId, userId,
orgId, orgId,
triesLeft: tkCfg?.triesLeft, triesLeft: tkCfg?.triesLeft,
aliasId aliasId,
payload
}, },
tx tx
); );
@@ -3,6 +3,7 @@ import { ProjectMembershipRole } from "@app/db/schemas";
export enum TokenType { export enum TokenType {
TOKEN_EMAIL_CONFIRMATION = "emailConfirmation", TOKEN_EMAIL_CONFIRMATION = "emailConfirmation",
TOKEN_EMAIL_VERIFICATION = "emailVerification", // unverified -> verified TOKEN_EMAIL_VERIFICATION = "emailVerification", // unverified -> verified
TOKEN_EMAIL_CHANGE_OTP = "emailChangeOtp",
TOKEN_EMAIL_MFA = "emailMfa", TOKEN_EMAIL_MFA = "emailMfa",
TOKEN_EMAIL_ORG_INVITATION = "organizationInvitation", TOKEN_EMAIL_ORG_INVITATION = "organizationInvitation",
TOKEN_EMAIL_PASSWORD_RESET = "passwordReset", TOKEN_EMAIL_PASSWORD_RESET = "passwordReset",
@@ -15,6 +16,7 @@ export type TCreateTokenForUserDTO = {
userId: string; userId: string;
orgId?: string; orgId?: string;
aliasId?: string; aliasId?: string;
payload?: string;
}; };
export type TCreateOrgInviteTokenDTO = { export type TCreateOrgInviteTokenDTO = {
@@ -52,6 +52,9 @@ export const constructPemChainFromCerts = (certificates: x509.X509Certificate[])
.join("\n") .join("\n")
.trim(); .trim();
export const prependCertToPemChain = (cert: x509.X509Certificate, pemChain: string) =>
`${cert.toString("pem")}\n${pemChain}`;
export const splitPemChain = (pemText: string) => { export const splitPemChain = (pemText: string) => {
const re2Pattern = new RE2("-----BEGIN CERTIFICATE-----[^-]+-----END CERTIFICATE-----", "g"); const re2Pattern = new RE2("-----BEGIN CERTIFICATE-----[^-]+-----END CERTIFICATE-----", "g");
@@ -661,7 +661,7 @@ describe("folderCommitServiceFactory", () => {
// Assert // Assert
expect(mockFolderCommitDAL.create).toHaveBeenCalled(); expect(mockFolderCommitDAL.create).toHaveBeenCalled();
expect(mockSecretV2BridgeDAL.invalidateSecretCacheByProjectId).toHaveBeenCalledWith(projectId); expect(mockSecretV2BridgeDAL.invalidateSecretCacheByProjectId).toHaveBeenCalledWith(projectId, {});
// Check that we got the right counts // Check that we got the right counts
expect(result.totalChanges).toEqual(2); expect(result.totalChanges).toEqual(2);
@@ -1386,7 +1386,7 @@ export const folderCommitServiceFactory = ({
); );
// Invalidate cache to reflect the changes // Invalidate cache to reflect the changes
await secretV2BridgeDAL.invalidateSecretCacheByProjectId(projectId); await secretV2BridgeDAL.invalidateSecretCacheByProjectId(projectId, tx);
return { return {
secretChangesCount: secretChanges.length, secretChangesCount: secretChanges.length,
@@ -6,6 +6,8 @@ import RE2 from "re2";
import { IdentityAuthMethod, TIdentityKubernetesAuthsUpdate } from "@app/db/schemas"; import { IdentityAuthMethod, TIdentityKubernetesAuthsUpdate } from "@app/db/schemas";
import { TGatewayDALFactory } from "@app/ee/services/gateway/gateway-dal"; import { TGatewayDALFactory } from "@app/ee/services/gateway/gateway-dal";
import { TGatewayServiceFactory } from "@app/ee/services/gateway/gateway-service"; import { TGatewayServiceFactory } from "@app/ee/services/gateway/gateway-service";
import { TGatewayV2DALFactory } from "@app/ee/services/gateway-v2/gateway-v2-dal";
import { TGatewayV2ServiceFactory } from "@app/ee/services/gateway-v2/gateway-v2-service";
import { TLicenseServiceFactory } from "@app/ee/services/license/license-service"; import { TLicenseServiceFactory } from "@app/ee/services/license/license-service";
import { import {
OrgPermissionGatewayActions, OrgPermissionGatewayActions,
@@ -21,6 +23,7 @@ import { getConfig } from "@app/lib/config/env";
import { crypto } from "@app/lib/crypto"; import { crypto } from "@app/lib/crypto";
import { BadRequestError, NotFoundError, PermissionBoundaryError, UnauthorizedError } from "@app/lib/errors"; import { BadRequestError, NotFoundError, PermissionBoundaryError, UnauthorizedError } from "@app/lib/errors";
import { GatewayHttpProxyActions, GatewayProxyProtocol, withGatewayProxy } from "@app/lib/gateway"; import { GatewayHttpProxyActions, GatewayProxyProtocol, withGatewayProxy } from "@app/lib/gateway";
import { withGatewayV2Proxy } from "@app/lib/gateway-v2/gateway-v2";
import { extractIPDetails, isValidIpOrCidr } from "@app/lib/ip"; import { extractIPDetails, isValidIpOrCidr } from "@app/lib/ip";
import { logger } from "@app/lib/logger"; import { logger } from "@app/lib/logger";
@@ -54,11 +57,15 @@ type TIdentityKubernetesAuthServiceFactoryDep = {
licenseService: Pick<TLicenseServiceFactory, "getPlan">; licenseService: Pick<TLicenseServiceFactory, "getPlan">;
kmsService: Pick<TKmsServiceFactory, "createCipherPairWithDataKey">; kmsService: Pick<TKmsServiceFactory, "createCipherPairWithDataKey">;
gatewayService: TGatewayServiceFactory; gatewayService: TGatewayServiceFactory;
gatewayV2Service: TGatewayV2ServiceFactory;
gatewayDAL: Pick<TGatewayDALFactory, "find">; gatewayDAL: Pick<TGatewayDALFactory, "find">;
gatewayV2DAL: Pick<TGatewayV2DALFactory, "find">;
}; };
export type TIdentityKubernetesAuthServiceFactory = ReturnType<typeof identityKubernetesAuthServiceFactory>; export type TIdentityKubernetesAuthServiceFactory = ReturnType<typeof identityKubernetesAuthServiceFactory>;
const GATEWAY_AUTH_DEFAULT_HOST = "https://kubernetes.default.svc.cluster.local";
export const identityKubernetesAuthServiceFactory = ({ export const identityKubernetesAuthServiceFactory = ({
identityKubernetesAuthDAL, identityKubernetesAuthDAL,
identityOrgMembershipDAL, identityOrgMembershipDAL,
@@ -66,7 +73,9 @@ export const identityKubernetesAuthServiceFactory = ({
permissionService, permissionService,
licenseService, licenseService,
gatewayService, gatewayService,
gatewayV2Service,
gatewayDAL, gatewayDAL,
gatewayV2DAL,
kmsService kmsService
}: TIdentityKubernetesAuthServiceFactoryDep) => { }: TIdentityKubernetesAuthServiceFactoryDep) => {
const $gatewayProxyWrapper = async <T>( const $gatewayProxyWrapper = async <T>(
@@ -79,6 +88,42 @@ export const identityKubernetesAuthServiceFactory = ({
}, },
gatewayCallback: (host: string, port: number, httpsAgent?: https.Agent) => Promise<T> gatewayCallback: (host: string, port: number, httpsAgent?: https.Agent) => Promise<T>
): Promise<T> => { ): Promise<T> => {
const gatewayV2ConnectionDetails = await gatewayV2Service.getPlatformConnectionDetailsByGatewayId({
gatewayId: inputs.gatewayId,
targetHost: inputs.targetHost ?? GATEWAY_AUTH_DEFAULT_HOST,
targetPort: inputs.targetPort ?? 443
});
if (gatewayV2ConnectionDetails) {
let httpsAgent: https.Agent | undefined;
if (!inputs.reviewTokenThroughGateway) {
httpsAgent = new https.Agent({
ca: inputs.caCert,
rejectUnauthorized: Boolean(inputs.caCert)
});
}
const callbackResult = await withGatewayV2Proxy(
async (port) => {
const res = await gatewayCallback(
inputs.reviewTokenThroughGateway ? "http://localhost" : "https://localhost",
port,
httpsAgent
);
return res;
},
{
protocol: inputs.reviewTokenThroughGateway ? GatewayProxyProtocol.Http : GatewayProxyProtocol.Tcp,
relayHost: gatewayV2ConnectionDetails.relayHost,
gateway: gatewayV2ConnectionDetails.gateway,
relay: gatewayV2ConnectionDetails.relay,
httpsAgent
}
);
return callbackResult;
}
const relayDetails = await gatewayService.fnGetGatewayClientTlsByGatewayId(inputs.gatewayId); const relayDetails = await gatewayService.fnGetGatewayClientTlsByGatewayId(inputs.gatewayId);
const [relayHost, relayPort] = relayDetails.relayAddress.split(":"); const [relayHost, relayPort] = relayDetails.relayAddress.split(":");
@@ -277,7 +322,7 @@ export const identityKubernetesAuthServiceFactory = ({
let data: TCreateTokenReviewResponse | undefined; let data: TCreateTokenReviewResponse | undefined;
if (identityKubernetesAuth.tokenReviewMode === IdentityKubernetesAuthTokenReviewMode.Gateway) { if (identityKubernetesAuth.tokenReviewMode === IdentityKubernetesAuthTokenReviewMode.Gateway) {
if (!identityKubernetesAuth.gatewayId) { if (!identityKubernetesAuth.gatewayId && !identityKubernetesAuth.gatewayV2Id) {
throw new BadRequestError({ throw new BadRequestError({
message: "Gateway ID is required when token review mode is set to Gateway" message: "Gateway ID is required when token review mode is set to Gateway"
}); });
@@ -285,7 +330,7 @@ export const identityKubernetesAuthServiceFactory = ({
data = await $gatewayProxyWrapper( data = await $gatewayProxyWrapper(
{ {
gatewayId: identityKubernetesAuth.gatewayId, gatewayId: (identityKubernetesAuth.gatewayV2Id ?? identityKubernetesAuth.gatewayId) as string,
reviewTokenThroughGateway: true reviewTokenThroughGateway: true
}, },
tokenReviewCallbackThroughGateway tokenReviewCallbackThroughGateway
@@ -304,17 +349,18 @@ export const identityKubernetesAuthServiceFactory = ({
const [k8sHost, k8sPort] = kubernetesHost.split(":"); const [k8sHost, k8sPort] = kubernetesHost.split(":");
data = identityKubernetesAuth.gatewayId data =
? await $gatewayProxyWrapper( identityKubernetesAuth.gatewayId || identityKubernetesAuth.gatewayV2Id
{ ? await $gatewayProxyWrapper(
gatewayId: identityKubernetesAuth.gatewayId, {
targetHost: k8sHost, gatewayId: (identityKubernetesAuth.gatewayV2Id ?? identityKubernetesAuth.gatewayId) as string,
targetPort: k8sPort ? Number(k8sPort) : 443, targetHost: k8sHost,
reviewTokenThroughGateway: false targetPort: k8sPort ? Number(k8sPort) : 443,
}, reviewTokenThroughGateway: false
tokenReviewCallbackRaw },
) tokenReviewCallbackRaw
: await tokenReviewCallbackRaw(); )
: await tokenReviewCallbackRaw();
} else { } else {
throw new BadRequestError({ throw new BadRequestError({
message: `Invalid token review mode: ${identityKubernetesAuth.tokenReviewMode}` message: `Invalid token review mode: ${identityKubernetesAuth.tokenReviewMode}`
@@ -490,14 +536,20 @@ export const identityKubernetesAuthServiceFactory = ({
return extractIPDetails(accessTokenTrustedIp.ipAddress); return extractIPDetails(accessTokenTrustedIp.ipAddress);
}); });
let isGatewayV1 = true;
if (gatewayId) { if (gatewayId) {
const [gateway] = await gatewayDAL.find({ id: gatewayId, orgId: identityMembershipOrg.orgId }); const [gateway] = await gatewayDAL.find({ id: gatewayId, orgId: identityMembershipOrg.orgId });
if (!gateway) { const [gatewayV2] = await gatewayV2DAL.find({ id: gatewayId, orgId: identityMembershipOrg.orgId });
if (!gateway && !gatewayV2) {
throw new NotFoundError({ throw new NotFoundError({
message: `Gateway with ID ${gatewayId} not found` message: `Gateway with ID ${gatewayId} not found`
}); });
} }
if (!gateway) {
isGatewayV1 = false;
}
const { permission: orgPermission } = await permissionService.getOrgPermission( const { permission: orgPermission } = await permissionService.getOrgPermission(
actor, actor,
actorId, actorId,
@@ -528,7 +580,8 @@ export const identityKubernetesAuthServiceFactory = ({
accessTokenMaxTTL, accessTokenMaxTTL,
accessTokenTTL, accessTokenTTL,
accessTokenNumUsesLimit, accessTokenNumUsesLimit,
gatewayId, gatewayId: isGatewayV1 ? gatewayId : null,
gatewayV2Id: isGatewayV1 ? null : gatewayId,
accessTokenTrustedIps: JSON.stringify(reformattedAccessTokenTrustedIps), accessTokenTrustedIps: JSON.stringify(reformattedAccessTokenTrustedIps),
encryptedKubernetesTokenReviewerJwt: tokenReviewerJwt encryptedKubernetesTokenReviewerJwt: tokenReviewerJwt
? encryptor({ plainText: Buffer.from(tokenReviewerJwt) }).cipherTextBlob ? encryptor({ plainText: Buffer.from(tokenReviewerJwt) }).cipherTextBlob
@@ -608,14 +661,21 @@ export const identityKubernetesAuthServiceFactory = ({
return extractIPDetails(accessTokenTrustedIp.ipAddress); return extractIPDetails(accessTokenTrustedIp.ipAddress);
}); });
let isGatewayV1 = true;
if (gatewayId) { if (gatewayId) {
const [gateway] = await gatewayDAL.find({ id: gatewayId, orgId: identityMembershipOrg.orgId }); const [gateway] = await gatewayDAL.find({ id: gatewayId, orgId: identityMembershipOrg.orgId });
if (!gateway) { const [gatewayV2] = await gatewayV2DAL.find({ id: gatewayId, orgId: identityMembershipOrg.orgId });
if (!gateway && !gatewayV2) {
throw new NotFoundError({ throw new NotFoundError({
message: `Gateway with ID ${gatewayId} not found` message: `Gateway with ID ${gatewayId} not found`
}); });
} }
if (!gateway) {
isGatewayV1 = false;
}
const { permission: orgPermission } = await permissionService.getOrgPermission( const { permission: orgPermission } = await permissionService.getOrgPermission(
actor, actor,
actorId, actorId,
@@ -629,13 +689,18 @@ export const identityKubernetesAuthServiceFactory = ({
); );
} }
const shouldUpdateGatewayId = Boolean(gatewayId);
const gatewayIdValue = isGatewayV1 ? gatewayId : null;
const gatewayV2IdValue = isGatewayV1 ? null : gatewayId;
const updateQuery: TIdentityKubernetesAuthsUpdate = { const updateQuery: TIdentityKubernetesAuthsUpdate = {
kubernetesHost, kubernetesHost,
tokenReviewMode, tokenReviewMode,
allowedNamespaces, allowedNamespaces,
allowedNames, allowedNames,
allowedAudience, allowedAudience,
gatewayId, gatewayId: shouldUpdateGatewayId ? gatewayIdValue : undefined,
gatewayV2Id: shouldUpdateGatewayId ? gatewayV2IdValue : undefined,
accessTokenMaxTTL, accessTokenMaxTTL,
accessTokenTTL, accessTokenTTL,
accessTokenNumUsesLimit, accessTokenNumUsesLimit,
@@ -730,7 +795,13 @@ export const identityKubernetesAuthServiceFactory = ({
}).toString(); }).toString();
} }
return { ...identityKubernetesAuth, caCert, tokenReviewerJwt, orgId: identityMembershipOrg.orgId }; return {
...identityKubernetesAuth,
caCert,
tokenReviewerJwt,
orgId: identityMembershipOrg.orgId,
gatewayId: identityKubernetesAuth.gatewayId ?? identityKubernetesAuth.gatewayV2Id
};
}; };
const revokeIdentityKubernetesAuth = async ({ const revokeIdentityKubernetesAuth = async ({
@@ -25,7 +25,7 @@ export const identityDALFactory = (db: TDbClient) => {
} as const; } as const;
const tableName = authMethodToTableName[authMethod]; const tableName = authMethodToTableName[authMethod];
if (!tableName) return; if (!tableName) return;
const data = await db(tableName).where({ identityId }).first(); const data = await db.replicaNode()(tableName).where({ identityId }).first();
if (!data) return; if (!data) return;
return data.accessTokenTrustedIps; return data.accessTokenTrustedIps;
}; };
@@ -30,7 +30,7 @@ export const integrationAuthDALFactory = (db: TDbClient) => {
const getByOrg = async (orgId: string, tx?: Knex) => { const getByOrg = async (orgId: string, tx?: Knex) => {
try { try {
const integrationAuths = await (tx || db)(TableName.IntegrationAuth) const integrationAuths = await (tx || db.replicaNode())(TableName.IntegrationAuth)
.join(TableName.Project, `${TableName.Project}.id`, `${TableName.IntegrationAuth}.projectId`) .join(TableName.Project, `${TableName.Project}.id`, `${TableName.IntegrationAuth}.projectId`)
.join(TableName.Organization, `${TableName.Organization}.id`, `${TableName.Project}.orgId`) .join(TableName.Organization, `${TableName.Organization}.id`, `${TableName.Project}.orgId`)
.where(`${TableName.Organization}.id`, "=", orgId) .where(`${TableName.Organization}.id`, "=", orgId)

Some files were not shown because too many files have changed in this diff Show More