Files
pangolin/server/lib/aiInferenceResource.ts
T
2026-08-04 15:54:24 -04:00

513 lines
15 KiB
TypeScript

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<typeof modelAccessModeSchema>;
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<number, ModelAccessMode>();
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<ResourceAiProviderAttachment[] | InferenceFieldsError> {
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<InferenceFieldsError | null> {
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<void> {
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<void> {
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<void> {
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<void> {
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<void> {
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<void> {
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<string | null> {
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<string | null> {
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<string | null> {
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<string | null> {
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;
}