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),
status: varchar("status")
.$type<"pending" | "approved">()
.default("approved"),
aiProviderId: integer("aiProviderId").references(
() => aiProviders.providerId,
{ onDelete: "set null" }
),
modelAccessMode: varchar("modelAccessMode").$type<
"passthrough" | "catalog" | "allowlist"
>()
.default("approved")
},
(t) => [
index("idx_resources_fulldomain")
@@ -229,6 +222,23 @@ export const resources = pgTable(
]
);
export const resourceAiProviders = pgTable(
"resourceAiProviders",
{
resourceId: integer("resourceId")
.notNull()
.references(() => resources.resourceId, { onDelete: "cascade" }),
providerId: integer("providerId")
.notNull()
.references(() => aiProviders.providerId, { onDelete: "cascade" }),
modelAccessMode: varchar("modelAccessMode")
.$type<"passthrough" | "catalog" | "allowlist">()
.notNull()
.default("passthrough")
},
(t) => [primaryKey({ columns: [t.resourceId, t.providerId] })]
);
export const resourceAiModels = pgTable(
"resourceAiModels",
{
@@ -496,18 +506,30 @@ export const siteResources = pgTable(
fullDomain: varchar("fullDomain"),
status: varchar("status")
.$type<"pending" | "approved">()
.default("approved"),
aiProviderId: integer("aiProviderId").references(
() => aiProviders.providerId,
{ onDelete: "set null" }
),
modelAccessMode: varchar("modelAccessMode").$type<
"passthrough" | "catalog" | "allowlist"
>()
.default("approved")
},
(t) => [index("idx_siteresources_orgid_niceid").on(t.orgId, t.niceId)]
);
export const siteResourceAiProviders = pgTable(
"siteResourceAiProviders",
{
siteResourceId: integer("siteResourceId")
.notNull()
.references(() => siteResources.siteResourceId, {
onDelete: "cascade"
}),
providerId: integer("providerId")
.notNull()
.references(() => aiProviders.providerId, { onDelete: "cascade" }),
modelAccessMode: varchar("modelAccessMode")
.$type<"passthrough" | "catalog" | "allowlist">()
.notNull()
.default("passthrough")
},
(t) => [primaryKey({ columns: [t.siteResourceId, t.providerId] })]
);
export const siteResourceAiModels = pgTable(
"siteResourceAiModels",
{
@@ -1806,5 +1828,9 @@ export type AiProvider = InferSelectModel<typeof aiProviders>;
export type AiModel = InferSelectModel<typeof aiModels>;
export type AiBudget = InferSelectModel<typeof aiBudgets>;
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 SiteResourceAiModel = InferSelectModel<typeof siteResourceAiModels>;
+42 -16
View File
@@ -216,16 +216,26 @@ export const resources = sqliteTable("resources", {
.$type<"site" | "remote" | "native">()
.default("site"),
authDaemonPort: integer("authDaemonPort").default(22123),
status: text("status").$type<"pending" | "approved">().default("approved"),
aiProviderId: integer("aiProviderId").references(
() => aiProviders.providerId,
{ onDelete: "set null" }
),
modelAccessMode: text("modelAccessMode").$type<
"passthrough" | "catalog" | "allowlist"
>()
status: text("status").$type<"pending" | "approved">().default("approved")
});
export const resourceAiProviders = sqliteTable(
"resourceAiProviders",
{
resourceId: integer("resourceId")
.notNull()
.references(() => resources.resourceId, { onDelete: "cascade" }),
providerId: integer("providerId")
.notNull()
.references(() => aiProviders.providerId, { onDelete: "cascade" }),
modelAccessMode: text("modelAccessMode")
.$type<"passthrough" | "catalog" | "allowlist">()
.notNull()
.default("passthrough")
},
(t) => [primaryKey({ columns: [t.resourceId, t.providerId] })]
);
export const resourceAiModels = sqliteTable(
"resourceAiModels",
{
@@ -483,16 +493,28 @@ export const siteResources = sqliteTable("siteResources", {
}),
subdomain: text("subdomain"),
fullDomain: text("fullDomain"),
status: text("status").$type<"pending" | "approved">().default("approved"),
aiProviderId: integer("aiProviderId").references(
() => aiProviders.providerId,
{ onDelete: "set null" }
),
modelAccessMode: text("modelAccessMode").$type<
"passthrough" | "catalog" | "allowlist"
>()
status: text("status").$type<"pending" | "approved">().default("approved")
});
export const siteResourceAiProviders = sqliteTable(
"siteResourceAiProviders",
{
siteResourceId: integer("siteResourceId")
.notNull()
.references(() => siteResources.siteResourceId, {
onDelete: "cascade"
}),
providerId: integer("providerId")
.notNull()
.references(() => aiProviders.providerId, { onDelete: "cascade" }),
modelAccessMode: text("modelAccessMode")
.$type<"passthrough" | "catalog" | "allowlist">()
.notNull()
.default("passthrough")
},
(t) => [primaryKey({ columns: [t.siteResourceId, t.providerId] })]
);
export const siteResourceAiModels = sqliteTable(
"siteResourceAiModels",
{
@@ -1787,5 +1809,9 @@ export type AiProvider = InferSelectModel<typeof aiProviders>;
export type AiModel = InferSelectModel<typeof aiModels>;
export type AiBudget = InferSelectModel<typeof aiBudgets>;
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 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 {
and,
eq,
@@ -214,7 +214,7 @@ export async function getTraefikConfig(
// central AI gateway), so they can't be reached via the targets->sites
// join above - query them separately and include them on every exit node.
const inferenceResources = await db
.select({
.selectDistinct({
resourceId: resources.resourceId,
resourceName: resources.name,
fullDomain: resources.fullDomain,
@@ -226,9 +226,13 @@ export async function getTraefikConfig(
preferWildcardCert: domains.preferWildcardCert
})
.from(resources)
.innerJoin(
resourceAiProviders,
eq(resources.resourceId, resourceAiProviders.resourceId)
)
.innerJoin(
aiProviders,
eq(resources.aiProviderId, aiProviders.providerId)
eq(resourceAiProviders.providerId, aiProviders.providerId)
)
.leftJoin(domains, eq(domains.domainId, resources.domainId))
.where(
+18 -5
View File
@@ -42,7 +42,9 @@ import {
siteResources,
Target,
targets,
aiProviders
aiProviders,
resourceAiProviders,
siteResourceAiProviders
} from "@server/db";
import {
sanitize,
@@ -402,7 +404,7 @@ export async function getTraefikConfig(
// so they can't be reached via the joins above - query them separately
// and include them on every exit node.
const inferenceResources = await db
.select({
.selectDistinct({
resourceId: resources.resourceId,
fullDomain: resources.fullDomain,
ssl: resources.ssl,
@@ -414,9 +416,13 @@ export async function getTraefikConfig(
preferWildcardCert: domains.preferWildcardCert
})
.from(resources)
.innerJoin(
resourceAiProviders,
eq(resources.resourceId, resourceAiProviders.resourceId)
)
.innerJoin(
aiProviders,
eq(resources.aiProviderId, aiProviders.providerId)
eq(resourceAiProviders.providerId, aiProviders.providerId)
)
.leftJoin(domains, eq(domains.domainId, resources.domainId))
.where(
@@ -435,16 +441,23 @@ export async function getTraefikConfig(
}[] = [];
if (build == "enterprise") {
siteResourcesInference = await db
.select({
.selectDistinct({
siteResourceId: siteResources.siteResourceId,
alias: siteResources.alias,
ssl: siteResources.ssl,
enabled: siteResources.enabled
})
.from(siteResources)
.innerJoin(
siteResourceAiProviders,
eq(
siteResources.siteResourceId,
siteResourceAiProviders.siteResourceId
)
)
.innerJoin(
aiProviders,
eq(siteResources.aiProviderId, aiProviders.providerId)
eq(siteResourceAiProviders.providerId, aiProviders.providerId)
)
.where(
and(
+218 -70
View File
@@ -8,8 +8,10 @@ import {
db,
exitNodes,
resourceAiModels,
resourceAiProviders,
resources,
siteResourceAiModels,
siteResourceAiProviders,
siteResources,
users
} from "@server/db";
@@ -21,7 +23,6 @@ import {
AiProviderType,
resolveAiProviderConfig
} from "@server/lib/aiProviderDefaults";
import { verifyResourceAccessToken } from "@server/auth/verifyResourceAccessToken";
import {
SESSION_COOKIE_NAME,
validateSessionToken
@@ -31,6 +32,7 @@ import { isIpInCidr } from "@server/lib/ip";
import { localCache } from "@server/lib/cache";
import logger from "@server/logger";
import HttpCode from "@server/types/HttpCode";
import type { ModelAccessMode } from "@server/lib/aiInferenceResource";
// Short-lived local caches so a burst of requests from the same IP/user
// doesn't hit the database on every single request. None of this is
@@ -85,14 +87,23 @@ async function findClientByIp(ip: string): Promise<CachedClient> {
return result;
}
type ProviderAttachment = {
provider: AiProvider;
modelAccessMode: ModelAccessMode;
};
type ResolvedTarget = {
resourceId: number | null;
orgId: string | null;
provider: AiProvider;
// null = no restriction; every enabled model on the provider is allowed
allowedModelIds: number[] | null;
attachments: ProviderAttachment[];
// model IDs on this resource's allowlist (resource-wide)
allowedModelIds: number[];
};
type ProviderSelection =
| { ok: true; provider: AiProvider }
| { ok: false; status: number; message: string };
export type RequestUser = {
userId: string;
username: string;
@@ -184,88 +195,245 @@ async function resolveTarget(host: string): Promise<ResolvedTarget | null> {
const [resourceRow] = await db
.select({
resourceId: resources.resourceId,
orgId: resources.orgId,
provider: aiProviders
orgId: resources.orgId
})
.from(resources)
.innerJoin(
aiProviders,
eq(resources.aiProviderId, aiProviders.providerId)
)
.where(
and(
eq(resources.fullDomain, host),
eq(resources.mode, "inference"),
eq(resources.enabled, true),
eq(aiProviders.enabled, true)
eq(resources.enabled, true)
)
)
.limit(1);
if (resourceRow) {
const restrictions = await db
.select({ modelId: resourceAiModels.modelId })
.from(resourceAiModels)
.where(eq(resourceAiModels.resourceId, resourceRow.resourceId));
const attachmentRows = await db
.select({
modelAccessMode: resourceAiProviders.modelAccessMode,
provider: aiProviders
})
.from(resourceAiProviders)
.innerJoin(
aiProviders,
eq(resourceAiProviders.providerId, aiProviders.providerId)
)
.where(
and(
eq(resourceAiProviders.resourceId, resourceRow.resourceId),
eq(aiProviders.enabled, true)
)
);
if (attachmentRows.length === 0) {
return null;
}
const hasAllowlist = attachmentRows.some(
(a) => a.modelAccessMode === "allowlist"
);
let allowedModelIds: number[] = [];
if (hasAllowlist) {
const restrictions = await db
.select({ modelId: resourceAiModels.modelId })
.from(resourceAiModels)
.where(eq(resourceAiModels.resourceId, resourceRow.resourceId));
allowedModelIds = restrictions.map((r) => r.modelId);
}
return {
resourceId: resourceRow.resourceId,
orgId: resourceRow.orgId,
provider: resourceRow.provider,
allowedModelIds: restrictions.length
? restrictions.map((r) => r.modelId)
: null
attachments: attachmentRows.map((a) => ({
provider: a.provider,
modelAccessMode: a.modelAccessMode as ModelAccessMode
})),
allowedModelIds
};
}
const [siteResourceRow] = await db
.select({
siteResourceId: siteResources.siteResourceId,
orgId: siteResources.orgId,
provider: aiProviders
orgId: siteResources.orgId
})
.from(siteResources)
.innerJoin(
aiProviders,
eq(siteResources.aiProviderId, aiProviders.providerId)
)
.where(
and(
eq(siteResources.alias, host),
eq(siteResources.mode, "inference"),
eq(siteResources.enabled, true),
eq(aiProviders.enabled, true)
eq(siteResources.enabled, true)
)
)
.limit(1);
if (siteResourceRow) {
const restrictions = await db
.select({ modelId: siteResourceAiModels.modelId })
.from(siteResourceAiModels)
const attachmentRows = await db
.select({
modelAccessMode: siteResourceAiProviders.modelAccessMode,
provider: aiProviders
})
.from(siteResourceAiProviders)
.innerJoin(
aiProviders,
eq(siteResourceAiProviders.providerId, aiProviders.providerId)
)
.where(
eq(
siteResourceAiModels.siteResourceId,
siteResourceRow.siteResourceId
and(
eq(
siteResourceAiProviders.siteResourceId,
siteResourceRow.siteResourceId
),
eq(aiProviders.enabled, true)
)
);
if (attachmentRows.length === 0) {
return null;
}
const hasAllowlist = attachmentRows.some(
(a) => a.modelAccessMode === "allowlist"
);
let allowedModelIds: number[] = [];
if (hasAllowlist) {
const restrictions = await db
.select({ modelId: siteResourceAiModels.modelId })
.from(siteResourceAiModels)
.where(
eq(
siteResourceAiModels.siteResourceId,
siteResourceRow.siteResourceId
)
);
allowedModelIds = restrictions.map((r) => r.modelId);
}
return {
// siteResources have no per-user auth/policy stack today (see
// the routing comment in getTraefikConfig.ts), so there's no
// resource access token scope to validate a user token against.
resourceId: null,
orgId: siteResourceRow.orgId,
provider: siteResourceRow.provider,
allowedModelIds: restrictions.length
? restrictions.map((r) => r.modelId)
: null
attachments: attachmentRows.map((a) => ({
provider: a.provider,
modelAccessMode: a.modelAccessMode as ModelAccessMode
})),
allowedModelIds
};
}
return null;
}
async function providerMatchesModel(
attachment: ProviderAttachment,
requestedModel: string,
allowedModelIds: number[]
): Promise<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
// different path/schema; everything else here is OpenAI-compatible today.
function getCompletionsPath(type: AiProviderType): string {
@@ -296,7 +464,7 @@ export async function chatCompletions(
});
}
const { provider, allowedModelIds, resourceId, orgId } = target;
const { attachments, allowedModelIds, resourceId, orgId } = target;
// Best-effort identity resolution - not yet enforced, but lets us
// start making per-user access decisions (e.g. model/role-based
@@ -311,39 +479,19 @@ export async function chatCompletions(
const requestedModel =
typeof req.body?.model === "string" ? req.body.model : undefined;
if (allowedModelIds) {
if (!requestedModel) {
return res.status(HttpCode.FORBIDDEN).json({
error: {
message:
"This resource restricts access to specific models; a model must be specified"
}
});
}
const [matchedModel] = await db
.select({ modelId: aiModels.modelId })
.from(aiModels)
.where(
and(
eq(aiModels.providerId, provider.providerId),
eq(aiModels.modelKey, requestedModel)
)
)
.limit(1);
if (
!matchedModel ||
!allowedModelIds.includes(matchedModel.modelId)
) {
return res.status(HttpCode.FORBIDDEN).json({
error: {
message: `Model "${requestedModel}" is not permitted on this resource`
}
});
}
const selection = await selectProvider(
attachments,
allowedModelIds,
requestedModel
);
if (!selection.ok) {
return res.status(selection.status).json({
error: { message: selection.message }
});
}
const { provider } = selection;
if (!provider.apiKey) {
return res.status(HttpCode.INTERNAL_SERVER_ERROR).json({
error: { message: "AI provider has no API key configured" }
+6 -19
View File
@@ -9,26 +9,16 @@ import { fromError } from "zod-validation-error";
import { OpenAPITags, registry } from "@server/openApi";
import { and, eq } from "drizzle-orm";
import type { CreateOrEditAiModelResponse } from "@server/routers/aiProvider/types";
import {
aiBudgetUnitSchema,
refineBudgetFields
} from "@server/routers/aiProvider/validation";
const paramsSchema = z.strictObject({
providerId: z.coerce.number().int().positive()
});
const bodySchema = z
.strictObject({
modelKey: z.string().nonempty(),
name: z.string().nonempty(),
budgetAmount: z.number().positive().optional().nullable(),
budgetUnit: aiBudgetUnitSchema.optional().nullable(),
enabled: z.boolean().optional()
})
.superRefine((data, ctx) => {
refineBudgetFields(data, ctx);
});
const bodySchema = z.strictObject({
modelKey: z.string().nonempty(),
name: z.string().nonempty(),
enabled: z.boolean().optional()
});
registry.registerPath({
method: "put",
@@ -79,8 +69,7 @@ export async function createAiModel(
}
const { providerId } = parsedParams.data;
const { modelKey, name, budgetAmount, budgetUnit, enabled } =
parsedBody.data;
const { modelKey, name, enabled } = parsedBody.data;
const [provider] =
req.aiProvider && req.aiProvider.providerId === providerId
@@ -127,8 +116,6 @@ export async function createAiModel(
providerId,
modelKey,
name,
budgetAmount: budgetAmount ?? null,
budgetUnit: budgetUnit ?? null,
enabled: enabled ?? true,
createdAt: now,
updatedAt: now
@@ -13,10 +13,8 @@ import type { CreateOrEditAiProviderResponse } from "@server/routers/aiProvider/
import { toPublicAiProvider } from "@server/routers/aiProvider/types";
import {
aiAuthTypeSchema,
aiBudgetUnitSchema,
aiProviderTypeSchema,
aiRoutingModeSchema,
refineBudgetFields,
refineProviderUpstreamFields
} from "@server/routers/aiProvider/validation";
@@ -33,13 +31,10 @@ const bodySchema = z
authType: aiAuthTypeSchema.optional().nullable(),
routingMode: aiRoutingModeSchema.optional(),
skipTlsVerification: z.boolean().optional(),
budgetAmount: z.number().positive().optional().nullable(),
budgetUnit: aiBudgetUnitSchema.optional().nullable(),
enabled: z.boolean().optional()
})
.superRefine((data, ctx) => {
refineProviderUpstreamFields(data, ctx);
refineBudgetFields(data, ctx);
});
registry.registerPath({
@@ -99,8 +94,6 @@ export async function createAiProvider(
authType,
routingMode,
skipTlsVerification,
budgetAmount,
budgetUnit,
enabled
} = parsedBody.data;
@@ -126,8 +119,6 @@ export async function createAiProvider(
authType: authType ?? null,
routingMode: resolvedRoutingMode,
skipTlsVerification: skipTlsVerification ?? false,
budgetAmount: budgetAmount ?? null,
budgetUnit: budgetUnit ?? null,
enabled: enabled ?? true,
createdAt: now,
updatedAt: now
+5 -21
View File
@@ -9,26 +9,16 @@ import { fromError } from "zod-validation-error";
import { OpenAPITags, registry } from "@server/openApi";
import { and, eq, ne } from "drizzle-orm";
import type { CreateOrEditAiModelResponse } from "@server/routers/aiProvider/types";
import {
aiBudgetUnitSchema,
refineBudgetFields
} from "@server/routers/aiProvider/validation";
const paramsSchema = z.strictObject({
modelId: z.coerce.number().int().positive()
});
const bodySchema = z
.strictObject({
modelKey: z.string().nonempty().optional(),
name: z.string().nonempty().optional(),
budgetAmount: z.number().positive().optional().nullable(),
budgetUnit: aiBudgetUnitSchema.optional().nullable(),
enabled: z.boolean().optional()
})
.superRefine((data, ctx) => {
refineBudgetFields(data, ctx);
});
const bodySchema = z.strictObject({
modelKey: z.string().nonempty().optional(),
name: z.string().nonempty().optional(),
enabled: z.boolean().optional()
});
registry.registerPath({
method: "post",
@@ -135,12 +125,6 @@ export async function updateAiModel(
if (body.name !== undefined) {
updateData.name = body.name;
}
if (body.budgetAmount !== undefined) {
updateData.budgetAmount = body.budgetAmount;
}
if (body.budgetUnit !== undefined) {
updateData.budgetUnit = body.budgetUnit;
}
if (body.enabled !== undefined) {
updateData.enabled = body.enabled;
}
+9 -23
View File
@@ -14,10 +14,8 @@ import type { CreateOrEditAiProviderResponse } from "@server/routers/aiProvider/
import { toPublicAiProvider } from "@server/routers/aiProvider/types";
import {
aiAuthTypeSchema,
aiBudgetUnitSchema,
aiProviderTypeSchema,
aiRoutingModeSchema,
refineBudgetFields,
refineProviderUpstreamFields
} from "@server/routers/aiProvider/validation";
import type {
@@ -29,21 +27,15 @@ const paramsSchema = z.strictObject({
providerId: z.coerce.number().int().positive()
});
const bodySchema = z
.strictObject({
name: z.string().nonempty().optional(),
upstreamUrl: z.url().optional().nullable(),
apiKey: z.string().optional(),
authType: aiAuthTypeSchema.optional().nullable(),
routingMode: aiRoutingModeSchema.optional(),
skipTlsVerification: z.boolean().optional(),
budgetAmount: z.number().positive().optional().nullable(),
budgetUnit: aiBudgetUnitSchema.optional().nullable(),
enabled: z.boolean().optional()
})
.superRefine((data, ctx) => {
refineBudgetFields(data, ctx);
});
const bodySchema = z.strictObject({
name: z.string().nonempty().optional(),
upstreamUrl: z.url().optional().nullable(),
apiKey: z.string().optional(),
authType: aiAuthTypeSchema.optional().nullable(),
routingMode: aiRoutingModeSchema.optional(),
skipTlsVerification: z.boolean().optional(),
enabled: z.boolean().optional()
});
registry.registerPath({
method: "post",
@@ -168,12 +160,6 @@ export async function updateAiProvider(
if (body.enabled !== undefined) {
updateData.enabled = body.enabled;
}
if (body.budgetAmount !== undefined) {
updateData.budgetAmount = body.budgetAmount;
}
if (body.budgetUnit !== undefined) {
updateData.budgetUnit = body.budgetUnit;
}
if (nextRoutingMode === "target") {
updateData.upstreamUrl = null;
} else if (body.upstreamUrl !== undefined) {
-24
View File
@@ -1,7 +1,6 @@
import { z } from "zod";
import {
providerRequiresUpstreamUrl,
type AiBudgetUnit,
type AiProviderRoutingMode,
type AiProviderType
} from "@server/lib/aiProviderDefaults";
@@ -18,33 +17,10 @@ export const aiProviderTypeSchema = z.enum([
"custom"
]);
export const aiBudgetUnitSchema = z.enum(["usd", "tokens"]);
export const aiAuthTypeSchema = z.enum(["bearer"]);
export const aiRoutingModeSchema = z.enum(["url", "target"]);
export function refineBudgetFields(
data: {
budgetAmount?: number | null;
budgetUnit?: AiBudgetUnit | null;
},
ctx: z.RefinementCtx
) {
const hasAmount =
data.budgetAmount !== undefined && data.budgetAmount !== null;
const hasUnit = data.budgetUnit !== undefined && data.budgetUnit !== null;
if (hasAmount !== hasUnit) {
ctx.addIssue({
code: "custom",
message:
"budgetAmount and budgetUnit must both be set or both omitted",
path: hasAmount ? ["budgetUnit"] : ["budgetAmount"]
});
}
}
export function refineProviderUpstreamFields(
data: {
type: AiProviderType;
+94
View File
@@ -414,6 +414,13 @@ authenticated.get(
siteResource.listSiteResourceAiModels
);
authenticated.get(
"/site-resource/:siteResourceId/ai-providers",
verifySiteResourceAccess,
verifyUserHasAction(ActionsEnum.listResourceAiModels),
siteResource.listSiteResourceAiProviders
);
authenticated.post(
"/site-resource/:siteResourceId/roles",
verifySiteResourceAccess,
@@ -432,6 +439,46 @@ authenticated.post(
siteResource.setSiteResourceAiModels
);
authenticated.post(
"/site-resource/:siteResourceId/ai-models/add",
verifySiteResourceAccess,
verifyUserHasAction(ActionsEnum.setResourceAiModels),
logActionAudit(ActionsEnum.setResourceAiModels),
siteResource.addAiModelToSiteResource
);
authenticated.post(
"/site-resource/:siteResourceId/ai-models/remove",
verifySiteResourceAccess,
verifyUserHasAction(ActionsEnum.setResourceAiModels),
logActionAudit(ActionsEnum.setResourceAiModels),
siteResource.removeAiModelFromSiteResource
);
authenticated.post(
"/site-resource/:siteResourceId/ai-providers",
verifySiteResourceAccess,
verifyUserHasAction(ActionsEnum.setResourceAiModels),
logActionAudit(ActionsEnum.setResourceAiModels),
siteResource.setSiteResourceAiProviders
);
authenticated.post(
"/site-resource/:siteResourceId/ai-providers/add",
verifySiteResourceAccess,
verifyUserHasAction(ActionsEnum.setResourceAiModels),
logActionAudit(ActionsEnum.setResourceAiModels),
siteResource.addAiProviderToSiteResource
);
authenticated.post(
"/site-resource/:siteResourceId/ai-providers/remove",
verifySiteResourceAccess,
verifyUserHasAction(ActionsEnum.setResourceAiModels),
logActionAudit(ActionsEnum.setResourceAiModels),
siteResource.removeAiProviderFromSiteResource
);
authenticated.post(
"/site-resource/:siteResourceId/users",
verifySiteResourceAccess,
@@ -673,6 +720,13 @@ authenticated.get(
resource.listResourceAiModels
);
authenticated.get(
"/resource/:resourceId/ai-providers",
verifyResourceAccess,
verifyUserHasAction(ActionsEnum.listResourceAiModels),
resource.listResourceAiProviders
);
authenticated.get(
"/resource/:resourceId",
verifyResourceAccess,
@@ -884,6 +938,46 @@ authenticated.post(
resource.setResourceAiModels
);
authenticated.post(
"/resource/:resourceId/ai-models/add",
verifyResourceAccess,
verifyUserHasAction(ActionsEnum.setResourceAiModels),
logActionAudit(ActionsEnum.setResourceAiModels),
resource.addAiModelToResource
);
authenticated.post(
"/resource/:resourceId/ai-models/remove",
verifyResourceAccess,
verifyUserHasAction(ActionsEnum.setResourceAiModels),
logActionAudit(ActionsEnum.setResourceAiModels),
resource.removeAiModelFromResource
);
authenticated.post(
"/resource/:resourceId/ai-providers",
verifyResourceAccess,
verifyUserHasAction(ActionsEnum.setResourceAiModels),
logActionAudit(ActionsEnum.setResourceAiModels),
resource.setResourceAiProviders
);
authenticated.post(
"/resource/:resourceId/ai-providers/add",
verifyResourceAccess,
verifyUserHasAction(ActionsEnum.setResourceAiModels),
logActionAudit(ActionsEnum.setResourceAiModels),
resource.addAiProviderToResource
);
authenticated.post(
"/resource/:resourceId/ai-providers/remove",
verifyResourceAccess,
verifyUserHasAction(ActionsEnum.setResourceAiModels),
logActionAudit(ActionsEnum.setResourceAiModels),
resource.removeAiProviderFromResource
);
authenticated.put(
"/resource-policy/:resourcePolicyId/access-control",
verifyResourcePolicyAccess,
+86
View File
@@ -256,6 +256,16 @@ authenticated.get(
siteResource.listSiteResourceAiModels
);
authenticated.get(
[
"/site-resource/:siteResourceId/ai-providers",
"/private-resource/:siteResourceId/ai-providers"
],
verifyApiKeySiteResourceAccess,
verifyApiKeyHasAction(ActionsEnum.listResourceAiModels),
siteResource.listSiteResourceAiProviders
);
authenticated.post(
[
"/site-resource/:siteResourceId/roles",
@@ -341,6 +351,39 @@ authenticated.post(
siteResource.removeAiModelFromSiteResource
);
authenticated.post(
[
"/site-resource/:siteResourceId/ai-providers",
"/private-resource/:siteResourceId/ai-providers"
],
verifyApiKeySiteResourceAccess,
verifyApiKeyHasAction(ActionsEnum.setResourceAiModels),
logActionAudit(ActionsEnum.setResourceAiModels),
siteResource.setSiteResourceAiProviders
);
authenticated.post(
[
"/site-resource/:siteResourceId/ai-providers/add",
"/private-resource/:siteResourceId/ai-providers/add"
],
verifyApiKeySiteResourceAccess,
verifyApiKeyHasAction(ActionsEnum.setResourceAiModels),
logActionAudit(ActionsEnum.setResourceAiModels),
siteResource.addAiProviderToSiteResource
);
authenticated.post(
[
"/site-resource/:siteResourceId/ai-providers/remove",
"/private-resource/:siteResourceId/ai-providers/remove"
],
verifyApiKeySiteResourceAccess,
verifyApiKeyHasAction(ActionsEnum.setResourceAiModels),
logActionAudit(ActionsEnum.setResourceAiModels),
siteResource.removeAiProviderFromSiteResource
);
authenticated.post(
[
"/site-resource/:siteResourceId/users/add",
@@ -563,6 +606,16 @@ authenticated.get(
resource.listResourceAiModels
);
authenticated.get(
[
"/resource/:resourceId/ai-providers",
"/public-resource/:resourceId/ai-providers"
],
verifyApiKeyResourceAccess,
verifyApiKeyHasAction(ActionsEnum.listResourceAiModels),
resource.listResourceAiProviders
);
authenticated.get(
["/resource/:resourceId", "/public-resource/:resourceId"],
verifyApiKeyResourceAccess,
@@ -775,6 +828,17 @@ authenticated.post(
resource.setResourceAiModels
);
authenticated.post(
[
"/resource/:resourceId/ai-providers",
"/public-resource/:resourceId/ai-providers"
],
verifyApiKeyResourceAccess,
verifyApiKeyHasAction(ActionsEnum.setResourceAiModels),
logActionAudit(ActionsEnum.setResourceAiModels),
resource.setResourceAiProviders
);
authenticated.post(
["/resource/:resourceId/users", "/public-resource/:resourceId/users"],
verifyApiKeyResourceAccess,
@@ -989,6 +1053,28 @@ authenticated.post(
resource.removeAiModelFromResource
);
authenticated.post(
[
"/resource/:resourceId/ai-providers/add",
"/public-resource/:resourceId/ai-providers/add"
],
verifyApiKeyResourceAccess,
verifyApiKeyHasAction(ActionsEnum.setResourceAiModels),
logActionAudit(ActionsEnum.setResourceAiModels),
resource.addAiProviderToResource
);
authenticated.post(
[
"/resource/:resourceId/ai-providers/remove",
"/public-resource/:resourceId/ai-providers/remove"
],
verifyApiKeyResourceAccess,
verifyApiKeyHasAction(ActionsEnum.setResourceAiModels),
logActionAudit(ActionsEnum.setResourceAiModels),
resource.removeAiProviderFromResource
);
authenticated.post(
[
"/resource/:resourceId/users/add",
+16 -28
View File
@@ -1,6 +1,6 @@
import { Request, Response, NextFunction } from "express";
import { z } from "zod";
import { db, resources, resourceAiModels, aiModels } from "@server/db";
import { db, resources, resourceAiModels } from "@server/db";
import { eq, and } from "drizzle-orm";
import response from "@server/lib/response";
import HttpCode from "@server/types/HttpCode";
@@ -8,7 +8,10 @@ import createHttpError from "http-errors";
import logger from "@server/logger";
import { fromError } from "zod-validation-error";
import { OpenAPITags, registry } from "@server/openApi";
import {
assertPublicAllowlistApiEligible,
assertModelsBelongToPublicAllowlistProviders
} from "@server/lib/aiInferenceResource";
const addAiModelToResourceBodySchema = z.strictObject({
modelId: z.int().positive()
});
@@ -21,7 +24,7 @@ registry.registerPath({
method: "post",
path: "/resource/{resourceId}/ai-models/add",
description:
"Add a single AI model to a resource's model restriction allow-list.",
"Add a single catalog model to an inference resource allowlist. Requires at least one attached AI provider in allowlist mode. The model must belong to a provider attached in allowlist mode.",
tags: [OpenAPITags.PublicResource],
request: {
params: addAiModelToResourceParamsSchema,
@@ -95,33 +98,18 @@ export async function addAiModelToResource(
);
}
if (!resource.aiProviderId) {
return next(
createHttpError(
HttpCode.BAD_REQUEST,
"Resource has no AI provider linked"
)
);
const eligibleError = await assertPublicAllowlistApiEligible(resource);
if (eligibleError) {
return next(createHttpError(HttpCode.BAD_REQUEST, eligibleError));
}
const [model] = await db
.select()
.from(aiModels)
.where(
and(
eq(aiModels.modelId, modelId),
eq(aiModels.providerId, resource.aiProviderId)
)
)
.limit(1);
if (!model) {
return next(
createHttpError(
HttpCode.NOT_FOUND,
"Model not found or does not belong to this resource's AI provider"
)
);
const modelError = await assertModelsBelongToPublicAllowlistProviders({
orgId: resource.orgId,
resourceId,
modelIds: [modelId]
});
if (modelError) {
return next(createHttpError(HttpCode.BAD_REQUEST, modelError));
}
const existingEntry = await db
@@ -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";
import { usageService } from "@server/lib/billing/usageService";
import { LimitId } from "@server/lib/billing";
import {
isInferenceFieldsError,
resolveProviderAttachments,
resourceAiProviderAttachmentSchema,
setPublicResourceAiProviders,
type ResourceAiProviderAttachment
} from "@server/lib/aiInferenceResource";
const createResourceParamsSchema = z.strictObject({
orgId: z.string()
@@ -98,13 +105,11 @@ const createHttpResourceSchema = z
authDaemonPort: z.int().positive().optional(),
authDaemonMode: z.enum(["site", "remote", "native"]).optional(),
// Inference settings
aiProviderId: z
.number()
.int()
.positive()
aiProviders: z
.array(resourceAiProviderAttachmentSchema)
.optional()
.describe(
"For inference-mode resources: the AI provider this resource proxies chat completions to."
"For inference-mode resources: AI providers to attach. Each entry may set modelAccessMode (passthrough, catalog, or allowlist); defaults to passthrough. At most one passthrough provider is allowed."
)
})
.refine(
@@ -377,11 +382,35 @@ async function createHttpResource(
authDaemonPort,
authDaemonMode,
pamMode,
aiProviderId
aiProviders: aiProviderInputs
} = parsedBody.data;
const subdomain = parsedBody.data.subdomain;
const stickySession = parsedBody.data.stickySession;
const effectiveMode = mode ?? "http";
let providerAttachments: ResourceAiProviderAttachment[] = [];
if (effectiveMode === "inference") {
const resolved = await resolveProviderAttachments({
orgId,
attachments: aiProviderInputs ?? [],
requireAtLeastOne: true
});
if (isInferenceFieldsError(resolved)) {
return next(
createHttpError(HttpCode.BAD_REQUEST, resolved.error)
);
}
providerAttachments = resolved;
} else if (aiProviderInputs && aiProviderInputs.length > 0) {
return next(
createHttpError(
HttpCode.BAD_REQUEST,
"AI providers can only be attached to inference-mode resources"
)
);
}
// Wildcard subdomains are a paid feature
if (subdomain && subdomain.includes("*")) {
const isLicensed = await isLicensedOrSubscribed(
@@ -422,7 +451,7 @@ async function createHttpResource(
}
if (
["ssh", "rdp", "vnc"].includes(mode!) &&
["ssh", "rdp", "vnc"].includes(effectiveMode) &&
!isLicensedOrSubscribed(
orgId!,
tierMatrix[TierFeature.AdvancedPublicResources]
@@ -555,7 +584,7 @@ async function createHttpResource(
orgId,
name,
subdomain: finalSubdomain,
mode: mode,
mode: effectiveMode,
pamMode: pamMode,
authDaemonMode: authDaemonMode,
authDaemonPort: authDaemonPort,
@@ -564,11 +593,18 @@ async function createHttpResource(
postAuthPath: postAuthPath,
wildcard,
health: "unknown",
defaultResourcePolicyId: defaultPolicy.resourcePolicyId,
aiProviderId: aiProviderId ?? null
defaultResourcePolicyId: defaultPolicy.resourcePolicyId
})
.returning();
if (providerAttachments.length > 0) {
await setPublicResourceAiProviders(
newResource[0].resourceId,
providerAttachments,
trx
);
}
await trx.insert(roleResources).values({
roleId: adminRole[0].roleId,
resourceId: newResource[0].resourceId
+4
View File
@@ -39,3 +39,7 @@ export * from "./listResourceAiModels";
export * from "./setResourceAiModels";
export * from "./addAiModelToResource";
export * from "./removeAiModelFromResource";
export * from "./listResourceAiProviders";
export * from "./setResourceAiProviders";
export * from "./addAiProviderToResource";
export * from "./removeAiProviderFromResource";
@@ -34,7 +34,7 @@ registry.registerPath({
method: "get",
path: "/resource/{resourceId}/ai-models",
description:
"List the AI models a resource is restricted to. An empty list means the resource is not restricted and every enabled model on its linked AI provider is allowed.",
"List catalog models on this resource's allowlist. Only enforced when modelAccessMode=allowlist; an empty allowlist denies all models.",
tags: [OpenAPITags.PublicResource],
request: {
params: listResourceAiModelsParamsSchema
@@ -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 { fromError } from "zod-validation-error";
import { OpenAPITags, registry } from "@server/openApi";
import { assertPublicAllowlistApiEligible } from "@server/lib/aiInferenceResource";
const removeAiModelFromResourceBodySchema = z.strictObject({
modelId: z.int().positive()
@@ -21,7 +22,7 @@ registry.registerPath({
method: "post",
path: "/resource/{resourceId}/ai-models/remove",
description:
"Remove a single AI model from a resource's model restriction allow-list.",
"Remove a single catalog model from an inference resource allowlist. Requires at least one attached AI provider in allowlist mode.",
tags: [OpenAPITags.PublicResource],
request: {
params: removeAiModelFromResourceParamsSchema,
@@ -97,6 +98,11 @@ export async function removeAiModelFromResource(
);
}
const eligibleError = await assertPublicAllowlistApiEligible(resource);
if (eligibleError) {
return next(createHttpError(HttpCode.BAD_REQUEST, eligibleError));
}
const existingEntry = await db
.select()
.from(resourceAiModels)
@@ -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 { z } from "zod";
import { db, resources, resourceAiModels, aiModels } from "@server/db";
import { eq, and, inArray } from "drizzle-orm";
import { db, resources, resourceAiModels } from "@server/db";
import { eq } from "drizzle-orm";
import response from "@server/lib/response";
import HttpCode from "@server/types/HttpCode";
import createHttpError from "http-errors";
import logger from "@server/logger";
import { fromError } from "zod-validation-error";
import { OpenAPITags, registry } from "@server/openApi";
import {
assertPublicAllowlistApiEligible,
assertModelsBelongToPublicAllowlistProviders
} from "@server/lib/aiInferenceResource";
const setResourceAiModelsBodySchema = z.strictObject({
modelIds: z.array(z.int().positive())
@@ -21,7 +25,7 @@ registry.registerPath({
method: "post",
path: "/resource/{resourceId}/ai-models",
description:
"Set the AI models a resource is restricted to. This replaces all existing restrictions. Pass an empty array to remove the restriction (allow every enabled model on the linked provider).",
"Replace the allowlist of catalog models for an inference resource. Requires at least one attached AI provider in allowlist mode. Models must belong to a provider attached in allowlist mode. An empty array denies all models.",
tags: [OpenAPITags.PublicResource],
request: {
params: setResourceAiModelsParamsSchema,
@@ -95,34 +99,18 @@ export async function setResourceAiModels(
);
}
if (modelIds.length > 0) {
if (!resource.aiProviderId) {
return next(
createHttpError(
HttpCode.BAD_REQUEST,
"Resource has no AI provider linked"
)
);
}
const eligibleError = await assertPublicAllowlistApiEligible(resource);
if (eligibleError) {
return next(createHttpError(HttpCode.BAD_REQUEST, eligibleError));
}
const validModels = await db
.select({ modelId: aiModels.modelId })
.from(aiModels)
.where(
and(
inArray(aiModels.modelId, modelIds),
eq(aiModels.providerId, resource.aiProviderId)
)
);
if (validModels.length !== new Set(modelIds).size) {
return next(
createHttpError(
HttpCode.BAD_REQUEST,
"One or more model IDs do not exist or do not belong to this resource's AI provider"
)
);
}
const modelError = await assertModelsBelongToPublicAllowlistProviders({
orgId: resource.orgId,
resourceId,
modelIds
});
if (modelError) {
return next(createHttpError(HttpCode.BAD_REQUEST, modelError));
}
await db.transaction(async (trx) => {
@@ -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()
.describe(
"ID of the resource policy to apply to this resource. Set to null to remove the resource policy and fall back to the inline policy settings."
),
aiProviderId: z
.number()
.int()
.positive()
.nullable()
.optional()
.describe(
"For inference-mode resources: the AI provider this resource proxies chat completions to. Set to null to unlink."
)
})
.refine((data) => Object.keys(data).length > 0, {
@@ -354,8 +345,10 @@ export async function updateResource(
);
}
if (["http", "ssh", "rdp", "vnc"].includes(resource.mode)) {
// HANDLE UPDATING HTTP RESOURCES
if (
["http", "ssh", "rdp", "vnc", "inference"].includes(resource.mode)
) {
// HANDLE UPDATING HTTP / BROWSER / INFERENCE RESOURCES
return await updateHttpResource(
{
req,
@@ -1,11 +1,6 @@
import { Request, Response, NextFunction } from "express";
import { z } from "zod";
import {
db,
siteResources,
siteResourceAiModels,
aiModels
} from "@server/db";
import { db, siteResources, siteResourceAiModels } from "@server/db";
import { eq, and } from "drizzle-orm";
import response from "@server/lib/response";
import HttpCode from "@server/types/HttpCode";
@@ -13,6 +8,10 @@ import createHttpError from "http-errors";
import logger from "@server/logger";
import { fromError } from "zod-validation-error";
import { OpenAPITags, registry } from "@server/openApi";
import {
assertSiteAllowlistApiEligible,
assertModelsBelongToSiteAllowlistProviders
} from "@server/lib/aiInferenceResource";
const addAiModelToSiteResourceBodySchema = z.strictObject({
modelId: z.int().positive()
@@ -26,7 +25,7 @@ registry.registerPath({
method: "post",
path: "/site-resource/{siteResourceId}/ai-models/add",
description:
"Add a single AI model to a site resource's model restriction allow-list.",
"Add a single catalog model to an inference site resource allowlist. Requires at least one attached AI provider in allowlist mode. The model must belong to a provider attached in allowlist mode.",
tags: [OpenAPITags.PrivateResource],
request: {
params: addAiModelToSiteResourceParamsSchema,
@@ -102,33 +101,19 @@ export async function addAiModelToSiteResource(
);
}
if (!siteResource.aiProviderId) {
return next(
createHttpError(
HttpCode.BAD_REQUEST,
"Site resource has no AI provider linked"
)
);
const eligibleError =
await assertSiteAllowlistApiEligible(siteResource);
if (eligibleError) {
return next(createHttpError(HttpCode.BAD_REQUEST, eligibleError));
}
const [model] = await db
.select()
.from(aiModels)
.where(
and(
eq(aiModels.modelId, modelId),
eq(aiModels.providerId, siteResource.aiProviderId)
)
)
.limit(1);
if (!model) {
return next(
createHttpError(
HttpCode.NOT_FOUND,
"Model not found or does not belong to this site resource's AI provider"
)
);
const modelError = await assertModelsBelongToSiteAllowlistProviders({
orgId: siteResource.orgId,
siteResourceId,
modelIds: [modelId]
});
if (modelError) {
return next(createHttpError(HttpCode.BAD_REQUEST, modelError));
}
const existingEntry = await db
@@ -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 { usageService } from "@server/lib/billing/usageService";
import { LimitId } from "@server/lib/billing";
import {
isInferenceFieldsError,
resolveProviderAttachments,
resourceAiProviderAttachmentSchema,
setSiteResourceAiProviders,
type ResourceAiProviderAttachment
} from "@server/lib/aiInferenceResource";
const createSiteResourceParamsSchema = z.strictObject({
orgId: z.string()
@@ -79,13 +86,11 @@ const createSiteResourceSchema = z
pamMode: z.enum(["passthrough", "push"]).optional(),
domainId: z.string().optional(), // only used for http mode, we need this to verify the alias is unique within the org
subdomain: z.string().optional(), // only used for http mode, we need this to verify the alias is unique within the org
aiProviderId: z
.number()
.int()
.positive()
aiProviders: z
.array(resourceAiProviderAttachmentSchema)
.optional()
.describe(
"For inference-mode site resources: the AI provider this resource proxies chat completions to."
"For inference-mode site resources: AI providers to attach. Each entry may set modelAccessMode (passthrough, catalog, or allowlist); defaults to passthrough. At most one passthrough provider is allowed."
)
})
.strict()
@@ -334,7 +339,7 @@ export async function createSiteResource(
pamMode,
domainId,
subdomain,
aiProviderId
aiProviders: aiProviderInputs
} = parsedBody.data;
// Backward compatibility: merge deprecated siteId into siteIds array
@@ -343,6 +348,28 @@ export async function createSiteResource(
siteIds.push(siteId);
}
let providerAttachments: ResourceAiProviderAttachment[] = [];
if (mode === "inference") {
const resolved = await resolveProviderAttachments({
orgId,
attachments: aiProviderInputs ?? [],
requireAtLeastOne: true
});
if (isInferenceFieldsError(resolved)) {
return next(
createHttpError(HttpCode.BAD_REQUEST, resolved.error)
);
}
providerAttachments = resolved;
} else if (aiProviderInputs && aiProviderInputs.length > 0) {
return next(
createHttpError(
HttpCode.BAD_REQUEST,
"AI providers can only be attached to inference-mode resources"
)
);
}
if (build == "saas") {
const usage = await usageService.getUsage(
orgId,
@@ -608,8 +635,7 @@ export async function createSiteResource(
domainId,
subdomain: finalSubdomain,
fullDomain,
requiresExitNodeConnection: mode === "inference", // in the future we might want to have different modes that do this
aiProviderId: aiProviderId ?? null
requiresExitNodeConnection: mode === "inference" // in the future we might want to have different modes that do this
};
if (isLicensedSshPam) {
if (authDaemonPort !== undefined)
@@ -625,6 +651,14 @@ export async function createSiteResource(
const siteResourceId = newSiteResource.siteResourceId;
if (providerAttachments.length > 0) {
await setSiteResourceAiProviders(
siteResourceId,
providerAttachments,
trx
);
}
//////////////////// update the associations ////////////////////
if (network) {
+4
View File
@@ -21,3 +21,7 @@ export * from "./listSiteResourceAiModels";
export * from "./setSiteResourceAiModels";
export * from "./addAiModelToSiteResource";
export * from "./removeAiModelFromSiteResource";
export * from "./listSiteResourceAiProviders";
export * from "./setSiteResourceAiProviders";
export * from "./addAiProviderToSiteResource";
export * from "./removeAiProviderFromSiteResource";
@@ -1,11 +1,6 @@
import { Request, Response, NextFunction } from "express";
import { z } from "zod";
import {
db,
siteResources,
siteResourceAiModels,
aiModels
} from "@server/db";
import { db, siteResources, siteResourceAiModels, aiModels } from "@server/db";
import { eq } from "drizzle-orm";
import response from "@server/lib/response";
import HttpCode from "@server/types/HttpCode";
@@ -27,10 +22,7 @@ async function query(siteResourceId: number) {
enabled: aiModels.enabled
})
.from(siteResourceAiModels)
.innerJoin(
aiModels,
eq(siteResourceAiModels.modelId, aiModels.modelId)
)
.innerJoin(aiModels, eq(siteResourceAiModels.modelId, aiModels.modelId))
.where(eq(siteResourceAiModels.siteResourceId, siteResourceId));
}
@@ -42,7 +34,7 @@ registry.registerPath({
method: "get",
path: "/site-resource/{siteResourceId}/ai-models",
description:
"List the AI models a site resource is restricted to. An empty list means the site resource is not restricted and every enabled model on its linked AI provider is allowed.",
"List catalog models on this site resource's allowlist. Only enforced when modelAccessMode=allowlist; an empty allowlist denies all models.",
tags: [OpenAPITags.PrivateResource],
request: {
params: listSiteResourceAiModelsParamsSchema
@@ -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 { fromError } from "zod-validation-error";
import { OpenAPITags, registry } from "@server/openApi";
import { assertSiteAllowlistApiEligible } from "@server/lib/aiInferenceResource";
const removeAiModelFromSiteResourceBodySchema = z.strictObject({
modelId: z.int().positive()
@@ -21,7 +22,7 @@ registry.registerPath({
method: "post",
path: "/site-resource/{siteResourceId}/ai-models/remove",
description:
"Remove a single AI model from a site resource's model restriction allow-list.",
"Remove a single catalog model from an inference site resource allowlist. Requires at least one attached AI provider in allowlist mode.",
tags: [OpenAPITags.PrivateResource],
request: {
params: removeAiModelFromSiteResourceParamsSchema,
@@ -96,6 +97,12 @@ export async function removeAiModelFromSiteResource(
);
}
const eligibleError =
await assertSiteAllowlistApiEligible(siteResource);
if (eligibleError) {
return next(createHttpError(HttpCode.BAD_REQUEST, eligibleError));
}
const existingEntry = await db
.select()
.from(siteResourceAiModels)
@@ -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 { z } from "zod";
import {
db,
siteResources,
siteResourceAiModels,
aiModels
} from "@server/db";
import { eq, and, inArray } from "drizzle-orm";
import { db, siteResources, siteResourceAiModels } from "@server/db";
import { eq } from "drizzle-orm";
import response from "@server/lib/response";
import HttpCode from "@server/types/HttpCode";
import createHttpError from "http-errors";
import logger from "@server/logger";
import { fromError } from "zod-validation-error";
import { OpenAPITags, registry } from "@server/openApi";
import {
assertSiteAllowlistApiEligible,
assertModelsBelongToSiteAllowlistProviders
} from "@server/lib/aiInferenceResource";
const setSiteResourceAiModelsBodySchema = z.strictObject({
modelIds: z.array(z.int().positive())
@@ -26,7 +25,7 @@ registry.registerPath({
method: "post",
path: "/site-resource/{siteResourceId}/ai-models",
description:
"Set the AI models a site resource is restricted to. This replaces all existing restrictions. Pass an empty array to remove the restriction (allow every enabled model on the linked provider).",
"Replace the allowlist of catalog models for an inference site resource. Requires at least one attached AI provider in allowlist mode. Models must belong to a provider attached in allowlist mode. An empty array denies all models.",
tags: [OpenAPITags.PrivateResource],
request: {
params: setSiteResourceAiModelsParamsSchema,
@@ -102,52 +101,33 @@ export async function setSiteResourceAiModels(
);
}
if (modelIds.length > 0) {
if (!siteResource.aiProviderId) {
return next(
createHttpError(
HttpCode.BAD_REQUEST,
"Site resource has no AI provider linked"
)
);
}
const eligibleError =
await assertSiteAllowlistApiEligible(siteResource);
if (eligibleError) {
return next(createHttpError(HttpCode.BAD_REQUEST, eligibleError));
}
const validModels = await db
.select({ modelId: aiModels.modelId })
.from(aiModels)
.where(
and(
inArray(aiModels.modelId, modelIds),
eq(aiModels.providerId, siteResource.aiProviderId)
)
);
if (validModels.length !== new Set(modelIds).size) {
return next(
createHttpError(
HttpCode.BAD_REQUEST,
"One or more model IDs do not exist or do not belong to this site resource's AI provider"
)
);
}
const modelError = await assertModelsBelongToSiteAllowlistProviders({
orgId: siteResource.orgId,
siteResourceId,
modelIds
});
if (modelError) {
return next(createHttpError(HttpCode.BAD_REQUEST, modelError));
}
await db.transaction(async (trx) => {
await trx
.delete(siteResourceAiModels)
.where(
eq(siteResourceAiModels.siteResourceId, siteResourceId)
);
.where(eq(siteResourceAiModels.siteResourceId, siteResourceId));
if (modelIds.length > 0) {
await trx
.insert(siteResourceAiModels)
.values(
modelIds.map((modelId) => ({
siteResourceId,
modelId
}))
);
await trx.insert(siteResourceAiModels).values(
modelIds.map((modelId) => ({
siteResourceId,
modelId
}))
);
}
});
@@ -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 { z } from "zod";
import { fromError } from "zod-validation-error";
import { clearSiteResourceAiConfig } from "@server/lib/aiInferenceResource";
const updateSiteResourceParamsSchema = z.strictObject({
siteResourceId: z.coerce.number().int().positive()
@@ -78,16 +79,7 @@ const updateSiteResourceSchema = z
authDaemonMode: z.enum(["site", "remote", "native"]).optional(),
pamMode: z.enum(["passthrough", "push"]).optional(),
domainId: z.string().optional(),
subdomain: z.string().optional(),
aiProviderId: z
.number()
.int()
.positive()
.nullable()
.optional()
.describe(
"For inference-mode site resources: the AI provider this resource proxies chat completions to. Set to null to unlink."
)
subdomain: z.string().optional()
})
.strict()
.refine(
@@ -341,8 +333,7 @@ export async function updateSiteResource(
authDaemonMode,
pamMode,
domainId,
subdomain,
aiProviderId
subdomain
} = parsedBody.data;
// Backward compatibility: merge deprecated siteId into siteIds array
@@ -608,12 +599,19 @@ export async function updateSiteResource(
networkId: mode === "inference" ? null : undefined,
requiresExitNodeConnection:
mode !== undefined ? mode === "inference" : undefined,
aiProviderId: aiProviderId,
...sshPamSet
})
.where(and(eq(siteResources.siteResourceId, siteResourceId)))
.returning();
const effectiveMode = mode ?? existingSiteResource.mode;
if (
existingSiteResource.mode === "inference" &&
effectiveMode !== "inference"
) {
await clearSiteResourceAiConfig(siteResourceId, trx);
}
//////////////////// update the associations ////////////////////
if (mode === "inference") {