import { and, eq, inArray } from "drizzle-orm"; import { aiModels, aiProviders, db, resourceAiModels, resourceAiProviders, siteResourceAiModels, siteResourceAiProviders, type Transaction } from "@server/db"; import { z } from "zod"; type DbOrTrx = Transaction | typeof db; export const modelAccessModeSchema = z.enum([ "passthrough", "catalog", "allowlist" ]); export type ModelAccessMode = z.infer; export const resourceAiProviderAttachmentSchema = z.strictObject({ providerId: z.number().int().positive(), modelAccessMode: modelAccessModeSchema.optional() }); export type ResourceAiProviderInput = z.infer< typeof resourceAiProviderAttachmentSchema >; export type ResourceAiProviderAttachment = { providerId: number; modelAccessMode: ModelAccessMode; }; export type InferenceFieldsError = { error: string; }; export function isInferenceFieldsError( value: { error: string } | object ): value is InferenceFieldsError { return "error" in value; } function normalizeAttachments( inputs: ResourceAiProviderInput[] ): ResourceAiProviderAttachment[] { const byProvider = new Map(); for (const input of inputs) { byProvider.set( input.providerId, input.modelAccessMode ?? "passthrough" ); } return [...byProvider.entries()].map(([providerId, modelAccessMode]) => ({ providerId, modelAccessMode })); } /** * Validate provider attachments for an org. * At most one passthrough provider is allowed per resource. */ export async function resolveProviderAttachments(input: { orgId: string; attachments: ResourceAiProviderInput[]; requireAtLeastOne: boolean; }): Promise { const attachments = normalizeAttachments(input.attachments); if (input.requireAtLeastOne && attachments.length === 0) { return { error: "At least one AI provider is required for inference-mode resources" }; } const passthroughCount = attachments.filter( (a) => a.modelAccessMode === "passthrough" ).length; if (passthroughCount > 1) { return { error: "A resource may have at most one AI provider in passthrough mode" }; } if (attachments.length === 0) { return []; } const providerIds = attachments.map((a) => a.providerId); const providers = await db .select({ providerId: aiProviders.providerId, orgId: aiProviders.orgId, enabled: aiProviders.enabled }) .from(aiProviders) .where( and( inArray(aiProviders.providerId, providerIds), eq(aiProviders.orgId, input.orgId) ) ); if (providers.length !== providerIds.length) { return { error: "One or more AI providers were not found in this organization" }; } const disabled = providers.find((p) => !p.enabled); if (disabled) { return { error: `AI provider with ID ${disabled.providerId} is disabled` }; } return attachments; } export async function assertInferenceModeAllowsProviderFields(input: { mode: string; hasProviderAttachments: boolean; }): Promise { if (input.mode === "inference") { return null; } if (input.hasProviderAttachments) { return { error: "AI providers can only be attached to inference-mode resources" }; } return null; } export async function setPublicResourceAiProviders( resourceId: number, attachments: ResourceAiProviderAttachment[], trx: DbOrTrx = db ): Promise { await trx .delete(resourceAiProviders) .where(eq(resourceAiProviders.resourceId, resourceId)); if (attachments.length > 0) { await trx.insert(resourceAiProviders).values( attachments.map((a) => ({ resourceId, providerId: a.providerId, modelAccessMode: a.modelAccessMode })) ); } await prunePublicResourceAllowlistToAllowlistProviders(resourceId, trx); } export async function setSiteResourceAiProviders( siteResourceId: number, attachments: ResourceAiProviderAttachment[], trx: DbOrTrx = db ): Promise { await trx .delete(siteResourceAiProviders) .where(eq(siteResourceAiProviders.siteResourceId, siteResourceId)); if (attachments.length > 0) { await trx.insert(siteResourceAiProviders).values( attachments.map((a) => ({ siteResourceId, providerId: a.providerId, modelAccessMode: a.modelAccessMode })) ); } await pruneSiteResourceAllowlistToAllowlistProviders(siteResourceId, trx); } export async function clearPublicResourceAiConfig( resourceId: number, trx: DbOrTrx = db ): Promise { await trx .delete(resourceAiModels) .where(eq(resourceAiModels.resourceId, resourceId)); await trx .delete(resourceAiProviders) .where(eq(resourceAiProviders.resourceId, resourceId)); } export async function clearSiteResourceAiConfig( siteResourceId: number, trx: DbOrTrx = db ): Promise { await trx .delete(siteResourceAiModels) .where(eq(siteResourceAiModels.siteResourceId, siteResourceId)); await trx .delete(siteResourceAiProviders) .where(eq(siteResourceAiProviders.siteResourceId, siteResourceId)); } async function prunePublicResourceAllowlistToAllowlistProviders( resourceId: number, trx: DbOrTrx = db ): Promise { const allowlistProviders = await trx .select({ providerId: resourceAiProviders.providerId }) .from(resourceAiProviders) .where( and( eq(resourceAiProviders.resourceId, resourceId), eq(resourceAiProviders.modelAccessMode, "allowlist") ) ); if (allowlistProviders.length === 0) { await trx .delete(resourceAiModels) .where(eq(resourceAiModels.resourceId, resourceId)); return; } const validModels = await trx .select({ modelId: aiModels.modelId }) .from(aiModels) .where( inArray( aiModels.providerId, allowlistProviders.map((p) => p.providerId) ) ); const validIds = validModels.map((m) => m.modelId); const existing = await trx .select({ modelId: resourceAiModels.modelId }) .from(resourceAiModels) .where(eq(resourceAiModels.resourceId, resourceId)); const toRemove = existing .map((e) => e.modelId) .filter((id) => !validIds.includes(id)); if (toRemove.length > 0) { await trx .delete(resourceAiModels) .where( and( eq(resourceAiModels.resourceId, resourceId), inArray(resourceAiModels.modelId, toRemove) ) ); } } async function pruneSiteResourceAllowlistToAllowlistProviders( siteResourceId: number, trx: DbOrTrx = db ): Promise { const allowlistProviders = await trx .select({ providerId: siteResourceAiProviders.providerId }) .from(siteResourceAiProviders) .where( and( eq(siteResourceAiProviders.siteResourceId, siteResourceId), eq(siteResourceAiProviders.modelAccessMode, "allowlist") ) ); if (allowlistProviders.length === 0) { await trx .delete(siteResourceAiModels) .where(eq(siteResourceAiModels.siteResourceId, siteResourceId)); return; } const validModels = await trx .select({ modelId: aiModels.modelId }) .from(aiModels) .where( inArray( aiModels.providerId, allowlistProviders.map((p) => p.providerId) ) ); const validIds = validModels.map((m) => m.modelId); const existing = await trx .select({ modelId: siteResourceAiModels.modelId }) .from(siteResourceAiModels) .where(eq(siteResourceAiModels.siteResourceId, siteResourceId)); const toRemove = existing .map((e) => e.modelId) .filter((id) => !validIds.includes(id)); if (toRemove.length > 0) { await trx .delete(siteResourceAiModels) .where( and( eq(siteResourceAiModels.siteResourceId, siteResourceId), inArray(siteResourceAiModels.modelId, toRemove) ) ); } } export async function listPublicResourceAiProviders(resourceId: number) { return db .select({ providerId: resourceAiProviders.providerId, modelAccessMode: resourceAiProviders.modelAccessMode, name: aiProviders.name, type: aiProviders.type, enabled: aiProviders.enabled }) .from(resourceAiProviders) .innerJoin( aiProviders, eq(resourceAiProviders.providerId, aiProviders.providerId) ) .where(eq(resourceAiProviders.resourceId, resourceId)); } export async function listSiteResourceAiProviders(siteResourceId: number) { return db .select({ providerId: siteResourceAiProviders.providerId, modelAccessMode: siteResourceAiProviders.modelAccessMode, name: aiProviders.name, type: aiProviders.type, enabled: aiProviders.enabled }) .from(siteResourceAiProviders) .innerJoin( aiProviders, eq(siteResourceAiProviders.providerId, aiProviders.providerId) ) .where(eq(siteResourceAiProviders.siteResourceId, siteResourceId)); } /** * Allowlist APIs require an inference resource with at least one * attached provider in allowlist mode. */ export async function assertPublicAllowlistApiEligible(resource: { resourceId: number; mode: string; }): Promise { if (resource.mode !== "inference") { return "AI model allowlists are only supported on inference-mode resources"; } const [row] = await db .select({ providerId: resourceAiProviders.providerId }) .from(resourceAiProviders) .where( and( eq(resourceAiProviders.resourceId, resource.resourceId), eq(resourceAiProviders.modelAccessMode, "allowlist") ) ) .limit(1); if (!row) { return "Attach at least one AI provider with modelAccessMode=allowlist before managing allowed models"; } return null; } export async function assertSiteAllowlistApiEligible(siteResource: { siteResourceId: number; mode: string; }): Promise { if (siteResource.mode !== "inference") { return "AI model allowlists are only supported on inference-mode resources"; } const [row] = await db .select({ providerId: siteResourceAiProviders.providerId }) .from(siteResourceAiProviders) .where( and( eq( siteResourceAiProviders.siteResourceId, siteResource.siteResourceId ), eq(siteResourceAiProviders.modelAccessMode, "allowlist") ) ) .limit(1); if (!row) { return "Attach at least one AI provider with modelAccessMode=allowlist before managing allowed models"; } return null; } /** * Models must belong to providers attached to this resource in allowlist mode, * and those providers must belong to the resource's org. */ export async function assertModelsBelongToPublicAllowlistProviders(input: { orgId: string; resourceId: number; modelIds: number[]; }): Promise { const uniqueIds = [...new Set(input.modelIds)]; if (uniqueIds.length === 0) { return null; } const allowlistProviders = await db .select({ providerId: resourceAiProviders.providerId }) .from(resourceAiProviders) .innerJoin( aiProviders, eq(resourceAiProviders.providerId, aiProviders.providerId) ) .where( and( eq(resourceAiProviders.resourceId, input.resourceId), eq(resourceAiProviders.modelAccessMode, "allowlist"), eq(aiProviders.orgId, input.orgId) ) ); if (allowlistProviders.length === 0) { return "No allowlist AI providers are attached to this resource"; } const validModels = await db .select({ modelId: aiModels.modelId }) .from(aiModels) .innerJoin(aiProviders, eq(aiModels.providerId, aiProviders.providerId)) .where( and( inArray(aiModels.modelId, uniqueIds), inArray( aiModels.providerId, allowlistProviders.map((p) => p.providerId) ), eq(aiProviders.orgId, input.orgId) ) ); if (validModels.length !== uniqueIds.length) { return "One or more model IDs do not exist or do not belong to an allowlist provider on this resource"; } return null; } export async function assertModelsBelongToSiteAllowlistProviders(input: { orgId: string; siteResourceId: number; modelIds: number[]; }): Promise { const uniqueIds = [...new Set(input.modelIds)]; if (uniqueIds.length === 0) { return null; } const allowlistProviders = await db .select({ providerId: siteResourceAiProviders.providerId }) .from(siteResourceAiProviders) .innerJoin( aiProviders, eq(siteResourceAiProviders.providerId, aiProviders.providerId) ) .where( and( eq( siteResourceAiProviders.siteResourceId, input.siteResourceId ), eq(siteResourceAiProviders.modelAccessMode, "allowlist"), eq(aiProviders.orgId, input.orgId) ) ); if (allowlistProviders.length === 0) { return "No allowlist AI providers are attached to this site resource"; } const validModels = await db .select({ modelId: aiModels.modelId }) .from(aiModels) .innerJoin(aiProviders, eq(aiModels.providerId, aiProviders.providerId)) .where( and( inArray(aiModels.modelId, uniqueIds), inArray( aiModels.providerId, allowlistProviders.map((p) => p.providerId) ), eq(aiProviders.orgId, input.orgId) ) ); if (validModels.length !== uniqueIds.length) { return "One or more model IDs do not exist or do not belong to an allowlist provider on this site resource"; } return null; }