add crud for adding providers and models to resources

This commit is contained in:
miloschwartz
2026-08-04 15:54:24 -04:00
parent ed8545f8a2
commit e38359c74f
39 changed files with 2343 additions and 438 deletions
+42 -16
View File
@@ -211,14 +211,7 @@ export const resources = pgTable(
authDaemonPort: integer("authDaemonPort").default(22123), authDaemonPort: integer("authDaemonPort").default(22123),
status: varchar("status") status: varchar("status")
.$type<"pending" | "approved">() .$type<"pending" | "approved">()
.default("approved"), .default("approved")
aiProviderId: integer("aiProviderId").references(
() => aiProviders.providerId,
{ onDelete: "set null" }
),
modelAccessMode: varchar("modelAccessMode").$type<
"passthrough" | "catalog" | "allowlist"
>()
}, },
(t) => [ (t) => [
index("idx_resources_fulldomain") 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( export const resourceAiModels = pgTable(
"resourceAiModels", "resourceAiModels",
{ {
@@ -496,18 +506,30 @@ export const siteResources = pgTable(
fullDomain: varchar("fullDomain"), fullDomain: varchar("fullDomain"),
status: varchar("status") status: varchar("status")
.$type<"pending" | "approved">() .$type<"pending" | "approved">()
.default("approved"), .default("approved")
aiProviderId: integer("aiProviderId").references(
() => aiProviders.providerId,
{ onDelete: "set null" }
),
modelAccessMode: varchar("modelAccessMode").$type<
"passthrough" | "catalog" | "allowlist"
>()
}, },
(t) => [index("idx_siteresources_orgid_niceid").on(t.orgId, t.niceId)] (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( export const siteResourceAiModels = pgTable(
"siteResourceAiModels", "siteResourceAiModels",
{ {
@@ -1806,5 +1828,9 @@ export type AiProvider = InferSelectModel<typeof aiProviders>;
export type AiModel = InferSelectModel<typeof aiModels>; export type AiModel = InferSelectModel<typeof aiModels>;
export type AiBudget = InferSelectModel<typeof aiBudgets>; export type AiBudget = InferSelectModel<typeof aiBudgets>;
export type AiBudgetPeriod = InferSelectModel<typeof aiBudgetPeriods>; export type AiBudgetPeriod = InferSelectModel<typeof aiBudgetPeriods>;
export type ResourceAiProvider = InferSelectModel<typeof resourceAiProviders>;
export type SiteResourceAiProvider = InferSelectModel<
typeof siteResourceAiProviders
>;
export type ResourceAiModel = InferSelectModel<typeof resourceAiModels>; export type ResourceAiModel = InferSelectModel<typeof resourceAiModels>;
export type SiteResourceAiModel = InferSelectModel<typeof siteResourceAiModels>; export type SiteResourceAiModel = InferSelectModel<typeof siteResourceAiModels>;
+42 -16
View File
@@ -216,16 +216,26 @@ export const resources = sqliteTable("resources", {
.$type<"site" | "remote" | "native">() .$type<"site" | "remote" | "native">()
.default("site"), .default("site"),
authDaemonPort: integer("authDaemonPort").default(22123), authDaemonPort: integer("authDaemonPort").default(22123),
status: text("status").$type<"pending" | "approved">().default("approved"), status: text("status").$type<"pending" | "approved">().default("approved")
aiProviderId: integer("aiProviderId").references(
() => aiProviders.providerId,
{ onDelete: "set null" }
),
modelAccessMode: text("modelAccessMode").$type<
"passthrough" | "catalog" | "allowlist"
>()
}); });
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( export const resourceAiModels = sqliteTable(
"resourceAiModels", "resourceAiModels",
{ {
@@ -483,16 +493,28 @@ export const siteResources = sqliteTable("siteResources", {
}), }),
subdomain: text("subdomain"), subdomain: text("subdomain"),
fullDomain: text("fullDomain"), fullDomain: text("fullDomain"),
status: text("status").$type<"pending" | "approved">().default("approved"), status: text("status").$type<"pending" | "approved">().default("approved")
aiProviderId: integer("aiProviderId").references(
() => aiProviders.providerId,
{ onDelete: "set null" }
),
modelAccessMode: text("modelAccessMode").$type<
"passthrough" | "catalog" | "allowlist"
>()
}); });
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( export const siteResourceAiModels = sqliteTable(
"siteResourceAiModels", "siteResourceAiModels",
{ {
@@ -1787,5 +1809,9 @@ export type AiProvider = InferSelectModel<typeof aiProviders>;
export type AiModel = InferSelectModel<typeof aiModels>; export type AiModel = InferSelectModel<typeof aiModels>;
export type AiBudget = InferSelectModel<typeof aiBudgets>; export type AiBudget = InferSelectModel<typeof aiBudgets>;
export type AiBudgetPeriod = InferSelectModel<typeof aiBudgetPeriods>; export type AiBudgetPeriod = InferSelectModel<typeof aiBudgetPeriods>;
export type ResourceAiProvider = InferSelectModel<typeof resourceAiProviders>;
export type SiteResourceAiProvider = InferSelectModel<
typeof siteResourceAiProviders
>;
export type ResourceAiModel = InferSelectModel<typeof resourceAiModels>; export type ResourceAiModel = InferSelectModel<typeof resourceAiModels>;
export type SiteResourceAiModel = InferSelectModel<typeof siteResourceAiModels>; export type SiteResourceAiModel = InferSelectModel<typeof siteResourceAiModels>;
+512
View File
@@ -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<typeof modelAccessModeSchema>;
export const resourceAiProviderAttachmentSchema = z.strictObject({
providerId: z.number().int().positive(),
modelAccessMode: modelAccessModeSchema.optional()
});
export type ResourceAiProviderInput = z.infer<
typeof resourceAiProviderAttachmentSchema
>;
export type ResourceAiProviderAttachment = {
providerId: number;
modelAccessMode: ModelAccessMode;
};
export type InferenceFieldsError = {
error: string;
};
export function isInferenceFieldsError(
value: { error: string } | object
): value is InferenceFieldsError {
return "error" in value;
}
function normalizeAttachments(
inputs: ResourceAiProviderInput[]
): ResourceAiProviderAttachment[] {
const byProvider = new Map<number, ModelAccessMode>();
for (const input of inputs) {
byProvider.set(
input.providerId,
input.modelAccessMode ?? "passthrough"
);
}
return [...byProvider.entries()].map(([providerId, modelAccessMode]) => ({
providerId,
modelAccessMode
}));
}
/**
* Validate provider attachments for an org.
* At most one passthrough provider is allowed per resource.
*/
export async function resolveProviderAttachments(input: {
orgId: string;
attachments: ResourceAiProviderInput[];
requireAtLeastOne: boolean;
}): Promise<ResourceAiProviderAttachment[] | InferenceFieldsError> {
const attachments = normalizeAttachments(input.attachments);
if (input.requireAtLeastOne && attachments.length === 0) {
return {
error: "At least one AI provider is required for inference-mode resources"
};
}
const passthroughCount = attachments.filter(
(a) => a.modelAccessMode === "passthrough"
).length;
if (passthroughCount > 1) {
return {
error: "A resource may have at most one AI provider in passthrough mode"
};
}
if (attachments.length === 0) {
return [];
}
const providerIds = attachments.map((a) => a.providerId);
const providers = await db
.select({
providerId: aiProviders.providerId,
orgId: aiProviders.orgId,
enabled: aiProviders.enabled
})
.from(aiProviders)
.where(
and(
inArray(aiProviders.providerId, providerIds),
eq(aiProviders.orgId, input.orgId)
)
);
if (providers.length !== providerIds.length) {
return {
error: "One or more AI providers were not found in this organization"
};
}
const disabled = providers.find((p) => !p.enabled);
if (disabled) {
return {
error: `AI provider with ID ${disabled.providerId} is disabled`
};
}
return attachments;
}
export async function assertInferenceModeAllowsProviderFields(input: {
mode: string;
hasProviderAttachments: boolean;
}): Promise<InferenceFieldsError | null> {
if (input.mode === "inference") {
return null;
}
if (input.hasProviderAttachments) {
return {
error: "AI providers can only be attached to inference-mode resources"
};
}
return null;
}
export async function setPublicResourceAiProviders(
resourceId: number,
attachments: ResourceAiProviderAttachment[],
trx: DbOrTrx = db
): Promise<void> {
await trx
.delete(resourceAiProviders)
.where(eq(resourceAiProviders.resourceId, resourceId));
if (attachments.length > 0) {
await trx.insert(resourceAiProviders).values(
attachments.map((a) => ({
resourceId,
providerId: a.providerId,
modelAccessMode: a.modelAccessMode
}))
);
}
await prunePublicResourceAllowlistToAllowlistProviders(resourceId, trx);
}
export async function setSiteResourceAiProviders(
siteResourceId: number,
attachments: ResourceAiProviderAttachment[],
trx: DbOrTrx = db
): Promise<void> {
await trx
.delete(siteResourceAiProviders)
.where(eq(siteResourceAiProviders.siteResourceId, siteResourceId));
if (attachments.length > 0) {
await trx.insert(siteResourceAiProviders).values(
attachments.map((a) => ({
siteResourceId,
providerId: a.providerId,
modelAccessMode: a.modelAccessMode
}))
);
}
await pruneSiteResourceAllowlistToAllowlistProviders(siteResourceId, trx);
}
export async function clearPublicResourceAiConfig(
resourceId: number,
trx: DbOrTrx = db
): Promise<void> {
await trx
.delete(resourceAiModels)
.where(eq(resourceAiModels.resourceId, resourceId));
await trx
.delete(resourceAiProviders)
.where(eq(resourceAiProviders.resourceId, resourceId));
}
export async function clearSiteResourceAiConfig(
siteResourceId: number,
trx: DbOrTrx = db
): Promise<void> {
await trx
.delete(siteResourceAiModels)
.where(eq(siteResourceAiModels.siteResourceId, siteResourceId));
await trx
.delete(siteResourceAiProviders)
.where(eq(siteResourceAiProviders.siteResourceId, siteResourceId));
}
async function prunePublicResourceAllowlistToAllowlistProviders(
resourceId: number,
trx: DbOrTrx = db
): Promise<void> {
const allowlistProviders = await trx
.select({ providerId: resourceAiProviders.providerId })
.from(resourceAiProviders)
.where(
and(
eq(resourceAiProviders.resourceId, resourceId),
eq(resourceAiProviders.modelAccessMode, "allowlist")
)
);
if (allowlistProviders.length === 0) {
await trx
.delete(resourceAiModels)
.where(eq(resourceAiModels.resourceId, resourceId));
return;
}
const validModels = await trx
.select({ modelId: aiModels.modelId })
.from(aiModels)
.where(
inArray(
aiModels.providerId,
allowlistProviders.map((p) => p.providerId)
)
);
const validIds = validModels.map((m) => m.modelId);
const existing = await trx
.select({ modelId: resourceAiModels.modelId })
.from(resourceAiModels)
.where(eq(resourceAiModels.resourceId, resourceId));
const toRemove = existing
.map((e) => e.modelId)
.filter((id) => !validIds.includes(id));
if (toRemove.length > 0) {
await trx
.delete(resourceAiModels)
.where(
and(
eq(resourceAiModels.resourceId, resourceId),
inArray(resourceAiModels.modelId, toRemove)
)
);
}
}
async function pruneSiteResourceAllowlistToAllowlistProviders(
siteResourceId: number,
trx: DbOrTrx = db
): Promise<void> {
const allowlistProviders = await trx
.select({ providerId: siteResourceAiProviders.providerId })
.from(siteResourceAiProviders)
.where(
and(
eq(siteResourceAiProviders.siteResourceId, siteResourceId),
eq(siteResourceAiProviders.modelAccessMode, "allowlist")
)
);
if (allowlistProviders.length === 0) {
await trx
.delete(siteResourceAiModels)
.where(eq(siteResourceAiModels.siteResourceId, siteResourceId));
return;
}
const validModels = await trx
.select({ modelId: aiModels.modelId })
.from(aiModels)
.where(
inArray(
aiModels.providerId,
allowlistProviders.map((p) => p.providerId)
)
);
const validIds = validModels.map((m) => m.modelId);
const existing = await trx
.select({ modelId: siteResourceAiModels.modelId })
.from(siteResourceAiModels)
.where(eq(siteResourceAiModels.siteResourceId, siteResourceId));
const toRemove = existing
.map((e) => e.modelId)
.filter((id) => !validIds.includes(id));
if (toRemove.length > 0) {
await trx
.delete(siteResourceAiModels)
.where(
and(
eq(siteResourceAiModels.siteResourceId, siteResourceId),
inArray(siteResourceAiModels.modelId, toRemove)
)
);
}
}
export async function listPublicResourceAiProviders(resourceId: number) {
return db
.select({
providerId: resourceAiProviders.providerId,
modelAccessMode: resourceAiProviders.modelAccessMode,
name: aiProviders.name,
type: aiProviders.type,
enabled: aiProviders.enabled
})
.from(resourceAiProviders)
.innerJoin(
aiProviders,
eq(resourceAiProviders.providerId, aiProviders.providerId)
)
.where(eq(resourceAiProviders.resourceId, resourceId));
}
export async function listSiteResourceAiProviders(siteResourceId: number) {
return db
.select({
providerId: siteResourceAiProviders.providerId,
modelAccessMode: siteResourceAiProviders.modelAccessMode,
name: aiProviders.name,
type: aiProviders.type,
enabled: aiProviders.enabled
})
.from(siteResourceAiProviders)
.innerJoin(
aiProviders,
eq(siteResourceAiProviders.providerId, aiProviders.providerId)
)
.where(eq(siteResourceAiProviders.siteResourceId, siteResourceId));
}
/**
* Allowlist APIs require an inference resource with at least one
* attached provider in allowlist mode.
*/
export async function assertPublicAllowlistApiEligible(resource: {
resourceId: number;
mode: string;
}): Promise<string | null> {
if (resource.mode !== "inference") {
return "AI model allowlists are only supported on inference-mode resources";
}
const [row] = await db
.select({ providerId: resourceAiProviders.providerId })
.from(resourceAiProviders)
.where(
and(
eq(resourceAiProviders.resourceId, resource.resourceId),
eq(resourceAiProviders.modelAccessMode, "allowlist")
)
)
.limit(1);
if (!row) {
return "Attach at least one AI provider with modelAccessMode=allowlist before managing allowed models";
}
return null;
}
export async function assertSiteAllowlistApiEligible(siteResource: {
siteResourceId: number;
mode: string;
}): Promise<string | null> {
if (siteResource.mode !== "inference") {
return "AI model allowlists are only supported on inference-mode resources";
}
const [row] = await db
.select({ providerId: siteResourceAiProviders.providerId })
.from(siteResourceAiProviders)
.where(
and(
eq(
siteResourceAiProviders.siteResourceId,
siteResource.siteResourceId
),
eq(siteResourceAiProviders.modelAccessMode, "allowlist")
)
)
.limit(1);
if (!row) {
return "Attach at least one AI provider with modelAccessMode=allowlist before managing allowed models";
}
return null;
}
/**
* Models must belong to providers attached to this resource in allowlist mode,
* and those providers must belong to the resource's org.
*/
export async function assertModelsBelongToPublicAllowlistProviders(input: {
orgId: string;
resourceId: number;
modelIds: number[];
}): Promise<string | null> {
const uniqueIds = [...new Set(input.modelIds)];
if (uniqueIds.length === 0) {
return null;
}
const allowlistProviders = await db
.select({ providerId: resourceAiProviders.providerId })
.from(resourceAiProviders)
.innerJoin(
aiProviders,
eq(resourceAiProviders.providerId, aiProviders.providerId)
)
.where(
and(
eq(resourceAiProviders.resourceId, input.resourceId),
eq(resourceAiProviders.modelAccessMode, "allowlist"),
eq(aiProviders.orgId, input.orgId)
)
);
if (allowlistProviders.length === 0) {
return "No allowlist AI providers are attached to this resource";
}
const validModels = await db
.select({ modelId: aiModels.modelId })
.from(aiModels)
.innerJoin(aiProviders, eq(aiModels.providerId, aiProviders.providerId))
.where(
and(
inArray(aiModels.modelId, uniqueIds),
inArray(
aiModels.providerId,
allowlistProviders.map((p) => p.providerId)
),
eq(aiProviders.orgId, input.orgId)
)
);
if (validModels.length !== uniqueIds.length) {
return "One or more model IDs do not exist or do not belong to an allowlist provider on this resource";
}
return null;
}
export async function assertModelsBelongToSiteAllowlistProviders(input: {
orgId: string;
siteResourceId: number;
modelIds: number[];
}): Promise<string | null> {
const uniqueIds = [...new Set(input.modelIds)];
if (uniqueIds.length === 0) {
return null;
}
const allowlistProviders = await db
.select({ providerId: siteResourceAiProviders.providerId })
.from(siteResourceAiProviders)
.innerJoin(
aiProviders,
eq(siteResourceAiProviders.providerId, aiProviders.providerId)
)
.where(
and(
eq(
siteResourceAiProviders.siteResourceId,
input.siteResourceId
),
eq(siteResourceAiProviders.modelAccessMode, "allowlist"),
eq(aiProviders.orgId, input.orgId)
)
);
if (allowlistProviders.length === 0) {
return "No allowlist AI providers are attached to this site resource";
}
const validModels = await db
.select({ modelId: aiModels.modelId })
.from(aiModels)
.innerJoin(aiProviders, eq(aiModels.providerId, aiProviders.providerId))
.where(
and(
inArray(aiModels.modelId, uniqueIds),
inArray(
aiModels.providerId,
allowlistProviders.map((p) => p.providerId)
),
eq(aiProviders.orgId, input.orgId)
)
);
if (validModels.length !== uniqueIds.length) {
return "One or more model IDs do not exist or do not belong to an allowlist provider on this site resource";
}
return null;
}
+7 -3
View File
@@ -1,4 +1,4 @@
import { db, targetHealthCheck, domains, aiProviders } from "@server/db"; import { db, targetHealthCheck, domains, aiProviders, resourceAiProviders } from "@server/db";
import { import {
and, and,
eq, eq,
@@ -214,7 +214,7 @@ export async function getTraefikConfig(
// central AI gateway), so they can't be reached via the targets->sites // 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. // join above - query them separately and include them on every exit node.
const inferenceResources = await db const inferenceResources = await db
.select({ .selectDistinct({
resourceId: resources.resourceId, resourceId: resources.resourceId,
resourceName: resources.name, resourceName: resources.name,
fullDomain: resources.fullDomain, fullDomain: resources.fullDomain,
@@ -226,9 +226,13 @@ export async function getTraefikConfig(
preferWildcardCert: domains.preferWildcardCert preferWildcardCert: domains.preferWildcardCert
}) })
.from(resources) .from(resources)
.innerJoin(
resourceAiProviders,
eq(resources.resourceId, resourceAiProviders.resourceId)
)
.innerJoin( .innerJoin(
aiProviders, aiProviders,
eq(resources.aiProviderId, aiProviders.providerId) eq(resourceAiProviders.providerId, aiProviders.providerId)
) )
.leftJoin(domains, eq(domains.domainId, resources.domainId)) .leftJoin(domains, eq(domains.domainId, resources.domainId))
.where( .where(
+18 -5
View File
@@ -42,7 +42,9 @@ import {
siteResources, siteResources,
Target, Target,
targets, targets,
aiProviders aiProviders,
resourceAiProviders,
siteResourceAiProviders
} from "@server/db"; } from "@server/db";
import { import {
sanitize, sanitize,
@@ -402,7 +404,7 @@ export async function getTraefikConfig(
// so they can't be reached via the joins above - query them separately // so they can't be reached via the joins above - query them separately
// and include them on every exit node. // and include them on every exit node.
const inferenceResources = await db const inferenceResources = await db
.select({ .selectDistinct({
resourceId: resources.resourceId, resourceId: resources.resourceId,
fullDomain: resources.fullDomain, fullDomain: resources.fullDomain,
ssl: resources.ssl, ssl: resources.ssl,
@@ -414,9 +416,13 @@ export async function getTraefikConfig(
preferWildcardCert: domains.preferWildcardCert preferWildcardCert: domains.preferWildcardCert
}) })
.from(resources) .from(resources)
.innerJoin(
resourceAiProviders,
eq(resources.resourceId, resourceAiProviders.resourceId)
)
.innerJoin( .innerJoin(
aiProviders, aiProviders,
eq(resources.aiProviderId, aiProviders.providerId) eq(resourceAiProviders.providerId, aiProviders.providerId)
) )
.leftJoin(domains, eq(domains.domainId, resources.domainId)) .leftJoin(domains, eq(domains.domainId, resources.domainId))
.where( .where(
@@ -435,16 +441,23 @@ export async function getTraefikConfig(
}[] = []; }[] = [];
if (build == "enterprise") { if (build == "enterprise") {
siteResourcesInference = await db siteResourcesInference = await db
.select({ .selectDistinct({
siteResourceId: siteResources.siteResourceId, siteResourceId: siteResources.siteResourceId,
alias: siteResources.alias, alias: siteResources.alias,
ssl: siteResources.ssl, ssl: siteResources.ssl,
enabled: siteResources.enabled enabled: siteResources.enabled
}) })
.from(siteResources) .from(siteResources)
.innerJoin(
siteResourceAiProviders,
eq(
siteResources.siteResourceId,
siteResourceAiProviders.siteResourceId
)
)
.innerJoin( .innerJoin(
aiProviders, aiProviders,
eq(siteResources.aiProviderId, aiProviders.providerId) eq(siteResourceAiProviders.providerId, aiProviders.providerId)
) )
.where( .where(
and( and(
+218 -70
View File
@@ -8,8 +8,10 @@ import {
db, db,
exitNodes, exitNodes,
resourceAiModels, resourceAiModels,
resourceAiProviders,
resources, resources,
siteResourceAiModels, siteResourceAiModels,
siteResourceAiProviders,
siteResources, siteResources,
users users
} from "@server/db"; } from "@server/db";
@@ -21,7 +23,6 @@ import {
AiProviderType, AiProviderType,
resolveAiProviderConfig resolveAiProviderConfig
} from "@server/lib/aiProviderDefaults"; } from "@server/lib/aiProviderDefaults";
import { verifyResourceAccessToken } from "@server/auth/verifyResourceAccessToken";
import { import {
SESSION_COOKIE_NAME, SESSION_COOKIE_NAME,
validateSessionToken validateSessionToken
@@ -31,6 +32,7 @@ import { isIpInCidr } from "@server/lib/ip";
import { localCache } from "@server/lib/cache"; import { localCache } from "@server/lib/cache";
import logger from "@server/logger"; import logger from "@server/logger";
import HttpCode from "@server/types/HttpCode"; 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 // 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 // doesn't hit the database on every single request. None of this is
@@ -85,14 +87,23 @@ async function findClientByIp(ip: string): Promise<CachedClient> {
return result; return result;
} }
type ProviderAttachment = {
provider: AiProvider;
modelAccessMode: ModelAccessMode;
};
type ResolvedTarget = { type ResolvedTarget = {
resourceId: number | null; resourceId: number | null;
orgId: string | null; orgId: string | null;
provider: AiProvider; attachments: ProviderAttachment[];
// null = no restriction; every enabled model on the provider is allowed // model IDs on this resource's allowlist (resource-wide)
allowedModelIds: number[] | null; allowedModelIds: number[];
}; };
type ProviderSelection =
| { ok: true; provider: AiProvider }
| { ok: false; status: number; message: string };
export type RequestUser = { export type RequestUser = {
userId: string; userId: string;
username: string; username: string;
@@ -184,88 +195,245 @@ async function resolveTarget(host: string): Promise<ResolvedTarget | null> {
const [resourceRow] = await db const [resourceRow] = await db
.select({ .select({
resourceId: resources.resourceId, resourceId: resources.resourceId,
orgId: resources.orgId, orgId: resources.orgId
provider: aiProviders
}) })
.from(resources) .from(resources)
.innerJoin(
aiProviders,
eq(resources.aiProviderId, aiProviders.providerId)
)
.where( .where(
and( and(
eq(resources.fullDomain, host), eq(resources.fullDomain, host),
eq(resources.mode, "inference"), eq(resources.mode, "inference"),
eq(resources.enabled, true), eq(resources.enabled, true)
eq(aiProviders.enabled, true)
) )
) )
.limit(1); .limit(1);
if (resourceRow) { if (resourceRow) {
const restrictions = await db const attachmentRows = await db
.select({ modelId: resourceAiModels.modelId }) .select({
.from(resourceAiModels) modelAccessMode: resourceAiProviders.modelAccessMode,
.where(eq(resourceAiModels.resourceId, resourceRow.resourceId)); 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 { return {
resourceId: resourceRow.resourceId, resourceId: resourceRow.resourceId,
orgId: resourceRow.orgId, orgId: resourceRow.orgId,
provider: resourceRow.provider, attachments: attachmentRows.map((a) => ({
allowedModelIds: restrictions.length provider: a.provider,
? restrictions.map((r) => r.modelId) modelAccessMode: a.modelAccessMode as ModelAccessMode
: null })),
allowedModelIds
}; };
} }
const [siteResourceRow] = await db const [siteResourceRow] = await db
.select({ .select({
siteResourceId: siteResources.siteResourceId, siteResourceId: siteResources.siteResourceId,
orgId: siteResources.orgId, orgId: siteResources.orgId
provider: aiProviders
}) })
.from(siteResources) .from(siteResources)
.innerJoin(
aiProviders,
eq(siteResources.aiProviderId, aiProviders.providerId)
)
.where( .where(
and( and(
eq(siteResources.alias, host), eq(siteResources.alias, host),
eq(siteResources.mode, "inference"), eq(siteResources.mode, "inference"),
eq(siteResources.enabled, true), eq(siteResources.enabled, true)
eq(aiProviders.enabled, true)
) )
) )
.limit(1); .limit(1);
if (siteResourceRow) { if (siteResourceRow) {
const restrictions = await db const attachmentRows = await db
.select({ modelId: siteResourceAiModels.modelId }) .select({
.from(siteResourceAiModels) modelAccessMode: siteResourceAiProviders.modelAccessMode,
provider: aiProviders
})
.from(siteResourceAiProviders)
.innerJoin(
aiProviders,
eq(siteResourceAiProviders.providerId, aiProviders.providerId)
)
.where( .where(
eq( and(
siteResourceAiModels.siteResourceId, eq(
siteResourceRow.siteResourceId 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 { return {
// siteResources have no per-user auth/policy stack today (see // siteResources have no per-user auth/policy stack today (see
// the routing comment in getTraefikConfig.ts), so there's no // the routing comment in getTraefikConfig.ts), so there's no
// resource access token scope to validate a user token against. // resource access token scope to validate a user token against.
resourceId: null, resourceId: null,
orgId: siteResourceRow.orgId, orgId: siteResourceRow.orgId,
provider: siteResourceRow.provider, attachments: attachmentRows.map((a) => ({
allowedModelIds: restrictions.length provider: a.provider,
? restrictions.map((r) => r.modelId) modelAccessMode: a.modelAccessMode as ModelAccessMode
: null })),
allowedModelIds
}; };
} }
return null; return null;
} }
async function providerMatchesModel(
attachment: ProviderAttachment,
requestedModel: string,
allowedModelIds: number[]
): Promise<boolean> {
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<ProviderSelection> {
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 // Generic OpenAI-wire-compatible passthrough. Anthropic's native API uses a
// different path/schema; everything else here is OpenAI-compatible today. // different path/schema; everything else here is OpenAI-compatible today.
function getCompletionsPath(type: AiProviderType): string { 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 // Best-effort identity resolution - not yet enforced, but lets us
// start making per-user access decisions (e.g. model/role-based // start making per-user access decisions (e.g. model/role-based
@@ -311,39 +479,19 @@ export async function chatCompletions(
const requestedModel = const requestedModel =
typeof req.body?.model === "string" ? req.body.model : undefined; typeof req.body?.model === "string" ? req.body.model : undefined;
if (allowedModelIds) { const selection = await selectProvider(
if (!requestedModel) { attachments,
return res.status(HttpCode.FORBIDDEN).json({ allowedModelIds,
error: { requestedModel
message: );
"This resource restricts access to specific models; a model must be specified" if (!selection.ok) {
} return res.status(selection.status).json({
}); error: { message: selection.message }
} });
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 { provider } = selection;
if (!provider.apiKey) { if (!provider.apiKey) {
return res.status(HttpCode.INTERNAL_SERVER_ERROR).json({ return res.status(HttpCode.INTERNAL_SERVER_ERROR).json({
error: { message: "AI provider has no API key configured" } error: { message: "AI provider has no API key configured" }
+6 -19
View File
@@ -9,26 +9,16 @@ import { fromError } from "zod-validation-error";
import { OpenAPITags, registry } from "@server/openApi"; import { OpenAPITags, registry } from "@server/openApi";
import { and, eq } from "drizzle-orm"; import { and, eq } from "drizzle-orm";
import type { CreateOrEditAiModelResponse } from "@server/routers/aiProvider/types"; import type { CreateOrEditAiModelResponse } from "@server/routers/aiProvider/types";
import {
aiBudgetUnitSchema,
refineBudgetFields
} from "@server/routers/aiProvider/validation";
const paramsSchema = z.strictObject({ const paramsSchema = z.strictObject({
providerId: z.coerce.number().int().positive() providerId: z.coerce.number().int().positive()
}); });
const bodySchema = z const bodySchema = z.strictObject({
.strictObject({ modelKey: z.string().nonempty(),
modelKey: z.string().nonempty(), name: z.string().nonempty(),
name: z.string().nonempty(), enabled: z.boolean().optional()
budgetAmount: z.number().positive().optional().nullable(), });
budgetUnit: aiBudgetUnitSchema.optional().nullable(),
enabled: z.boolean().optional()
})
.superRefine((data, ctx) => {
refineBudgetFields(data, ctx);
});
registry.registerPath({ registry.registerPath({
method: "put", method: "put",
@@ -79,8 +69,7 @@ export async function createAiModel(
} }
const { providerId } = parsedParams.data; const { providerId } = parsedParams.data;
const { modelKey, name, budgetAmount, budgetUnit, enabled } = const { modelKey, name, enabled } = parsedBody.data;
parsedBody.data;
const [provider] = const [provider] =
req.aiProvider && req.aiProvider.providerId === providerId req.aiProvider && req.aiProvider.providerId === providerId
@@ -127,8 +116,6 @@ export async function createAiModel(
providerId, providerId,
modelKey, modelKey,
name, name,
budgetAmount: budgetAmount ?? null,
budgetUnit: budgetUnit ?? null,
enabled: enabled ?? true, enabled: enabled ?? true,
createdAt: now, createdAt: now,
updatedAt: now updatedAt: now
@@ -13,10 +13,8 @@ import type { CreateOrEditAiProviderResponse } from "@server/routers/aiProvider/
import { toPublicAiProvider } from "@server/routers/aiProvider/types"; import { toPublicAiProvider } from "@server/routers/aiProvider/types";
import { import {
aiAuthTypeSchema, aiAuthTypeSchema,
aiBudgetUnitSchema,
aiProviderTypeSchema, aiProviderTypeSchema,
aiRoutingModeSchema, aiRoutingModeSchema,
refineBudgetFields,
refineProviderUpstreamFields refineProviderUpstreamFields
} from "@server/routers/aiProvider/validation"; } from "@server/routers/aiProvider/validation";
@@ -33,13 +31,10 @@ const bodySchema = z
authType: aiAuthTypeSchema.optional().nullable(), authType: aiAuthTypeSchema.optional().nullable(),
routingMode: aiRoutingModeSchema.optional(), routingMode: aiRoutingModeSchema.optional(),
skipTlsVerification: z.boolean().optional(), skipTlsVerification: z.boolean().optional(),
budgetAmount: z.number().positive().optional().nullable(),
budgetUnit: aiBudgetUnitSchema.optional().nullable(),
enabled: z.boolean().optional() enabled: z.boolean().optional()
}) })
.superRefine((data, ctx) => { .superRefine((data, ctx) => {
refineProviderUpstreamFields(data, ctx); refineProviderUpstreamFields(data, ctx);
refineBudgetFields(data, ctx);
}); });
registry.registerPath({ registry.registerPath({
@@ -99,8 +94,6 @@ export async function createAiProvider(
authType, authType,
routingMode, routingMode,
skipTlsVerification, skipTlsVerification,
budgetAmount,
budgetUnit,
enabled enabled
} = parsedBody.data; } = parsedBody.data;
@@ -126,8 +119,6 @@ export async function createAiProvider(
authType: authType ?? null, authType: authType ?? null,
routingMode: resolvedRoutingMode, routingMode: resolvedRoutingMode,
skipTlsVerification: skipTlsVerification ?? false, skipTlsVerification: skipTlsVerification ?? false,
budgetAmount: budgetAmount ?? null,
budgetUnit: budgetUnit ?? null,
enabled: enabled ?? true, enabled: enabled ?? true,
createdAt: now, createdAt: now,
updatedAt: now updatedAt: now
+5 -21
View File
@@ -9,26 +9,16 @@ import { fromError } from "zod-validation-error";
import { OpenAPITags, registry } from "@server/openApi"; import { OpenAPITags, registry } from "@server/openApi";
import { and, eq, ne } from "drizzle-orm"; import { and, eq, ne } from "drizzle-orm";
import type { CreateOrEditAiModelResponse } from "@server/routers/aiProvider/types"; import type { CreateOrEditAiModelResponse } from "@server/routers/aiProvider/types";
import {
aiBudgetUnitSchema,
refineBudgetFields
} from "@server/routers/aiProvider/validation";
const paramsSchema = z.strictObject({ const paramsSchema = z.strictObject({
modelId: z.coerce.number().int().positive() modelId: z.coerce.number().int().positive()
}); });
const bodySchema = z const bodySchema = z.strictObject({
.strictObject({ modelKey: z.string().nonempty().optional(),
modelKey: z.string().nonempty().optional(), name: z.string().nonempty().optional(),
name: z.string().nonempty().optional(), enabled: z.boolean().optional()
budgetAmount: z.number().positive().optional().nullable(), });
budgetUnit: aiBudgetUnitSchema.optional().nullable(),
enabled: z.boolean().optional()
})
.superRefine((data, ctx) => {
refineBudgetFields(data, ctx);
});
registry.registerPath({ registry.registerPath({
method: "post", method: "post",
@@ -135,12 +125,6 @@ export async function updateAiModel(
if (body.name !== undefined) { if (body.name !== undefined) {
updateData.name = body.name; updateData.name = body.name;
} }
if (body.budgetAmount !== undefined) {
updateData.budgetAmount = body.budgetAmount;
}
if (body.budgetUnit !== undefined) {
updateData.budgetUnit = body.budgetUnit;
}
if (body.enabled !== undefined) { if (body.enabled !== undefined) {
updateData.enabled = body.enabled; updateData.enabled = body.enabled;
} }
+9 -23
View File
@@ -14,10 +14,8 @@ import type { CreateOrEditAiProviderResponse } from "@server/routers/aiProvider/
import { toPublicAiProvider } from "@server/routers/aiProvider/types"; import { toPublicAiProvider } from "@server/routers/aiProvider/types";
import { import {
aiAuthTypeSchema, aiAuthTypeSchema,
aiBudgetUnitSchema,
aiProviderTypeSchema, aiProviderTypeSchema,
aiRoutingModeSchema, aiRoutingModeSchema,
refineBudgetFields,
refineProviderUpstreamFields refineProviderUpstreamFields
} from "@server/routers/aiProvider/validation"; } from "@server/routers/aiProvider/validation";
import type { import type {
@@ -29,21 +27,15 @@ const paramsSchema = z.strictObject({
providerId: z.coerce.number().int().positive() providerId: z.coerce.number().int().positive()
}); });
const bodySchema = z const bodySchema = z.strictObject({
.strictObject({ name: z.string().nonempty().optional(),
name: z.string().nonempty().optional(), upstreamUrl: z.url().optional().nullable(),
upstreamUrl: z.url().optional().nullable(), apiKey: z.string().optional(),
apiKey: z.string().optional(), authType: aiAuthTypeSchema.optional().nullable(),
authType: aiAuthTypeSchema.optional().nullable(), routingMode: aiRoutingModeSchema.optional(),
routingMode: aiRoutingModeSchema.optional(), skipTlsVerification: z.boolean().optional(),
skipTlsVerification: z.boolean().optional(), enabled: z.boolean().optional()
budgetAmount: z.number().positive().optional().nullable(), });
budgetUnit: aiBudgetUnitSchema.optional().nullable(),
enabled: z.boolean().optional()
})
.superRefine((data, ctx) => {
refineBudgetFields(data, ctx);
});
registry.registerPath({ registry.registerPath({
method: "post", method: "post",
@@ -168,12 +160,6 @@ export async function updateAiProvider(
if (body.enabled !== undefined) { if (body.enabled !== undefined) {
updateData.enabled = body.enabled; updateData.enabled = body.enabled;
} }
if (body.budgetAmount !== undefined) {
updateData.budgetAmount = body.budgetAmount;
}
if (body.budgetUnit !== undefined) {
updateData.budgetUnit = body.budgetUnit;
}
if (nextRoutingMode === "target") { if (nextRoutingMode === "target") {
updateData.upstreamUrl = null; updateData.upstreamUrl = null;
} else if (body.upstreamUrl !== undefined) { } else if (body.upstreamUrl !== undefined) {
-24
View File
@@ -1,7 +1,6 @@
import { z } from "zod"; import { z } from "zod";
import { import {
providerRequiresUpstreamUrl, providerRequiresUpstreamUrl,
type AiBudgetUnit,
type AiProviderRoutingMode, type AiProviderRoutingMode,
type AiProviderType type AiProviderType
} from "@server/lib/aiProviderDefaults"; } from "@server/lib/aiProviderDefaults";
@@ -18,33 +17,10 @@ export const aiProviderTypeSchema = z.enum([
"custom" "custom"
]); ]);
export const aiBudgetUnitSchema = z.enum(["usd", "tokens"]);
export const aiAuthTypeSchema = z.enum(["bearer"]); export const aiAuthTypeSchema = z.enum(["bearer"]);
export const aiRoutingModeSchema = z.enum(["url", "target"]); 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( export function refineProviderUpstreamFields(
data: { data: {
type: AiProviderType; type: AiProviderType;
+94
View File
@@ -414,6 +414,13 @@ authenticated.get(
siteResource.listSiteResourceAiModels siteResource.listSiteResourceAiModels
); );
authenticated.get(
"/site-resource/:siteResourceId/ai-providers",
verifySiteResourceAccess,
verifyUserHasAction(ActionsEnum.listResourceAiModels),
siteResource.listSiteResourceAiProviders
);
authenticated.post( authenticated.post(
"/site-resource/:siteResourceId/roles", "/site-resource/:siteResourceId/roles",
verifySiteResourceAccess, verifySiteResourceAccess,
@@ -432,6 +439,46 @@ authenticated.post(
siteResource.setSiteResourceAiModels 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( authenticated.post(
"/site-resource/:siteResourceId/users", "/site-resource/:siteResourceId/users",
verifySiteResourceAccess, verifySiteResourceAccess,
@@ -673,6 +720,13 @@ authenticated.get(
resource.listResourceAiModels resource.listResourceAiModels
); );
authenticated.get(
"/resource/:resourceId/ai-providers",
verifyResourceAccess,
verifyUserHasAction(ActionsEnum.listResourceAiModels),
resource.listResourceAiProviders
);
authenticated.get( authenticated.get(
"/resource/:resourceId", "/resource/:resourceId",
verifyResourceAccess, verifyResourceAccess,
@@ -884,6 +938,46 @@ authenticated.post(
resource.setResourceAiModels 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( authenticated.put(
"/resource-policy/:resourcePolicyId/access-control", "/resource-policy/:resourcePolicyId/access-control",
verifyResourcePolicyAccess, verifyResourcePolicyAccess,
+86
View File
@@ -256,6 +256,16 @@ authenticated.get(
siteResource.listSiteResourceAiModels siteResource.listSiteResourceAiModels
); );
authenticated.get(
[
"/site-resource/:siteResourceId/ai-providers",
"/private-resource/:siteResourceId/ai-providers"
],
verifyApiKeySiteResourceAccess,
verifyApiKeyHasAction(ActionsEnum.listResourceAiModels),
siteResource.listSiteResourceAiProviders
);
authenticated.post( authenticated.post(
[ [
"/site-resource/:siteResourceId/roles", "/site-resource/:siteResourceId/roles",
@@ -341,6 +351,39 @@ authenticated.post(
siteResource.removeAiModelFromSiteResource 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( authenticated.post(
[ [
"/site-resource/:siteResourceId/users/add", "/site-resource/:siteResourceId/users/add",
@@ -563,6 +606,16 @@ authenticated.get(
resource.listResourceAiModels resource.listResourceAiModels
); );
authenticated.get(
[
"/resource/:resourceId/ai-providers",
"/public-resource/:resourceId/ai-providers"
],
verifyApiKeyResourceAccess,
verifyApiKeyHasAction(ActionsEnum.listResourceAiModels),
resource.listResourceAiProviders
);
authenticated.get( authenticated.get(
["/resource/:resourceId", "/public-resource/:resourceId"], ["/resource/:resourceId", "/public-resource/:resourceId"],
verifyApiKeyResourceAccess, verifyApiKeyResourceAccess,
@@ -775,6 +828,17 @@ authenticated.post(
resource.setResourceAiModels resource.setResourceAiModels
); );
authenticated.post(
[
"/resource/:resourceId/ai-providers",
"/public-resource/:resourceId/ai-providers"
],
verifyApiKeyResourceAccess,
verifyApiKeyHasAction(ActionsEnum.setResourceAiModels),
logActionAudit(ActionsEnum.setResourceAiModels),
resource.setResourceAiProviders
);
authenticated.post( authenticated.post(
["/resource/:resourceId/users", "/public-resource/:resourceId/users"], ["/resource/:resourceId/users", "/public-resource/:resourceId/users"],
verifyApiKeyResourceAccess, verifyApiKeyResourceAccess,
@@ -989,6 +1053,28 @@ authenticated.post(
resource.removeAiModelFromResource 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( authenticated.post(
[ [
"/resource/:resourceId/users/add", "/resource/:resourceId/users/add",
+16 -28
View File
@@ -1,6 +1,6 @@
import { Request, Response, NextFunction } from "express"; import { Request, Response, NextFunction } from "express";
import { z } from "zod"; 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 { eq, and } from "drizzle-orm";
import response from "@server/lib/response"; import response from "@server/lib/response";
import HttpCode from "@server/types/HttpCode"; import HttpCode from "@server/types/HttpCode";
@@ -8,7 +8,10 @@ import createHttpError from "http-errors";
import logger from "@server/logger"; import logger from "@server/logger";
import { fromError } from "zod-validation-error"; import { fromError } from "zod-validation-error";
import { OpenAPITags, registry } from "@server/openApi"; import { OpenAPITags, registry } from "@server/openApi";
import {
assertPublicAllowlistApiEligible,
assertModelsBelongToPublicAllowlistProviders
} from "@server/lib/aiInferenceResource";
const addAiModelToResourceBodySchema = z.strictObject({ const addAiModelToResourceBodySchema = z.strictObject({
modelId: z.int().positive() modelId: z.int().positive()
}); });
@@ -21,7 +24,7 @@ registry.registerPath({
method: "post", method: "post",
path: "/resource/{resourceId}/ai-models/add", path: "/resource/{resourceId}/ai-models/add",
description: 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], tags: [OpenAPITags.PublicResource],
request: { request: {
params: addAiModelToResourceParamsSchema, params: addAiModelToResourceParamsSchema,
@@ -95,33 +98,18 @@ export async function addAiModelToResource(
); );
} }
if (!resource.aiProviderId) { const eligibleError = await assertPublicAllowlistApiEligible(resource);
return next( if (eligibleError) {
createHttpError( return next(createHttpError(HttpCode.BAD_REQUEST, eligibleError));
HttpCode.BAD_REQUEST,
"Resource has no AI provider linked"
)
);
} }
const [model] = await db const modelError = await assertModelsBelongToPublicAllowlistProviders({
.select() orgId: resource.orgId,
.from(aiModels) resourceId,
.where( modelIds: [modelId]
and( });
eq(aiModels.modelId, modelId), if (modelError) {
eq(aiModels.providerId, resource.aiProviderId) return next(createHttpError(HttpCode.BAD_REQUEST, modelError));
)
)
.limit(1);
if (!model) {
return next(
createHttpError(
HttpCode.NOT_FOUND,
"Model not found or does not belong to this resource's AI provider"
)
);
} }
const existingEntry = await db const existingEntry = await db
@@ -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<any> {
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")
);
}
}
+46 -10
View File
@@ -38,6 +38,13 @@ import {
} from "@server/db/names"; } from "@server/db/names";
import { usageService } from "@server/lib/billing/usageService"; import { usageService } from "@server/lib/billing/usageService";
import { LimitId } from "@server/lib/billing"; import { LimitId } from "@server/lib/billing";
import {
isInferenceFieldsError,
resolveProviderAttachments,
resourceAiProviderAttachmentSchema,
setPublicResourceAiProviders,
type ResourceAiProviderAttachment
} from "@server/lib/aiInferenceResource";
const createResourceParamsSchema = z.strictObject({ const createResourceParamsSchema = z.strictObject({
orgId: z.string() orgId: z.string()
@@ -98,13 +105,11 @@ const createHttpResourceSchema = z
authDaemonPort: z.int().positive().optional(), authDaemonPort: z.int().positive().optional(),
authDaemonMode: z.enum(["site", "remote", "native"]).optional(), authDaemonMode: z.enum(["site", "remote", "native"]).optional(),
// Inference settings // Inference settings
aiProviderId: z aiProviders: z
.number() .array(resourceAiProviderAttachmentSchema)
.int()
.positive()
.optional() .optional()
.describe( .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( .refine(
@@ -377,11 +382,35 @@ async function createHttpResource(
authDaemonPort, authDaemonPort,
authDaemonMode, authDaemonMode,
pamMode, pamMode,
aiProviderId aiProviders: aiProviderInputs
} = parsedBody.data; } = parsedBody.data;
const subdomain = parsedBody.data.subdomain; const subdomain = parsedBody.data.subdomain;
const stickySession = parsedBody.data.stickySession; 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 // Wildcard subdomains are a paid feature
if (subdomain && subdomain.includes("*")) { if (subdomain && subdomain.includes("*")) {
const isLicensed = await isLicensedOrSubscribed( const isLicensed = await isLicensedOrSubscribed(
@@ -422,7 +451,7 @@ async function createHttpResource(
} }
if ( if (
["ssh", "rdp", "vnc"].includes(mode!) && ["ssh", "rdp", "vnc"].includes(effectiveMode) &&
!isLicensedOrSubscribed( !isLicensedOrSubscribed(
orgId!, orgId!,
tierMatrix[TierFeature.AdvancedPublicResources] tierMatrix[TierFeature.AdvancedPublicResources]
@@ -555,7 +584,7 @@ async function createHttpResource(
orgId, orgId,
name, name,
subdomain: finalSubdomain, subdomain: finalSubdomain,
mode: mode, mode: effectiveMode,
pamMode: pamMode, pamMode: pamMode,
authDaemonMode: authDaemonMode, authDaemonMode: authDaemonMode,
authDaemonPort: authDaemonPort, authDaemonPort: authDaemonPort,
@@ -564,11 +593,18 @@ async function createHttpResource(
postAuthPath: postAuthPath, postAuthPath: postAuthPath,
wildcard, wildcard,
health: "unknown", health: "unknown",
defaultResourcePolicyId: defaultPolicy.resourcePolicyId, defaultResourcePolicyId: defaultPolicy.resourcePolicyId
aiProviderId: aiProviderId ?? null
}) })
.returning(); .returning();
if (providerAttachments.length > 0) {
await setPublicResourceAiProviders(
newResource[0].resourceId,
providerAttachments,
trx
);
}
await trx.insert(roleResources).values({ await trx.insert(roleResources).values({
roleId: adminRole[0].roleId, roleId: adminRole[0].roleId,
resourceId: newResource[0].resourceId resourceId: newResource[0].resourceId
+4
View File
@@ -39,3 +39,7 @@ export * from "./listResourceAiModels";
export * from "./setResourceAiModels"; export * from "./setResourceAiModels";
export * from "./addAiModelToResource"; export * from "./addAiModelToResource";
export * from "./removeAiModelFromResource"; export * from "./removeAiModelFromResource";
export * from "./listResourceAiProviders";
export * from "./setResourceAiProviders";
export * from "./addAiProviderToResource";
export * from "./removeAiProviderFromResource";
@@ -34,7 +34,7 @@ registry.registerPath({
method: "get", method: "get",
path: "/resource/{resourceId}/ai-models", path: "/resource/{resourceId}/ai-models",
description: 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], tags: [OpenAPITags.PublicResource],
request: { request: {
params: listResourceAiModelsParamsSchema params: listResourceAiModelsParamsSchema
@@ -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<ReturnType<typeof listPublicResourceAiProviders>>;
};
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<any> {
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<ListResourceAiProvidersResponse>(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")
);
}
}
@@ -8,6 +8,7 @@ import createHttpError from "http-errors";
import logger from "@server/logger"; import logger from "@server/logger";
import { fromError } from "zod-validation-error"; import { fromError } from "zod-validation-error";
import { OpenAPITags, registry } from "@server/openApi"; import { OpenAPITags, registry } from "@server/openApi";
import { assertPublicAllowlistApiEligible } from "@server/lib/aiInferenceResource";
const removeAiModelFromResourceBodySchema = z.strictObject({ const removeAiModelFromResourceBodySchema = z.strictObject({
modelId: z.int().positive() modelId: z.int().positive()
@@ -21,7 +22,7 @@ registry.registerPath({
method: "post", method: "post",
path: "/resource/{resourceId}/ai-models/remove", path: "/resource/{resourceId}/ai-models/remove",
description: 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], tags: [OpenAPITags.PublicResource],
request: { request: {
params: removeAiModelFromResourceParamsSchema, 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 const existingEntry = await db
.select() .select()
.from(resourceAiModels) .from(resourceAiModels)
@@ -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<any> {
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")
);
}
}
+18 -30
View File
@@ -1,13 +1,17 @@
import { Request, Response, NextFunction } from "express"; import { Request, Response, NextFunction } from "express";
import { z } from "zod"; import { z } from "zod";
import { db, resources, resourceAiModels, aiModels } from "@server/db"; import { db, resources, resourceAiModels } from "@server/db";
import { eq, and, inArray } from "drizzle-orm"; import { eq } from "drizzle-orm";
import response from "@server/lib/response"; import response from "@server/lib/response";
import HttpCode from "@server/types/HttpCode"; import HttpCode from "@server/types/HttpCode";
import createHttpError from "http-errors"; import createHttpError from "http-errors";
import logger from "@server/logger"; import logger from "@server/logger";
import { fromError } from "zod-validation-error"; import { fromError } from "zod-validation-error";
import { OpenAPITags, registry } from "@server/openApi"; import { OpenAPITags, registry } from "@server/openApi";
import {
assertPublicAllowlistApiEligible,
assertModelsBelongToPublicAllowlistProviders
} from "@server/lib/aiInferenceResource";
const setResourceAiModelsBodySchema = z.strictObject({ const setResourceAiModelsBodySchema = z.strictObject({
modelIds: z.array(z.int().positive()) modelIds: z.array(z.int().positive())
@@ -21,7 +25,7 @@ registry.registerPath({
method: "post", method: "post",
path: "/resource/{resourceId}/ai-models", path: "/resource/{resourceId}/ai-models",
description: 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], tags: [OpenAPITags.PublicResource],
request: { request: {
params: setResourceAiModelsParamsSchema, params: setResourceAiModelsParamsSchema,
@@ -95,34 +99,18 @@ export async function setResourceAiModels(
); );
} }
if (modelIds.length > 0) { const eligibleError = await assertPublicAllowlistApiEligible(resource);
if (!resource.aiProviderId) { if (eligibleError) {
return next( return next(createHttpError(HttpCode.BAD_REQUEST, eligibleError));
createHttpError( }
HttpCode.BAD_REQUEST,
"Resource has no AI provider linked"
)
);
}
const validModels = await db const modelError = await assertModelsBelongToPublicAllowlistProviders({
.select({ modelId: aiModels.modelId }) orgId: resource.orgId,
.from(aiModels) resourceId,
.where( modelIds
and( });
inArray(aiModels.modelId, modelIds), if (modelError) {
eq(aiModels.providerId, resource.aiProviderId) return next(createHttpError(HttpCode.BAD_REQUEST, modelError));
)
);
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"
)
);
}
} }
await db.transaction(async (trx) => { await db.transaction(async (trx) => {
@@ -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<any> {
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")
);
}
}
+4 -11
View File
@@ -120,15 +120,6 @@ const updateHttpResourceBodySchema = z
.optional() .optional()
.describe( .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." "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, { .refine((data) => Object.keys(data).length > 0, {
@@ -354,8 +345,10 @@ export async function updateResource(
); );
} }
if (["http", "ssh", "rdp", "vnc"].includes(resource.mode)) { if (
// HANDLE UPDATING HTTP RESOURCES ["http", "ssh", "rdp", "vnc", "inference"].includes(resource.mode)
) {
// HANDLE UPDATING HTTP / BROWSER / INFERENCE RESOURCES
return await updateHttpResource( return await updateHttpResource(
{ {
req, req,
@@ -1,11 +1,6 @@
import { Request, Response, NextFunction } from "express"; import { Request, Response, NextFunction } from "express";
import { z } from "zod"; import { z } from "zod";
import { import { db, siteResources, siteResourceAiModels } from "@server/db";
db,
siteResources,
siteResourceAiModels,
aiModels
} from "@server/db";
import { eq, and } from "drizzle-orm"; import { eq, and } from "drizzle-orm";
import response from "@server/lib/response"; import response from "@server/lib/response";
import HttpCode from "@server/types/HttpCode"; import HttpCode from "@server/types/HttpCode";
@@ -13,6 +8,10 @@ import createHttpError from "http-errors";
import logger from "@server/logger"; import logger from "@server/logger";
import { fromError } from "zod-validation-error"; import { fromError } from "zod-validation-error";
import { OpenAPITags, registry } from "@server/openApi"; import { OpenAPITags, registry } from "@server/openApi";
import {
assertSiteAllowlistApiEligible,
assertModelsBelongToSiteAllowlistProviders
} from "@server/lib/aiInferenceResource";
const addAiModelToSiteResourceBodySchema = z.strictObject({ const addAiModelToSiteResourceBodySchema = z.strictObject({
modelId: z.int().positive() modelId: z.int().positive()
@@ -26,7 +25,7 @@ registry.registerPath({
method: "post", method: "post",
path: "/site-resource/{siteResourceId}/ai-models/add", path: "/site-resource/{siteResourceId}/ai-models/add",
description: 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], tags: [OpenAPITags.PrivateResource],
request: { request: {
params: addAiModelToSiteResourceParamsSchema, params: addAiModelToSiteResourceParamsSchema,
@@ -102,33 +101,19 @@ export async function addAiModelToSiteResource(
); );
} }
if (!siteResource.aiProviderId) { const eligibleError =
return next( await assertSiteAllowlistApiEligible(siteResource);
createHttpError( if (eligibleError) {
HttpCode.BAD_REQUEST, return next(createHttpError(HttpCode.BAD_REQUEST, eligibleError));
"Site resource has no AI provider linked"
)
);
} }
const [model] = await db const modelError = await assertModelsBelongToSiteAllowlistProviders({
.select() orgId: siteResource.orgId,
.from(aiModels) siteResourceId,
.where( modelIds: [modelId]
and( });
eq(aiModels.modelId, modelId), if (modelError) {
eq(aiModels.providerId, siteResource.aiProviderId) return next(createHttpError(HttpCode.BAD_REQUEST, modelError));
)
)
.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 existingEntry = await db const existingEntry = await db
@@ -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<any> {
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")
);
}
}
@@ -39,6 +39,13 @@ import { createCertificate } from "#dynamic/routers/certificates/createCertifica
import { build } from "@server/build"; import { build } from "@server/build";
import { usageService } from "@server/lib/billing/usageService"; import { usageService } from "@server/lib/billing/usageService";
import { LimitId } from "@server/lib/billing"; import { LimitId } from "@server/lib/billing";
import {
isInferenceFieldsError,
resolveProviderAttachments,
resourceAiProviderAttachmentSchema,
setSiteResourceAiProviders,
type ResourceAiProviderAttachment
} from "@server/lib/aiInferenceResource";
const createSiteResourceParamsSchema = z.strictObject({ const createSiteResourceParamsSchema = z.strictObject({
orgId: z.string() orgId: z.string()
@@ -79,13 +86,11 @@ const createSiteResourceSchema = z
pamMode: z.enum(["passthrough", "push"]).optional(), 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 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 subdomain: z.string().optional(), // only used for http mode, we need this to verify the alias is unique within the org
aiProviderId: z aiProviders: z
.number() .array(resourceAiProviderAttachmentSchema)
.int()
.positive()
.optional() .optional()
.describe( .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() .strict()
@@ -334,7 +339,7 @@ export async function createSiteResource(
pamMode, pamMode,
domainId, domainId,
subdomain, subdomain,
aiProviderId aiProviders: aiProviderInputs
} = parsedBody.data; } = parsedBody.data;
// Backward compatibility: merge deprecated siteId into siteIds array // Backward compatibility: merge deprecated siteId into siteIds array
@@ -343,6 +348,28 @@ export async function createSiteResource(
siteIds.push(siteId); 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") { if (build == "saas") {
const usage = await usageService.getUsage( const usage = await usageService.getUsage(
orgId, orgId,
@@ -608,8 +635,7 @@ export async function createSiteResource(
domainId, domainId,
subdomain: finalSubdomain, subdomain: finalSubdomain,
fullDomain, fullDomain,
requiresExitNodeConnection: mode === "inference", // in the future we might want to have different modes that do this requiresExitNodeConnection: mode === "inference" // in the future we might want to have different modes that do this
aiProviderId: aiProviderId ?? null
}; };
if (isLicensedSshPam) { if (isLicensedSshPam) {
if (authDaemonPort !== undefined) if (authDaemonPort !== undefined)
@@ -625,6 +651,14 @@ export async function createSiteResource(
const siteResourceId = newSiteResource.siteResourceId; const siteResourceId = newSiteResource.siteResourceId;
if (providerAttachments.length > 0) {
await setSiteResourceAiProviders(
siteResourceId,
providerAttachments,
trx
);
}
//////////////////// update the associations //////////////////// //////////////////// update the associations ////////////////////
if (network) { if (network) {
+4
View File
@@ -21,3 +21,7 @@ export * from "./listSiteResourceAiModels";
export * from "./setSiteResourceAiModels"; export * from "./setSiteResourceAiModels";
export * from "./addAiModelToSiteResource"; export * from "./addAiModelToSiteResource";
export * from "./removeAiModelFromSiteResource"; export * from "./removeAiModelFromSiteResource";
export * from "./listSiteResourceAiProviders";
export * from "./setSiteResourceAiProviders";
export * from "./addAiProviderToSiteResource";
export * from "./removeAiProviderFromSiteResource";
@@ -1,11 +1,6 @@
import { Request, Response, NextFunction } from "express"; import { Request, Response, NextFunction } from "express";
import { z } from "zod"; import { z } from "zod";
import { import { db, siteResources, siteResourceAiModels, aiModels } from "@server/db";
db,
siteResources,
siteResourceAiModels,
aiModels
} from "@server/db";
import { eq } from "drizzle-orm"; import { eq } from "drizzle-orm";
import response from "@server/lib/response"; import response from "@server/lib/response";
import HttpCode from "@server/types/HttpCode"; import HttpCode from "@server/types/HttpCode";
@@ -27,10 +22,7 @@ async function query(siteResourceId: number) {
enabled: aiModels.enabled enabled: aiModels.enabled
}) })
.from(siteResourceAiModels) .from(siteResourceAiModels)
.innerJoin( .innerJoin(aiModels, eq(siteResourceAiModels.modelId, aiModels.modelId))
aiModels,
eq(siteResourceAiModels.modelId, aiModels.modelId)
)
.where(eq(siteResourceAiModels.siteResourceId, siteResourceId)); .where(eq(siteResourceAiModels.siteResourceId, siteResourceId));
} }
@@ -42,7 +34,7 @@ registry.registerPath({
method: "get", method: "get",
path: "/site-resource/{siteResourceId}/ai-models", path: "/site-resource/{siteResourceId}/ai-models",
description: 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], tags: [OpenAPITags.PrivateResource],
request: { request: {
params: listSiteResourceAiModelsParamsSchema params: listSiteResourceAiModelsParamsSchema
@@ -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<ReturnType<typeof listAttachments>>;
};
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<any> {
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<ListSiteResourceAiProvidersResponse>(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")
);
}
}
@@ -8,6 +8,7 @@ import createHttpError from "http-errors";
import logger from "@server/logger"; import logger from "@server/logger";
import { fromError } from "zod-validation-error"; import { fromError } from "zod-validation-error";
import { OpenAPITags, registry } from "@server/openApi"; import { OpenAPITags, registry } from "@server/openApi";
import { assertSiteAllowlistApiEligible } from "@server/lib/aiInferenceResource";
const removeAiModelFromSiteResourceBodySchema = z.strictObject({ const removeAiModelFromSiteResourceBodySchema = z.strictObject({
modelId: z.int().positive() modelId: z.int().positive()
@@ -21,7 +22,7 @@ registry.registerPath({
method: "post", method: "post",
path: "/site-resource/{siteResourceId}/ai-models/remove", path: "/site-resource/{siteResourceId}/ai-models/remove",
description: 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], tags: [OpenAPITags.PrivateResource],
request: { request: {
params: removeAiModelFromSiteResourceParamsSchema, 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 const existingEntry = await db
.select() .select()
.from(siteResourceAiModels) .from(siteResourceAiModels)
@@ -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<any> {
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")
);
}
}
@@ -1,18 +1,17 @@
import { Request, Response, NextFunction } from "express"; import { Request, Response, NextFunction } from "express";
import { z } from "zod"; import { z } from "zod";
import { import { db, siteResources, siteResourceAiModels } from "@server/db";
db, import { eq } from "drizzle-orm";
siteResources,
siteResourceAiModels,
aiModels
} from "@server/db";
import { eq, and, inArray } from "drizzle-orm";
import response from "@server/lib/response"; import response from "@server/lib/response";
import HttpCode from "@server/types/HttpCode"; import HttpCode from "@server/types/HttpCode";
import createHttpError from "http-errors"; import createHttpError from "http-errors";
import logger from "@server/logger"; import logger from "@server/logger";
import { fromError } from "zod-validation-error"; import { fromError } from "zod-validation-error";
import { OpenAPITags, registry } from "@server/openApi"; import { OpenAPITags, registry } from "@server/openApi";
import {
assertSiteAllowlistApiEligible,
assertModelsBelongToSiteAllowlistProviders
} from "@server/lib/aiInferenceResource";
const setSiteResourceAiModelsBodySchema = z.strictObject({ const setSiteResourceAiModelsBodySchema = z.strictObject({
modelIds: z.array(z.int().positive()) modelIds: z.array(z.int().positive())
@@ -26,7 +25,7 @@ registry.registerPath({
method: "post", method: "post",
path: "/site-resource/{siteResourceId}/ai-models", path: "/site-resource/{siteResourceId}/ai-models",
description: 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], tags: [OpenAPITags.PrivateResource],
request: { request: {
params: setSiteResourceAiModelsParamsSchema, params: setSiteResourceAiModelsParamsSchema,
@@ -102,52 +101,33 @@ export async function setSiteResourceAiModels(
); );
} }
if (modelIds.length > 0) { const eligibleError =
if (!siteResource.aiProviderId) { await assertSiteAllowlistApiEligible(siteResource);
return next( if (eligibleError) {
createHttpError( return next(createHttpError(HttpCode.BAD_REQUEST, eligibleError));
HttpCode.BAD_REQUEST, }
"Site resource has no AI provider linked"
)
);
}
const validModels = await db const modelError = await assertModelsBelongToSiteAllowlistProviders({
.select({ modelId: aiModels.modelId }) orgId: siteResource.orgId,
.from(aiModels) siteResourceId,
.where( modelIds
and( });
inArray(aiModels.modelId, modelIds), if (modelError) {
eq(aiModels.providerId, siteResource.aiProviderId) return next(createHttpError(HttpCode.BAD_REQUEST, modelError));
)
);
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"
)
);
}
} }
await db.transaction(async (trx) => { await db.transaction(async (trx) => {
await trx await trx
.delete(siteResourceAiModels) .delete(siteResourceAiModels)
.where( .where(eq(siteResourceAiModels.siteResourceId, siteResourceId));
eq(siteResourceAiModels.siteResourceId, siteResourceId)
);
if (modelIds.length > 0) { if (modelIds.length > 0) {
await trx await trx.insert(siteResourceAiModels).values(
.insert(siteResourceAiModels) modelIds.map((modelId) => ({
.values( siteResourceId,
modelIds.map((modelId) => ({ modelId
siteResourceId, }))
modelId );
}))
);
} }
}); });
@@ -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<any> {
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")
);
}
}
@@ -29,6 +29,7 @@ import { NextFunction, Request, Response } from "express";
import createHttpError from "http-errors"; import createHttpError from "http-errors";
import { z } from "zod"; import { z } from "zod";
import { fromError } from "zod-validation-error"; import { fromError } from "zod-validation-error";
import { clearSiteResourceAiConfig } from "@server/lib/aiInferenceResource";
const updateSiteResourceParamsSchema = z.strictObject({ const updateSiteResourceParamsSchema = z.strictObject({
siteResourceId: z.coerce.number().int().positive() siteResourceId: z.coerce.number().int().positive()
@@ -78,16 +79,7 @@ const updateSiteResourceSchema = z
authDaemonMode: z.enum(["site", "remote", "native"]).optional(), authDaemonMode: z.enum(["site", "remote", "native"]).optional(),
pamMode: z.enum(["passthrough", "push"]).optional(), pamMode: z.enum(["passthrough", "push"]).optional(),
domainId: z.string().optional(), domainId: z.string().optional(),
subdomain: 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."
)
}) })
.strict() .strict()
.refine( .refine(
@@ -341,8 +333,7 @@ export async function updateSiteResource(
authDaemonMode, authDaemonMode,
pamMode, pamMode,
domainId, domainId,
subdomain, subdomain
aiProviderId
} = parsedBody.data; } = parsedBody.data;
// Backward compatibility: merge deprecated siteId into siteIds array // Backward compatibility: merge deprecated siteId into siteIds array
@@ -608,12 +599,19 @@ export async function updateSiteResource(
networkId: mode === "inference" ? null : undefined, networkId: mode === "inference" ? null : undefined,
requiresExitNodeConnection: requiresExitNodeConnection:
mode !== undefined ? mode === "inference" : undefined, mode !== undefined ? mode === "inference" : undefined,
aiProviderId: aiProviderId,
...sshPamSet ...sshPamSet
}) })
.where(and(eq(siteResources.siteResourceId, siteResourceId))) .where(and(eq(siteResources.siteResourceId, siteResourceId)))
.returning(); .returning();
const effectiveMode = mode ?? existingSiteResource.mode;
if (
existingSiteResource.mode === "inference" &&
effectiveMode !== "inference"
) {
await clearSiteResourceAiConfig(siteResourceId, trx);
}
//////////////////// update the associations //////////////////// //////////////////// update the associations ////////////////////
if (mode === "inference") { if (mode === "inference") {
@@ -66,8 +66,6 @@ export default function AiProviderAuthenticationPage() {
authType: (provider.authType as "bearer" | null) ?? "bearer", authType: (provider.authType as "bearer" | null) ?? "bearer",
routingMode: (provider.routingMode as "url" | "target") ?? "url", routingMode: (provider.routingMode as "url" | "target") ?? "url",
skipTlsVerification: provider.skipTlsVerification, skipTlsVerification: provider.skipTlsVerification,
budgetAmount: provider.budgetAmount,
budgetUnit: provider.budgetUnit as "usd" | "tokens" | null,
enabled: provider.enabled enabled: provider.enabled
} }
}); });
@@ -96,8 +94,6 @@ export default function AiProviderAuthenticationPage() {
authType: (updated.authType as "bearer" | null) ?? "bearer", authType: (updated.authType as "bearer" | null) ?? "bearer",
routingMode: (updated.routingMode as "url" | "target") ?? "url", routingMode: (updated.routingMode as "url" | "target") ?? "url",
skipTlsVerification: updated.skipTlsVerification, skipTlsVerification: updated.skipTlsVerification,
budgetAmount: updated.budgetAmount,
budgetUnit: updated.budgetUnit as "usd" | "tokens" | null,
enabled: updated.enabled enabled: updated.enabled
}); });
toast({ toast({
@@ -75,8 +75,6 @@ export default function AiProviderNetworkPage() {
authType: (provider.authType as "bearer" | null) ?? "bearer", authType: (provider.authType as "bearer" | null) ?? "bearer",
routingMode: (provider.routingMode as "url" | "target") ?? "url", routingMode: (provider.routingMode as "url" | "target") ?? "url",
skipTlsVerification: provider.skipTlsVerification, skipTlsVerification: provider.skipTlsVerification,
budgetAmount: provider.budgetAmount,
budgetUnit: provider.budgetUnit as "usd" | "tokens" | null,
enabled: provider.enabled enabled: provider.enabled
} }
}); });
@@ -120,8 +118,6 @@ export default function AiProviderNetworkPage() {
authType: (updated.authType as "bearer" | null) ?? "bearer", authType: (updated.authType as "bearer" | null) ?? "bearer",
routingMode: (updated.routingMode as "url" | "target") ?? "url", routingMode: (updated.routingMode as "url" | "target") ?? "url",
skipTlsVerification: updated.skipTlsVerification, skipTlsVerification: updated.skipTlsVerification,
budgetAmount: updated.budgetAmount,
budgetUnit: updated.budgetUnit as "usd" | "tokens" | null,
enabled: updated.enabled enabled: updated.enabled
}); });
@@ -79,8 +79,6 @@ export default function CreateAiProviderPage() {
authType: "bearer", authType: "bearer",
routingMode: "url", routingMode: "url",
skipTlsVerification: false, skipTlsVerification: false,
budgetAmount: null,
budgetUnit: null,
enabled: true enabled: true
} }
}); });
-30
View File
@@ -26,8 +26,6 @@ export const aiProviderFormSchema = z
authType: z.enum(["bearer"]).optional().nullable(), authType: z.enum(["bearer"]).optional().nullable(),
routingMode: z.enum(["url", "target"]).optional(), routingMode: z.enum(["url", "target"]).optional(),
skipTlsVerification: z.boolean().optional(), skipTlsVerification: z.boolean().optional(),
budgetAmount: z.number().positive().nullable().optional(),
budgetUnit: z.enum(["usd", "tokens"]).optional().nullable(),
enabled: z.boolean().optional() enabled: z.boolean().optional()
}) })
.superRefine((data, ctx) => { .superRefine((data, ctx) => {
@@ -78,20 +76,6 @@ export const aiProviderFormSchema = z
path: ["authType"] 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<typeof aiProviderFormSchema>; export type AiProviderFormValues = z.infer<typeof aiProviderFormSchema>;
@@ -145,11 +129,6 @@ export function toAiProviderCreatePayload(values: AiProviderFormValues) {
? upstreamRaw ? upstreamRaw
: null; : null;
const hasBudget =
values.budgetAmount !== undefined &&
values.budgetAmount !== null &&
values.budgetUnit;
return { return {
name: values.name.trim(), name: values.name.trim(),
type: values.type, type: values.type,
@@ -161,8 +140,6 @@ export function toAiProviderCreatePayload(values: AiProviderFormValues) {
? (values.authType ?? "bearer") ? (values.authType ?? "bearer")
: (values.authType ?? undefined), : (values.authType ?? undefined),
skipTlsVerification: values.skipTlsVerification, skipTlsVerification: values.skipTlsVerification,
budgetAmount: hasBudget ? values.budgetAmount : null,
budgetUnit: hasBudget ? values.budgetUnit : null,
enabled: values.enabled ?? true enabled: values.enabled ?? true
}; };
} }
@@ -178,11 +155,6 @@ export function toAiProviderUpdatePayload(values: AiProviderFormValues) {
? upstreamRaw ? upstreamRaw
: null; : null;
const hasBudget =
values.budgetAmount !== undefined &&
values.budgetAmount !== null &&
values.budgetUnit;
const payload: Record<string, unknown> = { const payload: Record<string, unknown> = {
name: values.name.trim(), name: values.name.trim(),
routingMode: values.type === "custom" ? routingMode : "url", routingMode: values.type === "custom" ? routingMode : "url",
@@ -192,8 +164,6 @@ export function toAiProviderUpdatePayload(values: AiProviderFormValues) {
? (values.authType ?? "bearer") ? (values.authType ?? "bearer")
: (values.authType ?? null), : (values.authType ?? null),
skipTlsVerification: values.skipTlsVerification ?? false, skipTlsVerification: values.skipTlsVerification ?? false,
budgetAmount: hasBudget ? values.budgetAmount : null,
budgetUnit: hasBudget ? values.budgetUnit : null,
enabled: values.enabled ?? true enabled: values.enabled ?? true
}; };