diff --git a/server/lib/blueprints/findOrgUser.ts b/server/lib/blueprints/findOrgUser.ts new file mode 100644 index 000000000..1fc054fd1 --- /dev/null +++ b/server/lib/blueprints/findOrgUser.ts @@ -0,0 +1,38 @@ +import { and, asc, eq, or } from "drizzle-orm"; +import { Transaction, User, userOrgs, users } from "@server/db"; + +export async function findOrgUserByIdentifier( + trx: Transaction, + orgId: string, + identifier: string +): Promise { + const [match] = await trx + .select() + .from(users) + .innerJoin(userOrgs, eq(users.userId, userOrgs.userId)) + .where( + and( + or(eq(users.username, identifier), eq(users.email, identifier)), + eq(userOrgs.orgId, orgId) + ) + ) + .orderBy(asc(users.dateCreated), asc(users.userId)) + .limit(1); + + return match?.user ?? null; +} + +export async function resolveOrgUserIds( + trx: Transaction, + orgId: string, + identifiers: string[] +): Promise { + const userIds = new Set(); + for (const identifier of identifiers) { + const user = await findOrgUserByIdentifier(trx, orgId, identifier); + if (user) { + userIds.add(user.userId); + } + } + return [...userIds]; +} diff --git a/server/lib/blueprints/privateResources.ts b/server/lib/blueprints/privateResources.ts index 65cf199d6..085fed836 100644 --- a/server/lib/blueprints/privateResources.ts +++ b/server/lib/blueprints/privateResources.ts @@ -11,15 +11,14 @@ import { siteNetworks, siteResources, Transaction, - userOrgs, - users, userSiteResources, networks } from "@server/db"; import { sites } from "@server/db"; -import { eq, and, ne, inArray, or, isNotNull } from "drizzle-orm"; +import { eq, and, ne, inArray, isNotNull } from "drizzle-orm"; import { Config } from "./types"; import { getOrCreateLabelIds, syncSiteResourceLabels } from "./labels"; +import { resolveOrgUserIds } from "./findOrgUser"; import logger from "@server/logger"; import { defaultRoleAllowedActions } from "@server/routers/role/createRole"; import { getNextAvailableAliasAddress } from "../ip"; @@ -389,28 +388,22 @@ export async function updatePrivateResources( .where(eq(userSiteResources.siteResourceId, siteResourceId)); if (resourceData.users.length > 0) { - // get userIds from username - const usersToUpdate = await trx - .select() - .from(users) - .innerJoin(userOrgs, eq(users.userId, userOrgs.userId)) - .where( - and( - or( - inArray(users.username, resourceData.users), - inArray(users.email, resourceData.users) - ), - eq(userOrgs.orgId, orgId) - ) - ); + const userIds = await resolveOrgUserIds( + trx, + orgId, + resourceData.users + ); - const userIds = usersToUpdate.map((user) => user.user.userId); - - await trx - .insert(userSiteResources) - .values( - userIds.map((userId) => ({ userId, siteResourceId })) - ); + if (userIds.length > 0) { + await trx + .insert(userSiteResources) + .values( + userIds.map((userId) => ({ + userId, + siteResourceId + })) + ); + } } // Get all admin role IDs for this org to exclude from deletion @@ -721,28 +714,22 @@ export async function updatePrivateResources( } if (resourceData.users.length > 0) { - // get userIds from username - const usersToUpdate = await trx - .select() - .from(users) - .innerJoin(userOrgs, eq(users.userId, userOrgs.userId)) - .where( - and( - or( - inArray(users.username, resourceData.users), - inArray(users.email, resourceData.users) - ), - eq(userOrgs.orgId, orgId) - ) - ); + const userIds = await resolveOrgUserIds( + trx, + orgId, + resourceData.users + ); - const userIds = usersToUpdate.map((user) => user.user.userId); - - await trx - .insert(userSiteResources) - .values( - userIds.map((userId) => ({ userId, siteResourceId })) - ); + if (userIds.length > 0) { + await trx + .insert(userSiteResources) + .values( + userIds.map((userId) => ({ + userId, + siteResourceId + })) + ); + } } if (resourceData.machines.length > 0) { diff --git a/server/lib/blueprints/publicResources.ts b/server/lib/blueprints/publicResources.ts index a76bcc26c..7822c2b8a 100644 --- a/server/lib/blueprints/publicResources.ts +++ b/server/lib/blueprints/publicResources.ts @@ -46,11 +46,12 @@ import { encrypt } from "@server/lib/crypto"; import logger from "@server/logger"; import { defaultRoleAllowedActions } from "@server/routers/role/createRole"; import { pickPort } from "@server/routers/target/helpers"; -import { and, asc, eq, isNotNull, ne, or } from "drizzle-orm"; +import { and, asc, eq, isNotNull, ne } from "drizzle-orm"; import { tierMatrix } from "../billing/tierMatrix"; import { isValidCIDR, isValidIP, isValidUrlGlobPattern } from "../validators"; import { Config, isTargetsOnlyResource, TargetData } from "./types"; import { getOrCreateLabelIds, syncResourceLabels } from "./labels"; +import { findOrgUserByIdentifier } from "./findOrgUser"; import { LimitId } from "../billing"; import { usageService } from "../billing/usageService"; import { syncInferenceAiConfig } from "./aiProviders"; @@ -1563,29 +1564,19 @@ async function syncUserResources( .where(eq(userResources.resourceId, resourceId)); for (const username of ssoUsers) { - const [user] = await trx - .select() - .from(users) - .innerJoin(userOrgs, eq(users.userId, userOrgs.userId)) - .where( - and( - or(eq(users.username, username), eq(users.email, username)), - eq(userOrgs.orgId, orgId) - ) - ) - .limit(1); + const user = await findOrgUserByIdentifier(trx, orgId, username); if (!user) { throw new Error(`User not found: ${username} in org ${orgId}`); } const existingUserResource = existingUserResources.find( - (rr) => rr.userId === user.user.userId + (rr) => rr.userId === user.userId ); if (!existingUserResource) { await trx.insert(userResources).values({ - userId: user.user.userId, + userId: user.userId, resourceId: resourceId }); } @@ -1955,29 +1946,19 @@ async function syncUserPolicies( .where(eq(userPolicies.resourcePolicyId, policyId)); for (const username of ssoUsers) { - const [user] = await trx - .select() - .from(users) - .innerJoin(userOrgs, eq(users.userId, userOrgs.userId)) - .where( - and( - or(eq(users.username, username), eq(users.email, username)), - eq(userOrgs.orgId, orgId) - ) - ) - .limit(1); + const user = await findOrgUserByIdentifier(trx, orgId, username); if (!user) { throw new Error(`User not found: ${username} in org ${orgId}`); } const existingUserPolicy = existingUserPoliciesList.find( - (up) => up.userId === user.user.userId + (up) => up.userId === user.userId ); if (!existingUserPolicy) { await trx.insert(userPolicies).values({ - userId: user.user.userId, + userId: user.userId, resourcePolicyId: policyId }); } diff --git a/server/lib/blueprints/resourcePolicies.ts b/server/lib/blueprints/resourcePolicies.ts index 87e9101cc..4a883f64b 100644 --- a/server/lib/blueprints/resourcePolicies.ts +++ b/server/lib/blueprints/resourcePolicies.ts @@ -13,7 +13,7 @@ import { userPolicies, users } from "@server/db"; -import { eq, and, or } from "drizzle-orm"; +import { eq, and } from "drizzle-orm"; import { Config, ResourcePolicyData } from "./types"; import logger from "@server/logger"; import { getUniqueResourcePolicyName } from "@server/db/names"; @@ -22,6 +22,7 @@ import { idpExistsForOrg } from "@server/lib/idp/idpExistsForOrg"; import { isValidCIDR, isValidIP, isValidUrlGlobPattern } from "../validators"; import { isLicensedOrSubscribed } from "#dynamic/lib/isLicencedOrSubscribed"; import { tierMatrix } from "../billing/tierMatrix"; +import { findOrgUserByIdentifier } from "./findOrgUser"; export type ResourcePoliciesResults = { resourcePolicyId: number; @@ -466,17 +467,7 @@ async function syncUserPolicies( .where(eq(userPolicies.resourcePolicyId, policyId)); for (const username of ssoUsers) { - const [user] = await trx - .select() - .from(users) - .innerJoin(userOrgs, eq(users.userId, userOrgs.userId)) - .where( - and( - or(eq(users.username, username), eq(users.email, username)), - eq(userOrgs.orgId, orgId) - ) - ) - .limit(1); + const user = await findOrgUserByIdentifier(trx, orgId, username); if (!user) { logger.warn( @@ -486,12 +477,12 @@ async function syncUserPolicies( } const alreadyExists = existingUserPolicies.some( - (up) => up.userId === user.user.userId + (up) => up.userId === user.userId ); if (!alreadyExists) { await trx.insert(userPolicies).values({ - userId: user.user.userId, + userId: user.userId, resourcePolicyId: policyId }); } @@ -536,17 +527,7 @@ async function addUserPolicies( trx: Transaction ) { for (const username of ssoUsers) { - const [user] = await trx - .select() - .from(users) - .innerJoin(userOrgs, eq(users.userId, userOrgs.userId)) - .where( - and( - or(eq(users.username, username), eq(users.email, username)), - eq(userOrgs.orgId, orgId) - ) - ) - .limit(1); + const user = await findOrgUserByIdentifier(trx, orgId, username); if (!user) { logger.warn( @@ -556,7 +537,7 @@ async function addUserPolicies( } await trx.insert(userPolicies).values({ - userId: user.user.userId, + userId: user.userId, resourcePolicyId: policyId }); }