diff --git a/messages/en-US.json b/messages/en-US.json index e994c823c..8cc780829 100644 --- a/messages/en-US.json +++ b/messages/en-US.json @@ -1667,7 +1667,7 @@ "aiProviderTypeCustom": "Custom", "aiProviderTypeOpenaiDescription": "OpenAI API with default upstream URL", "aiProviderTypeAnthropicDescription": "Anthropic API with default upstream URL", - "aiProviderTypeGoogleGeminiDescription": "Google Gemini OpenAI-compatible endpoint", + "aiProviderTypeGoogleGeminiDescription": "Google Gemini generateContent API", "aiProviderTypeVertexAiDescription": "Google Vertex AI; upstream URL required", "aiProviderTypeBedrockDescription": "Amazon Bedrock Runtime", "aiProviderTypeMicrosoftFoundryDescription": "Microsoft Foundry; upstream URL required", @@ -1759,13 +1759,20 @@ "aiProviderMessageRemove": "This will permanently delete the provider and its models and targets. This cannot be undone.", "aiProviderErrorNoUpdate": "AI provider is not available to update", "aiProviderModels": "Models", - "aiProviderModelsDescription": "Define model names available on this provider. Requests must match one of these keys. Use * and ? as wildcards (for example gpt-4* or claude-?).", + "aiProviderModelsDescription": "Define allow and block patterns for this provider. Requests must match an allow pattern and must not match a block pattern. Use * and ? as wildcards (for example gpt-4* or claude-?). An empty allow list denies all models.", "aiProviderModelsPlaceholder": "Model name or pattern (e.g. gpt-4*)", + "aiProviderModelsAllow": "Allow List", + "aiProviderModelsAllowDescription": "Models that may be used through this provider. Empty means deny all.", + "aiProviderModelsAllowPlaceholder": "Allowed model or pattern (e.g. gpt-4*)", + "aiProviderModelsBlock": "Block List", + "aiProviderModelsBlockDescription": "Models to deny even if they match an allow pattern.", + "aiProviderModelsBlockPlaceholder": "Blocked model or pattern (e.g. gpt-4o-mini)", + "aiProviderModelsOverlapError": "These patterns cannot be on both lists: {keys}", "aiProviderModelsUpdated": "Models updated", "aiProviderModelsErrorUpdate": "Failed to update models", "aiResourceProviders": "Providers", "aiResourceProvidersDescription": "Choose which AI providers this inference resource can use", - "aiResourceProvidersHelp": "Models must be defined on each provider. Exact names and patterns that conflict (identical keys, or an exact key matching another provider's pattern) are not allowed across selected providers.", + "aiResourceProvidersHelp": "Each attached provider uses its own allow and block lists. Allow patterns that conflict across attached providers are not allowed.", "aiResourceProvidersSelect": "Select providers", "aiResourceProvidersEmpty": "No AI providers found", "aiResourceProvidersUpdated": "Providers updated", diff --git a/server/db/pg/schema/schema.ts b/server/db/pg/schema/schema.ts index 6bc038660..90dd5abd9 100644 --- a/server/db/pg/schema/schema.ts +++ b/server/db/pg/schema/schema.ts @@ -231,10 +231,10 @@ export const resourceAiProviders = pgTable( providerId: integer("providerId") .notNull() .references(() => aiProviders.providerId, { onDelete: "cascade" }), - modelAccessMode: varchar("modelAccessMode") - .$type<"catalog" | "allowlist">() + accessMode: varchar("accessMode") + .$type<"inherit" | "select">() .notNull() - .default("catalog") + .default("inherit") }, (t) => [primaryKey({ columns: [t.resourceId, t.providerId] })] ); @@ -247,7 +247,11 @@ export const resourceAiModels = pgTable( .references(() => resources.resourceId, { onDelete: "cascade" }), modelId: integer("modelId") .notNull() - .references(() => aiModels.modelId, { onDelete: "cascade" }) + .references(() => aiModels.modelId, { onDelete: "cascade" }), + listType: varchar("listType") + .$type<"allow" | "block">() + .notNull() + .default("allow") }, (t) => [primaryKey({ columns: [t.resourceId, t.modelId] })] ); @@ -522,10 +526,10 @@ export const siteResourceAiProviders = pgTable( providerId: integer("providerId") .notNull() .references(() => aiProviders.providerId, { onDelete: "cascade" }), - modelAccessMode: varchar("modelAccessMode") - .$type<"catalog" | "allowlist">() + accessMode: varchar("accessMode") + .$type<"inherit" | "select">() .notNull() - .default("catalog") + .default("inherit") }, (t) => [primaryKey({ columns: [t.siteResourceId, t.providerId] })] ); @@ -540,7 +544,11 @@ export const siteResourceAiModels = pgTable( }), modelId: integer("modelId") .notNull() - .references(() => aiModels.modelId, { onDelete: "cascade" }) + .references(() => aiModels.modelId, { onDelete: "cascade" }), + listType: varchar("listType") + .$type<"allow" | "block">() + .notNull() + .default("allow") }, (t) => [primaryKey({ columns: [t.siteResourceId, t.modelId] })] ); @@ -1678,6 +1686,10 @@ export const aiModels = pgTable( .references(() => aiProviders.providerId, { onDelete: "cascade" }), modelKey: varchar("modelKey").notNull(), name: varchar("name").notNull(), + listType: varchar("listType") + .$type<"allow" | "block">() + .notNull() + .default("allow"), enabled: boolean("enabled").notNull().default(true), createdAt: bigint("createdAt", { mode: "number" }).notNull(), updatedAt: bigint("updatedAt", { mode: "number" }).notNull() diff --git a/server/db/sqlite/schema/schema.ts b/server/db/sqlite/schema/schema.ts index 8130af9ea..dab445282 100644 --- a/server/db/sqlite/schema/schema.ts +++ b/server/db/sqlite/schema/schema.ts @@ -228,10 +228,10 @@ export const resourceAiProviders = sqliteTable( providerId: integer("providerId") .notNull() .references(() => aiProviders.providerId, { onDelete: "cascade" }), - modelAccessMode: text("modelAccessMode") - .$type<"catalog" | "allowlist">() + accessMode: text("accessMode") + .$type<"inherit" | "select">() .notNull() - .default("catalog") + .default("inherit") }, (t) => [primaryKey({ columns: [t.resourceId, t.providerId] })] ); @@ -244,7 +244,11 @@ export const resourceAiModels = sqliteTable( .references(() => resources.resourceId, { onDelete: "cascade" }), modelId: integer("modelId") .notNull() - .references(() => aiModels.modelId, { onDelete: "cascade" }) + .references(() => aiModels.modelId, { onDelete: "cascade" }), + listType: text("listType") + .$type<"allow" | "block">() + .notNull() + .default("allow") }, (t) => [primaryKey({ columns: [t.resourceId, t.modelId] })] ); @@ -507,10 +511,10 @@ export const siteResourceAiProviders = sqliteTable( providerId: integer("providerId") .notNull() .references(() => aiProviders.providerId, { onDelete: "cascade" }), - modelAccessMode: text("modelAccessMode") - .$type<"catalog" | "allowlist">() + accessMode: text("accessMode") + .$type<"inherit" | "select">() .notNull() - .default("catalog") + .default("inherit") }, (t) => [primaryKey({ columns: [t.siteResourceId, t.providerId] })] ); @@ -525,7 +529,11 @@ export const siteResourceAiModels = sqliteTable( }), modelId: integer("modelId") .notNull() - .references(() => aiModels.modelId, { onDelete: "cascade" }) + .references(() => aiModels.modelId, { onDelete: "cascade" }), + listType: text("listType") + .$type<"allow" | "block">() + .notNull() + .default("allow") }, (t) => [primaryKey({ columns: [t.siteResourceId, t.modelId] })] ); @@ -1660,6 +1668,10 @@ export const aiModels = sqliteTable( .references(() => aiProviders.providerId, { onDelete: "cascade" }), modelKey: text("modelKey").notNull(), name: text("name").notNull(), + listType: text("listType") + .$type<"allow" | "block">() + .notNull() + .default("allow"), enabled: integer("enabled", { mode: "boolean" }) .notNull() .default(true), diff --git a/server/lib/aiCapabilities.ts b/server/lib/aiCapabilities.ts index e899f83c3..f7b370691 100644 --- a/server/lib/aiCapabilities.ts +++ b/server/lib/aiCapabilities.ts @@ -40,18 +40,37 @@ function paramModel(req: Request): string | undefined { } /** - * Join base URL with a path, avoiding double slashes and a duplicated trailing - * /v1 when the inbound path already starts with /v1 and the base ends with /v1. + * Join a provider base URL with an inbound request path. */ export function joinUpstreamUrl(baseUrl: string, path: string): string { const base = baseUrl.replace(/\/+$/, ""); let suffix = path.startsWith("/") ? path : `/${path}`; - if ( - base.endsWith("/v1") && - (suffix === "/v1" || suffix.startsWith("/v1/")) - ) { - suffix = suffix.slice("/v1".length) || "/"; + let basePathname = "/"; + try { + basePathname = new URL(base).pathname.replace(/\/+$/, "") || "/"; + } catch { + // Fall through with "/" non-absolute bases are not expected in + // production, but keep joining usable for malformed input. + } + + if (basePathname !== "/") { + const baseSegs = basePathname.split("/").filter(Boolean); + const pathSegs = suffix.split("/").filter(Boolean); + const max = Math.min(baseSegs.length, pathSegs.length); + let overlap = 0; + for (let n = max; n >= 1; n--) { + const baseSuffix = baseSegs.slice(-n); + const pathPrefix = pathSegs.slice(0, n); + if (baseSuffix.every((seg, i) => seg === pathPrefix[i])) { + overlap = n; + break; + } + } + if (overlap > 0) { + const remaining = pathSegs.slice(overlap); + suffix = remaining.length > 0 ? `/${remaining.join("/")}` : "/"; + } } if (suffix === "/") { @@ -175,7 +194,7 @@ export const AI_PROVIDER_CAPABILITY_DEFAULTS: Record< > = { openai: ["openai_chat"], anthropic: ["anthropic_messages"], - googleGemini: ["openai_chat"], + googleGemini: ["gemini_generate_content"], vertexAi: ["google_generate_content"], bedrock: ["bedrock_converse"], microsoftFoundry: ["openai_chat"], diff --git a/server/lib/aiInferenceResource.ts b/server/lib/aiInferenceResource.ts index d01b30dd9..51333473a 100644 --- a/server/lib/aiInferenceResource.ts +++ b/server/lib/aiInferenceResource.ts @@ -14,13 +14,17 @@ import { modelKeysConflict } from "@server/lib/aiModelKeyMatch"; type DbOrTrx = Transaction | typeof db; -export const modelAccessModeSchema = z.enum(["catalog", "allowlist"]); +export const modelListTypeSchema = z.enum(["allow", "block"]); -export type ModelAccessMode = z.infer; +export type ModelListType = z.infer; + +export const accessModeSchema = z.enum(["inherit", "select"]); + +export type AccessMode = z.infer; export const resourceAiProviderAttachmentSchema = z.strictObject({ providerId: z.number().int().positive(), - modelAccessMode: modelAccessModeSchema.optional() + accessMode: accessModeSchema.optional().default("inherit") }); export type ResourceAiProviderInput = z.infer< @@ -29,9 +33,16 @@ export type ResourceAiProviderInput = z.infer< export type ResourceAiProviderAttachment = { providerId: number; - modelAccessMode: ModelAccessMode; + accessMode: AccessMode; }; +export const resourceAiModelEntrySchema = z.strictObject({ + modelId: z.number().int().positive(), + listType: modelListTypeSchema +}); + +export type ResourceAiModelEntry = z.infer; + export type InferenceFieldsError = { error: string; }; @@ -42,58 +53,153 @@ export function isInferenceFieldsError( return "error" in value; } +/** + * Resolve which allow/block patterns apply for an attachment. + * inherit → provider lists; select → resource-selected lists (replace). + */ +export function resolveEffectiveLists(input: { + accessMode: AccessMode; + providerAllows: string[]; + providerBlocks: string[]; + resourceAllows: string[]; + resourceBlocks: string[]; +}): { allows: string[]; blocks: string[] } { + if (input.accessMode === "select") { + return { + allows: input.resourceAllows, + blocks: input.resourceBlocks + }; + } + return { + allows: input.providerAllows, + blocks: input.providerBlocks + }; +} + function normalizeAttachments( inputs: ResourceAiProviderInput[] ): ResourceAiProviderAttachment[] { - const byProvider = new Map(); + const byProviderId = new Map(); for (const input of inputs) { - byProvider.set(input.providerId, input.modelAccessMode ?? "catalog"); + byProviderId.set(input.providerId, input.accessMode ?? "inherit"); } - return [...byProvider.entries()].map(([providerId, modelAccessMode]) => ({ + return [...byProviderId.entries()].map(([providerId, accessMode]) => ({ providerId, - modelAccessMode + accessMode })); } +type EffectiveAllowRow = { + providerId: number; + modelKey: string; +}; + /** - * Ensure enabled catalog modelKeys do not conflict across attached providers. - * Catalog attachments contribute all enabled models on the provider. - * Allowlist attachments contribute nothing until models are allowlisted - * (those are checked when the allowlist is set). - * - * Conflicts: identical keys, or an exact key that matches another provider's - * pattern. Full glob intersections are left to runtime ambiguity errors. + * Ensure effective allow modelKeys do not conflict across attached providers. + * inherit uses provider allows; select uses resource-selected allows (or the + * optional override map). Block patterns are ignored for overlap checks. */ export async function assertNoOverlappingModelKeys( attachments: ResourceAiProviderAttachment[], - trx: DbOrTrx = db + options: { + trx?: DbOrTrx; + resourceId?: number; + siteResourceId?: number; + selectedAllowsByProvider?: Map; + } = {} ): Promise { - const catalogProviderIds = attachments - .filter((a) => a.modelAccessMode === "catalog") - .map((a) => a.providerId); + const trx = options.trx ?? db; - if (catalogProviderIds.length < 2) { + if (attachments.length < 2) { return null; } - const models = await trx - .select({ - providerId: aiModels.providerId, - modelKey: aiModels.modelKey - }) - .from(aiModels) - .where( - and( - inArray(aiModels.providerId, catalogProviderIds), - eq(aiModels.enabled, true) - ) - ); + const inheritProviderIds = attachments + .filter((a) => a.accessMode === "inherit") + .map((a) => a.providerId); + const selectProviderIds = attachments + .filter((a) => a.accessMode === "select") + .map((a) => a.providerId); + + const effectiveAllows: EffectiveAllowRow[] = []; + + if (inheritProviderIds.length > 0) { + const providerAllows = await trx + .select({ + providerId: aiModels.providerId, + modelKey: aiModels.modelKey + }) + .from(aiModels) + .where( + and( + inArray(aiModels.providerId, inheritProviderIds), + eq(aiModels.enabled, true), + eq(aiModels.listType, "allow") + ) + ); + effectiveAllows.push(...providerAllows); + } + + if (selectProviderIds.length > 0) { + if (options.selectedAllowsByProvider) { + for (const providerId of selectProviderIds) { + const keys = + options.selectedAllowsByProvider.get(providerId) ?? []; + for (const modelKey of keys) { + effectiveAllows.push({ providerId, modelKey }); + } + } + } else if (options.resourceId !== undefined) { + const rows = await trx + .select({ + providerId: aiModels.providerId, + modelKey: aiModels.modelKey + }) + .from(resourceAiModels) + .innerJoin( + aiModels, + eq(resourceAiModels.modelId, aiModels.modelId) + ) + .where( + and( + eq(resourceAiModels.resourceId, options.resourceId), + inArray(aiModels.providerId, selectProviderIds), + eq(resourceAiModels.listType, "allow"), + eq(aiModels.enabled, true) + ) + ); + effectiveAllows.push(...rows); + } else if (options.siteResourceId !== undefined) { + const rows = await trx + .select({ + providerId: aiModels.providerId, + modelKey: aiModels.modelKey + }) + .from(siteResourceAiModels) + .innerJoin( + aiModels, + eq(siteResourceAiModels.modelId, aiModels.modelId) + ) + .where( + and( + eq( + siteResourceAiModels.siteResourceId, + options.siteResourceId + ), + inArray(aiModels.providerId, selectProviderIds), + eq(siteResourceAiModels.listType, "allow"), + eq(aiModels.enabled, true) + ) + ); + effectiveAllows.push(...rows); + } + } const conflictPairs: string[] = []; - for (let i = 0; i < models.length; i++) { - for (let j = i + 1; j < models.length; j++) { - const left = models[i]; - const right = models[j]; + for (let i = 0; i < effectiveAllows.length; i++) { + for (let j = i + 1; j < effectiveAllows.length; j++) { + const left = effectiveAllows[i]; + const right = effectiveAllows[j]; if (left.providerId === right.providerId) { continue; } @@ -124,6 +230,8 @@ export async function resolveProviderAttachments(input: { orgId: string; attachments: ResourceAiProviderInput[]; requireAtLeastOne: boolean; + resourceId?: number; + siteResourceId?: number; }): Promise { const attachments = normalizeAttachments(input.attachments); @@ -165,7 +273,10 @@ export async function resolveProviderAttachments(input: { }; } - const overlapError = await assertNoOverlappingModelKeys(attachments); + const overlapError = await assertNoOverlappingModelKeys(attachments, { + resourceId: input.resourceId, + siteResourceId: input.siteResourceId + }); if (overlapError) { return overlapError; } @@ -188,6 +299,11 @@ export async function assertInferenceModeAllowsProviderFields(input: { return null; } +/** + * Attach providers to a resource. Inherit attachments use the provider lists + * as-is (resource model rows for those providers are pruned). Select + * attachments keep resource-selected allow/block subsets. + */ export async function setPublicResourceAiProviders( resourceId: number, attachments: ResourceAiProviderAttachment[], @@ -202,12 +318,16 @@ export async function setPublicResourceAiProviders( attachments.map((a) => ({ resourceId, providerId: a.providerId, - modelAccessMode: a.modelAccessMode + accessMode: a.accessMode })) ); } - await prunePublicResourceAllowlistToAllowlistProviders(resourceId, trx); + await prunePublicResourceModelsToSelectProviders( + resourceId, + attachments, + trx + ); } export async function setSiteResourceAiProviders( @@ -224,12 +344,103 @@ export async function setSiteResourceAiProviders( attachments.map((a) => ({ siteResourceId, providerId: a.providerId, - modelAccessMode: a.modelAccessMode + accessMode: a.accessMode })) ); } - await pruneSiteResourceAllowlistToAllowlistProviders(siteResourceId, trx); + await pruneSiteResourceModelsToSelectProviders( + siteResourceId, + attachments, + trx + ); +} + +/** + * Keep resource model rows only for providers in select mode. + */ +async function prunePublicResourceModelsToSelectProviders( + resourceId: number, + attachments: ResourceAiProviderAttachment[], + trx: DbOrTrx +): Promise { + const selectProviderIds = attachments + .filter((a) => a.accessMode === "select") + .map((a) => a.providerId); + + if (selectProviderIds.length === 0) { + await trx + .delete(resourceAiModels) + .where(eq(resourceAiModels.resourceId, resourceId)); + return; + } + + const existing = await trx + .select({ + modelId: resourceAiModels.modelId, + providerId: aiModels.providerId + }) + .from(resourceAiModels) + .innerJoin(aiModels, eq(resourceAiModels.modelId, aiModels.modelId)) + .where(eq(resourceAiModels.resourceId, resourceId)); + + const allowed = new Set(selectProviderIds); + const toRemove = existing + .filter((row) => !allowed.has(row.providerId)) + .map((row) => row.modelId); + + if (toRemove.length > 0) { + await trx + .delete(resourceAiModels) + .where( + and( + eq(resourceAiModels.resourceId, resourceId), + inArray(resourceAiModels.modelId, toRemove) + ) + ); + } +} + +async function pruneSiteResourceModelsToSelectProviders( + siteResourceId: number, + attachments: ResourceAiProviderAttachment[], + trx: DbOrTrx +): Promise { + const selectProviderIds = attachments + .filter((a) => a.accessMode === "select") + .map((a) => a.providerId); + + if (selectProviderIds.length === 0) { + await trx + .delete(siteResourceAiModels) + .where(eq(siteResourceAiModels.siteResourceId, siteResourceId)); + return; + } + + const existing = await trx + .select({ + modelId: siteResourceAiModels.modelId, + providerId: aiModels.providerId + }) + .from(siteResourceAiModels) + .innerJoin(aiModels, eq(siteResourceAiModels.modelId, aiModels.modelId)) + .where(eq(siteResourceAiModels.siteResourceId, siteResourceId)); + + const allowed = new Set(selectProviderIds); + const toRemove = existing + .filter((row) => !allowed.has(row.providerId)) + .map((row) => row.modelId); + + if (toRemove.length > 0) { + await trx + .delete(siteResourceAiModels) + .where( + and( + eq(siteResourceAiModels.siteResourceId, siteResourceId), + inArray(siteResourceAiModels.modelId, toRemove) + ) + ); + } } export async function clearPublicResourceAiConfig( @@ -256,120 +467,14 @@ export async function clearSiteResourceAiConfig( .where(eq(siteResourceAiProviders.siteResourceId, siteResourceId)); } -async function prunePublicResourceAllowlistToAllowlistProviders( - resourceId: number, - trx: DbOrTrx = db -): Promise { - const allowlistProviders = await trx - .select({ providerId: resourceAiProviders.providerId }) - .from(resourceAiProviders) - .where( - and( - eq(resourceAiProviders.resourceId, resourceId), - eq(resourceAiProviders.modelAccessMode, "allowlist") - ) - ); - - if (allowlistProviders.length === 0) { - await trx - .delete(resourceAiModels) - .where(eq(resourceAiModels.resourceId, resourceId)); - return; - } - - const validModels = await trx - .select({ modelId: aiModels.modelId }) - .from(aiModels) - .where( - inArray( - aiModels.providerId, - allowlistProviders.map((p) => p.providerId) - ) - ); - const validIds = validModels.map((m) => m.modelId); - - const existing = await trx - .select({ modelId: resourceAiModels.modelId }) - .from(resourceAiModels) - .where(eq(resourceAiModels.resourceId, resourceId)); - - const toRemove = existing - .map((e) => e.modelId) - .filter((id) => !validIds.includes(id)); - - if (toRemove.length > 0) { - await trx - .delete(resourceAiModels) - .where( - and( - eq(resourceAiModels.resourceId, resourceId), - inArray(resourceAiModels.modelId, toRemove) - ) - ); - } -} - -async function pruneSiteResourceAllowlistToAllowlistProviders( - siteResourceId: number, - trx: DbOrTrx = db -): Promise { - const allowlistProviders = await trx - .select({ providerId: siteResourceAiProviders.providerId }) - .from(siteResourceAiProviders) - .where( - and( - eq(siteResourceAiProviders.siteResourceId, siteResourceId), - eq(siteResourceAiProviders.modelAccessMode, "allowlist") - ) - ); - - if (allowlistProviders.length === 0) { - await trx - .delete(siteResourceAiModels) - .where(eq(siteResourceAiModels.siteResourceId, siteResourceId)); - return; - } - - const validModels = await trx - .select({ modelId: aiModels.modelId }) - .from(aiModels) - .where( - inArray( - aiModels.providerId, - allowlistProviders.map((p) => p.providerId) - ) - ); - const validIds = validModels.map((m) => m.modelId); - - const existing = await trx - .select({ modelId: siteResourceAiModels.modelId }) - .from(siteResourceAiModels) - .where(eq(siteResourceAiModels.siteResourceId, siteResourceId)); - - const toRemove = existing - .map((e) => e.modelId) - .filter((id) => !validIds.includes(id)); - - if (toRemove.length > 0) { - await trx - .delete(siteResourceAiModels) - .where( - and( - eq(siteResourceAiModels.siteResourceId, siteResourceId), - inArray(siteResourceAiModels.modelId, toRemove) - ) - ); - } -} - export async function listPublicResourceAiProviders(resourceId: number) { return db .select({ providerId: resourceAiProviders.providerId, - modelAccessMode: resourceAiProviders.modelAccessMode, name: aiProviders.name, type: aiProviders.type, - enabled: aiProviders.enabled + enabled: aiProviders.enabled, + accessMode: resourceAiProviders.accessMode }) .from(resourceAiProviders) .innerJoin( @@ -383,10 +488,10 @@ export async function listSiteResourceAiProviders(siteResourceId: number) { return db .select({ providerId: siteResourceAiProviders.providerId, - modelAccessMode: siteResourceAiProviders.modelAccessMode, name: aiProviders.name, type: aiProviders.type, - enabled: aiProviders.enabled + enabled: aiProviders.enabled, + accessMode: siteResourceAiProviders.accessMode }) .from(siteResourceAiProviders) .innerJoin( @@ -397,15 +502,15 @@ export async function listSiteResourceAiProviders(siteResourceId: number) { } /** - * Allowlist APIs require an inference resource with at least one - * attached provider in allowlist mode. + * Model list APIs require an inference resource with at least one select-mode + * attached provider. */ -export async function assertPublicAllowlistApiEligible(resource: { +export async function assertPublicModelListApiEligible(resource: { resourceId: number; mode: string; }): Promise { if (resource.mode !== "inference") { - return "AI model allowlists are only supported on inference-mode resources"; + return "AI model lists are only supported on inference-mode resources"; } const [row] = await db @@ -414,23 +519,23 @@ export async function assertPublicAllowlistApiEligible(resource: { .where( and( eq(resourceAiProviders.resourceId, resource.resourceId), - eq(resourceAiProviders.modelAccessMode, "allowlist") + eq(resourceAiProviders.accessMode, "select") ) ) .limit(1); if (!row) { - return "Attach at least one AI provider with modelAccessMode=allowlist before managing allowed models"; + return "Set at least one attached AI provider to select mode before managing model lists"; } return null; } -export async function assertSiteAllowlistApiEligible(siteResource: { +export async function assertSiteModelListApiEligible(siteResource: { siteResourceId: number; mode: string; }): Promise { if (siteResource.mode !== "inference") { - return "AI model allowlists are only supported on inference-mode resources"; + return "AI model lists are only supported on inference-mode resources"; } const [row] = await db @@ -442,33 +547,36 @@ export async function assertSiteAllowlistApiEligible(siteResource: { siteResourceAiProviders.siteResourceId, siteResource.siteResourceId ), - eq(siteResourceAiProviders.modelAccessMode, "allowlist") + eq(siteResourceAiProviders.accessMode, "select") ) ) .limit(1); if (!row) { - return "Attach at least one AI provider with modelAccessMode=allowlist before managing allowed models"; + return "Set at least one attached AI provider to select mode before managing model lists"; } return null; } /** - * Models must belong to providers attached to this resource in allowlist mode, - * and those providers must belong to the resource's org. + * Resource model entries must belong to select-mode attached providers, and + * listType must match the provider catalog entry (allow→allow, block→block). */ -export async function assertModelsBelongToPublicAllowlistProviders(input: { +export async function assertPublicResourceModelEntriesValid(input: { orgId: string; resourceId: number; - modelIds: number[]; + models: ResourceAiModelEntry[]; }): Promise { - const uniqueIds = [...new Set(input.modelIds)]; - if (uniqueIds.length === 0) { + const uniqueModels = dedupeModelEntries(input.models); + if (uniqueModels.length === 0) { return null; } - const allowlistProviders = await db - .select({ providerId: resourceAiProviders.providerId }) + const attachments = await db + .select({ + providerId: resourceAiProviders.providerId, + accessMode: resourceAiProviders.accessMode + }) .from(resourceAiProviders) .innerJoin( aiProviders, @@ -477,49 +585,33 @@ export async function assertModelsBelongToPublicAllowlistProviders(input: { .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; + return assertModelEntriesValid({ + orgId: input.orgId, + modelEntries: uniqueModels, + attachments, + resourceLabel: "resource" + }); } -export async function assertModelsBelongToSiteAllowlistProviders(input: { +export async function assertSiteResourceModelEntriesValid(input: { orgId: string; siteResourceId: number; - modelIds: number[]; + models: ResourceAiModelEntry[]; }): Promise { - const uniqueIds = [...new Set(input.modelIds)]; - if (uniqueIds.length === 0) { + const uniqueModels = dedupeModelEntries(input.models); + if (uniqueModels.length === 0) { return null; } - const allowlistProviders = await db - .select({ providerId: siteResourceAiProviders.providerId }) + const attachments = await db + .select({ + providerId: siteResourceAiProviders.providerId, + accessMode: siteResourceAiProviders.accessMode + }) .from(siteResourceAiProviders) .innerJoin( aiProviders, @@ -531,32 +623,92 @@ export async function assertModelsBelongToSiteAllowlistProviders(input: { 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"; + return assertModelEntriesValid({ + orgId: input.orgId, + modelEntries: uniqueModels, + attachments, + resourceLabel: "site resource" + }); +} + +function dedupeModelEntries( + models: ResourceAiModelEntry[] +): ResourceAiModelEntry[] { + const byModelId = new Map( + models.map((m) => [m.modelId, m.listType] as const) + ); + return [...byModelId.entries()].map(([modelId, listType]) => ({ + modelId, + listType + })); +} + +async function assertModelEntriesValid(input: { + orgId: string; + modelEntries: ResourceAiModelEntry[]; + attachments: ResourceAiProviderAttachment[]; + resourceLabel: string; +}): Promise { + const selectProviderIds = input.attachments + .filter((a) => a.accessMode === "select") + .map((a) => a.providerId); + + if (selectProviderIds.length === 0) { + return "Set at least one attached AI provider to select mode before managing model lists"; } - const validModels = await db - .select({ modelId: aiModels.modelId }) + const modelIds = input.modelEntries.map((m) => m.modelId); + const catalogRows = await db + .select({ + modelId: aiModels.modelId, + modelKey: aiModels.modelKey, + listType: aiModels.listType, + providerId: aiModels.providerId, + enabled: aiModels.enabled + }) .from(aiModels) .innerJoin(aiProviders, eq(aiModels.providerId, aiProviders.providerId)) .where( and( - inArray(aiModels.modelId, uniqueIds), - inArray( - aiModels.providerId, - allowlistProviders.map((p) => p.providerId) - ), + inArray(aiModels.modelId, modelIds), + inArray(aiModels.providerId, selectProviderIds), 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"; + if (catalogRows.length !== modelIds.length) { + return `One or more model IDs do not exist or do not belong to a select-mode provider on this ${input.resourceLabel}`; + } + + const catalogById = new Map(catalogRows.map((row) => [row.modelId, row])); + const selectedAllowsByProvider = new Map(); + for (const entry of input.modelEntries) { + const catalog = catalogById.get(entry.modelId); + if (!catalog) { + return `One or more model IDs do not exist or do not belong to a select-mode provider on this ${input.resourceLabel}`; + } + if (catalog.listType !== entry.listType) { + return `Model ${entry.modelId} must use listType "${catalog.listType}" to match the provider catalog entry`; + } + if (!catalog.enabled) { + return `Model ${entry.modelId} is disabled on its provider`; + } + if (entry.listType === "allow") { + const keys = selectedAllowsByProvider.get(catalog.providerId) ?? []; + keys.push(catalog.modelKey); + selectedAllowsByProvider.set(catalog.providerId, keys); + } + } + + const overlapError = await assertNoOverlappingModelKeys(input.attachments, { + selectedAllowsByProvider + }); + if (overlapError) { + return overlapError.error; } return null; diff --git a/server/lib/aiModelKeyMatch.ts b/server/lib/aiModelKeyMatch.ts index fa27be56f..38184c3b0 100644 --- a/server/lib/aiModelKeyMatch.ts +++ b/server/lib/aiModelKeyMatch.ts @@ -83,3 +83,42 @@ export function modelKeysConflict(a: string, b: string): boolean { return modelKeyMatches(b, a); } + +/** + * Provider-layer policy: empty allowlist denies all. Blocklist only applies + * after an allow match. + */ +export function isAllowedByLists( + requested: string, + allows: string[], + blocks: string[] +): boolean { + if (allows.length === 0) { + return false; + } + if (!allows.some((pattern) => modelKeyMatches(pattern, requested))) { + return false; + } + if (blocks.some((pattern) => modelKeyMatches(pattern, requested))) { + return false; + } + return true; +} + +/** + * Among allow patterns that match `requested`, return the most specific one, + * or null if none match. + */ +export function mostSpecificMatchingAllow( + requested: string, + allows: string[] +): string | null { + const matching = allows.filter((pattern) => + modelKeyMatches(pattern, requested) + ); + if (matching.length === 0) { + return null; + } + matching.sort(compareModelKeySpecificity); + return matching[0]; +} diff --git a/server/lib/aiProviderDefaults.ts b/server/lib/aiProviderDefaults.ts index e48c6f877..f24330ec8 100644 --- a/server/lib/aiProviderDefaults.ts +++ b/server/lib/aiProviderDefaults.ts @@ -43,7 +43,7 @@ export const AI_PROVIDER_DEFAULTS: Record< authType: "x-api-key" }, googleGemini: { - upstreamUrl: "https://generativelanguage.googleapis.com/v1beta/openai/", + upstreamUrl: "https://generativelanguage.googleapis.com", authType: "x-goog-api-key" }, vertexAi: { diff --git a/server/routers/aiGateway/pipeline.ts b/server/routers/aiGateway/pipeline.ts index 04d318004..d881b27ca 100644 --- a/server/routers/aiGateway/pipeline.ts +++ b/server/routers/aiGateway/pipeline.ts @@ -38,10 +38,15 @@ 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"; +import { + resolveEffectiveLists, + type AccessMode, + type ModelListType +} from "@server/lib/aiInferenceResource"; import { compareModelKeySpecificity, - modelKeyMatches + isAllowedByLists, + mostSpecificMatchingAllow } from "@server/lib/aiModelKeyMatch"; import { aiGatewayUpstreamFetch } from "@server/lib/aiGatewayUpstreamFetch"; @@ -96,7 +101,19 @@ async function findClientByIp(ip: string): Promise { type ProviderAttachment = { provider: AiProvider; - modelAccessMode: ModelAccessMode; + accessMode: AccessMode; +}; + +type ResourceModelPattern = { + providerId: number; + modelKey: string; + listType: ModelListType; + enabled: boolean; +}; + +type ProviderPatternLists = { + allows: string[]; + blocks: string[]; }; type ResolvedTarget = { @@ -104,7 +121,7 @@ type ResolvedTarget = { siteResourceId: number | null; orgId: string | null; attachments: ProviderAttachment[]; - allowlistedModelIds: Set; + resourceListsByProvider: Map; }; type ProviderSelection = @@ -231,8 +248,8 @@ async function resolveTarget(host: string): Promise { if (resourceRow) { const attachmentRows = await db .select({ - modelAccessMode: resourceAiProviders.modelAccessMode, - provider: aiProviders + provider: aiProviders, + accessMode: resourceAiProviders.accessMode }) .from(resourceAiProviders) .innerJoin( @@ -252,38 +269,26 @@ async function resolveTarget(host: string): Promise { const attachments: ProviderAttachment[] = attachmentRows.map((a) => ({ provider: a.provider, - modelAccessMode: a.modelAccessMode as ModelAccessMode + accessMode: a.accessMode })); - const allowlistProviderIds = attachments - .filter((a) => a.modelAccessMode === "allowlist") - .map((a) => a.provider.providerId); - const allowlistedModelIds = new Set(); - if (allowlistProviderIds.length > 0) { - const restrictions = await db - .select({ modelId: resourceAiModels.modelId }) - .from(resourceAiModels) - .innerJoin( - aiModels, - eq(resourceAiModels.modelId, aiModels.modelId) - ) - .where( - and( - eq(resourceAiModels.resourceId, resourceRow.resourceId), - inArray(aiModels.providerId, allowlistProviderIds) - ) - ); - for (const row of restrictions) { - allowlistedModelIds.add(row.modelId); - } - } + const resourcePatterns = await db + .select({ + providerId: aiModels.providerId, + modelKey: aiModels.modelKey, + listType: resourceAiModels.listType, + enabled: aiModels.enabled + }) + .from(resourceAiModels) + .innerJoin(aiModels, eq(resourceAiModels.modelId, aiModels.modelId)) + .where(eq(resourceAiModels.resourceId, resourceRow.resourceId)); return { resourceId: resourceRow.resourceId, siteResourceId: null, orgId: resourceRow.orgId, attachments, - allowlistedModelIds + resourceListsByProvider: groupPatternsByProvider(resourcePatterns) }; } @@ -305,8 +310,8 @@ async function resolveTarget(host: string): Promise { if (siteResourceRow) { const attachmentRows = await db .select({ - modelAccessMode: siteResourceAiProviders.modelAccessMode, - provider: aiProviders + provider: aiProviders, + accessMode: siteResourceAiProviders.accessMode }) .from(siteResourceAiProviders) .innerJoin( @@ -329,50 +334,65 @@ async function resolveTarget(host: string): Promise { const attachments: ProviderAttachment[] = attachmentRows.map((a) => ({ provider: a.provider, - modelAccessMode: a.modelAccessMode as ModelAccessMode + accessMode: a.accessMode })); - const allowlistProviderIds = attachments - .filter((a) => a.modelAccessMode === "allowlist") - .map((a) => a.provider.providerId); - const allowlistedModelIds = new Set(); - if (allowlistProviderIds.length > 0) { - const restrictions = await db - .select({ modelId: siteResourceAiModels.modelId }) - .from(siteResourceAiModels) - .innerJoin( - aiModels, - eq(siteResourceAiModels.modelId, aiModels.modelId) + const resourcePatterns = await db + .select({ + providerId: aiModels.providerId, + modelKey: aiModels.modelKey, + listType: siteResourceAiModels.listType, + enabled: aiModels.enabled + }) + .from(siteResourceAiModels) + .innerJoin( + aiModels, + eq(siteResourceAiModels.modelId, aiModels.modelId) + ) + .where( + eq( + siteResourceAiModels.siteResourceId, + siteResourceRow.siteResourceId ) - .where( - and( - eq( - siteResourceAiModels.siteResourceId, - siteResourceRow.siteResourceId - ), - inArray(aiModels.providerId, allowlistProviderIds) - ) - ); - for (const row of restrictions) { - allowlistedModelIds.add(row.modelId); - } - } + ); return { resourceId: null, siteResourceId: siteResourceRow.siteResourceId, orgId: siteResourceRow.orgId, attachments, - allowlistedModelIds + resourceListsByProvider: groupPatternsByProvider(resourcePatterns) }; } return null; } +function groupPatternsByProvider( + patterns: ResourceModelPattern[] +): Map { + const byProvider = new Map(); + for (const pattern of patterns) { + if (!pattern.enabled) { + continue; + } + let lists = byProvider.get(pattern.providerId); + if (!lists) { + lists = { allows: [], blocks: [] }; + byProvider.set(pattern.providerId, lists); + } + if (pattern.listType === "allow") { + lists.allows.push(pattern.modelKey); + } else { + lists.blocks.push(pattern.modelKey); + } + } + return byProvider; +} + async function selectProvider( attachments: ProviderAttachment[], - allowlistedModelIds: Set, + resourceListsByProvider: Map, requestedModel: string | undefined ): Promise { if (!requestedModel) { @@ -383,10 +403,10 @@ async function selectProvider( }; } - const providerById = new Map( + const attachmentByProviderId = new Map( attachments.map((a) => [a.provider.providerId, a]) ); - const providerIds = [...providerById.keys()]; + const providerIds = [...attachmentByProviderId.keys()]; if (providerIds.length === 0) { return { ok: false, @@ -397,48 +417,54 @@ async function selectProvider( const providerModels = await db .select({ - modelId: aiModels.modelId, providerId: aiModels.providerId, modelKey: aiModels.modelKey, + listType: aiModels.listType, enabled: aiModels.enabled }) .from(aiModels) .where(inArray(aiModels.providerId, providerIds)); + const allowsByProvider = new Map(); + const blocksByProvider = new Map(); + for (const model of providerModels) { + if (!model.enabled) { + continue; + } + const targetMap = + model.listType === "allow" ? allowsByProvider : blocksByProvider; + const existing = targetMap.get(model.providerId) ?? []; + existing.push(model.modelKey); + targetMap.set(model.providerId, existing); + } + type ModelCandidate = { provider: AiProvider; modelKey: string; }; const candidates: ModelCandidate[] = []; - for (const model of providerModels) { - if (!model.enabled) { + for (const [providerId, attachment] of attachmentByProviderId) { + const resourceLists = resourceListsByProvider.get(providerId); + const { allows, blocks } = resolveEffectiveLists({ + accessMode: attachment.accessMode, + providerAllows: allowsByProvider.get(providerId) ?? [], + providerBlocks: blocksByProvider.get(providerId) ?? [], + resourceAllows: resourceLists?.allows ?? [], + resourceBlocks: resourceLists?.blocks ?? [] + }); + + if (!isAllowedByLists(requestedModel, allows, blocks)) { continue; } - - if (!modelKeyMatches(model.modelKey, requestedModel)) { + const matchingAllow = mostSpecificMatchingAllow(requestedModel, allows); + if (!matchingAllow) { continue; } - - const attachment = providerById.get(model.providerId); - if (!attachment) { - continue; - } - - if (attachment.modelAccessMode === "catalog") { - candidates.push({ - provider: attachment.provider, - modelKey: model.modelKey - }); - continue; - } - - if (allowlistedModelIds.has(model.modelId)) { - candidates.push({ - provider: attachment.provider, - modelKey: model.modelKey - }); - } + candidates.push({ + provider: attachment.provider, + modelKey: matchingAllow + }); } if (candidates.length === 0) { @@ -504,7 +530,8 @@ export async function handleAiGatewayProxy( }); } - const { attachments, allowlistedModelIds, resourceId, orgId } = target; + const { attachments, resourceListsByProvider, resourceId, orgId } = + target; const requestUser = await resolveRequestUser(req, resourceId, orgId); if (requestUser) { @@ -529,7 +556,7 @@ export async function handleAiGatewayProxy( const selection = await selectProvider( capableAttachments, - allowlistedModelIds, + resourceListsByProvider, requestedModel ); if (!selection.ok) { diff --git a/server/routers/aiProvider/createAiModel.ts b/server/routers/aiProvider/createAiModel.ts index dfe2a62e1..4a2734bbd 100644 --- a/server/routers/aiProvider/createAiModel.ts +++ b/server/routers/aiProvider/createAiModel.ts @@ -9,6 +9,7 @@ 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 { modelListTypeSchema } from "@server/lib/aiInferenceResource"; const paramsSchema = z.strictObject({ providerId: z.coerce.number().int().positive() @@ -17,7 +18,8 @@ const paramsSchema = z.strictObject({ const bodySchema = z.strictObject({ modelKey: z.string().nonempty(), name: z.string().nonempty(), - enabled: z.boolean().optional() + enabled: z.boolean().optional(), + listType: modelListTypeSchema.optional().default("allow") }); registry.registerPath({ @@ -69,7 +71,7 @@ export async function createAiModel( } const { providerId } = parsedParams.data; - const { modelKey, name, enabled } = parsedBody.data; + const { modelKey, name, enabled, listType } = parsedBody.data; const [provider] = req.aiProvider && req.aiProvider.providerId === providerId @@ -116,6 +118,7 @@ export async function createAiModel( providerId, modelKey, name, + listType, enabled: enabled ?? true, createdAt: now, updatedAt: now diff --git a/server/routers/aiProvider/updateAiModel.ts b/server/routers/aiProvider/updateAiModel.ts index fe9ccf273..c0d1514d2 100644 --- a/server/routers/aiProvider/updateAiModel.ts +++ b/server/routers/aiProvider/updateAiModel.ts @@ -9,6 +9,7 @@ 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 { modelListTypeSchema } from "@server/lib/aiInferenceResource"; const paramsSchema = z.strictObject({ modelId: z.coerce.number().int().positive() @@ -17,7 +18,8 @@ const paramsSchema = z.strictObject({ const bodySchema = z.strictObject({ modelKey: z.string().nonempty().optional(), name: z.string().nonempty().optional(), - enabled: z.boolean().optional() + enabled: z.boolean().optional(), + listType: modelListTypeSchema.optional() }); registry.registerPath({ @@ -128,6 +130,9 @@ export async function updateAiModel( if (body.enabled !== undefined) { updateData.enabled = body.enabled; } + if (body.listType !== undefined) { + updateData.listType = body.listType; + } const [model] = await db .update(aiModels) diff --git a/server/routers/resource/addAiModelToResource.ts b/server/routers/resource/addAiModelToResource.ts index 3121c9ffa..4777bfe4d 100644 --- a/server/routers/resource/addAiModelToResource.ts +++ b/server/routers/resource/addAiModelToResource.ts @@ -9,11 +9,14 @@ import logger from "@server/logger"; import { fromError } from "zod-validation-error"; import { OpenAPITags, registry } from "@server/openApi"; import { - assertPublicAllowlistApiEligible, - assertModelsBelongToPublicAllowlistProviders + assertPublicModelListApiEligible, + assertPublicResourceModelEntriesValid, + modelListTypeSchema } from "@server/lib/aiInferenceResource"; + const addAiModelToResourceBodySchema = z.strictObject({ - modelId: z.int().positive() + modelId: z.number().int().positive(), + listType: modelListTypeSchema.optional().default("allow") }); const addAiModelToResourceParamsSchema = z.strictObject({ @@ -24,7 +27,7 @@ registry.registerPath({ method: "post", path: "/resource/{resourceId}/ai-models/add", description: - "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.", + "Add a single model to an inference resource allow/block selection. Requires at least one attached AI provider in select mode. The model must belong to a select-mode provider and its listType must match the provider catalog entry. listType defaults to allow.", tags: [OpenAPITags.PublicResource], request: { params: addAiModelToResourceParamsSchema, @@ -70,7 +73,7 @@ export async function addAiModelToResource( ); } - const { modelId } = parsedBody.data; + const { modelId, listType } = parsedBody.data; const parsedParams = addAiModelToResourceParamsSchema.safeParse( req.params @@ -98,15 +101,15 @@ export async function addAiModelToResource( ); } - const eligibleError = await assertPublicAllowlistApiEligible(resource); + const eligibleError = await assertPublicModelListApiEligible(resource); if (eligibleError) { return next(createHttpError(HttpCode.BAD_REQUEST, eligibleError)); } - const modelError = await assertModelsBelongToPublicAllowlistProviders({ + const modelError = await assertPublicResourceModelEntriesValid({ orgId: resource.orgId, resourceId, - modelIds: [modelId] + models: [{ modelId, listType }] }); if (modelError) { return next(createHttpError(HttpCode.BAD_REQUEST, modelError)); @@ -131,7 +134,9 @@ export async function addAiModelToResource( ); } - await db.insert(resourceAiModels).values({ resourceId, modelId }); + await db + .insert(resourceAiModels) + .values({ resourceId, modelId, listType }); return response(res, { data: {}, diff --git a/server/routers/resource/addAiProviderToResource.ts b/server/routers/resource/addAiProviderToResource.ts index d51cf244c..bcc175c18 100644 --- a/server/routers/resource/addAiProviderToResource.ts +++ b/server/routers/resource/addAiProviderToResource.ts @@ -11,14 +11,12 @@ 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() + providerId: z.number().int().positive() }); const addAiProviderToResourceParamsSchema = z.strictObject({ @@ -29,7 +27,7 @@ registry.registerPath({ method: "post", path: "/resource/{resourceId}/ai-providers/add", description: - "Add or replace a single AI provider attachment on an inference resource.", + "Add or replace a single AI provider attachment on an inference resource. The provider is attached in inherit mode, using its own allow/block lists.", tags: [OpenAPITags.PublicResource], request: { params: addAiProviderToResourceParamsSchema, @@ -77,7 +75,7 @@ export async function addAiProviderToResource( ); } - const { providerId, modelAccessMode } = parsedBody.data; + const { providerId } = parsedBody.data; const parsedParams = addAiProviderToResourceParamsSchema.safeParse( req.params @@ -120,18 +118,21 @@ export async function addAiProviderToResource( .filter((a) => a.providerId !== providerId) .map((a) => ({ providerId: a.providerId, - modelAccessMode: a.modelAccessMode + accessMode: a.accessMode })), - { providerId, modelAccessMode } + { providerId, accessMode: "inherit" as const } ]; const attachments = await resolveProviderAttachments({ orgId: resource.orgId, attachments: nextAttachments, - requireAtLeastOne: true + requireAtLeastOne: true, + resourceId }); if (isInferenceFieldsError(attachments)) { - return next(createHttpError(HttpCode.BAD_REQUEST, attachments.error)); + return next( + createHttpError(HttpCode.BAD_REQUEST, attachments.error) + ); } await setPublicResourceAiProviders(resourceId, attachments); diff --git a/server/routers/resource/createResource.ts b/server/routers/resource/createResource.ts index 47e70140f..e225458d1 100644 --- a/server/routers/resource/createResource.ts +++ b/server/routers/resource/createResource.ts @@ -109,7 +109,7 @@ const createHttpResourceSchema = z .array(resourceAiProviderAttachmentSchema) .optional() .describe( - "For inference-mode resources: AI providers to attach. Each entry may set modelAccessMode (catalog or allowlist); defaults to catalog. Model keys must be unique across attached catalog providers." + "For inference-mode resources: AI providers to attach. Providers are attached in inherit mode, using each provider's own allow/block lists. Effective allow model keys must be unique across attached providers." ) }) .refine( @@ -391,9 +391,14 @@ async function createHttpResource( let providerAttachments: ResourceAiProviderAttachment[] = []; if (effectiveMode === "inference") { + // A new resource has no model selections yet, so providers always start + // in inherit mode; select can be enabled afterwards. const resolved = await resolveProviderAttachments({ orgId, - attachments: aiProviderInputs ?? [], + attachments: (aiProviderInputs ?? []).map((p) => ({ + providerId: p.providerId, + accessMode: "inherit" as const + })), requireAtLeastOne: false }); if (isInferenceFieldsError(resolved)) { diff --git a/server/routers/resource/listResourceAiModels.ts b/server/routers/resource/listResourceAiModels.ts index 0a5dea778..f16e0c2da 100644 --- a/server/routers/resource/listResourceAiModels.ts +++ b/server/routers/resource/listResourceAiModels.ts @@ -19,7 +19,9 @@ async function query(resourceId: number) { modelId: aiModels.modelId, modelKey: aiModels.modelKey, name: aiModels.name, - enabled: aiModels.enabled + providerId: aiModels.providerId, + enabled: aiModels.enabled, + listType: resourceAiModels.listType }) .from(resourceAiModels) .innerJoin(aiModels, eq(resourceAiModels.modelId, aiModels.modelId)) @@ -34,7 +36,7 @@ registry.registerPath({ method: "get", path: "/resource/{resourceId}/ai-models", description: - "List catalog models on this resource's allowlist. Only enforced when modelAccessMode=allowlist; an empty allowlist denies all models.", + "List the models this resource has selected from its select-mode providers' allow/block lists. Providers in inherit mode are not represented here; they use their own lists.", tags: [OpenAPITags.PublicResource], request: { params: listResourceAiModelsParamsSchema diff --git a/server/routers/resource/listResourceAiProviders.ts b/server/routers/resource/listResourceAiProviders.ts index 2c81c45d5..f48ad71fe 100644 --- a/server/routers/resource/listResourceAiProviders.ts +++ b/server/routers/resource/listResourceAiProviders.ts @@ -21,7 +21,8 @@ export type ListResourceAiProvidersResponse = { registry.registerPath({ method: "get", path: "/resource/{resourceId}/ai-providers", - description: "List AI providers attached to an inference resource.", + description: + "List AI providers attached to an inference resource, including each attachment's accessMode.", tags: [OpenAPITags.PublicResource], request: { params: listResourceAiProvidersParamsSchema diff --git a/server/routers/resource/removeAiModelFromResource.ts b/server/routers/resource/removeAiModelFromResource.ts index 7b5cb012f..626b94666 100644 --- a/server/routers/resource/removeAiModelFromResource.ts +++ b/server/routers/resource/removeAiModelFromResource.ts @@ -8,7 +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"; +import { assertPublicModelListApiEligible } from "@server/lib/aiInferenceResource"; const removeAiModelFromResourceBodySchema = z.strictObject({ modelId: z.int().positive() @@ -22,7 +22,7 @@ registry.registerPath({ method: "post", path: "/resource/{resourceId}/ai-models/remove", description: - "Remove a single catalog model from an inference resource allowlist. Requires at least one attached AI provider in allowlist mode.", + "Remove a single model from an inference resource allow/block list. Requires at least one attached AI provider.", tags: [OpenAPITags.PublicResource], request: { params: removeAiModelFromResourceParamsSchema, @@ -98,7 +98,7 @@ export async function removeAiModelFromResource( ); } - const eligibleError = await assertPublicAllowlistApiEligible(resource); + const eligibleError = await assertPublicModelListApiEligible(resource); if (eligibleError) { return next(createHttpError(HttpCode.BAD_REQUEST, eligibleError)); } diff --git a/server/routers/resource/removeAiProviderFromResource.ts b/server/routers/resource/removeAiProviderFromResource.ts index e74b0c886..a7f9298ee 100644 --- a/server/routers/resource/removeAiProviderFromResource.ts +++ b/server/routers/resource/removeAiProviderFromResource.ts @@ -127,13 +127,14 @@ export async function removeAiProviderFromResource( .filter((a) => a.providerId !== providerId) .map((a) => ({ providerId: a.providerId, - modelAccessMode: a.modelAccessMode + accessMode: a.accessMode })); const attachments = await resolveProviderAttachments({ orgId: resource.orgId, attachments: remaining, - requireAtLeastOne: false + requireAtLeastOne: false, + resourceId }); if (isInferenceFieldsError(attachments)) { return next( diff --git a/server/routers/resource/setResourceAiModels.ts b/server/routers/resource/setResourceAiModels.ts index 0c3323ebe..78a67fdfb 100644 --- a/server/routers/resource/setResourceAiModels.ts +++ b/server/routers/resource/setResourceAiModels.ts @@ -9,12 +9,13 @@ import logger from "@server/logger"; import { fromError } from "zod-validation-error"; import { OpenAPITags, registry } from "@server/openApi"; import { - assertPublicAllowlistApiEligible, - assertModelsBelongToPublicAllowlistProviders + assertPublicModelListApiEligible, + assertPublicResourceModelEntriesValid, + resourceAiModelEntrySchema } from "@server/lib/aiInferenceResource"; const setResourceAiModelsBodySchema = z.strictObject({ - modelIds: z.array(z.int().positive()) + models: z.array(resourceAiModelEntrySchema) }); const setResourceAiModelsParamsSchema = z.strictObject({ @@ -25,7 +26,7 @@ registry.registerPath({ method: "post", path: "/resource/{resourceId}/ai-models", description: - "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.", + "Replace the allow/block model selection for an inference resource. Requires at least one attached AI provider in select mode. Models must belong to a select-mode provider and their listType must match the provider catalog entry. An empty array clears the selection, which denies all models for select-mode providers.", tags: [OpenAPITags.PublicResource], request: { params: setResourceAiModelsParamsSchema, @@ -71,7 +72,7 @@ export async function setResourceAiModels( ); } - const { modelIds } = parsedBody.data; + const { models } = parsedBody.data; const parsedParams = setResourceAiModelsParamsSchema.safeParse( req.params @@ -99,15 +100,22 @@ export async function setResourceAiModels( ); } - const eligibleError = await assertPublicAllowlistApiEligible(resource); + const eligibleError = await assertPublicModelListApiEligible(resource); if (eligibleError) { return next(createHttpError(HttpCode.BAD_REQUEST, eligibleError)); } - const modelError = await assertModelsBelongToPublicAllowlistProviders({ + const byModelId = new Map( + models.map((m) => [m.modelId, m.listType] as const) + ); + const uniqueModels = [...byModelId.entries()].map( + ([modelId, listType]) => ({ modelId, listType }) + ); + + const modelError = await assertPublicResourceModelEntriesValid({ orgId: resource.orgId, resourceId, - modelIds + models: uniqueModels }); if (modelError) { return next(createHttpError(HttpCode.BAD_REQUEST, modelError)); @@ -118,12 +126,14 @@ export async function setResourceAiModels( .delete(resourceAiModels) .where(eq(resourceAiModels.resourceId, resourceId)); - if (modelIds.length > 0) { - await trx - .insert(resourceAiModels) - .values( - modelIds.map((modelId) => ({ resourceId, modelId })) - ); + if (uniqueModels.length > 0) { + await trx.insert(resourceAiModels).values( + uniqueModels.map((m) => ({ + resourceId, + modelId: m.modelId, + listType: m.listType + })) + ); } }); diff --git a/server/routers/resource/setResourceAiProviders.ts b/server/routers/resource/setResourceAiProviders.ts index 88552954b..ef312ac70 100644 --- a/server/routers/resource/setResourceAiProviders.ts +++ b/server/routers/resource/setResourceAiProviders.ts @@ -27,7 +27,7 @@ registry.registerPath({ method: "post", path: "/resource/{resourceId}/ai-providers", description: - "Replace the AI providers attached to an inference resource. An empty list clears all providers. Model keys must be unique across attached catalog providers.", + "Replace the AI providers attached to an inference resource. Each provider uses accessMode inherit (default, uses the provider's own allow/block lists) or select (uses the resource's selected subset of that provider's catalog). An empty list clears all providers. Effective allow model keys must be unique across attached providers.", tags: [OpenAPITags.PublicResource], request: { params: setResourceAiProvidersParamsSchema, @@ -113,7 +113,8 @@ export async function setResourceAiProviders( const attachments = await resolveProviderAttachments({ orgId: resource.orgId, attachments: providers, - requireAtLeastOne: false + requireAtLeastOne: false, + resourceId }); if (isInferenceFieldsError(attachments)) { return next( diff --git a/server/routers/siteResource/addAiModelToSiteResource.ts b/server/routers/siteResource/addAiModelToSiteResource.ts index 0b6a004d7..8e1f8935e 100644 --- a/server/routers/siteResource/addAiModelToSiteResource.ts +++ b/server/routers/siteResource/addAiModelToSiteResource.ts @@ -9,12 +9,14 @@ import logger from "@server/logger"; import { fromError } from "zod-validation-error"; import { OpenAPITags, registry } from "@server/openApi"; import { - assertSiteAllowlistApiEligible, - assertModelsBelongToSiteAllowlistProviders + assertSiteModelListApiEligible, + assertSiteResourceModelEntriesValid, + modelListTypeSchema } from "@server/lib/aiInferenceResource"; const addAiModelToSiteResourceBodySchema = z.strictObject({ - modelId: z.int().positive() + modelId: z.number().int().positive(), + listType: modelListTypeSchema.optional().default("allow") }); const addAiModelToSiteResourceParamsSchema = z.strictObject({ @@ -25,7 +27,7 @@ registry.registerPath({ method: "post", path: "/site-resource/{siteResourceId}/ai-models/add", description: - "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.", + "Add a single model to an inference site resource allow/block selection. Requires at least one attached AI provider in select mode. The model must belong to a select-mode provider and its listType must match the provider catalog entry. listType defaults to allow.", tags: [OpenAPITags.PrivateResource], request: { params: addAiModelToSiteResourceParamsSchema, @@ -73,7 +75,7 @@ export async function addAiModelToSiteResource( ); } - const { modelId } = parsedBody.data; + const { modelId, listType } = parsedBody.data; const parsedParams = addAiModelToSiteResourceParamsSchema.safeParse( req.params @@ -102,15 +104,15 @@ export async function addAiModelToSiteResource( } const eligibleError = - await assertSiteAllowlistApiEligible(siteResource); + await assertSiteModelListApiEligible(siteResource); if (eligibleError) { return next(createHttpError(HttpCode.BAD_REQUEST, eligibleError)); } - const modelError = await assertModelsBelongToSiteAllowlistProviders({ + const modelError = await assertSiteResourceModelEntriesValid({ orgId: siteResource.orgId, siteResourceId, - modelIds: [modelId] + models: [{ modelId, listType }] }); if (modelError) { return next(createHttpError(HttpCode.BAD_REQUEST, modelError)); @@ -135,9 +137,11 @@ export async function addAiModelToSiteResource( ); } - await db - .insert(siteResourceAiModels) - .values({ siteResourceId, modelId }); + await db.insert(siteResourceAiModels).values({ + siteResourceId, + modelId, + listType + }); return response(res, { data: {}, diff --git a/server/routers/siteResource/addAiProviderToSiteResource.ts b/server/routers/siteResource/addAiProviderToSiteResource.ts index bae346347..79a0ae093 100644 --- a/server/routers/siteResource/addAiProviderToSiteResource.ts +++ b/server/routers/siteResource/addAiProviderToSiteResource.ts @@ -11,14 +11,12 @@ 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() + providerId: z.number().int().positive() }); const addAiProviderToSiteResourceParamsSchema = z.strictObject({ @@ -29,7 +27,7 @@ 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.", + "Add or replace a single AI provider attachment on an inference site resource. The provider is attached in inherit mode, using its own allow/block lists.", tags: [OpenAPITags.PrivateResource], request: { params: addAiProviderToSiteResourceParamsSchema, @@ -77,7 +75,7 @@ export async function addAiProviderToSiteResource( ); } - const { providerId, modelAccessMode } = parsedBody.data; + const { providerId } = parsedBody.data; const parsedParams = addAiProviderToSiteResourceParamsSchema.safeParse( req.params @@ -120,18 +118,21 @@ export async function addAiProviderToSiteResource( .filter((a) => a.providerId !== providerId) .map((a) => ({ providerId: a.providerId, - modelAccessMode: a.modelAccessMode + accessMode: a.accessMode })), - { providerId, modelAccessMode } + { providerId, accessMode: "inherit" as const } ]; const attachments = await resolveProviderAttachments({ orgId: siteResource.orgId, attachments: nextAttachments, - requireAtLeastOne: true + requireAtLeastOne: true, + siteResourceId }); if (isInferenceFieldsError(attachments)) { - return next(createHttpError(HttpCode.BAD_REQUEST, attachments.error)); + return next( + createHttpError(HttpCode.BAD_REQUEST, attachments.error) + ); } await setSiteResourceAiProviders(siteResourceId, attachments); diff --git a/server/routers/siteResource/createSiteResource.ts b/server/routers/siteResource/createSiteResource.ts index 31f631ce2..15b82cb56 100644 --- a/server/routers/siteResource/createSiteResource.ts +++ b/server/routers/siteResource/createSiteResource.ts @@ -90,7 +90,7 @@ const createSiteResourceSchema = z .array(resourceAiProviderAttachmentSchema) .optional() .describe( - "For inference-mode site resources: AI providers to attach. Each entry may set modelAccessMode (catalog or allowlist); defaults to catalog. Model keys must be unique across attached catalog providers." + "For inference-mode site resources: AI providers to attach. Providers are attached in inherit mode, using each provider's own allow/block lists. Effective allow model keys must be unique across attached providers." ) }) .strict() @@ -350,9 +350,14 @@ export async function createSiteResource( let providerAttachments: ResourceAiProviderAttachment[] = []; if (mode === "inference") { + // A new site resource has no model selections yet, so providers + // always start in inherit mode; select can be enabled afterwards. const resolved = await resolveProviderAttachments({ orgId, - attachments: aiProviderInputs ?? [], + attachments: (aiProviderInputs ?? []).map((p) => ({ + providerId: p.providerId, + accessMode: "inherit" as const + })), requireAtLeastOne: false }); if (isInferenceFieldsError(resolved)) { diff --git a/server/routers/siteResource/listSiteResourceAiModels.ts b/server/routers/siteResource/listSiteResourceAiModels.ts index 26e709140..fd65d29ea 100644 --- a/server/routers/siteResource/listSiteResourceAiModels.ts +++ b/server/routers/siteResource/listSiteResourceAiModels.ts @@ -19,7 +19,9 @@ async function query(siteResourceId: number) { modelId: aiModels.modelId, modelKey: aiModels.modelKey, name: aiModels.name, - enabled: aiModels.enabled + providerId: aiModels.providerId, + enabled: aiModels.enabled, + listType: siteResourceAiModels.listType }) .from(siteResourceAiModels) .innerJoin(aiModels, eq(siteResourceAiModels.modelId, aiModels.modelId)) @@ -34,7 +36,7 @@ registry.registerPath({ method: "get", path: "/site-resource/{siteResourceId}/ai-models", description: - "List catalog models on this site resource's allowlist. Only enforced when modelAccessMode=allowlist; an empty allowlist denies all models.", + "List the models this site resource has selected from its select-mode providers' allow/block lists. Providers in inherit mode are not represented here; they use their own lists.", tags: [OpenAPITags.PrivateResource], request: { params: listSiteResourceAiModelsParamsSchema diff --git a/server/routers/siteResource/listSiteResourceAiProviders.ts b/server/routers/siteResource/listSiteResourceAiProviders.ts index bcbe37e8b..7d7671dad 100644 --- a/server/routers/siteResource/listSiteResourceAiProviders.ts +++ b/server/routers/siteResource/listSiteResourceAiProviders.ts @@ -21,7 +21,8 @@ export type ListSiteResourceAiProvidersResponse = { registry.registerPath({ method: "get", path: "/site-resource/{siteResourceId}/ai-providers", - description: "List AI providers attached to an inference site resource.", + description: + "List AI providers attached to an inference site resource, including each attachment's accessMode.", tags: [OpenAPITags.PrivateResource], request: { params: listSiteResourceAiProvidersParamsSchema diff --git a/server/routers/siteResource/removeAiModelFromSiteResource.ts b/server/routers/siteResource/removeAiModelFromSiteResource.ts index eb8d649a2..c52dfb43d 100644 --- a/server/routers/siteResource/removeAiModelFromSiteResource.ts +++ b/server/routers/siteResource/removeAiModelFromSiteResource.ts @@ -8,7 +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"; +import { assertSiteModelListApiEligible } from "@server/lib/aiInferenceResource"; const removeAiModelFromSiteResourceBodySchema = z.strictObject({ modelId: z.int().positive() @@ -22,7 +22,7 @@ registry.registerPath({ method: "post", path: "/site-resource/{siteResourceId}/ai-models/remove", description: - "Remove a single catalog model from an inference site resource allowlist. Requires at least one attached AI provider in allowlist mode.", + "Remove a single model from an inference site resource allow/block list. Requires at least one attached AI provider.", tags: [OpenAPITags.PrivateResource], request: { params: removeAiModelFromSiteResourceParamsSchema, @@ -98,7 +98,7 @@ export async function removeAiModelFromSiteResource( } const eligibleError = - await assertSiteAllowlistApiEligible(siteResource); + await assertSiteModelListApiEligible(siteResource); if (eligibleError) { return next(createHttpError(HttpCode.BAD_REQUEST, eligibleError)); } diff --git a/server/routers/siteResource/removeAiProviderFromSiteResource.ts b/server/routers/siteResource/removeAiProviderFromSiteResource.ts index 3d4392820..ffb4f1ff8 100644 --- a/server/routers/siteResource/removeAiProviderFromSiteResource.ts +++ b/server/routers/siteResource/removeAiProviderFromSiteResource.ts @@ -126,13 +126,14 @@ export async function removeAiProviderFromSiteResource( .filter((a) => a.providerId !== providerId) .map((a) => ({ providerId: a.providerId, - modelAccessMode: a.modelAccessMode + accessMode: a.accessMode })); const attachments = await resolveProviderAttachments({ orgId: siteResource.orgId, attachments: remaining, - requireAtLeastOne: false + requireAtLeastOne: false, + siteResourceId }); if (isInferenceFieldsError(attachments)) { return next( diff --git a/server/routers/siteResource/setSiteResourceAiModels.ts b/server/routers/siteResource/setSiteResourceAiModels.ts index 391634553..eeacbbd23 100644 --- a/server/routers/siteResource/setSiteResourceAiModels.ts +++ b/server/routers/siteResource/setSiteResourceAiModels.ts @@ -9,12 +9,13 @@ import logger from "@server/logger"; import { fromError } from "zod-validation-error"; import { OpenAPITags, registry } from "@server/openApi"; import { - assertSiteAllowlistApiEligible, - assertModelsBelongToSiteAllowlistProviders + assertSiteModelListApiEligible, + assertSiteResourceModelEntriesValid, + resourceAiModelEntrySchema } from "@server/lib/aiInferenceResource"; const setSiteResourceAiModelsBodySchema = z.strictObject({ - modelIds: z.array(z.int().positive()) + models: z.array(resourceAiModelEntrySchema) }); const setSiteResourceAiModelsParamsSchema = z.strictObject({ @@ -25,7 +26,7 @@ registry.registerPath({ method: "post", path: "/site-resource/{siteResourceId}/ai-models", description: - "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.", + "Replace the allow/block model selection for an inference site resource. Requires at least one attached AI provider in select mode. Models must belong to a select-mode provider and their listType must match the provider catalog entry. An empty array clears the selection, which denies all models for select-mode providers.", tags: [OpenAPITags.PrivateResource], request: { params: setSiteResourceAiModelsParamsSchema, @@ -73,7 +74,7 @@ export async function setSiteResourceAiModels( ); } - const { modelIds } = parsedBody.data; + const { models } = parsedBody.data; const parsedParams = setSiteResourceAiModelsParamsSchema.safeParse( req.params @@ -102,15 +103,22 @@ export async function setSiteResourceAiModels( } const eligibleError = - await assertSiteAllowlistApiEligible(siteResource); + await assertSiteModelListApiEligible(siteResource); if (eligibleError) { return next(createHttpError(HttpCode.BAD_REQUEST, eligibleError)); } - const modelError = await assertModelsBelongToSiteAllowlistProviders({ + const byModelId = new Map( + models.map((m) => [m.modelId, m.listType] as const) + ); + const uniqueModels = [...byModelId.entries()].map( + ([modelId, listType]) => ({ modelId, listType }) + ); + + const modelError = await assertSiteResourceModelEntriesValid({ orgId: siteResource.orgId, siteResourceId, - modelIds + models: uniqueModels }); if (modelError) { return next(createHttpError(HttpCode.BAD_REQUEST, modelError)); @@ -121,11 +129,12 @@ export async function setSiteResourceAiModels( .delete(siteResourceAiModels) .where(eq(siteResourceAiModels.siteResourceId, siteResourceId)); - if (modelIds.length > 0) { + if (uniqueModels.length > 0) { await trx.insert(siteResourceAiModels).values( - modelIds.map((modelId) => ({ + uniqueModels.map((m) => ({ siteResourceId, - modelId + modelId: m.modelId, + listType: m.listType })) ); } diff --git a/server/routers/siteResource/setSiteResourceAiProviders.ts b/server/routers/siteResource/setSiteResourceAiProviders.ts index fd685bb20..a6c8a52c7 100644 --- a/server/routers/siteResource/setSiteResourceAiProviders.ts +++ b/server/routers/siteResource/setSiteResourceAiProviders.ts @@ -27,7 +27,7 @@ registry.registerPath({ method: "post", path: "/site-resource/{siteResourceId}/ai-providers", description: - "Replace the AI providers attached to an inference site resource. An empty list clears all providers. Model keys must be unique across attached catalog providers.", + "Replace the AI providers attached to an inference site resource. Each provider uses accessMode inherit (default, uses the provider's own allow/block lists) or select (uses the site resource's selected subset of that provider's catalog). An empty list clears all providers. Effective allow model keys must be unique across attached providers.", tags: [OpenAPITags.PrivateResource], request: { params: setSiteResourceAiProvidersParamsSchema, @@ -115,7 +115,8 @@ export async function setSiteResourceAiProviders( const attachments = await resolveProviderAttachments({ orgId: siteResource.orgId, attachments: providers, - requireAtLeastOne: false + requireAtLeastOne: false, + siteResourceId }); if (isInferenceFieldsError(attachments)) { return next( diff --git a/src/app/[orgId]/settings/ai-providers/[providerId]/models/page.tsx b/src/app/[orgId]/settings/ai-providers/[providerId]/models/page.tsx index f90143492..10db41571 100644 --- a/src/app/[orgId]/settings/ai-providers/[providerId]/models/page.tsx +++ b/src/app/[orgId]/settings/ai-providers/[providerId]/models/page.tsx @@ -12,6 +12,7 @@ import { } from "@app/components/Settings"; import { TagInput, type Tag } from "@app/components/tags/tag-input"; import { Button } from "@app/components/ui/button"; +import { Label } from "@app/components/ui/label"; import { useAiProviderContext } from "@app/hooks/useAiProviderContext"; import { useEnvContext } from "@app/hooks/useEnvContext"; import { toast } from "@app/hooks/useToast"; @@ -21,6 +22,8 @@ import { useQuery, useQueryClient } from "@tanstack/react-query"; import { useTranslations } from "next-intl"; import { useEffect, useState } from "react"; +type ModelListType = "allow" | "block"; + export default function AiProviderModelsPage() { const { provider } = useAiProviderContext(); const { env } = useEnvContext(); @@ -28,8 +31,14 @@ export default function AiProviderModelsPage() { const queryClient = useQueryClient(); const t = useTranslations(); const [saveLoading, setSaveLoading] = useState(false); - const [tags, setTags] = useState([]); - const [activeTagIndex, setActiveTagIndex] = useState(null); + const [allowTags, setAllowTags] = useState([]); + const [blockTags, setBlockTags] = useState([]); + const [activeAllowTagIndex, setActiveAllowTagIndex] = useState< + number | null + >(null); + const [activeBlockTagIndex, setActiveBlockTagIndex] = useState< + number | null + >(null); const modelsQuery = useQuery( aiProviderQueries.providerModels({ providerId: provider.providerId }) @@ -37,11 +46,21 @@ export default function AiProviderModelsPage() { useEffect(() => { if (!modelsQuery.data) return; - setTags( - modelsQuery.data.map((model) => ({ - id: String(model.modelId), - text: model.modelKey - })) + setAllowTags( + modelsQuery.data + .filter((model) => (model.listType ?? "allow") === "allow") + .map((model) => ({ + id: String(model.modelId), + text: model.modelKey + })) + ); + setBlockTags( + modelsQuery.data + .filter((model) => model.listType === "block") + .map((model) => ({ + id: String(model.modelId), + text: model.modelKey + })) ); }, [modelsQuery.data]); @@ -52,27 +71,74 @@ export default function AiProviderModelsPage() { const existingByKey = new Map( existing.map((model) => [model.modelKey, model]) ); - const nextKeys = new Set( - tags.map((tag) => tag.text.trim()).filter(Boolean) + + const nextAllow = new Set( + allowTags.map((tag) => tag.text.trim()).filter(Boolean) + ); + const nextBlock = new Set( + blockTags.map((tag) => tag.text.trim()).filter(Boolean) ); - const toCreate = [...nextKeys].filter( - (key) => !existingByKey.has(key) - ); - const toDelete = existing.filter( - (model) => !nextKeys.has(model.modelKey) - ); + const overlap = [...nextAllow].filter((key) => nextBlock.has(key)); + if (overlap.length > 0) { + toast({ + variant: "destructive", + title: t("aiProviderModelsErrorUpdate"), + description: t("aiProviderModelsOverlapError", { + keys: overlap.join(", ") + }) + }); + return; + } + + const desired = new Map(); + for (const key of nextAllow) { + desired.set(key, "allow"); + } + for (const key of nextBlock) { + desired.set(key, "block"); + } + + const toCreate: { modelKey: string; listType: ModelListType }[] = + []; + const toUpdate: { + modelId: number; + listType: ModelListType; + }[] = []; + const toDelete: number[] = []; + + for (const [modelKey, listType] of desired) { + const existingModel = existingByKey.get(modelKey); + if (!existingModel) { + toCreate.push({ modelKey, listType }); + continue; + } + if ((existingModel.listType ?? "allow") !== listType) { + toUpdate.push({ + modelId: existingModel.modelId, + listType + }); + } + } + + for (const model of existing) { + if (!desired.has(model.modelKey)) { + toDelete.push(model.modelId); + } + } await Promise.all([ - ...toCreate.map((modelKey) => + ...toCreate.map(({ modelKey, listType }) => api.put(`/ai-provider/${provider.providerId}/model`, { modelKey, - name: modelKey + name: modelKey, + listType }) ), - ...toDelete.map((model) => - api.delete(`/ai-model/${model.modelId}`) - ) + ...toUpdate.map(({ modelId, listType }) => + api.post(`/ai-model/${modelId}`, { listType }) + ), + ...toDelete.map((modelId) => api.delete(`/ai-model/${modelId}`)) ]); await queryClient.invalidateQueries( @@ -113,24 +179,59 @@ export default function AiProviderModelsPage() { - { - const next = - typeof newTags === "function" - ? newTags(tags) - : newTags; - setTags(next as Tag[]); - }} - allowDuplicates={false} - sortTags - delimiterList={[",", "Enter"]} - disabled={modelsQuery.isLoading || saveLoading} - /> +
+ + { + const next = + typeof newTags === "function" + ? newTags(allowTags) + : newTags; + setAllowTags(next as Tag[]); + }} + allowDuplicates={false} + sortTags + delimiterList={[",", "Enter"]} + disabled={modelsQuery.isLoading || saveLoading} + /> +

+ {t("aiProviderModelsAllowDescription")} +

+
+ +
+ + { + const next = + typeof newTags === "function" + ? newTags(blockTags) + : newTags; + setBlockTags(next as Tag[]); + }} + allowDuplicates={false} + sortTags + delimiterList={[",", "Enter"]} + disabled={modelsQuery.isLoading || saveLoading} + /> +

+ {t("aiProviderModelsBlockDescription")} +

+
diff --git a/src/app/[orgId]/settings/resources/private/[niceId]/inference/page.tsx b/src/app/[orgId]/settings/resources/private/[niceId]/inference/page.tsx index 48304d350..f9ad664a0 100644 --- a/src/app/[orgId]/settings/resources/private/[niceId]/inference/page.tsx +++ b/src/app/[orgId]/settings/resources/private/[niceId]/inference/page.tsx @@ -130,8 +130,7 @@ export default function PrivateResourceInferencePage() { await api.post(`/site-resource/${siteResource.id}/ai-providers`, { providers: data.providerIds.map((providerId) => ({ - providerId, - modelAccessMode: "catalog" + providerId })) }); @@ -258,12 +257,10 @@ export default function PrivateResourceInferencePage() { cols={2} hideFreeDomain defaultSubdomain={ - httpConfigSubdomain ?? - undefined + httpConfigSubdomain ?? undefined } defaultDomainId={ - httpConfigDomainId ?? - undefined + httpConfigDomainId ?? undefined } defaultFullDomain={ httpConfigFullDomain ?? diff --git a/src/app/[orgId]/settings/resources/public/[niceId]/inference/page.tsx b/src/app/[orgId]/settings/resources/public/[niceId]/inference/page.tsx index 8741159b4..ad7230591 100644 --- a/src/app/[orgId]/settings/resources/public/[niceId]/inference/page.tsx +++ b/src/app/[orgId]/settings/resources/public/[niceId]/inference/page.tsx @@ -102,8 +102,7 @@ export default function PublicResourceInferencePage() { try { await api.post(`/resource/${resource.resourceId}/ai-providers`, { providers: data.providerIds.map((providerId) => ({ - providerId, - modelAccessMode: "catalog" + providerId })) }); diff --git a/src/app/[orgId]/settings/resources/public/create/page.tsx b/src/app/[orgId]/settings/resources/public/create/page.tsx index a5088d84e..4d1c24efb 100644 --- a/src/app/[orgId]/settings/resources/public/create/page.tsx +++ b/src/app/[orgId]/settings/resources/public/create/page.tsx @@ -497,8 +497,7 @@ export default function Page() { if (resourceType === "inference") { Object.assign(payload, { aiProviders: selectedProviders.map((provider) => ({ - providerId: parseInt(provider.id, 10), - modelAccessMode: "catalog" + providerId: parseInt(provider.id, 10) })) }); } else if (resourceType === "ssh") { diff --git a/src/lib/privateResourceForm.ts b/src/lib/privateResourceForm.ts index f186ffb9d..669a96363 100644 --- a/src/lib/privateResourceForm.ts +++ b/src/lib/privateResourceForm.ts @@ -217,8 +217,7 @@ export function buildCreateSiteResourcePayload( }), ...(data.mode === "inference" && { aiProviders: (data.providerIds ?? []).map((providerId) => ({ - providerId, - modelAccessMode: "catalog" as const + providerId })), ssl: data.ssl ?? false, domainId: data.httpConfigDomainId diff --git a/src/lib/queries.ts b/src/lib/queries.ts index 9a0cb2403..6b7070f5b 100644 --- a/src/lib/queries.ts +++ b/src/lib/queries.ts @@ -1288,10 +1288,10 @@ export const resourceQueries = { AxiosResponse<{ providers: Array<{ providerId: number; - modelAccessMode: "catalog" | "allowlist"; name: string; type: string; enabled: boolean; + accessMode: "inherit" | "select"; }>; }> >(`/site-resource/${siteResourceId}/ai-providers`, { @@ -1308,10 +1308,10 @@ export const resourceQueries = { AxiosResponse<{ providers: Array<{ providerId: number; - modelAccessMode: "catalog" | "allowlist"; name: string; type: string; enabled: boolean; + accessMode: "inherit" | "select"; }>; }> >(`/resource/${resourceId}/ai-providers`, {