mirror of
https://github.com/fosrl/pangolin.git
synced 2026-08-05 12:10:52 +02:00
add crud for adding providers and models to resources
This commit is contained in:
@@ -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>;
|
||||
|
||||
@@ -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>;
|
||||
|
||||
@@ -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;
|
||||
}
|
||||
@@ -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(
|
||||
|
||||
@@ -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(
|
||||
|
||||
@@ -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" }
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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;
|
||||
}
|
||||
|
||||
@@ -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) {
|
||||
|
||||
@@ -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;
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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",
|
||||
|
||||
@@ -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")
|
||||
);
|
||||
}
|
||||
}
|
||||
@@ -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
|
||||
|
||||
@@ -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")
|
||||
);
|
||||
}
|
||||
}
|
||||
@@ -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")
|
||||
);
|
||||
}
|
||||
}
|
||||
@@ -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) {
|
||||
|
||||
@@ -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") {
|
||||
|
||||
Reference in New Issue
Block a user