From 5daf334e32138ecefd06a2a545b3e64e43e67524 Mon Sep 17 00:00:00 2001 From: Owen Date: Sun, 2 Aug 2026 10:50:50 -0400 Subject: [PATCH] Endpoints to update models on the resource --- server/auth/actions.ts | 2 + server/routers/external.ts | 30 ++++ server/routers/integration.ts | 86 +++++++++ .../routers/resource/addAiModelToResource.ts | 161 +++++++++++++++++ server/routers/resource/index.ts | 4 + .../routers/resource/listResourceAiModels.ts | 107 +++++++++++ .../resource/removeAiModelFromResource.ts | 141 +++++++++++++++ .../routers/resource/setResourceAiModels.ts | 155 ++++++++++++++++ .../siteResource/addAiModelToSiteResource.ts | 170 ++++++++++++++++++ server/routers/siteResource/index.ts | 4 + .../siteResource/listSiteResourceAiModels.ts | 115 ++++++++++++ .../removeAiModelFromSiteResource.ts | 140 +++++++++++++++ .../siteResource/setSiteResourceAiModels.ts | 167 +++++++++++++++++ 13 files changed, 1282 insertions(+) create mode 100644 server/routers/resource/addAiModelToResource.ts create mode 100644 server/routers/resource/listResourceAiModels.ts create mode 100644 server/routers/resource/removeAiModelFromResource.ts create mode 100644 server/routers/resource/setResourceAiModels.ts create mode 100644 server/routers/siteResource/addAiModelToSiteResource.ts create mode 100644 server/routers/siteResource/listSiteResourceAiModels.ts create mode 100644 server/routers/siteResource/removeAiModelFromSiteResource.ts create mode 100644 server/routers/siteResource/setSiteResourceAiModels.ts diff --git a/server/auth/actions.ts b/server/auth/actions.ts index 944adbff4..1b4a33e89 100644 --- a/server/auth/actions.ts +++ b/server/auth/actions.ts @@ -50,6 +50,8 @@ export enum ActionsEnum { setResourceUsers = "setResourceUsers", setResourceRoles = "setResourceRoles", listResourceUsers = "listResourceUsers", + listResourceAiModels = "listResourceAiModels", + setResourceAiModels = "setResourceAiModels", // removeRoleSite = "removeRoleSite", // addRoleAction = "addRoleAction", // removeRoleAction = "removeRoleAction", diff --git a/server/routers/external.ts b/server/routers/external.ts index 99bca3a8e..4ebf58cae 100644 --- a/server/routers/external.ts +++ b/server/routers/external.ts @@ -407,6 +407,13 @@ authenticated.get( siteResource.listSiteResourceClients ); +authenticated.get( + "/site-resource/:siteResourceId/ai-models", + verifySiteResourceAccess, + verifyUserHasAction(ActionsEnum.listResourceAiModels), + siteResource.listSiteResourceAiModels +); + authenticated.post( "/site-resource/:siteResourceId/roles", verifySiteResourceAccess, @@ -417,6 +424,14 @@ authenticated.post( siteResource.setSiteResourceRoles ); +authenticated.post( + "/site-resource/:siteResourceId/ai-models", + verifySiteResourceAccess, + verifyUserHasAction(ActionsEnum.setResourceAiModels), + logActionAudit(ActionsEnum.setResourceAiModels), + siteResource.setSiteResourceAiModels +); + authenticated.post( "/site-resource/:siteResourceId/users", verifySiteResourceAccess, @@ -651,6 +666,13 @@ authenticated.get( resource.listResourceUsers ); +authenticated.get( + "/resource/:resourceId/ai-models", + verifyResourceAccess, + verifyUserHasAction(ActionsEnum.listResourceAiModels), + resource.listResourceAiModels +); + authenticated.get( "/resource/:resourceId", verifyResourceAccess, @@ -854,6 +876,14 @@ authenticated.post( resource.setResourceUsers ); +authenticated.post( + "/resource/:resourceId/ai-models", + verifyResourceAccess, + verifyUserHasAction(ActionsEnum.setResourceAiModels), + logActionAudit(ActionsEnum.setResourceAiModels), + resource.setResourceAiModels +); + authenticated.put( "/resource-policy/:resourcePolicyId/access-control", verifyResourcePolicyAccess, diff --git a/server/routers/integration.ts b/server/routers/integration.ts index 384214e4c..41757e7a0 100644 --- a/server/routers/integration.ts +++ b/server/routers/integration.ts @@ -246,6 +246,16 @@ authenticated.get( siteResource.listSiteResourceClients ); +authenticated.get( + [ + "/site-resource/:siteResourceId/ai-models", + "/private-resource/:siteResourceId/ai-models" + ], + verifyApiKeySiteResourceAccess, + verifyApiKeyHasAction(ActionsEnum.listResourceAiModels), + siteResource.listSiteResourceAiModels +); + authenticated.post( [ "/site-resource/:siteResourceId/roles", @@ -298,6 +308,39 @@ authenticated.post( siteResource.removeRoleFromSiteResource ); +authenticated.post( + [ + "/site-resource/:siteResourceId/ai-models", + "/private-resource/:siteResourceId/ai-models" + ], + verifyApiKeySiteResourceAccess, + verifyApiKeyHasAction(ActionsEnum.setResourceAiModels), + logActionAudit(ActionsEnum.setResourceAiModels), + siteResource.setSiteResourceAiModels +); + +authenticated.post( + [ + "/site-resource/:siteResourceId/ai-models/add", + "/private-resource/:siteResourceId/ai-models/add" + ], + verifyApiKeySiteResourceAccess, + verifyApiKeyHasAction(ActionsEnum.setResourceAiModels), + logActionAudit(ActionsEnum.setResourceAiModels), + siteResource.addAiModelToSiteResource +); + +authenticated.post( + [ + "/site-resource/:siteResourceId/ai-models/remove", + "/private-resource/:siteResourceId/ai-models/remove" + ], + verifyApiKeySiteResourceAccess, + verifyApiKeyHasAction(ActionsEnum.setResourceAiModels), + logActionAudit(ActionsEnum.setResourceAiModels), + siteResource.removeAiModelFromSiteResource +); + authenticated.post( [ "/site-resource/:siteResourceId/users/add", @@ -510,6 +553,16 @@ authenticated.get( resource.listResourceUsers ); +authenticated.get( + [ + "/resource/:resourceId/ai-models", + "/public-resource/:resourceId/ai-models" + ], + verifyApiKeyResourceAccess, + verifyApiKeyHasAction(ActionsEnum.listResourceAiModels), + resource.listResourceAiModels +); + authenticated.get( ["/resource/:resourceId", "/public-resource/:resourceId"], verifyApiKeyResourceAccess, @@ -711,6 +764,17 @@ authenticated.post( resource.setResourceRoles ); +authenticated.post( + [ + "/resource/:resourceId/ai-models", + "/public-resource/:resourceId/ai-models" + ], + verifyApiKeyResourceAccess, + verifyApiKeyHasAction(ActionsEnum.setResourceAiModels), + logActionAudit(ActionsEnum.setResourceAiModels), + resource.setResourceAiModels +); + authenticated.post( ["/resource/:resourceId/users", "/public-resource/:resourceId/users"], verifyApiKeyResourceAccess, @@ -903,6 +967,28 @@ authenticated.post( resource.removeRoleFromResource ); +authenticated.post( + [ + "/resource/:resourceId/ai-models/add", + "/public-resource/:resourceId/ai-models/add" + ], + verifyApiKeyResourceAccess, + verifyApiKeyHasAction(ActionsEnum.setResourceAiModels), + logActionAudit(ActionsEnum.setResourceAiModels), + resource.addAiModelToResource +); + +authenticated.post( + [ + "/resource/:resourceId/ai-models/remove", + "/public-resource/:resourceId/ai-models/remove" + ], + verifyApiKeyResourceAccess, + verifyApiKeyHasAction(ActionsEnum.setResourceAiModels), + logActionAudit(ActionsEnum.setResourceAiModels), + resource.removeAiModelFromResource +); + authenticated.post( [ "/resource/:resourceId/users/add", diff --git a/server/routers/resource/addAiModelToResource.ts b/server/routers/resource/addAiModelToResource.ts new file mode 100644 index 000000000..fd33a3e91 --- /dev/null +++ b/server/routers/resource/addAiModelToResource.ts @@ -0,0 +1,161 @@ +import { Request, Response, NextFunction } from "express"; +import { z } from "zod"; +import { db, resources, resourceAiModels, aiModels } from "@server/db"; +import { eq, and } 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"; + +const addAiModelToResourceBodySchema = z.strictObject({ + modelId: z.int().positive() +}); + +const addAiModelToResourceParamsSchema = z.strictObject({ + resourceId: z.coerce.number().int().positive() +}); + +registry.registerPath({ + method: "post", + path: "/resource/{resourceId}/ai-models/add", + description: + "Add a single AI model to a resource's model restriction allow-list.", + tags: [OpenAPITags.PublicResource], + request: { + params: addAiModelToResourceParamsSchema, + body: { + content: { + "application/json": { + schema: addAiModelToResourceBodySchema + } + } + } + }, + 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 addAiModelToResource( + req: Request, + res: Response, + next: NextFunction +): Promise { + try { + const parsedBody = addAiModelToResourceBodySchema.safeParse(req.body); + if (!parsedBody.success) { + return next( + createHttpError( + HttpCode.BAD_REQUEST, + fromError(parsedBody.error).toString() + ) + ); + } + + const { modelId } = parsedBody.data; + + const parsedParams = addAiModelToResourceParamsSchema.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.aiProviderId) { + return next( + createHttpError( + HttpCode.BAD_REQUEST, + "Resource has no AI provider linked" + ) + ); + } + + 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 existingEntry = await db + .select() + .from(resourceAiModels) + .where( + and( + eq(resourceAiModels.resourceId, resourceId), + eq(resourceAiModels.modelId, modelId) + ) + ); + + if (existingEntry.length > 0) { + return next( + createHttpError( + HttpCode.CONFLICT, + "Model already assigned to resource" + ) + ); + } + + await db.insert(resourceAiModels).values({ resourceId, modelId }); + + return response(res, { + data: {}, + success: true, + error: false, + message: "Model added to resource successfully", + status: HttpCode.CREATED + }); + } catch (error) { + logger.error(error); + return next( + createHttpError(HttpCode.INTERNAL_SERVER_ERROR, "An error occurred") + ); + } +} diff --git a/server/routers/resource/index.ts b/server/routers/resource/index.ts index 709bf7340..32992d364 100644 --- a/server/routers/resource/index.ts +++ b/server/routers/resource/index.ts @@ -35,3 +35,7 @@ export * from "./removeEmailFromResourceWhitelist"; export * from "./getStatusHistory"; export * from "./getBatchedStatusHistory"; export * from "./getResourcePolicies"; +export * from "./listResourceAiModels"; +export * from "./setResourceAiModels"; +export * from "./addAiModelToResource"; +export * from "./removeAiModelFromResource"; diff --git a/server/routers/resource/listResourceAiModels.ts b/server/routers/resource/listResourceAiModels.ts new file mode 100644 index 000000000..f6875faff --- /dev/null +++ b/server/routers/resource/listResourceAiModels.ts @@ -0,0 +1,107 @@ +import { Request, Response, NextFunction } from "express"; +import { z } from "zod"; +import { db, resources, resourceAiModels, aiModels } 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"; + +const listResourceAiModelsParamsSchema = z.strictObject({ + resourceId: z.coerce.number().int().positive() +}); + +async function query(resourceId: number) { + return await db + .select({ + modelId: aiModels.modelId, + modelKey: aiModels.modelKey, + name: aiModels.name, + enabled: aiModels.enabled + }) + .from(resourceAiModels) + .innerJoin(aiModels, eq(resourceAiModels.modelId, aiModels.modelId)) + .where(eq(resourceAiModels.resourceId, resourceId)); +} + +export type ListResourceAiModelsResponse = { + models: NonNullable>>; +}; + +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.", + tags: [OpenAPITags.PublicResource], + request: { + params: listResourceAiModelsParamsSchema + }, + 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 listResourceAiModels( + req: Request, + res: Response, + next: NextFunction +): Promise { + try { + const parsedParams = listResourceAiModelsParamsSchema.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 models = await query(resourceId); + + return response(res, { + data: { models }, + success: true, + error: false, + message: "Resource AI models retrieved successfully", + status: HttpCode.OK + }); + } catch (error) { + logger.error(error); + return next( + createHttpError(HttpCode.INTERNAL_SERVER_ERROR, "An error occurred") + ); + } +} diff --git a/server/routers/resource/removeAiModelFromResource.ts b/server/routers/resource/removeAiModelFromResource.ts new file mode 100644 index 000000000..47445bbf5 --- /dev/null +++ b/server/routers/resource/removeAiModelFromResource.ts @@ -0,0 +1,141 @@ +import { Request, Response, NextFunction } from "express"; +import { z } from "zod"; +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"; +import createHttpError from "http-errors"; +import logger from "@server/logger"; +import { fromError } from "zod-validation-error"; +import { OpenAPITags, registry } from "@server/openApi"; + +const removeAiModelFromResourceBodySchema = z.strictObject({ + modelId: z.int().positive() +}); + +const removeAiModelFromResourceParamsSchema = z.strictObject({ + resourceId: z.coerce.number().int().positive() +}); + +registry.registerPath({ + method: "post", + path: "/resource/{resourceId}/ai-models/remove", + description: + "Remove a single AI model from a resource's model restriction allow-list.", + tags: [OpenAPITags.PublicResource], + request: { + params: removeAiModelFromResourceParamsSchema, + body: { + content: { + "application/json": { + schema: removeAiModelFromResourceBodySchema + } + } + } + }, + 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 removeAiModelFromResource( + req: Request, + res: Response, + next: NextFunction +): Promise { + try { + const parsedBody = removeAiModelFromResourceBodySchema.safeParse( + req.body + ); + if (!parsedBody.success) { + return next( + createHttpError( + HttpCode.BAD_REQUEST, + fromError(parsedBody.error).toString() + ) + ); + } + + const { modelId } = parsedBody.data; + + const parsedParams = removeAiModelFromResourceParamsSchema.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 existingEntry = await db + .select() + .from(resourceAiModels) + .where( + and( + eq(resourceAiModels.resourceId, resourceId), + eq(resourceAiModels.modelId, modelId) + ) + ); + + if (existingEntry.length === 0) { + return next( + createHttpError( + HttpCode.NOT_FOUND, + "Model not found in resource's restriction list" + ) + ); + } + + await db + .delete(resourceAiModels) + .where( + and( + eq(resourceAiModels.resourceId, resourceId), + eq(resourceAiModels.modelId, modelId) + ) + ); + + return response(res, { + data: {}, + success: true, + error: false, + message: "Model removed from resource successfully", + status: HttpCode.OK + }); + } catch (error) { + logger.error(error); + return next( + createHttpError(HttpCode.INTERNAL_SERVER_ERROR, "An error occurred") + ); + } +} diff --git a/server/routers/resource/setResourceAiModels.ts b/server/routers/resource/setResourceAiModels.ts new file mode 100644 index 000000000..2ef56f4fc --- /dev/null +++ b/server/routers/resource/setResourceAiModels.ts @@ -0,0 +1,155 @@ +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 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"; + +const setResourceAiModelsBodySchema = z.strictObject({ + modelIds: z.array(z.int().positive()) +}); + +const setResourceAiModelsParamsSchema = z.strictObject({ + resourceId: z.coerce.number().int().positive() +}); + +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).", + tags: [OpenAPITags.PublicResource], + request: { + params: setResourceAiModelsParamsSchema, + body: { + content: { + "application/json": { + schema: setResourceAiModelsBodySchema + } + } + } + }, + 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 setResourceAiModels( + req: Request, + res: Response, + next: NextFunction +): Promise { + try { + const parsedBody = setResourceAiModelsBodySchema.safeParse(req.body); + if (!parsedBody.success) { + return next( + createHttpError( + HttpCode.BAD_REQUEST, + fromError(parsedBody.error).toString() + ) + ); + } + + const { modelIds } = parsedBody.data; + + const parsedParams = setResourceAiModelsParamsSchema.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 (modelIds.length > 0) { + if (!resource.aiProviderId) { + return next( + createHttpError( + HttpCode.BAD_REQUEST, + "Resource has no AI provider linked" + ) + ); + } + + 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" + ) + ); + } + } + + await db.transaction(async (trx) => { + await trx + .delete(resourceAiModels) + .where(eq(resourceAiModels.resourceId, resourceId)); + + if (modelIds.length > 0) { + await trx + .insert(resourceAiModels) + .values( + modelIds.map((modelId) => ({ resourceId, modelId })) + ); + } + }); + + return response(res, { + data: {}, + success: true, + error: false, + message: "AI models set for resource successfully", + status: HttpCode.CREATED + }); + } catch (error) { + logger.error(error); + return next( + createHttpError(HttpCode.INTERNAL_SERVER_ERROR, "An error occurred") + ); + } +} diff --git a/server/routers/siteResource/addAiModelToSiteResource.ts b/server/routers/siteResource/addAiModelToSiteResource.ts new file mode 100644 index 000000000..2e55f100b --- /dev/null +++ b/server/routers/siteResource/addAiModelToSiteResource.ts @@ -0,0 +1,170 @@ +import { Request, Response, NextFunction } from "express"; +import { z } from "zod"; +import { + db, + siteResources, + siteResourceAiModels, + aiModels +} from "@server/db"; +import { eq, and } 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"; + +const addAiModelToSiteResourceBodySchema = z.strictObject({ + modelId: z.int().positive() +}); + +const addAiModelToSiteResourceParamsSchema = z.strictObject({ + siteResourceId: z.coerce.number().int().positive() +}); + +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.", + tags: [OpenAPITags.PrivateResource], + request: { + params: addAiModelToSiteResourceParamsSchema, + body: { + content: { + "application/json": { + schema: addAiModelToSiteResourceBodySchema + } + } + } + }, + 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 addAiModelToSiteResource( + req: Request, + res: Response, + next: NextFunction +): Promise { + try { + const parsedBody = addAiModelToSiteResourceBodySchema.safeParse( + req.body + ); + if (!parsedBody.success) { + return next( + createHttpError( + HttpCode.BAD_REQUEST, + fromError(parsedBody.error).toString() + ) + ); + } + + const { modelId } = parsedBody.data; + + const parsedParams = addAiModelToSiteResourceParamsSchema.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.aiProviderId) { + return next( + createHttpError( + HttpCode.BAD_REQUEST, + "Site resource has no AI provider linked" + ) + ); + } + + 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 existingEntry = await db + .select() + .from(siteResourceAiModels) + .where( + and( + eq(siteResourceAiModels.siteResourceId, siteResourceId), + eq(siteResourceAiModels.modelId, modelId) + ) + ); + + if (existingEntry.length > 0) { + return next( + createHttpError( + HttpCode.CONFLICT, + "Model already assigned to site resource" + ) + ); + } + + await db + .insert(siteResourceAiModels) + .values({ siteResourceId, modelId }); + + return response(res, { + data: {}, + success: true, + error: false, + message: "Model added to site resource successfully", + status: HttpCode.CREATED + }); + } catch (error) { + logger.error(error); + return next( + createHttpError(HttpCode.INTERNAL_SERVER_ERROR, "An error occurred") + ); + } +} diff --git a/server/routers/siteResource/index.ts b/server/routers/siteResource/index.ts index 5c09d3883..2acaf1e33 100644 --- a/server/routers/siteResource/index.ts +++ b/server/routers/siteResource/index.ts @@ -17,3 +17,7 @@ export * from "./setSiteResourceClients"; export * from "./addClientToSiteResource"; export * from "./batchAddClientToSiteResources"; export * from "./removeClientFromSiteResource"; +export * from "./listSiteResourceAiModels"; +export * from "./setSiteResourceAiModels"; +export * from "./addAiModelToSiteResource"; +export * from "./removeAiModelFromSiteResource"; diff --git a/server/routers/siteResource/listSiteResourceAiModels.ts b/server/routers/siteResource/listSiteResourceAiModels.ts new file mode 100644 index 000000000..26dea7cd0 --- /dev/null +++ b/server/routers/siteResource/listSiteResourceAiModels.ts @@ -0,0 +1,115 @@ +import { Request, Response, NextFunction } from "express"; +import { z } from "zod"; +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"; +import createHttpError from "http-errors"; +import logger from "@server/logger"; +import { fromError } from "zod-validation-error"; +import { OpenAPITags, registry } from "@server/openApi"; + +const listSiteResourceAiModelsParamsSchema = z.strictObject({ + siteResourceId: z.coerce.number().int().positive() +}); + +async function query(siteResourceId: number) { + return await db + .select({ + modelId: aiModels.modelId, + modelKey: aiModels.modelKey, + name: aiModels.name, + enabled: aiModels.enabled + }) + .from(siteResourceAiModels) + .innerJoin( + aiModels, + eq(siteResourceAiModels.modelId, aiModels.modelId) + ) + .where(eq(siteResourceAiModels.siteResourceId, siteResourceId)); +} + +export type ListSiteResourceAiModelsResponse = { + models: NonNullable>>; +}; + +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.", + tags: [OpenAPITags.PrivateResource], + request: { + params: listSiteResourceAiModelsParamsSchema + }, + 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 listSiteResourceAiModels( + req: Request, + res: Response, + next: NextFunction +): Promise { + try { + const parsedParams = listSiteResourceAiModelsParamsSchema.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 models = await query(siteResourceId); + + return response(res, { + data: { models }, + success: true, + error: false, + message: "Site resource AI models retrieved successfully", + status: HttpCode.OK + }); + } catch (error) { + logger.error(error); + return next( + createHttpError(HttpCode.INTERNAL_SERVER_ERROR, "An error occurred") + ); + } +} diff --git a/server/routers/siteResource/removeAiModelFromSiteResource.ts b/server/routers/siteResource/removeAiModelFromSiteResource.ts new file mode 100644 index 000000000..5af2f4d87 --- /dev/null +++ b/server/routers/siteResource/removeAiModelFromSiteResource.ts @@ -0,0 +1,140 @@ +import { Request, Response, NextFunction } from "express"; +import { z } from "zod"; +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"; +import createHttpError from "http-errors"; +import logger from "@server/logger"; +import { fromError } from "zod-validation-error"; +import { OpenAPITags, registry } from "@server/openApi"; + +const removeAiModelFromSiteResourceBodySchema = z.strictObject({ + modelId: z.int().positive() +}); + +const removeAiModelFromSiteResourceParamsSchema = z.strictObject({ + siteResourceId: z.coerce.number().int().positive() +}); + +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.", + tags: [OpenAPITags.PrivateResource], + request: { + params: removeAiModelFromSiteResourceParamsSchema, + body: { + content: { + "application/json": { + schema: removeAiModelFromSiteResourceBodySchema + } + } + } + }, + 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 removeAiModelFromSiteResource( + req: Request, + res: Response, + next: NextFunction +): Promise { + try { + const parsedBody = removeAiModelFromSiteResourceBodySchema.safeParse( + req.body + ); + if (!parsedBody.success) { + return next( + createHttpError( + HttpCode.BAD_REQUEST, + fromError(parsedBody.error).toString() + ) + ); + } + + const { modelId } = parsedBody.data; + + const parsedParams = + removeAiModelFromSiteResourceParamsSchema.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 existingEntry = await db + .select() + .from(siteResourceAiModels) + .where( + and( + eq(siteResourceAiModels.siteResourceId, siteResourceId), + eq(siteResourceAiModels.modelId, modelId) + ) + ); + + if (existingEntry.length === 0) { + return next( + createHttpError( + HttpCode.NOT_FOUND, + "Model not found in site resource's restriction list" + ) + ); + } + + await db + .delete(siteResourceAiModels) + .where( + and( + eq(siteResourceAiModels.siteResourceId, siteResourceId), + eq(siteResourceAiModels.modelId, modelId) + ) + ); + + return response(res, { + data: {}, + success: true, + error: false, + message: "Model removed from site resource successfully", + status: HttpCode.OK + }); + } catch (error) { + logger.error(error); + return next( + createHttpError(HttpCode.INTERNAL_SERVER_ERROR, "An error occurred") + ); + } +} diff --git a/server/routers/siteResource/setSiteResourceAiModels.ts b/server/routers/siteResource/setSiteResourceAiModels.ts new file mode 100644 index 000000000..1366187af --- /dev/null +++ b/server/routers/siteResource/setSiteResourceAiModels.ts @@ -0,0 +1,167 @@ +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 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"; + +const setSiteResourceAiModelsBodySchema = z.strictObject({ + modelIds: z.array(z.int().positive()) +}); + +const setSiteResourceAiModelsParamsSchema = z.strictObject({ + siteResourceId: z.coerce.number().int().positive() +}); + +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).", + tags: [OpenAPITags.PrivateResource], + request: { + params: setSiteResourceAiModelsParamsSchema, + body: { + content: { + "application/json": { + schema: setSiteResourceAiModelsBodySchema + } + } + } + }, + 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 setSiteResourceAiModels( + req: Request, + res: Response, + next: NextFunction +): Promise { + try { + const parsedBody = setSiteResourceAiModelsBodySchema.safeParse( + req.body + ); + if (!parsedBody.success) { + return next( + createHttpError( + HttpCode.BAD_REQUEST, + fromError(parsedBody.error).toString() + ) + ); + } + + const { modelIds } = parsedBody.data; + + const parsedParams = setSiteResourceAiModelsParamsSchema.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 (modelIds.length > 0) { + if (!siteResource.aiProviderId) { + return next( + createHttpError( + HttpCode.BAD_REQUEST, + "Site resource has no AI provider linked" + ) + ); + } + + 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" + ) + ); + } + } + + await db.transaction(async (trx) => { + await trx + .delete(siteResourceAiModels) + .where( + eq(siteResourceAiModels.siteResourceId, siteResourceId) + ); + + if (modelIds.length > 0) { + await trx + .insert(siteResourceAiModels) + .values( + modelIds.map((modelId) => ({ + siteResourceId, + modelId + })) + ); + } + }); + + return response(res, { + data: {}, + success: true, + error: false, + message: "AI models set for site resource successfully", + status: HttpCode.CREATED + }); + } catch (error) { + logger.error(error); + return next( + createHttpError(HttpCode.INTERNAL_SERVER_ERROR, "An error occurred") + ); + } +}