From e38359c74f44fc0cdc0e9e0a29b40d46970ab83b Mon Sep 17 00:00:00 2001 From: miloschwartz Date: Tue, 4 Aug 2026 15:54:24 -0400 Subject: [PATCH] add crud for adding providers and models to resources --- server/db/pg/schema/schema.ts | 58 +- server/db/sqlite/schema/schema.ts | 58 +- server/lib/aiInferenceResource.ts | 512 ++++++++++++++++++ server/lib/traefik/getTraefikConfig.ts | 10 +- .../private/lib/traefik/getTraefikConfig.ts | 23 +- server/routers/aiGateway/chatCompletions.ts | 288 +++++++--- server/routers/aiProvider/createAiModel.ts | 25 +- server/routers/aiProvider/createAiProvider.ts | 9 - server/routers/aiProvider/updateAiModel.ts | 26 +- server/routers/aiProvider/updateAiProvider.ts | 32 +- server/routers/aiProvider/validation.ts | 24 - server/routers/external.ts | 94 ++++ server/routers/integration.ts | 86 +++ .../routers/resource/addAiModelToResource.ts | 44 +- .../resource/addAiProviderToResource.ts | 152 ++++++ server/routers/resource/createResource.ts | 56 +- server/routers/resource/index.ts | 4 + .../routers/resource/listResourceAiModels.ts | 2 +- .../resource/listResourceAiProviders.ts | 94 ++++ .../resource/removeAiModelFromResource.ts | 8 +- .../resource/removeAiProviderFromResource.ts | 165 ++++++ .../routers/resource/setResourceAiModels.ts | 48 +- .../resource/setResourceAiProviders.ts | 137 +++++ server/routers/resource/updateResource.ts | 15 +- .../siteResource/addAiModelToSiteResource.ts | 49 +- .../addAiProviderToSiteResource.ts | 152 ++++++ .../siteResource/createSiteResource.ts | 50 +- server/routers/siteResource/index.ts | 4 + .../siteResource/listSiteResourceAiModels.ts | 14 +- .../listSiteResourceAiProviders.ts | 94 ++++ .../removeAiModelFromSiteResource.ts | 9 +- .../removeAiProviderFromSiteResource.ts | 164 ++++++ .../siteResource/setSiteResourceAiModels.ts | 72 +-- .../setSiteResourceAiProviders.ts | 139 +++++ .../siteResource/updateSiteResource.ts | 24 +- .../[providerId]/authentication/page.tsx | 4 - .../[providerId]/network/page.tsx | 4 - .../settings/ai-providers/create/page.tsx | 2 - src/lib/aiProviderFormSchema.ts | 30 - 39 files changed, 2343 insertions(+), 438 deletions(-) create mode 100644 server/lib/aiInferenceResource.ts create mode 100644 server/routers/resource/addAiProviderToResource.ts create mode 100644 server/routers/resource/listResourceAiProviders.ts create mode 100644 server/routers/resource/removeAiProviderFromResource.ts create mode 100644 server/routers/resource/setResourceAiProviders.ts create mode 100644 server/routers/siteResource/addAiProviderToSiteResource.ts create mode 100644 server/routers/siteResource/listSiteResourceAiProviders.ts create mode 100644 server/routers/siteResource/removeAiProviderFromSiteResource.ts create mode 100644 server/routers/siteResource/setSiteResourceAiProviders.ts diff --git a/server/db/pg/schema/schema.ts b/server/db/pg/schema/schema.ts index 6a24064a6..b323ae411 100644 --- a/server/db/pg/schema/schema.ts +++ b/server/db/pg/schema/schema.ts @@ -211,14 +211,7 @@ export const resources = pgTable( authDaemonPort: integer("authDaemonPort").default(22123), status: varchar("status") .$type<"pending" | "approved">() - .default("approved"), - aiProviderId: integer("aiProviderId").references( - () => aiProviders.providerId, - { onDelete: "set null" } - ), - modelAccessMode: varchar("modelAccessMode").$type< - "passthrough" | "catalog" | "allowlist" - >() + .default("approved") }, (t) => [ index("idx_resources_fulldomain") @@ -229,6 +222,23 @@ export const resources = pgTable( ] ); +export const resourceAiProviders = pgTable( + "resourceAiProviders", + { + resourceId: integer("resourceId") + .notNull() + .references(() => resources.resourceId, { onDelete: "cascade" }), + providerId: integer("providerId") + .notNull() + .references(() => aiProviders.providerId, { onDelete: "cascade" }), + modelAccessMode: varchar("modelAccessMode") + .$type<"passthrough" | "catalog" | "allowlist">() + .notNull() + .default("passthrough") + }, + (t) => [primaryKey({ columns: [t.resourceId, t.providerId] })] +); + export const resourceAiModels = pgTable( "resourceAiModels", { @@ -496,18 +506,30 @@ export const siteResources = pgTable( fullDomain: varchar("fullDomain"), status: varchar("status") .$type<"pending" | "approved">() - .default("approved"), - aiProviderId: integer("aiProviderId").references( - () => aiProviders.providerId, - { onDelete: "set null" } - ), - modelAccessMode: varchar("modelAccessMode").$type< - "passthrough" | "catalog" | "allowlist" - >() + .default("approved") }, (t) => [index("idx_siteresources_orgid_niceid").on(t.orgId, t.niceId)] ); +export const siteResourceAiProviders = pgTable( + "siteResourceAiProviders", + { + siteResourceId: integer("siteResourceId") + .notNull() + .references(() => siteResources.siteResourceId, { + onDelete: "cascade" + }), + providerId: integer("providerId") + .notNull() + .references(() => aiProviders.providerId, { onDelete: "cascade" }), + modelAccessMode: varchar("modelAccessMode") + .$type<"passthrough" | "catalog" | "allowlist">() + .notNull() + .default("passthrough") + }, + (t) => [primaryKey({ columns: [t.siteResourceId, t.providerId] })] +); + export const siteResourceAiModels = pgTable( "siteResourceAiModels", { @@ -1806,5 +1828,9 @@ export type AiProvider = InferSelectModel; export type AiModel = InferSelectModel; export type AiBudget = InferSelectModel; export type AiBudgetPeriod = InferSelectModel; +export type ResourceAiProvider = InferSelectModel; +export type SiteResourceAiProvider = InferSelectModel< + typeof siteResourceAiProviders +>; export type ResourceAiModel = InferSelectModel; export type SiteResourceAiModel = InferSelectModel; diff --git a/server/db/sqlite/schema/schema.ts b/server/db/sqlite/schema/schema.ts index f453d36f7..c8f25dfed 100644 --- a/server/db/sqlite/schema/schema.ts +++ b/server/db/sqlite/schema/schema.ts @@ -216,16 +216,26 @@ export const resources = sqliteTable("resources", { .$type<"site" | "remote" | "native">() .default("site"), authDaemonPort: integer("authDaemonPort").default(22123), - status: text("status").$type<"pending" | "approved">().default("approved"), - aiProviderId: integer("aiProviderId").references( - () => aiProviders.providerId, - { onDelete: "set null" } - ), - modelAccessMode: text("modelAccessMode").$type< - "passthrough" | "catalog" | "allowlist" - >() + status: text("status").$type<"pending" | "approved">().default("approved") }); +export const resourceAiProviders = sqliteTable( + "resourceAiProviders", + { + resourceId: integer("resourceId") + .notNull() + .references(() => resources.resourceId, { onDelete: "cascade" }), + providerId: integer("providerId") + .notNull() + .references(() => aiProviders.providerId, { onDelete: "cascade" }), + modelAccessMode: text("modelAccessMode") + .$type<"passthrough" | "catalog" | "allowlist">() + .notNull() + .default("passthrough") + }, + (t) => [primaryKey({ columns: [t.resourceId, t.providerId] })] +); + export const resourceAiModels = sqliteTable( "resourceAiModels", { @@ -483,16 +493,28 @@ export const siteResources = sqliteTable("siteResources", { }), subdomain: text("subdomain"), fullDomain: text("fullDomain"), - status: text("status").$type<"pending" | "approved">().default("approved"), - aiProviderId: integer("aiProviderId").references( - () => aiProviders.providerId, - { onDelete: "set null" } - ), - modelAccessMode: text("modelAccessMode").$type< - "passthrough" | "catalog" | "allowlist" - >() + status: text("status").$type<"pending" | "approved">().default("approved") }); +export const siteResourceAiProviders = sqliteTable( + "siteResourceAiProviders", + { + siteResourceId: integer("siteResourceId") + .notNull() + .references(() => siteResources.siteResourceId, { + onDelete: "cascade" + }), + providerId: integer("providerId") + .notNull() + .references(() => aiProviders.providerId, { onDelete: "cascade" }), + modelAccessMode: text("modelAccessMode") + .$type<"passthrough" | "catalog" | "allowlist">() + .notNull() + .default("passthrough") + }, + (t) => [primaryKey({ columns: [t.siteResourceId, t.providerId] })] +); + export const siteResourceAiModels = sqliteTable( "siteResourceAiModels", { @@ -1787,5 +1809,9 @@ export type AiProvider = InferSelectModel; export type AiModel = InferSelectModel; export type AiBudget = InferSelectModel; export type AiBudgetPeriod = InferSelectModel; +export type ResourceAiProvider = InferSelectModel; +export type SiteResourceAiProvider = InferSelectModel< + typeof siteResourceAiProviders +>; export type ResourceAiModel = InferSelectModel; export type SiteResourceAiModel = InferSelectModel; diff --git a/server/lib/aiInferenceResource.ts b/server/lib/aiInferenceResource.ts new file mode 100644 index 000000000..4175a11e8 --- /dev/null +++ b/server/lib/aiInferenceResource.ts @@ -0,0 +1,512 @@ +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; +} diff --git a/server/lib/traefik/getTraefikConfig.ts b/server/lib/traefik/getTraefikConfig.ts index 4f5dd3860..574b8a374 100644 --- a/server/lib/traefik/getTraefikConfig.ts +++ b/server/lib/traefik/getTraefikConfig.ts @@ -1,4 +1,4 @@ -import { db, targetHealthCheck, domains, aiProviders } from "@server/db"; +import { db, targetHealthCheck, domains, aiProviders, resourceAiProviders } from "@server/db"; import { and, eq, @@ -214,7 +214,7 @@ export async function getTraefikConfig( // central AI gateway), so they can't be reached via the targets->sites // join above - query them separately and include them on every exit node. const inferenceResources = await db - .select({ + .selectDistinct({ resourceId: resources.resourceId, resourceName: resources.name, fullDomain: resources.fullDomain, @@ -226,9 +226,13 @@ export async function getTraefikConfig( preferWildcardCert: domains.preferWildcardCert }) .from(resources) + .innerJoin( + resourceAiProviders, + eq(resources.resourceId, resourceAiProviders.resourceId) + ) .innerJoin( aiProviders, - eq(resources.aiProviderId, aiProviders.providerId) + eq(resourceAiProviders.providerId, aiProviders.providerId) ) .leftJoin(domains, eq(domains.domainId, resources.domainId)) .where( diff --git a/server/private/lib/traefik/getTraefikConfig.ts b/server/private/lib/traefik/getTraefikConfig.ts index 8e1b1eca2..abe2fc691 100644 --- a/server/private/lib/traefik/getTraefikConfig.ts +++ b/server/private/lib/traefik/getTraefikConfig.ts @@ -42,7 +42,9 @@ import { siteResources, Target, targets, - aiProviders + aiProviders, + resourceAiProviders, + siteResourceAiProviders } from "@server/db"; import { sanitize, @@ -402,7 +404,7 @@ export async function getTraefikConfig( // so they can't be reached via the joins above - query them separately // and include them on every exit node. const inferenceResources = await db - .select({ + .selectDistinct({ resourceId: resources.resourceId, fullDomain: resources.fullDomain, ssl: resources.ssl, @@ -414,9 +416,13 @@ export async function getTraefikConfig( preferWildcardCert: domains.preferWildcardCert }) .from(resources) + .innerJoin( + resourceAiProviders, + eq(resources.resourceId, resourceAiProviders.resourceId) + ) .innerJoin( aiProviders, - eq(resources.aiProviderId, aiProviders.providerId) + eq(resourceAiProviders.providerId, aiProviders.providerId) ) .leftJoin(domains, eq(domains.domainId, resources.domainId)) .where( @@ -435,16 +441,23 @@ export async function getTraefikConfig( }[] = []; if (build == "enterprise") { siteResourcesInference = await db - .select({ + .selectDistinct({ siteResourceId: siteResources.siteResourceId, alias: siteResources.alias, ssl: siteResources.ssl, enabled: siteResources.enabled }) .from(siteResources) + .innerJoin( + siteResourceAiProviders, + eq( + siteResources.siteResourceId, + siteResourceAiProviders.siteResourceId + ) + ) .innerJoin( aiProviders, - eq(siteResources.aiProviderId, aiProviders.providerId) + eq(siteResourceAiProviders.providerId, aiProviders.providerId) ) .where( and( diff --git a/server/routers/aiGateway/chatCompletions.ts b/server/routers/aiGateway/chatCompletions.ts index c9c32c2ba..51508778a 100644 --- a/server/routers/aiGateway/chatCompletions.ts +++ b/server/routers/aiGateway/chatCompletions.ts @@ -8,8 +8,10 @@ import { db, exitNodes, resourceAiModels, + resourceAiProviders, resources, siteResourceAiModels, + siteResourceAiProviders, siteResources, users } from "@server/db"; @@ -21,7 +23,6 @@ import { AiProviderType, resolveAiProviderConfig } from "@server/lib/aiProviderDefaults"; -import { verifyResourceAccessToken } from "@server/auth/verifyResourceAccessToken"; import { SESSION_COOKIE_NAME, validateSessionToken @@ -31,6 +32,7 @@ import { isIpInCidr } from "@server/lib/ip"; import { localCache } from "@server/lib/cache"; import logger from "@server/logger"; import HttpCode from "@server/types/HttpCode"; +import type { ModelAccessMode } from "@server/lib/aiInferenceResource"; // Short-lived local caches so a burst of requests from the same IP/user // doesn't hit the database on every single request. None of this is @@ -85,14 +87,23 @@ async function findClientByIp(ip: string): Promise { return result; } +type ProviderAttachment = { + provider: AiProvider; + modelAccessMode: ModelAccessMode; +}; + type ResolvedTarget = { resourceId: number | null; orgId: string | null; - provider: AiProvider; - // null = no restriction; every enabled model on the provider is allowed - allowedModelIds: number[] | null; + attachments: ProviderAttachment[]; + // model IDs on this resource's allowlist (resource-wide) + allowedModelIds: number[]; }; +type ProviderSelection = + | { ok: true; provider: AiProvider } + | { ok: false; status: number; message: string }; + export type RequestUser = { userId: string; username: string; @@ -184,88 +195,245 @@ async function resolveTarget(host: string): Promise { const [resourceRow] = await db .select({ resourceId: resources.resourceId, - orgId: resources.orgId, - provider: aiProviders + orgId: resources.orgId }) .from(resources) - .innerJoin( - aiProviders, - eq(resources.aiProviderId, aiProviders.providerId) - ) .where( and( eq(resources.fullDomain, host), eq(resources.mode, "inference"), - eq(resources.enabled, true), - eq(aiProviders.enabled, true) + eq(resources.enabled, true) ) ) .limit(1); if (resourceRow) { - const restrictions = await db - .select({ modelId: resourceAiModels.modelId }) - .from(resourceAiModels) - .where(eq(resourceAiModels.resourceId, resourceRow.resourceId)); + const attachmentRows = await db + .select({ + modelAccessMode: resourceAiProviders.modelAccessMode, + provider: aiProviders + }) + .from(resourceAiProviders) + .innerJoin( + aiProviders, + eq(resourceAiProviders.providerId, aiProviders.providerId) + ) + .where( + and( + eq(resourceAiProviders.resourceId, resourceRow.resourceId), + eq(aiProviders.enabled, true) + ) + ); + + if (attachmentRows.length === 0) { + return null; + } + + const hasAllowlist = attachmentRows.some( + (a) => a.modelAccessMode === "allowlist" + ); + let allowedModelIds: number[] = []; + if (hasAllowlist) { + const restrictions = await db + .select({ modelId: resourceAiModels.modelId }) + .from(resourceAiModels) + .where(eq(resourceAiModels.resourceId, resourceRow.resourceId)); + allowedModelIds = restrictions.map((r) => r.modelId); + } return { resourceId: resourceRow.resourceId, orgId: resourceRow.orgId, - provider: resourceRow.provider, - allowedModelIds: restrictions.length - ? restrictions.map((r) => r.modelId) - : null + attachments: attachmentRows.map((a) => ({ + provider: a.provider, + modelAccessMode: a.modelAccessMode as ModelAccessMode + })), + allowedModelIds }; } const [siteResourceRow] = await db .select({ siteResourceId: siteResources.siteResourceId, - orgId: siteResources.orgId, - provider: aiProviders + orgId: siteResources.orgId }) .from(siteResources) - .innerJoin( - aiProviders, - eq(siteResources.aiProviderId, aiProviders.providerId) - ) .where( and( eq(siteResources.alias, host), eq(siteResources.mode, "inference"), - eq(siteResources.enabled, true), - eq(aiProviders.enabled, true) + eq(siteResources.enabled, true) ) ) .limit(1); if (siteResourceRow) { - const restrictions = await db - .select({ modelId: siteResourceAiModels.modelId }) - .from(siteResourceAiModels) + const attachmentRows = await db + .select({ + modelAccessMode: siteResourceAiProviders.modelAccessMode, + provider: aiProviders + }) + .from(siteResourceAiProviders) + .innerJoin( + aiProviders, + eq(siteResourceAiProviders.providerId, aiProviders.providerId) + ) .where( - eq( - siteResourceAiModels.siteResourceId, - siteResourceRow.siteResourceId + and( + eq( + siteResourceAiProviders.siteResourceId, + siteResourceRow.siteResourceId + ), + eq(aiProviders.enabled, true) ) ); + if (attachmentRows.length === 0) { + return null; + } + + const hasAllowlist = attachmentRows.some( + (a) => a.modelAccessMode === "allowlist" + ); + let allowedModelIds: number[] = []; + if (hasAllowlist) { + const restrictions = await db + .select({ modelId: siteResourceAiModels.modelId }) + .from(siteResourceAiModels) + .where( + eq( + siteResourceAiModels.siteResourceId, + siteResourceRow.siteResourceId + ) + ); + allowedModelIds = restrictions.map((r) => r.modelId); + } + return { // siteResources have no per-user auth/policy stack today (see // the routing comment in getTraefikConfig.ts), so there's no // resource access token scope to validate a user token against. resourceId: null, orgId: siteResourceRow.orgId, - provider: siteResourceRow.provider, - allowedModelIds: restrictions.length - ? restrictions.map((r) => r.modelId) - : null + attachments: attachmentRows.map((a) => ({ + provider: a.provider, + modelAccessMode: a.modelAccessMode as ModelAccessMode + })), + allowedModelIds }; } return null; } +async function providerMatchesModel( + attachment: ProviderAttachment, + requestedModel: string, + allowedModelIds: number[] +): Promise { + if (attachment.modelAccessMode === "passthrough") { + return true; + } + + const [matchedModel] = await db + .select({ + modelId: aiModels.modelId, + enabled: aiModels.enabled + }) + .from(aiModels) + .where( + and( + eq(aiModels.providerId, attachment.provider.providerId), + eq(aiModels.modelKey, requestedModel) + ) + ) + .limit(1); + + if (!matchedModel) { + return false; + } + + if (attachment.modelAccessMode === "catalog") { + return matchedModel.enabled; + } + + // allowlist + return allowedModelIds.includes(matchedModel.modelId); +} + +async function selectProvider( + attachments: ProviderAttachment[], + allowedModelIds: number[], + requestedModel: string | undefined +): Promise { + const passthroughAttachments = attachments.filter( + (a) => a.modelAccessMode === "passthrough" + ); + const hasRestricted = attachments.some( + (a) => + a.modelAccessMode === "catalog" || a.modelAccessMode === "allowlist" + ); + + if (!requestedModel) { + if (hasRestricted) { + return { + ok: false, + status: HttpCode.FORBIDDEN, + message: + "This resource restricts access to specific models; a model must be specified" + }; + } + if (passthroughAttachments.length === 1) { + return { ok: true, provider: passthroughAttachments[0].provider }; + } + return { + ok: false, + status: HttpCode.FORBIDDEN, + message: "A model must be specified for this resource" + }; + } + + const candidates: ProviderAttachment[] = []; + for (const attachment of attachments) { + if (attachment.modelAccessMode === "passthrough") { + candidates.push(attachment); + continue; + } + if ( + await providerMatchesModel( + attachment, + requestedModel, + allowedModelIds + ) + ) { + candidates.push(attachment); + } + } + + if (candidates.length === 1) { + return { ok: true, provider: candidates[0].provider }; + } + + if (candidates.length > 1) { + return { + ok: false, + status: HttpCode.FORBIDDEN, + message: `Model "${requestedModel}" is ambiguous across multiple AI providers on this resource` + }; + } + + // Zero candidates: fall back to a single passthrough attachment if present + if (passthroughAttachments.length === 1) { + return { ok: true, provider: passthroughAttachments[0].provider }; + } + + return { + ok: false, + status: HttpCode.FORBIDDEN, + message: `Model "${requestedModel}" is not permitted on this resource` + }; +} + // Generic OpenAI-wire-compatible passthrough. Anthropic's native API uses a // different path/schema; everything else here is OpenAI-compatible today. function getCompletionsPath(type: AiProviderType): string { @@ -296,7 +464,7 @@ export async function chatCompletions( }); } - const { provider, allowedModelIds, resourceId, orgId } = target; + const { attachments, allowedModelIds, resourceId, orgId } = target; // Best-effort identity resolution - not yet enforced, but lets us // start making per-user access decisions (e.g. model/role-based @@ -311,39 +479,19 @@ export async function chatCompletions( const requestedModel = typeof req.body?.model === "string" ? req.body.model : undefined; - if (allowedModelIds) { - if (!requestedModel) { - return res.status(HttpCode.FORBIDDEN).json({ - error: { - message: - "This resource restricts access to specific models; a model must be specified" - } - }); - } - - const [matchedModel] = await db - .select({ modelId: aiModels.modelId }) - .from(aiModels) - .where( - and( - eq(aiModels.providerId, provider.providerId), - eq(aiModels.modelKey, requestedModel) - ) - ) - .limit(1); - - if ( - !matchedModel || - !allowedModelIds.includes(matchedModel.modelId) - ) { - return res.status(HttpCode.FORBIDDEN).json({ - error: { - message: `Model "${requestedModel}" is not permitted on this resource` - } - }); - } + const selection = await selectProvider( + attachments, + allowedModelIds, + requestedModel + ); + if (!selection.ok) { + return res.status(selection.status).json({ + error: { message: selection.message } + }); } + const { provider } = selection; + if (!provider.apiKey) { return res.status(HttpCode.INTERNAL_SERVER_ERROR).json({ error: { message: "AI provider has no API key configured" } diff --git a/server/routers/aiProvider/createAiModel.ts b/server/routers/aiProvider/createAiModel.ts index 7fda989d2..dfe2a62e1 100644 --- a/server/routers/aiProvider/createAiModel.ts +++ b/server/routers/aiProvider/createAiModel.ts @@ -9,26 +9,16 @@ import { fromError } from "zod-validation-error"; import { OpenAPITags, registry } from "@server/openApi"; import { and, eq } from "drizzle-orm"; import type { CreateOrEditAiModelResponse } from "@server/routers/aiProvider/types"; -import { - aiBudgetUnitSchema, - refineBudgetFields -} from "@server/routers/aiProvider/validation"; const paramsSchema = z.strictObject({ providerId: z.coerce.number().int().positive() }); -const bodySchema = z - .strictObject({ - modelKey: z.string().nonempty(), - name: z.string().nonempty(), - budgetAmount: z.number().positive().optional().nullable(), - budgetUnit: aiBudgetUnitSchema.optional().nullable(), - enabled: z.boolean().optional() - }) - .superRefine((data, ctx) => { - refineBudgetFields(data, ctx); - }); +const bodySchema = z.strictObject({ + modelKey: z.string().nonempty(), + name: z.string().nonempty(), + enabled: z.boolean().optional() +}); registry.registerPath({ method: "put", @@ -79,8 +69,7 @@ export async function createAiModel( } const { providerId } = parsedParams.data; - const { modelKey, name, budgetAmount, budgetUnit, enabled } = - parsedBody.data; + const { modelKey, name, enabled } = parsedBody.data; const [provider] = req.aiProvider && req.aiProvider.providerId === providerId @@ -127,8 +116,6 @@ export async function createAiModel( providerId, modelKey, name, - budgetAmount: budgetAmount ?? null, - budgetUnit: budgetUnit ?? null, enabled: enabled ?? true, createdAt: now, updatedAt: now diff --git a/server/routers/aiProvider/createAiProvider.ts b/server/routers/aiProvider/createAiProvider.ts index 1d58344cb..6fb3f395a 100644 --- a/server/routers/aiProvider/createAiProvider.ts +++ b/server/routers/aiProvider/createAiProvider.ts @@ -13,10 +13,8 @@ import type { CreateOrEditAiProviderResponse } from "@server/routers/aiProvider/ import { toPublicAiProvider } from "@server/routers/aiProvider/types"; import { aiAuthTypeSchema, - aiBudgetUnitSchema, aiProviderTypeSchema, aiRoutingModeSchema, - refineBudgetFields, refineProviderUpstreamFields } from "@server/routers/aiProvider/validation"; @@ -33,13 +31,10 @@ const bodySchema = z authType: aiAuthTypeSchema.optional().nullable(), routingMode: aiRoutingModeSchema.optional(), skipTlsVerification: z.boolean().optional(), - budgetAmount: z.number().positive().optional().nullable(), - budgetUnit: aiBudgetUnitSchema.optional().nullable(), enabled: z.boolean().optional() }) .superRefine((data, ctx) => { refineProviderUpstreamFields(data, ctx); - refineBudgetFields(data, ctx); }); registry.registerPath({ @@ -99,8 +94,6 @@ export async function createAiProvider( authType, routingMode, skipTlsVerification, - budgetAmount, - budgetUnit, enabled } = parsedBody.data; @@ -126,8 +119,6 @@ export async function createAiProvider( authType: authType ?? null, routingMode: resolvedRoutingMode, skipTlsVerification: skipTlsVerification ?? false, - budgetAmount: budgetAmount ?? null, - budgetUnit: budgetUnit ?? null, enabled: enabled ?? true, createdAt: now, updatedAt: now diff --git a/server/routers/aiProvider/updateAiModel.ts b/server/routers/aiProvider/updateAiModel.ts index fa94e8da4..fe9ccf273 100644 --- a/server/routers/aiProvider/updateAiModel.ts +++ b/server/routers/aiProvider/updateAiModel.ts @@ -9,26 +9,16 @@ import { fromError } from "zod-validation-error"; import { OpenAPITags, registry } from "@server/openApi"; import { and, eq, ne } from "drizzle-orm"; import type { CreateOrEditAiModelResponse } from "@server/routers/aiProvider/types"; -import { - aiBudgetUnitSchema, - refineBudgetFields -} from "@server/routers/aiProvider/validation"; const paramsSchema = z.strictObject({ modelId: z.coerce.number().int().positive() }); -const bodySchema = z - .strictObject({ - modelKey: z.string().nonempty().optional(), - name: z.string().nonempty().optional(), - budgetAmount: z.number().positive().optional().nullable(), - budgetUnit: aiBudgetUnitSchema.optional().nullable(), - enabled: z.boolean().optional() - }) - .superRefine((data, ctx) => { - refineBudgetFields(data, ctx); - }); +const bodySchema = z.strictObject({ + modelKey: z.string().nonempty().optional(), + name: z.string().nonempty().optional(), + enabled: z.boolean().optional() +}); registry.registerPath({ method: "post", @@ -135,12 +125,6 @@ export async function updateAiModel( if (body.name !== undefined) { updateData.name = body.name; } - if (body.budgetAmount !== undefined) { - updateData.budgetAmount = body.budgetAmount; - } - if (body.budgetUnit !== undefined) { - updateData.budgetUnit = body.budgetUnit; - } if (body.enabled !== undefined) { updateData.enabled = body.enabled; } diff --git a/server/routers/aiProvider/updateAiProvider.ts b/server/routers/aiProvider/updateAiProvider.ts index 39787fb93..35e27fba2 100644 --- a/server/routers/aiProvider/updateAiProvider.ts +++ b/server/routers/aiProvider/updateAiProvider.ts @@ -14,10 +14,8 @@ import type { CreateOrEditAiProviderResponse } from "@server/routers/aiProvider/ import { toPublicAiProvider } from "@server/routers/aiProvider/types"; import { aiAuthTypeSchema, - aiBudgetUnitSchema, aiProviderTypeSchema, aiRoutingModeSchema, - refineBudgetFields, refineProviderUpstreamFields } from "@server/routers/aiProvider/validation"; import type { @@ -29,21 +27,15 @@ const paramsSchema = z.strictObject({ providerId: z.coerce.number().int().positive() }); -const bodySchema = z - .strictObject({ - name: z.string().nonempty().optional(), - upstreamUrl: z.url().optional().nullable(), - apiKey: z.string().optional(), - authType: aiAuthTypeSchema.optional().nullable(), - routingMode: aiRoutingModeSchema.optional(), - skipTlsVerification: z.boolean().optional(), - budgetAmount: z.number().positive().optional().nullable(), - budgetUnit: aiBudgetUnitSchema.optional().nullable(), - enabled: z.boolean().optional() - }) - .superRefine((data, ctx) => { - refineBudgetFields(data, ctx); - }); +const bodySchema = z.strictObject({ + name: z.string().nonempty().optional(), + upstreamUrl: z.url().optional().nullable(), + apiKey: z.string().optional(), + authType: aiAuthTypeSchema.optional().nullable(), + routingMode: aiRoutingModeSchema.optional(), + skipTlsVerification: z.boolean().optional(), + enabled: z.boolean().optional() +}); registry.registerPath({ method: "post", @@ -168,12 +160,6 @@ export async function updateAiProvider( if (body.enabled !== undefined) { updateData.enabled = body.enabled; } - if (body.budgetAmount !== undefined) { - updateData.budgetAmount = body.budgetAmount; - } - if (body.budgetUnit !== undefined) { - updateData.budgetUnit = body.budgetUnit; - } if (nextRoutingMode === "target") { updateData.upstreamUrl = null; } else if (body.upstreamUrl !== undefined) { diff --git a/server/routers/aiProvider/validation.ts b/server/routers/aiProvider/validation.ts index c830c08de..a5e2aa962 100644 --- a/server/routers/aiProvider/validation.ts +++ b/server/routers/aiProvider/validation.ts @@ -1,7 +1,6 @@ import { z } from "zod"; import { providerRequiresUpstreamUrl, - type AiBudgetUnit, type AiProviderRoutingMode, type AiProviderType } from "@server/lib/aiProviderDefaults"; @@ -18,33 +17,10 @@ export const aiProviderTypeSchema = z.enum([ "custom" ]); -export const aiBudgetUnitSchema = z.enum(["usd", "tokens"]); - export const aiAuthTypeSchema = z.enum(["bearer"]); export const aiRoutingModeSchema = z.enum(["url", "target"]); -export function refineBudgetFields( - data: { - budgetAmount?: number | null; - budgetUnit?: AiBudgetUnit | null; - }, - ctx: z.RefinementCtx -) { - const hasAmount = - data.budgetAmount !== undefined && data.budgetAmount !== null; - const hasUnit = data.budgetUnit !== undefined && data.budgetUnit !== null; - - if (hasAmount !== hasUnit) { - ctx.addIssue({ - code: "custom", - message: - "budgetAmount and budgetUnit must both be set or both omitted", - path: hasAmount ? ["budgetUnit"] : ["budgetAmount"] - }); - } -} - export function refineProviderUpstreamFields( data: { type: AiProviderType; diff --git a/server/routers/external.ts b/server/routers/external.ts index 4ebf58cae..94165cc8d 100644 --- a/server/routers/external.ts +++ b/server/routers/external.ts @@ -414,6 +414,13 @@ authenticated.get( siteResource.listSiteResourceAiModels ); +authenticated.get( + "/site-resource/:siteResourceId/ai-providers", + verifySiteResourceAccess, + verifyUserHasAction(ActionsEnum.listResourceAiModels), + siteResource.listSiteResourceAiProviders +); + authenticated.post( "/site-resource/:siteResourceId/roles", verifySiteResourceAccess, @@ -432,6 +439,46 @@ authenticated.post( siteResource.setSiteResourceAiModels ); +authenticated.post( + "/site-resource/:siteResourceId/ai-models/add", + verifySiteResourceAccess, + verifyUserHasAction(ActionsEnum.setResourceAiModels), + logActionAudit(ActionsEnum.setResourceAiModels), + siteResource.addAiModelToSiteResource +); + +authenticated.post( + "/site-resource/:siteResourceId/ai-models/remove", + verifySiteResourceAccess, + verifyUserHasAction(ActionsEnum.setResourceAiModels), + logActionAudit(ActionsEnum.setResourceAiModels), + siteResource.removeAiModelFromSiteResource +); + +authenticated.post( + "/site-resource/:siteResourceId/ai-providers", + verifySiteResourceAccess, + verifyUserHasAction(ActionsEnum.setResourceAiModels), + logActionAudit(ActionsEnum.setResourceAiModels), + siteResource.setSiteResourceAiProviders +); + +authenticated.post( + "/site-resource/:siteResourceId/ai-providers/add", + verifySiteResourceAccess, + verifyUserHasAction(ActionsEnum.setResourceAiModels), + logActionAudit(ActionsEnum.setResourceAiModels), + siteResource.addAiProviderToSiteResource +); + +authenticated.post( + "/site-resource/:siteResourceId/ai-providers/remove", + verifySiteResourceAccess, + verifyUserHasAction(ActionsEnum.setResourceAiModels), + logActionAudit(ActionsEnum.setResourceAiModels), + siteResource.removeAiProviderFromSiteResource +); + authenticated.post( "/site-resource/:siteResourceId/users", verifySiteResourceAccess, @@ -673,6 +720,13 @@ authenticated.get( resource.listResourceAiModels ); +authenticated.get( + "/resource/:resourceId/ai-providers", + verifyResourceAccess, + verifyUserHasAction(ActionsEnum.listResourceAiModels), + resource.listResourceAiProviders +); + authenticated.get( "/resource/:resourceId", verifyResourceAccess, @@ -884,6 +938,46 @@ authenticated.post( resource.setResourceAiModels ); +authenticated.post( + "/resource/:resourceId/ai-models/add", + verifyResourceAccess, + verifyUserHasAction(ActionsEnum.setResourceAiModels), + logActionAudit(ActionsEnum.setResourceAiModels), + resource.addAiModelToResource +); + +authenticated.post( + "/resource/:resourceId/ai-models/remove", + verifyResourceAccess, + verifyUserHasAction(ActionsEnum.setResourceAiModels), + logActionAudit(ActionsEnum.setResourceAiModels), + resource.removeAiModelFromResource +); + +authenticated.post( + "/resource/:resourceId/ai-providers", + verifyResourceAccess, + verifyUserHasAction(ActionsEnum.setResourceAiModels), + logActionAudit(ActionsEnum.setResourceAiModels), + resource.setResourceAiProviders +); + +authenticated.post( + "/resource/:resourceId/ai-providers/add", + verifyResourceAccess, + verifyUserHasAction(ActionsEnum.setResourceAiModels), + logActionAudit(ActionsEnum.setResourceAiModels), + resource.addAiProviderToResource +); + +authenticated.post( + "/resource/:resourceId/ai-providers/remove", + verifyResourceAccess, + verifyUserHasAction(ActionsEnum.setResourceAiModels), + logActionAudit(ActionsEnum.setResourceAiModels), + resource.removeAiProviderFromResource +); + authenticated.put( "/resource-policy/:resourcePolicyId/access-control", verifyResourcePolicyAccess, diff --git a/server/routers/integration.ts b/server/routers/integration.ts index 41757e7a0..e08c14920 100644 --- a/server/routers/integration.ts +++ b/server/routers/integration.ts @@ -256,6 +256,16 @@ authenticated.get( siteResource.listSiteResourceAiModels ); +authenticated.get( + [ + "/site-resource/:siteResourceId/ai-providers", + "/private-resource/:siteResourceId/ai-providers" + ], + verifyApiKeySiteResourceAccess, + verifyApiKeyHasAction(ActionsEnum.listResourceAiModels), + siteResource.listSiteResourceAiProviders +); + authenticated.post( [ "/site-resource/:siteResourceId/roles", @@ -341,6 +351,39 @@ authenticated.post( siteResource.removeAiModelFromSiteResource ); +authenticated.post( + [ + "/site-resource/:siteResourceId/ai-providers", + "/private-resource/:siteResourceId/ai-providers" + ], + verifyApiKeySiteResourceAccess, + verifyApiKeyHasAction(ActionsEnum.setResourceAiModels), + logActionAudit(ActionsEnum.setResourceAiModels), + siteResource.setSiteResourceAiProviders +); + +authenticated.post( + [ + "/site-resource/:siteResourceId/ai-providers/add", + "/private-resource/:siteResourceId/ai-providers/add" + ], + verifyApiKeySiteResourceAccess, + verifyApiKeyHasAction(ActionsEnum.setResourceAiModels), + logActionAudit(ActionsEnum.setResourceAiModels), + siteResource.addAiProviderToSiteResource +); + +authenticated.post( + [ + "/site-resource/:siteResourceId/ai-providers/remove", + "/private-resource/:siteResourceId/ai-providers/remove" + ], + verifyApiKeySiteResourceAccess, + verifyApiKeyHasAction(ActionsEnum.setResourceAiModels), + logActionAudit(ActionsEnum.setResourceAiModels), + siteResource.removeAiProviderFromSiteResource +); + authenticated.post( [ "/site-resource/:siteResourceId/users/add", @@ -563,6 +606,16 @@ authenticated.get( resource.listResourceAiModels ); +authenticated.get( + [ + "/resource/:resourceId/ai-providers", + "/public-resource/:resourceId/ai-providers" + ], + verifyApiKeyResourceAccess, + verifyApiKeyHasAction(ActionsEnum.listResourceAiModels), + resource.listResourceAiProviders +); + authenticated.get( ["/resource/:resourceId", "/public-resource/:resourceId"], verifyApiKeyResourceAccess, @@ -775,6 +828,17 @@ authenticated.post( resource.setResourceAiModels ); +authenticated.post( + [ + "/resource/:resourceId/ai-providers", + "/public-resource/:resourceId/ai-providers" + ], + verifyApiKeyResourceAccess, + verifyApiKeyHasAction(ActionsEnum.setResourceAiModels), + logActionAudit(ActionsEnum.setResourceAiModels), + resource.setResourceAiProviders +); + authenticated.post( ["/resource/:resourceId/users", "/public-resource/:resourceId/users"], verifyApiKeyResourceAccess, @@ -989,6 +1053,28 @@ authenticated.post( resource.removeAiModelFromResource ); +authenticated.post( + [ + "/resource/:resourceId/ai-providers/add", + "/public-resource/:resourceId/ai-providers/add" + ], + verifyApiKeyResourceAccess, + verifyApiKeyHasAction(ActionsEnum.setResourceAiModels), + logActionAudit(ActionsEnum.setResourceAiModels), + resource.addAiProviderToResource +); + +authenticated.post( + [ + "/resource/:resourceId/ai-providers/remove", + "/public-resource/:resourceId/ai-providers/remove" + ], + verifyApiKeyResourceAccess, + verifyApiKeyHasAction(ActionsEnum.setResourceAiModels), + logActionAudit(ActionsEnum.setResourceAiModels), + resource.removeAiProviderFromResource +); + authenticated.post( [ "/resource/:resourceId/users/add", diff --git a/server/routers/resource/addAiModelToResource.ts b/server/routers/resource/addAiModelToResource.ts index fd33a3e91..3121c9ffa 100644 --- a/server/routers/resource/addAiModelToResource.ts +++ b/server/routers/resource/addAiModelToResource.ts @@ -1,6 +1,6 @@ import { Request, Response, NextFunction } from "express"; import { z } from "zod"; -import { db, resources, resourceAiModels, aiModels } from "@server/db"; +import { db, resources, resourceAiModels } from "@server/db"; import { eq, and } from "drizzle-orm"; import response from "@server/lib/response"; import HttpCode from "@server/types/HttpCode"; @@ -8,7 +8,10 @@ import createHttpError from "http-errors"; import logger from "@server/logger"; import { fromError } from "zod-validation-error"; import { OpenAPITags, registry } from "@server/openApi"; - +import { + assertPublicAllowlistApiEligible, + assertModelsBelongToPublicAllowlistProviders +} from "@server/lib/aiInferenceResource"; const addAiModelToResourceBodySchema = z.strictObject({ modelId: z.int().positive() }); @@ -21,7 +24,7 @@ registry.registerPath({ method: "post", path: "/resource/{resourceId}/ai-models/add", description: - "Add a single AI model to a resource's model restriction allow-list.", + "Add a single catalog model to an inference resource allowlist. Requires at least one attached AI provider in allowlist mode. The model must belong to a provider attached in allowlist mode.", tags: [OpenAPITags.PublicResource], request: { params: addAiModelToResourceParamsSchema, @@ -95,33 +98,18 @@ export async function addAiModelToResource( ); } - if (!resource.aiProviderId) { - return next( - createHttpError( - HttpCode.BAD_REQUEST, - "Resource has no AI provider linked" - ) - ); + const eligibleError = await assertPublicAllowlistApiEligible(resource); + if (eligibleError) { + return next(createHttpError(HttpCode.BAD_REQUEST, eligibleError)); } - const [model] = await db - .select() - .from(aiModels) - .where( - and( - eq(aiModels.modelId, modelId), - eq(aiModels.providerId, resource.aiProviderId) - ) - ) - .limit(1); - - if (!model) { - return next( - createHttpError( - HttpCode.NOT_FOUND, - "Model not found or does not belong to this resource's AI provider" - ) - ); + const modelError = await assertModelsBelongToPublicAllowlistProviders({ + orgId: resource.orgId, + resourceId, + modelIds: [modelId] + }); + if (modelError) { + return next(createHttpError(HttpCode.BAD_REQUEST, modelError)); } const existingEntry = await db diff --git a/server/routers/resource/addAiProviderToResource.ts b/server/routers/resource/addAiProviderToResource.ts new file mode 100644 index 000000000..d51cf244c --- /dev/null +++ b/server/routers/resource/addAiProviderToResource.ts @@ -0,0 +1,152 @@ +import { Request, Response, NextFunction } from "express"; +import { z } from "zod"; +import { db, resources } from "@server/db"; +import { eq } from "drizzle-orm"; +import response from "@server/lib/response"; +import HttpCode from "@server/types/HttpCode"; +import createHttpError from "http-errors"; +import logger from "@server/logger"; +import { fromError } from "zod-validation-error"; +import { OpenAPITags, registry } from "@server/openApi"; +import { + isInferenceFieldsError, + listPublicResourceAiProviders, + modelAccessModeSchema, + resolveProviderAttachments, + setPublicResourceAiProviders +} from "@server/lib/aiInferenceResource"; + +const addAiProviderToResourceBodySchema = z.strictObject({ + providerId: z.number().int().positive(), + modelAccessMode: modelAccessModeSchema.optional() +}); + +const addAiProviderToResourceParamsSchema = z.strictObject({ + resourceId: z.coerce.number().int().positive() +}); + +registry.registerPath({ + method: "post", + path: "/resource/{resourceId}/ai-providers/add", + description: + "Add or replace a single AI provider attachment on an inference resource.", + tags: [OpenAPITags.PublicResource], + request: { + params: addAiProviderToResourceParamsSchema, + body: { + content: { + "application/json": { + schema: addAiProviderToResourceBodySchema + } + } + } + }, + responses: { + 200: { + description: "Successful response", + content: { + "application/json": { + schema: z.object({ + data: z.record(z.string(), z.any()).nullable(), + success: z.boolean(), + error: z.boolean(), + message: z.string(), + status: z.number() + }) + } + } + } + } +}); + +export async function addAiProviderToResource( + req: Request, + res: Response, + next: NextFunction +): Promise { + try { + const parsedBody = addAiProviderToResourceBodySchema.safeParse( + req.body + ); + if (!parsedBody.success) { + return next( + createHttpError( + HttpCode.BAD_REQUEST, + fromError(parsedBody.error).toString() + ) + ); + } + + const { providerId, modelAccessMode } = parsedBody.data; + + const parsedParams = addAiProviderToResourceParamsSchema.safeParse( + req.params + ); + if (!parsedParams.success) { + return next( + createHttpError( + HttpCode.BAD_REQUEST, + fromError(parsedParams.error).toString() + ) + ); + } + + const { resourceId } = parsedParams.data; + + const [resource] = await db + .select() + .from(resources) + .where(eq(resources.resourceId, resourceId)) + .limit(1); + + if (!resource) { + return next( + createHttpError(HttpCode.NOT_FOUND, "Resource not found") + ); + } + + if (resource.mode !== "inference") { + return next( + createHttpError( + HttpCode.BAD_REQUEST, + "AI providers can only be attached to inference-mode resources" + ) + ); + } + + const existing = await listPublicResourceAiProviders(resourceId); + const nextAttachments = [ + ...existing + .filter((a) => a.providerId !== providerId) + .map((a) => ({ + providerId: a.providerId, + modelAccessMode: a.modelAccessMode + })), + { providerId, modelAccessMode } + ]; + + const attachments = await resolveProviderAttachments({ + orgId: resource.orgId, + attachments: nextAttachments, + requireAtLeastOne: true + }); + if (isInferenceFieldsError(attachments)) { + return next(createHttpError(HttpCode.BAD_REQUEST, attachments.error)); + } + + await setPublicResourceAiProviders(resourceId, attachments); + + return response(res, { + data: {}, + success: true, + error: false, + message: "AI provider added to resource successfully", + status: HttpCode.CREATED + }); + } catch (error) { + logger.error(error); + return next( + createHttpError(HttpCode.INTERNAL_SERVER_ERROR, "An error occurred") + ); + } +} diff --git a/server/routers/resource/createResource.ts b/server/routers/resource/createResource.ts index e56896e55..4c880c3af 100644 --- a/server/routers/resource/createResource.ts +++ b/server/routers/resource/createResource.ts @@ -38,6 +38,13 @@ import { } from "@server/db/names"; import { usageService } from "@server/lib/billing/usageService"; import { LimitId } from "@server/lib/billing"; +import { + isInferenceFieldsError, + resolveProviderAttachments, + resourceAiProviderAttachmentSchema, + setPublicResourceAiProviders, + type ResourceAiProviderAttachment +} from "@server/lib/aiInferenceResource"; const createResourceParamsSchema = z.strictObject({ orgId: z.string() @@ -98,13 +105,11 @@ const createHttpResourceSchema = z authDaemonPort: z.int().positive().optional(), authDaemonMode: z.enum(["site", "remote", "native"]).optional(), // Inference settings - aiProviderId: z - .number() - .int() - .positive() + aiProviders: z + .array(resourceAiProviderAttachmentSchema) .optional() .describe( - "For inference-mode resources: the AI provider this resource proxies chat completions to." + "For inference-mode resources: AI providers to attach. Each entry may set modelAccessMode (passthrough, catalog, or allowlist); defaults to passthrough. At most one passthrough provider is allowed." ) }) .refine( @@ -377,11 +382,35 @@ async function createHttpResource( authDaemonPort, authDaemonMode, pamMode, - aiProviderId + aiProviders: aiProviderInputs } = parsedBody.data; const subdomain = parsedBody.data.subdomain; const stickySession = parsedBody.data.stickySession; + const effectiveMode = mode ?? "http"; + + let providerAttachments: ResourceAiProviderAttachment[] = []; + if (effectiveMode === "inference") { + const resolved = await resolveProviderAttachments({ + orgId, + attachments: aiProviderInputs ?? [], + requireAtLeastOne: true + }); + if (isInferenceFieldsError(resolved)) { + return next( + createHttpError(HttpCode.BAD_REQUEST, resolved.error) + ); + } + providerAttachments = resolved; + } else if (aiProviderInputs && aiProviderInputs.length > 0) { + return next( + createHttpError( + HttpCode.BAD_REQUEST, + "AI providers can only be attached to inference-mode resources" + ) + ); + } + // Wildcard subdomains are a paid feature if (subdomain && subdomain.includes("*")) { const isLicensed = await isLicensedOrSubscribed( @@ -422,7 +451,7 @@ async function createHttpResource( } if ( - ["ssh", "rdp", "vnc"].includes(mode!) && + ["ssh", "rdp", "vnc"].includes(effectiveMode) && !isLicensedOrSubscribed( orgId!, tierMatrix[TierFeature.AdvancedPublicResources] @@ -555,7 +584,7 @@ async function createHttpResource( orgId, name, subdomain: finalSubdomain, - mode: mode, + mode: effectiveMode, pamMode: pamMode, authDaemonMode: authDaemonMode, authDaemonPort: authDaemonPort, @@ -564,11 +593,18 @@ async function createHttpResource( postAuthPath: postAuthPath, wildcard, health: "unknown", - defaultResourcePolicyId: defaultPolicy.resourcePolicyId, - aiProviderId: aiProviderId ?? null + defaultResourcePolicyId: defaultPolicy.resourcePolicyId }) .returning(); + if (providerAttachments.length > 0) { + await setPublicResourceAiProviders( + newResource[0].resourceId, + providerAttachments, + trx + ); + } + await trx.insert(roleResources).values({ roleId: adminRole[0].roleId, resourceId: newResource[0].resourceId diff --git a/server/routers/resource/index.ts b/server/routers/resource/index.ts index 32992d364..a09cf0452 100644 --- a/server/routers/resource/index.ts +++ b/server/routers/resource/index.ts @@ -39,3 +39,7 @@ export * from "./listResourceAiModels"; export * from "./setResourceAiModels"; export * from "./addAiModelToResource"; export * from "./removeAiModelFromResource"; +export * from "./listResourceAiProviders"; +export * from "./setResourceAiProviders"; +export * from "./addAiProviderToResource"; +export * from "./removeAiProviderFromResource"; diff --git a/server/routers/resource/listResourceAiModels.ts b/server/routers/resource/listResourceAiModels.ts index f6875faff..0a5dea778 100644 --- a/server/routers/resource/listResourceAiModels.ts +++ b/server/routers/resource/listResourceAiModels.ts @@ -34,7 +34,7 @@ registry.registerPath({ method: "get", path: "/resource/{resourceId}/ai-models", description: - "List the AI models a resource is restricted to. An empty list means the resource is not restricted and every enabled model on its linked AI provider is allowed.", + "List catalog models on this resource's allowlist. Only enforced when modelAccessMode=allowlist; an empty allowlist denies all models.", tags: [OpenAPITags.PublicResource], request: { params: listResourceAiModelsParamsSchema diff --git a/server/routers/resource/listResourceAiProviders.ts b/server/routers/resource/listResourceAiProviders.ts new file mode 100644 index 000000000..2c81c45d5 --- /dev/null +++ b/server/routers/resource/listResourceAiProviders.ts @@ -0,0 +1,94 @@ +import { Request, Response, NextFunction } from "express"; +import { z } from "zod"; +import { db, resources } from "@server/db"; +import { eq } from "drizzle-orm"; +import response from "@server/lib/response"; +import HttpCode from "@server/types/HttpCode"; +import createHttpError from "http-errors"; +import logger from "@server/logger"; +import { fromError } from "zod-validation-error"; +import { OpenAPITags, registry } from "@server/openApi"; +import { listPublicResourceAiProviders } from "@server/lib/aiInferenceResource"; + +const listResourceAiProvidersParamsSchema = z.strictObject({ + resourceId: z.coerce.number().int().positive() +}); + +export type ListResourceAiProvidersResponse = { + providers: Awaited>; +}; + +registry.registerPath({ + method: "get", + path: "/resource/{resourceId}/ai-providers", + description: "List AI providers attached to an inference resource.", + tags: [OpenAPITags.PublicResource], + request: { + params: listResourceAiProvidersParamsSchema + }, + responses: { + 200: { + description: "Successful response", + content: { + "application/json": { + schema: z.object({ + data: z.record(z.string(), z.any()).nullable(), + success: z.boolean(), + error: z.boolean(), + message: z.string(), + status: z.number() + }) + } + } + } + } +}); + +export async function listResourceAiProviders( + req: Request, + res: Response, + next: NextFunction +): Promise { + try { + const parsedParams = listResourceAiProvidersParamsSchema.safeParse( + req.params + ); + if (!parsedParams.success) { + return next( + createHttpError( + HttpCode.BAD_REQUEST, + fromError(parsedParams.error).toString() + ) + ); + } + + const { resourceId } = parsedParams.data; + + const [resource] = await db + .select() + .from(resources) + .where(eq(resources.resourceId, resourceId)) + .limit(1); + + if (!resource) { + return next( + createHttpError(HttpCode.NOT_FOUND, "Resource not found") + ); + } + + const providers = await listPublicResourceAiProviders(resourceId); + + return response(res, { + data: { providers }, + success: true, + error: false, + message: "Resource AI providers retrieved successfully", + status: HttpCode.OK + }); + } catch (error) { + logger.error(error); + return next( + createHttpError(HttpCode.INTERNAL_SERVER_ERROR, "An error occurred") + ); + } +} diff --git a/server/routers/resource/removeAiModelFromResource.ts b/server/routers/resource/removeAiModelFromResource.ts index 47445bbf5..7b5cb012f 100644 --- a/server/routers/resource/removeAiModelFromResource.ts +++ b/server/routers/resource/removeAiModelFromResource.ts @@ -8,6 +8,7 @@ import createHttpError from "http-errors"; import logger from "@server/logger"; import { fromError } from "zod-validation-error"; import { OpenAPITags, registry } from "@server/openApi"; +import { assertPublicAllowlistApiEligible } from "@server/lib/aiInferenceResource"; const removeAiModelFromResourceBodySchema = z.strictObject({ modelId: z.int().positive() @@ -21,7 +22,7 @@ registry.registerPath({ method: "post", path: "/resource/{resourceId}/ai-models/remove", description: - "Remove a single AI model from a resource's model restriction allow-list.", + "Remove a single catalog model from an inference resource allowlist. Requires at least one attached AI provider in allowlist mode.", tags: [OpenAPITags.PublicResource], request: { params: removeAiModelFromResourceParamsSchema, @@ -97,6 +98,11 @@ export async function removeAiModelFromResource( ); } + const eligibleError = await assertPublicAllowlistApiEligible(resource); + if (eligibleError) { + return next(createHttpError(HttpCode.BAD_REQUEST, eligibleError)); + } + const existingEntry = await db .select() .from(resourceAiModels) diff --git a/server/routers/resource/removeAiProviderFromResource.ts b/server/routers/resource/removeAiProviderFromResource.ts new file mode 100644 index 000000000..d3a4294ad --- /dev/null +++ b/server/routers/resource/removeAiProviderFromResource.ts @@ -0,0 +1,165 @@ +import { Request, Response, NextFunction } from "express"; +import { z } from "zod"; +import { db, resources } from "@server/db"; +import { eq } from "drizzle-orm"; +import response from "@server/lib/response"; +import HttpCode from "@server/types/HttpCode"; +import createHttpError from "http-errors"; +import logger from "@server/logger"; +import { fromError } from "zod-validation-error"; +import { OpenAPITags, registry } from "@server/openApi"; +import { + isInferenceFieldsError, + listPublicResourceAiProviders, + resolveProviderAttachments, + setPublicResourceAiProviders +} from "@server/lib/aiInferenceResource"; + +const removeAiProviderFromResourceBodySchema = z.strictObject({ + providerId: z.number().int().positive() +}); + +const removeAiProviderFromResourceParamsSchema = z.strictObject({ + resourceId: z.coerce.number().int().positive() +}); + +registry.registerPath({ + method: "post", + path: "/resource/{resourceId}/ai-providers/remove", + description: + "Remove an AI provider attachment from an inference resource. At least one provider must remain.", + tags: [OpenAPITags.PublicResource], + request: { + params: removeAiProviderFromResourceParamsSchema, + body: { + content: { + "application/json": { + schema: removeAiProviderFromResourceBodySchema + } + } + } + }, + responses: { + 200: { + description: "Successful response", + content: { + "application/json": { + schema: z.object({ + data: z.record(z.string(), z.any()).nullable(), + success: z.boolean(), + error: z.boolean(), + message: z.string(), + status: z.number() + }) + } + } + } + } +}); + +export async function removeAiProviderFromResource( + req: Request, + res: Response, + next: NextFunction +): Promise { + try { + const parsedBody = removeAiProviderFromResourceBodySchema.safeParse( + req.body + ); + if (!parsedBody.success) { + return next( + createHttpError( + HttpCode.BAD_REQUEST, + fromError(parsedBody.error).toString() + ) + ); + } + + const { providerId } = parsedBody.data; + + const parsedParams = + removeAiProviderFromResourceParamsSchema.safeParse(req.params); + if (!parsedParams.success) { + return next( + createHttpError( + HttpCode.BAD_REQUEST, + fromError(parsedParams.error).toString() + ) + ); + } + + const { resourceId } = parsedParams.data; + + const [resource] = await db + .select() + .from(resources) + .where(eq(resources.resourceId, resourceId)) + .limit(1); + + if (!resource) { + return next( + createHttpError(HttpCode.NOT_FOUND, "Resource not found") + ); + } + + if (resource.mode !== "inference") { + return next( + createHttpError( + HttpCode.BAD_REQUEST, + "AI providers can only be attached to inference-mode resources" + ) + ); + } + + const existing = await listPublicResourceAiProviders(resourceId); + const found = existing.find((a) => a.providerId === providerId); + if (!found) { + return next( + createHttpError( + HttpCode.NOT_FOUND, + "AI provider is not attached to this resource" + ) + ); + } + + const remaining = existing + .filter((a) => a.providerId !== providerId) + .map((a) => ({ + providerId: a.providerId, + modelAccessMode: a.modelAccessMode + })); + + if (remaining.length === 0) { + return next( + createHttpError( + HttpCode.BAD_REQUEST, + "At least one AI provider is required for inference-mode resources" + ) + ); + } + + const attachments = await resolveProviderAttachments({ + orgId: resource.orgId, + attachments: remaining, + requireAtLeastOne: true + }); + if (isInferenceFieldsError(attachments)) { + return next(createHttpError(HttpCode.BAD_REQUEST, attachments.error)); + } + + await setPublicResourceAiProviders(resourceId, attachments); + + return response(res, { + data: {}, + success: true, + error: false, + message: "AI provider removed from resource successfully", + status: HttpCode.OK + }); + } catch (error) { + logger.error(error); + return next( + createHttpError(HttpCode.INTERNAL_SERVER_ERROR, "An error occurred") + ); + } +} diff --git a/server/routers/resource/setResourceAiModels.ts b/server/routers/resource/setResourceAiModels.ts index 2ef56f4fc..0c3323ebe 100644 --- a/server/routers/resource/setResourceAiModels.ts +++ b/server/routers/resource/setResourceAiModels.ts @@ -1,13 +1,17 @@ import { Request, Response, NextFunction } from "express"; import { z } from "zod"; -import { db, resources, resourceAiModels, aiModels } from "@server/db"; -import { eq, and, inArray } from "drizzle-orm"; +import { db, resources, resourceAiModels } from "@server/db"; +import { eq } from "drizzle-orm"; import response from "@server/lib/response"; import HttpCode from "@server/types/HttpCode"; import createHttpError from "http-errors"; import logger from "@server/logger"; import { fromError } from "zod-validation-error"; import { OpenAPITags, registry } from "@server/openApi"; +import { + assertPublicAllowlistApiEligible, + assertModelsBelongToPublicAllowlistProviders +} from "@server/lib/aiInferenceResource"; const setResourceAiModelsBodySchema = z.strictObject({ modelIds: z.array(z.int().positive()) @@ -21,7 +25,7 @@ registry.registerPath({ method: "post", path: "/resource/{resourceId}/ai-models", description: - "Set the AI models a resource is restricted to. This replaces all existing restrictions. Pass an empty array to remove the restriction (allow every enabled model on the linked provider).", + "Replace the allowlist of catalog models for an inference resource. Requires at least one attached AI provider in allowlist mode. Models must belong to a provider attached in allowlist mode. An empty array denies all models.", tags: [OpenAPITags.PublicResource], request: { params: setResourceAiModelsParamsSchema, @@ -95,34 +99,18 @@ export async function setResourceAiModels( ); } - if (modelIds.length > 0) { - if (!resource.aiProviderId) { - return next( - createHttpError( - HttpCode.BAD_REQUEST, - "Resource has no AI provider linked" - ) - ); - } + const eligibleError = await assertPublicAllowlistApiEligible(resource); + if (eligibleError) { + return next(createHttpError(HttpCode.BAD_REQUEST, eligibleError)); + } - const validModels = await db - .select({ modelId: aiModels.modelId }) - .from(aiModels) - .where( - and( - inArray(aiModels.modelId, modelIds), - eq(aiModels.providerId, resource.aiProviderId) - ) - ); - - if (validModels.length !== new Set(modelIds).size) { - return next( - createHttpError( - HttpCode.BAD_REQUEST, - "One or more model IDs do not exist or do not belong to this resource's AI provider" - ) - ); - } + const modelError = await assertModelsBelongToPublicAllowlistProviders({ + orgId: resource.orgId, + resourceId, + modelIds + }); + if (modelError) { + return next(createHttpError(HttpCode.BAD_REQUEST, modelError)); } await db.transaction(async (trx) => { diff --git a/server/routers/resource/setResourceAiProviders.ts b/server/routers/resource/setResourceAiProviders.ts new file mode 100644 index 000000000..a1de8f176 --- /dev/null +++ b/server/routers/resource/setResourceAiProviders.ts @@ -0,0 +1,137 @@ +import { Request, Response, NextFunction } from "express"; +import { z } from "zod"; +import { db, resources } from "@server/db"; +import { eq } from "drizzle-orm"; +import response from "@server/lib/response"; +import HttpCode from "@server/types/HttpCode"; +import createHttpError from "http-errors"; +import logger from "@server/logger"; +import { fromError } from "zod-validation-error"; +import { OpenAPITags, registry } from "@server/openApi"; +import { + isInferenceFieldsError, + resolveProviderAttachments, + resourceAiProviderAttachmentSchema, + setPublicResourceAiProviders +} from "@server/lib/aiInferenceResource"; + +const setResourceAiProvidersBodySchema = z.strictObject({ + providers: z.array(resourceAiProviderAttachmentSchema) +}); + +const setResourceAiProvidersParamsSchema = z.strictObject({ + resourceId: z.coerce.number().int().positive() +}); + +registry.registerPath({ + method: "post", + path: "/resource/{resourceId}/ai-providers", + description: + "Replace the AI providers attached to an inference resource. At least one provider is required. At most one may use passthrough mode.", + tags: [OpenAPITags.PublicResource], + request: { + params: setResourceAiProvidersParamsSchema, + body: { + content: { + "application/json": { + schema: setResourceAiProvidersBodySchema + } + } + } + }, + responses: { + 200: { + description: "Successful response", + content: { + "application/json": { + schema: z.object({ + data: z.record(z.string(), z.any()).nullable(), + success: z.boolean(), + error: z.boolean(), + message: z.string(), + status: z.number() + }) + } + } + } + } +}); + +export async function setResourceAiProviders( + req: Request, + res: Response, + next: NextFunction +): Promise { + try { + const parsedBody = setResourceAiProvidersBodySchema.safeParse(req.body); + if (!parsedBody.success) { + return next( + createHttpError( + HttpCode.BAD_REQUEST, + fromError(parsedBody.error).toString() + ) + ); + } + + const { providers } = parsedBody.data; + + const parsedParams = setResourceAiProvidersParamsSchema.safeParse( + req.params + ); + if (!parsedParams.success) { + return next( + createHttpError( + HttpCode.BAD_REQUEST, + fromError(parsedParams.error).toString() + ) + ); + } + + const { resourceId } = parsedParams.data; + + const [resource] = await db + .select() + .from(resources) + .where(eq(resources.resourceId, resourceId)) + .limit(1); + + if (!resource) { + return next( + createHttpError(HttpCode.NOT_FOUND, "Resource not found") + ); + } + + if (resource.mode !== "inference") { + return next( + createHttpError( + HttpCode.BAD_REQUEST, + "AI providers can only be attached to inference-mode resources" + ) + ); + } + + const attachments = await resolveProviderAttachments({ + orgId: resource.orgId, + attachments: providers, + requireAtLeastOne: true + }); + if (isInferenceFieldsError(attachments)) { + return next(createHttpError(HttpCode.BAD_REQUEST, attachments.error)); + } + + await setPublicResourceAiProviders(resourceId, attachments); + + return response(res, { + data: {}, + success: true, + error: false, + message: "AI providers set for resource successfully", + status: HttpCode.CREATED + }); + } catch (error) { + logger.error(error); + return next( + createHttpError(HttpCode.INTERNAL_SERVER_ERROR, "An error occurred") + ); + } +} diff --git a/server/routers/resource/updateResource.ts b/server/routers/resource/updateResource.ts index c67cb2795..0d6f8f8aa 100644 --- a/server/routers/resource/updateResource.ts +++ b/server/routers/resource/updateResource.ts @@ -120,15 +120,6 @@ const updateHttpResourceBodySchema = z .optional() .describe( "ID of the resource policy to apply to this resource. Set to null to remove the resource policy and fall back to the inline policy settings." - ), - aiProviderId: z - .number() - .int() - .positive() - .nullable() - .optional() - .describe( - "For inference-mode resources: the AI provider this resource proxies chat completions to. Set to null to unlink." ) }) .refine((data) => Object.keys(data).length > 0, { @@ -354,8 +345,10 @@ export async function updateResource( ); } - if (["http", "ssh", "rdp", "vnc"].includes(resource.mode)) { - // HANDLE UPDATING HTTP RESOURCES + if ( + ["http", "ssh", "rdp", "vnc", "inference"].includes(resource.mode) + ) { + // HANDLE UPDATING HTTP / BROWSER / INFERENCE RESOURCES return await updateHttpResource( { req, diff --git a/server/routers/siteResource/addAiModelToSiteResource.ts b/server/routers/siteResource/addAiModelToSiteResource.ts index 2e55f100b..0b6a004d7 100644 --- a/server/routers/siteResource/addAiModelToSiteResource.ts +++ b/server/routers/siteResource/addAiModelToSiteResource.ts @@ -1,11 +1,6 @@ import { Request, Response, NextFunction } from "express"; import { z } from "zod"; -import { - db, - siteResources, - siteResourceAiModels, - aiModels -} from "@server/db"; +import { db, siteResources, siteResourceAiModels } from "@server/db"; import { eq, and } from "drizzle-orm"; import response from "@server/lib/response"; import HttpCode from "@server/types/HttpCode"; @@ -13,6 +8,10 @@ import createHttpError from "http-errors"; import logger from "@server/logger"; import { fromError } from "zod-validation-error"; import { OpenAPITags, registry } from "@server/openApi"; +import { + assertSiteAllowlistApiEligible, + assertModelsBelongToSiteAllowlistProviders +} from "@server/lib/aiInferenceResource"; const addAiModelToSiteResourceBodySchema = z.strictObject({ modelId: z.int().positive() @@ -26,7 +25,7 @@ registry.registerPath({ method: "post", path: "/site-resource/{siteResourceId}/ai-models/add", description: - "Add a single AI model to a site resource's model restriction allow-list.", + "Add a single catalog model to an inference site resource allowlist. Requires at least one attached AI provider in allowlist mode. The model must belong to a provider attached in allowlist mode.", tags: [OpenAPITags.PrivateResource], request: { params: addAiModelToSiteResourceParamsSchema, @@ -102,33 +101,19 @@ export async function addAiModelToSiteResource( ); } - if (!siteResource.aiProviderId) { - return next( - createHttpError( - HttpCode.BAD_REQUEST, - "Site resource has no AI provider linked" - ) - ); + const eligibleError = + await assertSiteAllowlistApiEligible(siteResource); + if (eligibleError) { + return next(createHttpError(HttpCode.BAD_REQUEST, eligibleError)); } - const [model] = await db - .select() - .from(aiModels) - .where( - and( - eq(aiModels.modelId, modelId), - eq(aiModels.providerId, siteResource.aiProviderId) - ) - ) - .limit(1); - - if (!model) { - return next( - createHttpError( - HttpCode.NOT_FOUND, - "Model not found or does not belong to this site resource's AI provider" - ) - ); + const modelError = await assertModelsBelongToSiteAllowlistProviders({ + orgId: siteResource.orgId, + siteResourceId, + modelIds: [modelId] + }); + if (modelError) { + return next(createHttpError(HttpCode.BAD_REQUEST, modelError)); } const existingEntry = await db diff --git a/server/routers/siteResource/addAiProviderToSiteResource.ts b/server/routers/siteResource/addAiProviderToSiteResource.ts new file mode 100644 index 000000000..bae346347 --- /dev/null +++ b/server/routers/siteResource/addAiProviderToSiteResource.ts @@ -0,0 +1,152 @@ +import { Request, Response, NextFunction } from "express"; +import { z } from "zod"; +import { db, siteResources } from "@server/db"; +import { eq } from "drizzle-orm"; +import response from "@server/lib/response"; +import HttpCode from "@server/types/HttpCode"; +import createHttpError from "http-errors"; +import logger from "@server/logger"; +import { fromError } from "zod-validation-error"; +import { OpenAPITags, registry } from "@server/openApi"; +import { + isInferenceFieldsError, + listSiteResourceAiProviders, + modelAccessModeSchema, + resolveProviderAttachments, + setSiteResourceAiProviders +} from "@server/lib/aiInferenceResource"; + +const addAiProviderToSiteResourceBodySchema = z.strictObject({ + providerId: z.number().int().positive(), + modelAccessMode: modelAccessModeSchema.optional() +}); + +const addAiProviderToSiteResourceParamsSchema = z.strictObject({ + siteResourceId: z.coerce.number().int().positive() +}); + +registry.registerPath({ + method: "post", + path: "/site-resource/{siteResourceId}/ai-providers/add", + description: + "Add or replace a single AI provider attachment on an inference site resource.", + tags: [OpenAPITags.PrivateResource], + request: { + params: addAiProviderToSiteResourceParamsSchema, + body: { + content: { + "application/json": { + schema: addAiProviderToSiteResourceBodySchema + } + } + } + }, + responses: { + 200: { + description: "Successful response", + content: { + "application/json": { + schema: z.object({ + data: z.record(z.string(), z.any()).nullable(), + success: z.boolean(), + error: z.boolean(), + message: z.string(), + status: z.number() + }) + } + } + } + } +}); + +export async function addAiProviderToSiteResource( + req: Request, + res: Response, + next: NextFunction +): Promise { + try { + const parsedBody = addAiProviderToSiteResourceBodySchema.safeParse( + req.body + ); + if (!parsedBody.success) { + return next( + createHttpError( + HttpCode.BAD_REQUEST, + fromError(parsedBody.error).toString() + ) + ); + } + + const { providerId, modelAccessMode } = parsedBody.data; + + const parsedParams = addAiProviderToSiteResourceParamsSchema.safeParse( + req.params + ); + if (!parsedParams.success) { + return next( + createHttpError( + HttpCode.BAD_REQUEST, + fromError(parsedParams.error).toString() + ) + ); + } + + const { siteResourceId } = parsedParams.data; + + const [siteResource] = await db + .select() + .from(siteResources) + .where(eq(siteResources.siteResourceId, siteResourceId)) + .limit(1); + + if (!siteResource) { + return next( + createHttpError(HttpCode.NOT_FOUND, "Site resource not found") + ); + } + + if (siteResource.mode !== "inference") { + return next( + createHttpError( + HttpCode.BAD_REQUEST, + "AI providers can only be attached to inference-mode resources" + ) + ); + } + + const existing = await listSiteResourceAiProviders(siteResourceId); + const nextAttachments = [ + ...existing + .filter((a) => a.providerId !== providerId) + .map((a) => ({ + providerId: a.providerId, + modelAccessMode: a.modelAccessMode + })), + { providerId, modelAccessMode } + ]; + + const attachments = await resolveProviderAttachments({ + orgId: siteResource.orgId, + attachments: nextAttachments, + requireAtLeastOne: true + }); + if (isInferenceFieldsError(attachments)) { + return next(createHttpError(HttpCode.BAD_REQUEST, attachments.error)); + } + + await setSiteResourceAiProviders(siteResourceId, attachments); + + return response(res, { + data: {}, + success: true, + error: false, + message: "AI provider added to site resource successfully", + status: HttpCode.CREATED + }); + } catch (error) { + logger.error(error); + return next( + createHttpError(HttpCode.INTERNAL_SERVER_ERROR, "An error occurred") + ); + } +} diff --git a/server/routers/siteResource/createSiteResource.ts b/server/routers/siteResource/createSiteResource.ts index 3244668f5..47dcfc59b 100644 --- a/server/routers/siteResource/createSiteResource.ts +++ b/server/routers/siteResource/createSiteResource.ts @@ -39,6 +39,13 @@ import { createCertificate } from "#dynamic/routers/certificates/createCertifica import { build } from "@server/build"; import { usageService } from "@server/lib/billing/usageService"; import { LimitId } from "@server/lib/billing"; +import { + isInferenceFieldsError, + resolveProviderAttachments, + resourceAiProviderAttachmentSchema, + setSiteResourceAiProviders, + type ResourceAiProviderAttachment +} from "@server/lib/aiInferenceResource"; const createSiteResourceParamsSchema = z.strictObject({ orgId: z.string() @@ -79,13 +86,11 @@ const createSiteResourceSchema = z pamMode: z.enum(["passthrough", "push"]).optional(), domainId: z.string().optional(), // only used for http mode, we need this to verify the alias is unique within the org subdomain: z.string().optional(), // only used for http mode, we need this to verify the alias is unique within the org - aiProviderId: z - .number() - .int() - .positive() + aiProviders: z + .array(resourceAiProviderAttachmentSchema) .optional() .describe( - "For inference-mode site resources: the AI provider this resource proxies chat completions to." + "For inference-mode site resources: AI providers to attach. Each entry may set modelAccessMode (passthrough, catalog, or allowlist); defaults to passthrough. At most one passthrough provider is allowed." ) }) .strict() @@ -334,7 +339,7 @@ export async function createSiteResource( pamMode, domainId, subdomain, - aiProviderId + aiProviders: aiProviderInputs } = parsedBody.data; // Backward compatibility: merge deprecated siteId into siteIds array @@ -343,6 +348,28 @@ export async function createSiteResource( siteIds.push(siteId); } + let providerAttachments: ResourceAiProviderAttachment[] = []; + if (mode === "inference") { + const resolved = await resolveProviderAttachments({ + orgId, + attachments: aiProviderInputs ?? [], + requireAtLeastOne: true + }); + if (isInferenceFieldsError(resolved)) { + return next( + createHttpError(HttpCode.BAD_REQUEST, resolved.error) + ); + } + providerAttachments = resolved; + } else if (aiProviderInputs && aiProviderInputs.length > 0) { + return next( + createHttpError( + HttpCode.BAD_REQUEST, + "AI providers can only be attached to inference-mode resources" + ) + ); + } + if (build == "saas") { const usage = await usageService.getUsage( orgId, @@ -608,8 +635,7 @@ export async function createSiteResource( domainId, subdomain: finalSubdomain, fullDomain, - requiresExitNodeConnection: mode === "inference", // in the future we might want to have different modes that do this - aiProviderId: aiProviderId ?? null + requiresExitNodeConnection: mode === "inference" // in the future we might want to have different modes that do this }; if (isLicensedSshPam) { if (authDaemonPort !== undefined) @@ -625,6 +651,14 @@ export async function createSiteResource( const siteResourceId = newSiteResource.siteResourceId; + if (providerAttachments.length > 0) { + await setSiteResourceAiProviders( + siteResourceId, + providerAttachments, + trx + ); + } + //////////////////// update the associations //////////////////// if (network) { diff --git a/server/routers/siteResource/index.ts b/server/routers/siteResource/index.ts index 2acaf1e33..046ed0a13 100644 --- a/server/routers/siteResource/index.ts +++ b/server/routers/siteResource/index.ts @@ -21,3 +21,7 @@ export * from "./listSiteResourceAiModels"; export * from "./setSiteResourceAiModels"; export * from "./addAiModelToSiteResource"; export * from "./removeAiModelFromSiteResource"; +export * from "./listSiteResourceAiProviders"; +export * from "./setSiteResourceAiProviders"; +export * from "./addAiProviderToSiteResource"; +export * from "./removeAiProviderFromSiteResource"; diff --git a/server/routers/siteResource/listSiteResourceAiModels.ts b/server/routers/siteResource/listSiteResourceAiModels.ts index 26dea7cd0..26e709140 100644 --- a/server/routers/siteResource/listSiteResourceAiModels.ts +++ b/server/routers/siteResource/listSiteResourceAiModels.ts @@ -1,11 +1,6 @@ import { Request, Response, NextFunction } from "express"; import { z } from "zod"; -import { - db, - siteResources, - siteResourceAiModels, - aiModels -} from "@server/db"; +import { db, siteResources, siteResourceAiModels, aiModels } from "@server/db"; import { eq } from "drizzle-orm"; import response from "@server/lib/response"; import HttpCode from "@server/types/HttpCode"; @@ -27,10 +22,7 @@ async function query(siteResourceId: number) { enabled: aiModels.enabled }) .from(siteResourceAiModels) - .innerJoin( - aiModels, - eq(siteResourceAiModels.modelId, aiModels.modelId) - ) + .innerJoin(aiModels, eq(siteResourceAiModels.modelId, aiModels.modelId)) .where(eq(siteResourceAiModels.siteResourceId, siteResourceId)); } @@ -42,7 +34,7 @@ registry.registerPath({ method: "get", path: "/site-resource/{siteResourceId}/ai-models", description: - "List the AI models a site resource is restricted to. An empty list means the site resource is not restricted and every enabled model on its linked AI provider is allowed.", + "List catalog models on this site resource's allowlist. Only enforced when modelAccessMode=allowlist; an empty allowlist denies all models.", tags: [OpenAPITags.PrivateResource], request: { params: listSiteResourceAiModelsParamsSchema diff --git a/server/routers/siteResource/listSiteResourceAiProviders.ts b/server/routers/siteResource/listSiteResourceAiProviders.ts new file mode 100644 index 000000000..bcbe37e8b --- /dev/null +++ b/server/routers/siteResource/listSiteResourceAiProviders.ts @@ -0,0 +1,94 @@ +import { Request, Response, NextFunction } from "express"; +import { z } from "zod"; +import { db, siteResources } from "@server/db"; +import { eq } from "drizzle-orm"; +import response from "@server/lib/response"; +import HttpCode from "@server/types/HttpCode"; +import createHttpError from "http-errors"; +import logger from "@server/logger"; +import { fromError } from "zod-validation-error"; +import { OpenAPITags, registry } from "@server/openApi"; +import { listSiteResourceAiProviders as listAttachments } from "@server/lib/aiInferenceResource"; + +const listSiteResourceAiProvidersParamsSchema = z.strictObject({ + siteResourceId: z.coerce.number().int().positive() +}); + +export type ListSiteResourceAiProvidersResponse = { + providers: Awaited>; +}; + +registry.registerPath({ + method: "get", + path: "/site-resource/{siteResourceId}/ai-providers", + description: "List AI providers attached to an inference site resource.", + tags: [OpenAPITags.PrivateResource], + request: { + params: listSiteResourceAiProvidersParamsSchema + }, + responses: { + 200: { + description: "Successful response", + content: { + "application/json": { + schema: z.object({ + data: z.record(z.string(), z.any()).nullable(), + success: z.boolean(), + error: z.boolean(), + message: z.string(), + status: z.number() + }) + } + } + } + } +}); + +export async function listSiteResourceAiProviders( + req: Request, + res: Response, + next: NextFunction +): Promise { + try { + const parsedParams = listSiteResourceAiProvidersParamsSchema.safeParse( + req.params + ); + if (!parsedParams.success) { + return next( + createHttpError( + HttpCode.BAD_REQUEST, + fromError(parsedParams.error).toString() + ) + ); + } + + const { siteResourceId } = parsedParams.data; + + const [siteResource] = await db + .select() + .from(siteResources) + .where(eq(siteResources.siteResourceId, siteResourceId)) + .limit(1); + + if (!siteResource) { + return next( + createHttpError(HttpCode.NOT_FOUND, "Site resource not found") + ); + } + + const providers = await listAttachments(siteResourceId); + + return response(res, { + data: { providers }, + success: true, + error: false, + message: "Site resource AI providers retrieved successfully", + status: HttpCode.OK + }); + } catch (error) { + logger.error(error); + return next( + createHttpError(HttpCode.INTERNAL_SERVER_ERROR, "An error occurred") + ); + } +} diff --git a/server/routers/siteResource/removeAiModelFromSiteResource.ts b/server/routers/siteResource/removeAiModelFromSiteResource.ts index 5af2f4d87..eb8d649a2 100644 --- a/server/routers/siteResource/removeAiModelFromSiteResource.ts +++ b/server/routers/siteResource/removeAiModelFromSiteResource.ts @@ -8,6 +8,7 @@ import createHttpError from "http-errors"; import logger from "@server/logger"; import { fromError } from "zod-validation-error"; import { OpenAPITags, registry } from "@server/openApi"; +import { assertSiteAllowlistApiEligible } from "@server/lib/aiInferenceResource"; const removeAiModelFromSiteResourceBodySchema = z.strictObject({ modelId: z.int().positive() @@ -21,7 +22,7 @@ registry.registerPath({ method: "post", path: "/site-resource/{siteResourceId}/ai-models/remove", description: - "Remove a single AI model from a site resource's model restriction allow-list.", + "Remove a single catalog model from an inference site resource allowlist. Requires at least one attached AI provider in allowlist mode.", tags: [OpenAPITags.PrivateResource], request: { params: removeAiModelFromSiteResourceParamsSchema, @@ -96,6 +97,12 @@ export async function removeAiModelFromSiteResource( ); } + const eligibleError = + await assertSiteAllowlistApiEligible(siteResource); + if (eligibleError) { + return next(createHttpError(HttpCode.BAD_REQUEST, eligibleError)); + } + const existingEntry = await db .select() .from(siteResourceAiModels) diff --git a/server/routers/siteResource/removeAiProviderFromSiteResource.ts b/server/routers/siteResource/removeAiProviderFromSiteResource.ts new file mode 100644 index 000000000..65781859c --- /dev/null +++ b/server/routers/siteResource/removeAiProviderFromSiteResource.ts @@ -0,0 +1,164 @@ +import { Request, Response, NextFunction } from "express"; +import { z } from "zod"; +import { db, siteResources } from "@server/db"; +import { eq } from "drizzle-orm"; +import response from "@server/lib/response"; +import HttpCode from "@server/types/HttpCode"; +import createHttpError from "http-errors"; +import logger from "@server/logger"; +import { fromError } from "zod-validation-error"; +import { OpenAPITags, registry } from "@server/openApi"; +import { + isInferenceFieldsError, + listSiteResourceAiProviders, + resolveProviderAttachments, + setSiteResourceAiProviders +} from "@server/lib/aiInferenceResource"; + +const removeAiProviderFromSiteResourceBodySchema = z.strictObject({ + providerId: z.number().int().positive() +}); + +const removeAiProviderFromSiteResourceParamsSchema = z.strictObject({ + siteResourceId: z.coerce.number().int().positive() +}); + +registry.registerPath({ + method: "post", + path: "/site-resource/{siteResourceId}/ai-providers/remove", + description: + "Remove an AI provider attachment from an inference site resource. At least one provider must remain.", + tags: [OpenAPITags.PrivateResource], + request: { + params: removeAiProviderFromSiteResourceParamsSchema, + body: { + content: { + "application/json": { + schema: removeAiProviderFromSiteResourceBodySchema + } + } + } + }, + responses: { + 200: { + description: "Successful response", + content: { + "application/json": { + schema: z.object({ + data: z.record(z.string(), z.any()).nullable(), + success: z.boolean(), + error: z.boolean(), + message: z.string(), + status: z.number() + }) + } + } + } + } +}); + +export async function removeAiProviderFromSiteResource( + req: Request, + res: Response, + next: NextFunction +): Promise { + try { + const parsedBody = + removeAiProviderFromSiteResourceBodySchema.safeParse(req.body); + if (!parsedBody.success) { + return next( + createHttpError( + HttpCode.BAD_REQUEST, + fromError(parsedBody.error).toString() + ) + ); + } + + const { providerId } = parsedBody.data; + + const parsedParams = + removeAiProviderFromSiteResourceParamsSchema.safeParse(req.params); + if (!parsedParams.success) { + return next( + createHttpError( + HttpCode.BAD_REQUEST, + fromError(parsedParams.error).toString() + ) + ); + } + + const { siteResourceId } = parsedParams.data; + + const [siteResource] = await db + .select() + .from(siteResources) + .where(eq(siteResources.siteResourceId, siteResourceId)) + .limit(1); + + if (!siteResource) { + return next( + createHttpError(HttpCode.NOT_FOUND, "Site resource not found") + ); + } + + if (siteResource.mode !== "inference") { + return next( + createHttpError( + HttpCode.BAD_REQUEST, + "AI providers can only be attached to inference-mode resources" + ) + ); + } + + const existing = await listSiteResourceAiProviders(siteResourceId); + const found = existing.find((a) => a.providerId === providerId); + if (!found) { + return next( + createHttpError( + HttpCode.NOT_FOUND, + "AI provider is not attached to this site resource" + ) + ); + } + + const remaining = existing + .filter((a) => a.providerId !== providerId) + .map((a) => ({ + providerId: a.providerId, + modelAccessMode: a.modelAccessMode + })); + + if (remaining.length === 0) { + return next( + createHttpError( + HttpCode.BAD_REQUEST, + "At least one AI provider is required for inference-mode resources" + ) + ); + } + + const attachments = await resolveProviderAttachments({ + orgId: siteResource.orgId, + attachments: remaining, + requireAtLeastOne: true + }); + if (isInferenceFieldsError(attachments)) { + return next(createHttpError(HttpCode.BAD_REQUEST, attachments.error)); + } + + await setSiteResourceAiProviders(siteResourceId, attachments); + + return response(res, { + data: {}, + success: true, + error: false, + message: "AI provider removed from site resource successfully", + status: HttpCode.OK + }); + } catch (error) { + logger.error(error); + return next( + createHttpError(HttpCode.INTERNAL_SERVER_ERROR, "An error occurred") + ); + } +} diff --git a/server/routers/siteResource/setSiteResourceAiModels.ts b/server/routers/siteResource/setSiteResourceAiModels.ts index 1366187af..391634553 100644 --- a/server/routers/siteResource/setSiteResourceAiModels.ts +++ b/server/routers/siteResource/setSiteResourceAiModels.ts @@ -1,18 +1,17 @@ import { Request, Response, NextFunction } from "express"; import { z } from "zod"; -import { - db, - siteResources, - siteResourceAiModels, - aiModels -} from "@server/db"; -import { eq, and, inArray } from "drizzle-orm"; +import { db, siteResources, siteResourceAiModels } from "@server/db"; +import { eq } from "drizzle-orm"; import response from "@server/lib/response"; import HttpCode from "@server/types/HttpCode"; import createHttpError from "http-errors"; import logger from "@server/logger"; import { fromError } from "zod-validation-error"; import { OpenAPITags, registry } from "@server/openApi"; +import { + assertSiteAllowlistApiEligible, + assertModelsBelongToSiteAllowlistProviders +} from "@server/lib/aiInferenceResource"; const setSiteResourceAiModelsBodySchema = z.strictObject({ modelIds: z.array(z.int().positive()) @@ -26,7 +25,7 @@ registry.registerPath({ method: "post", path: "/site-resource/{siteResourceId}/ai-models", description: - "Set the AI models a site resource is restricted to. This replaces all existing restrictions. Pass an empty array to remove the restriction (allow every enabled model on the linked provider).", + "Replace the allowlist of catalog models for an inference site resource. Requires at least one attached AI provider in allowlist mode. Models must belong to a provider attached in allowlist mode. An empty array denies all models.", tags: [OpenAPITags.PrivateResource], request: { params: setSiteResourceAiModelsParamsSchema, @@ -102,52 +101,33 @@ export async function setSiteResourceAiModels( ); } - if (modelIds.length > 0) { - if (!siteResource.aiProviderId) { - return next( - createHttpError( - HttpCode.BAD_REQUEST, - "Site resource has no AI provider linked" - ) - ); - } + const eligibleError = + await assertSiteAllowlistApiEligible(siteResource); + if (eligibleError) { + return next(createHttpError(HttpCode.BAD_REQUEST, eligibleError)); + } - const validModels = await db - .select({ modelId: aiModels.modelId }) - .from(aiModels) - .where( - and( - inArray(aiModels.modelId, modelIds), - eq(aiModels.providerId, siteResource.aiProviderId) - ) - ); - - if (validModels.length !== new Set(modelIds).size) { - return next( - createHttpError( - HttpCode.BAD_REQUEST, - "One or more model IDs do not exist or do not belong to this site resource's AI provider" - ) - ); - } + const modelError = await assertModelsBelongToSiteAllowlistProviders({ + orgId: siteResource.orgId, + siteResourceId, + modelIds + }); + if (modelError) { + return next(createHttpError(HttpCode.BAD_REQUEST, modelError)); } await db.transaction(async (trx) => { await trx .delete(siteResourceAiModels) - .where( - eq(siteResourceAiModels.siteResourceId, siteResourceId) - ); + .where(eq(siteResourceAiModels.siteResourceId, siteResourceId)); if (modelIds.length > 0) { - await trx - .insert(siteResourceAiModels) - .values( - modelIds.map((modelId) => ({ - siteResourceId, - modelId - })) - ); + await trx.insert(siteResourceAiModels).values( + modelIds.map((modelId) => ({ + siteResourceId, + modelId + })) + ); } }); diff --git a/server/routers/siteResource/setSiteResourceAiProviders.ts b/server/routers/siteResource/setSiteResourceAiProviders.ts new file mode 100644 index 000000000..93bf8a866 --- /dev/null +++ b/server/routers/siteResource/setSiteResourceAiProviders.ts @@ -0,0 +1,139 @@ +import { Request, Response, NextFunction } from "express"; +import { z } from "zod"; +import { db, siteResources } from "@server/db"; +import { eq } from "drizzle-orm"; +import response from "@server/lib/response"; +import HttpCode from "@server/types/HttpCode"; +import createHttpError from "http-errors"; +import logger from "@server/logger"; +import { fromError } from "zod-validation-error"; +import { OpenAPITags, registry } from "@server/openApi"; +import { + isInferenceFieldsError, + resolveProviderAttachments, + resourceAiProviderAttachmentSchema, + setSiteResourceAiProviders as replaceAttachments +} from "@server/lib/aiInferenceResource"; + +const setSiteResourceAiProvidersBodySchema = z.strictObject({ + providers: z.array(resourceAiProviderAttachmentSchema) +}); + +const setSiteResourceAiProvidersParamsSchema = z.strictObject({ + siteResourceId: z.coerce.number().int().positive() +}); + +registry.registerPath({ + method: "post", + path: "/site-resource/{siteResourceId}/ai-providers", + description: + "Replace the AI providers attached to an inference site resource. At least one provider is required. At most one may use passthrough mode.", + tags: [OpenAPITags.PrivateResource], + request: { + params: setSiteResourceAiProvidersParamsSchema, + body: { + content: { + "application/json": { + schema: setSiteResourceAiProvidersBodySchema + } + } + } + }, + responses: { + 200: { + description: "Successful response", + content: { + "application/json": { + schema: z.object({ + data: z.record(z.string(), z.any()).nullable(), + success: z.boolean(), + error: z.boolean(), + message: z.string(), + status: z.number() + }) + } + } + } + } +}); + +export async function setSiteResourceAiProviders( + req: Request, + res: Response, + next: NextFunction +): Promise { + try { + const parsedBody = setSiteResourceAiProvidersBodySchema.safeParse( + req.body + ); + if (!parsedBody.success) { + return next( + createHttpError( + HttpCode.BAD_REQUEST, + fromError(parsedBody.error).toString() + ) + ); + } + + const { providers } = parsedBody.data; + + const parsedParams = setSiteResourceAiProvidersParamsSchema.safeParse( + req.params + ); + if (!parsedParams.success) { + return next( + createHttpError( + HttpCode.BAD_REQUEST, + fromError(parsedParams.error).toString() + ) + ); + } + + const { siteResourceId } = parsedParams.data; + + const [siteResource] = await db + .select() + .from(siteResources) + .where(eq(siteResources.siteResourceId, siteResourceId)) + .limit(1); + + if (!siteResource) { + return next( + createHttpError(HttpCode.NOT_FOUND, "Site resource not found") + ); + } + + if (siteResource.mode !== "inference") { + return next( + createHttpError( + HttpCode.BAD_REQUEST, + "AI providers can only be attached to inference-mode resources" + ) + ); + } + + const attachments = await resolveProviderAttachments({ + orgId: siteResource.orgId, + attachments: providers, + requireAtLeastOne: true + }); + if (isInferenceFieldsError(attachments)) { + return next(createHttpError(HttpCode.BAD_REQUEST, attachments.error)); + } + + await replaceAttachments(siteResourceId, attachments); + + return response(res, { + data: {}, + success: true, + error: false, + message: "AI providers set for site resource successfully", + status: HttpCode.CREATED + }); + } catch (error) { + logger.error(error); + return next( + createHttpError(HttpCode.INTERNAL_SERVER_ERROR, "An error occurred") + ); + } +} diff --git a/server/routers/siteResource/updateSiteResource.ts b/server/routers/siteResource/updateSiteResource.ts index caf524804..e99d60e1e 100644 --- a/server/routers/siteResource/updateSiteResource.ts +++ b/server/routers/siteResource/updateSiteResource.ts @@ -29,6 +29,7 @@ import { NextFunction, Request, Response } from "express"; import createHttpError from "http-errors"; import { z } from "zod"; import { fromError } from "zod-validation-error"; +import { clearSiteResourceAiConfig } from "@server/lib/aiInferenceResource"; const updateSiteResourceParamsSchema = z.strictObject({ siteResourceId: z.coerce.number().int().positive() @@ -78,16 +79,7 @@ const updateSiteResourceSchema = z authDaemonMode: z.enum(["site", "remote", "native"]).optional(), pamMode: z.enum(["passthrough", "push"]).optional(), domainId: z.string().optional(), - subdomain: z.string().optional(), - aiProviderId: z - .number() - .int() - .positive() - .nullable() - .optional() - .describe( - "For inference-mode site resources: the AI provider this resource proxies chat completions to. Set to null to unlink." - ) + subdomain: z.string().optional() }) .strict() .refine( @@ -341,8 +333,7 @@ export async function updateSiteResource( authDaemonMode, pamMode, domainId, - subdomain, - aiProviderId + subdomain } = parsedBody.data; // Backward compatibility: merge deprecated siteId into siteIds array @@ -608,12 +599,19 @@ export async function updateSiteResource( networkId: mode === "inference" ? null : undefined, requiresExitNodeConnection: mode !== undefined ? mode === "inference" : undefined, - aiProviderId: aiProviderId, ...sshPamSet }) .where(and(eq(siteResources.siteResourceId, siteResourceId))) .returning(); + const effectiveMode = mode ?? existingSiteResource.mode; + if ( + existingSiteResource.mode === "inference" && + effectiveMode !== "inference" + ) { + await clearSiteResourceAiConfig(siteResourceId, trx); + } + //////////////////// update the associations //////////////////// if (mode === "inference") { diff --git a/src/app/[orgId]/settings/ai-providers/[providerId]/authentication/page.tsx b/src/app/[orgId]/settings/ai-providers/[providerId]/authentication/page.tsx index f054c7901..02d1d2142 100644 --- a/src/app/[orgId]/settings/ai-providers/[providerId]/authentication/page.tsx +++ b/src/app/[orgId]/settings/ai-providers/[providerId]/authentication/page.tsx @@ -66,8 +66,6 @@ export default function AiProviderAuthenticationPage() { authType: (provider.authType as "bearer" | null) ?? "bearer", routingMode: (provider.routingMode as "url" | "target") ?? "url", skipTlsVerification: provider.skipTlsVerification, - budgetAmount: provider.budgetAmount, - budgetUnit: provider.budgetUnit as "usd" | "tokens" | null, enabled: provider.enabled } }); @@ -96,8 +94,6 @@ export default function AiProviderAuthenticationPage() { authType: (updated.authType as "bearer" | null) ?? "bearer", routingMode: (updated.routingMode as "url" | "target") ?? "url", skipTlsVerification: updated.skipTlsVerification, - budgetAmount: updated.budgetAmount, - budgetUnit: updated.budgetUnit as "usd" | "tokens" | null, enabled: updated.enabled }); toast({ diff --git a/src/app/[orgId]/settings/ai-providers/[providerId]/network/page.tsx b/src/app/[orgId]/settings/ai-providers/[providerId]/network/page.tsx index 90ef74d10..2aa06d01d 100644 --- a/src/app/[orgId]/settings/ai-providers/[providerId]/network/page.tsx +++ b/src/app/[orgId]/settings/ai-providers/[providerId]/network/page.tsx @@ -75,8 +75,6 @@ export default function AiProviderNetworkPage() { authType: (provider.authType as "bearer" | null) ?? "bearer", routingMode: (provider.routingMode as "url" | "target") ?? "url", skipTlsVerification: provider.skipTlsVerification, - budgetAmount: provider.budgetAmount, - budgetUnit: provider.budgetUnit as "usd" | "tokens" | null, enabled: provider.enabled } }); @@ -120,8 +118,6 @@ export default function AiProviderNetworkPage() { authType: (updated.authType as "bearer" | null) ?? "bearer", routingMode: (updated.routingMode as "url" | "target") ?? "url", skipTlsVerification: updated.skipTlsVerification, - budgetAmount: updated.budgetAmount, - budgetUnit: updated.budgetUnit as "usd" | "tokens" | null, enabled: updated.enabled }); diff --git a/src/app/[orgId]/settings/ai-providers/create/page.tsx b/src/app/[orgId]/settings/ai-providers/create/page.tsx index ec6eee6c7..3111f0b3e 100644 --- a/src/app/[orgId]/settings/ai-providers/create/page.tsx +++ b/src/app/[orgId]/settings/ai-providers/create/page.tsx @@ -79,8 +79,6 @@ export default function CreateAiProviderPage() { authType: "bearer", routingMode: "url", skipTlsVerification: false, - budgetAmount: null, - budgetUnit: null, enabled: true } }); diff --git a/src/lib/aiProviderFormSchema.ts b/src/lib/aiProviderFormSchema.ts index 8861dcdbf..26245620c 100644 --- a/src/lib/aiProviderFormSchema.ts +++ b/src/lib/aiProviderFormSchema.ts @@ -26,8 +26,6 @@ export const aiProviderFormSchema = z authType: z.enum(["bearer"]).optional().nullable(), routingMode: z.enum(["url", "target"]).optional(), skipTlsVerification: z.boolean().optional(), - budgetAmount: z.number().positive().nullable().optional(), - budgetUnit: z.enum(["usd", "tokens"]).optional().nullable(), enabled: z.boolean().optional() }) .superRefine((data, ctx) => { @@ -78,20 +76,6 @@ export const aiProviderFormSchema = z path: ["authType"] }); } - - const hasAmount = - data.budgetAmount !== undefined && data.budgetAmount !== null; - const hasUnit = - data.budgetUnit !== undefined && data.budgetUnit !== null; - - if (hasAmount !== hasUnit) { - ctx.addIssue({ - code: "custom", - message: - "budgetAmount and budgetUnit must both be set or both omitted", - path: hasAmount ? ["budgetUnit"] : ["budgetAmount"] - }); - } }); export type AiProviderFormValues = z.infer; @@ -145,11 +129,6 @@ export function toAiProviderCreatePayload(values: AiProviderFormValues) { ? upstreamRaw : null; - const hasBudget = - values.budgetAmount !== undefined && - values.budgetAmount !== null && - values.budgetUnit; - return { name: values.name.trim(), type: values.type, @@ -161,8 +140,6 @@ export function toAiProviderCreatePayload(values: AiProviderFormValues) { ? (values.authType ?? "bearer") : (values.authType ?? undefined), skipTlsVerification: values.skipTlsVerification, - budgetAmount: hasBudget ? values.budgetAmount : null, - budgetUnit: hasBudget ? values.budgetUnit : null, enabled: values.enabled ?? true }; } @@ -178,11 +155,6 @@ export function toAiProviderUpdatePayload(values: AiProviderFormValues) { ? upstreamRaw : null; - const hasBudget = - values.budgetAmount !== undefined && - values.budgetAmount !== null && - values.budgetUnit; - const payload: Record = { name: values.name.trim(), routingMode: values.type === "custom" ? routingMode : "url", @@ -192,8 +164,6 @@ export function toAiProviderUpdatePayload(values: AiProviderFormValues) { ? (values.authType ?? "bearer") : (values.authType ?? null), skipTlsVerification: values.skipTlsVerification ?? false, - budgetAmount: hasBudget ? values.budgetAmount : null, - budgetUnit: hasBudget ? values.budgetUnit : null, enabled: values.enabled ?? true };