add targets and refactor endpoints

This commit is contained in:
miloschwartz
2026-07-31 17:36:04 -04:00
committed by Owen
parent 42c0abedb7
commit 6fa0009ebf
31 changed files with 741 additions and 490 deletions
+16 -4
View File
@@ -323,11 +323,18 @@ export const targets = pgTable(
"targets", "targets",
{ {
targetId: serial("targetId").primaryKey(), targetId: serial("targetId").primaryKey(),
resourceId: integer("resourceId") resourceId: integer("resourceId").references(
.references(() => resources.resourceId, { () => resources.resourceId,
{
onDelete: "cascade" onDelete: "cascade"
}) }
.notNull(), ),
providerId: integer("providerId").references(
() => aiProviders.providerId,
{
onDelete: "cascade"
}
),
siteId: integer("siteId") siteId: integer("siteId")
.references(() => sites.siteId, { .references(() => sites.siteId, {
onDelete: "cascade" onDelete: "cascade"
@@ -351,6 +358,7 @@ export const targets = pgTable(
}, },
(t) => [ (t) => [
index("idx_targets_resourceid_siteid").on(t.resourceId, t.siteId), index("idx_targets_resourceid_siteid").on(t.resourceId, t.siteId),
index("idx_targets_providerid_siteid").on(t.providerId, t.siteId),
index("idx_targets_site_enabled_priority_target_resource") index("idx_targets_site_enabled_priority_target_resource")
.on(t.siteId, t.priority.desc(), t.targetId, t.resourceId) .on(t.siteId, t.priority.desc(), t.targetId, t.resourceId)
.where(sql`${t.enabled} = true`) .where(sql`${t.enabled} = true`)
@@ -1572,6 +1580,10 @@ export const aiProviders = pgTable("aiProviders", {
apiKey: text("apiKey"), apiKey: text("apiKey"),
apiKeyLastChars: varchar("apiKeyLastChars"), apiKeyLastChars: varchar("apiKeyLastChars"),
authType: varchar("authType").$type<"bearer">(), authType: varchar("authType").$type<"bearer">(),
routingMode: varchar("routingMode")
.$type<"url" | "target">()
.notNull()
.default("url"),
skipTlsVerification: boolean("skipTlsVerification") skipTlsVerification: boolean("skipTlsVerification")
.notNull() .notNull()
.default(false), .default(false),
+10 -5
View File
@@ -326,11 +326,12 @@ export const clientLabels = sqliteTable(
export const targets = sqliteTable("targets", { export const targets = sqliteTable("targets", {
targetId: integer("targetId").primaryKey({ autoIncrement: true }), targetId: integer("targetId").primaryKey({ autoIncrement: true }),
resourceId: integer("resourceId") resourceId: integer("resourceId").references(() => resources.resourceId, {
.references(() => resources.resourceId, { onDelete: "cascade"
onDelete: "cascade" }),
}) providerId: integer("providerId").references(() => aiProviders.providerId, {
.notNull(), onDelete: "cascade"
}),
siteId: integer("siteId") siteId: integer("siteId")
.references(() => sites.siteId, { .references(() => sites.siteId, {
onDelete: "cascade" onDelete: "cascade"
@@ -1561,6 +1562,10 @@ export const aiProviders = sqliteTable("aiProviders", {
apiKey: text("apiKey"), apiKey: text("apiKey"),
apiKeyLastChars: text("apiKeyLastChars"), apiKeyLastChars: text("apiKeyLastChars"),
authType: text("authType").$type<"bearer">(), authType: text("authType").$type<"bearer">(),
routingMode: text("routingMode")
.$type<"url" | "target">()
.notNull()
.default("url"),
skipTlsVerification: integer("skipTlsVerification", { mode: "boolean" }) skipTlsVerification: integer("skipTlsVerification", { mode: "boolean" })
.notNull() .notNull()
.default(false), .default(false),
+24 -3
View File
@@ -11,6 +11,7 @@ export type AiProviderType =
export type AiProviderAuthType = "bearer"; export type AiProviderAuthType = "bearer";
export type AiBudgetUnit = "usd" | "tokens"; export type AiBudgetUnit = "usd" | "tokens";
export type AiProviderRoutingMode = "url" | "target";
type AiProviderDefaults = { type AiProviderDefaults = {
upstreamUrl: string | null; upstreamUrl: string | null;
@@ -55,7 +56,13 @@ export const AI_PROVIDER_DEFAULTS: Record<
} }
}; };
export function providerRequiresUpstreamUrl(type: AiProviderType): boolean { export function providerRequiresUpstreamUrl(
type: AiProviderType,
routingMode: AiProviderRoutingMode = "url"
): boolean {
if (routingMode === "target") {
return false;
}
if (type === "custom") { if (type === "custom") {
return true; return true;
} }
@@ -66,20 +73,34 @@ export function resolveAiProviderConfig(input: {
type: AiProviderType; type: AiProviderType;
upstreamUrl: string | null; upstreamUrl: string | null;
authType: AiProviderAuthType | null; authType: AiProviderAuthType | null;
routingMode?: AiProviderRoutingMode | null;
}): { }): {
upstreamUrl: string | null; upstreamUrl: string | null;
authType: AiProviderAuthType | null; authType: AiProviderAuthType | null;
routingMode: AiProviderRoutingMode;
} { } {
const routingMode = input.routingMode ?? "url";
if (routingMode === "target") {
return {
upstreamUrl: null,
authType: input.authType ?? "bearer",
routingMode
};
}
if (input.type === "custom") { if (input.type === "custom") {
return { return {
upstreamUrl: input.upstreamUrl, upstreamUrl: input.upstreamUrl,
authType: input.authType authType: input.authType,
routingMode
}; };
} }
const defaults = AI_PROVIDER_DEFAULTS[input.type]; const defaults = AI_PROVIDER_DEFAULTS[input.type];
return { return {
upstreamUrl: input.upstreamUrl ?? defaults.upstreamUrl, upstreamUrl: input.upstreamUrl ?? defaults.upstreamUrl,
authType: input.authType ?? defaults.authType authType: input.authType ?? defaults.authType,
routingMode
}; };
} }
@@ -202,6 +202,10 @@ async function handleResource(
return; return;
} }
if (!target.resourceId) {
return;
}
const [resource] = await trx const [resource] = await trx
.select() .select()
.from(resources) .from(resources)
@@ -227,9 +231,7 @@ async function handleResource(
let health = "healthy"; let health = "healthy";
const allUnknown = monitoredTargets.length === 0; const allUnknown = monitoredTargets.length === 0;
const allHealthy = monitoredTargets.every( const allHealthy = monitoredTargets.every((t) => t.hcHealth === "healthy");
(t) => t.hcHealth === "healthy"
);
const allUnhealthy = monitoredTargets.every( const allUnhealthy = monitoredTargets.every(
(t) => t.hcHealth === "unhealthy" (t) => t.hcHealth === "unhealthy"
); );
+8 -1
View File
@@ -64,13 +64,20 @@ export async function performDeleteResources(
const targetsByResourceId = new Map<number, Target[]>(); const targetsByResourceId = new Map<number, Target[]>();
for (const target of targetsToBeRemoved) { for (const target of targetsToBeRemoved) {
if (target.resourceId == null) {
continue;
}
const existing = targetsByResourceId.get(target.resourceId) ?? []; const existing = targetsByResourceId.get(target.resourceId) ?? [];
existing.push(target); existing.push(target);
targetsByResourceId.set(target.resourceId, existing); targetsByResourceId.set(target.resourceId, existing);
} }
const targetIdToResourceId = new Map( const targetIdToResourceId = new Map(
targetsToBeRemoved.map((target) => [target.targetId, target.resourceId]) targetsToBeRemoved.flatMap((target) =>
target.resourceId == null
? []
: [[target.targetId, target.resourceId] as const]
)
); );
const healthChecksByResourceId = new Map<number, TargetHealthCheck[]>(); const healthChecksByResourceId = new Map<number, TargetHealthCheck[]>();
+5 -3
View File
@@ -1,4 +1,4 @@
import { and, eq, inArray, sql } from "drizzle-orm"; import { and, eq, inArray, isNotNull, sql } from "drizzle-orm";
import { import {
db, db,
resources, resources,
@@ -33,9 +33,11 @@ export async function getResourceIdsForSite(
const rows = await trx const rows = await trx
.selectDistinct({ resourceId: targets.resourceId }) .selectDistinct({ resourceId: targets.resourceId })
.from(targets) .from(targets)
.where(eq(targets.siteId, siteId)); .where(and(eq(targets.siteId, siteId), isNotNull(targets.resourceId)));
return rows.map((row) => row.resourceId); return rows
.map((row) => row.resourceId)
.filter((resourceId): resourceId is number => resourceId != null);
} }
export async function getSiteResourceIdsForSite( export async function getSiteResourceIdsForSite(
@@ -14,9 +14,6 @@ export async function verifyApiKeyAiModelAccess(
const apiKey = req.apiKey; const apiKey = req.apiKey;
const modelIdRaw = getFirstString(req.params.modelId); const modelIdRaw = getFirstString(req.params.modelId);
const modelId = Number.parseInt(modelIdRaw ?? "", 10); const modelId = Number.parseInt(modelIdRaw ?? "", 10);
const providerIdRaw = getFirstString(req.params.providerId);
const providerId = Number.parseInt(providerIdRaw ?? "", 10);
const orgId = getFirstString(req.params.orgId);
if (!apiKey) { if (!apiKey) {
return next( return next(
@@ -24,18 +21,6 @@ export async function verifyApiKeyAiModelAccess(
); );
} }
if (!orgId) {
return next(
createHttpError(HttpCode.BAD_REQUEST, "Invalid organization ID")
);
}
if (Number.isNaN(providerId)) {
return next(
createHttpError(HttpCode.BAD_REQUEST, "Invalid provider ID")
);
}
if (Number.isNaN(modelId)) { if (Number.isNaN(modelId)) {
return next( return next(
createHttpError(HttpCode.BAD_REQUEST, "Invalid model ID") createHttpError(HttpCode.BAD_REQUEST, "Invalid model ID")
@@ -52,13 +37,7 @@ export async function verifyApiKeyAiModelAccess(
aiProviders, aiProviders,
eq(aiModels.providerId, aiProviders.providerId) eq(aiModels.providerId, aiProviders.providerId)
) )
.where( .where(eq(aiModels.modelId, modelId))
and(
eq(aiModels.modelId, modelId),
eq(aiModels.providerId, providerId),
eq(aiProviders.orgId, orgId)
)
)
.limit(1); .limit(1);
if (!row) { if (!row) {
@@ -76,7 +55,9 @@ export async function verifyApiKeyAiModelAccess(
return next(); return next();
} }
if (!req.apiKeyOrg) { const orgId = row.provider.orgId;
if (!req.apiKeyOrg || req.apiKeyOrg.orgId !== orgId) {
const apiKeyOrgRes = await db const apiKeyOrgRes = await db
.select() .select()
.from(apiKeyOrg) .from(apiKeyOrg)
@@ -14,7 +14,6 @@ export async function verifyApiKeyAiProviderAccess(
const apiKey = req.apiKey; const apiKey = req.apiKey;
const providerIdRaw = getFirstString(req.params.providerId); const providerIdRaw = getFirstString(req.params.providerId);
const providerId = Number.parseInt(providerIdRaw ?? "", 10); const providerId = Number.parseInt(providerIdRaw ?? "", 10);
const orgId = getFirstString(req.params.orgId);
if (!apiKey) { if (!apiKey) {
return next( return next(
@@ -22,52 +21,16 @@ export async function verifyApiKeyAiProviderAccess(
); );
} }
if (!orgId) {
return next(
createHttpError(HttpCode.BAD_REQUEST, "Invalid organization ID")
);
}
if (Number.isNaN(providerId)) { if (Number.isNaN(providerId)) {
return next( return next(
createHttpError(HttpCode.BAD_REQUEST, "Invalid provider ID") createHttpError(HttpCode.BAD_REQUEST, "Invalid provider ID")
); );
} }
if (apiKey.isRoot) {
const [provider] = await db
.select()
.from(aiProviders)
.where(
and(
eq(aiProviders.providerId, providerId),
eq(aiProviders.orgId, orgId)
)
)
.limit(1);
if (!provider) {
return next(
createHttpError(
HttpCode.NOT_FOUND,
`AI provider with ID ${providerId} not found`
)
);
}
req.aiProvider = provider;
return next();
}
const [provider] = await db const [provider] = await db
.select() .select()
.from(aiProviders) .from(aiProviders)
.where( .where(eq(aiProviders.providerId, providerId))
and(
eq(aiProviders.providerId, providerId),
eq(aiProviders.orgId, orgId)
)
)
.limit(1); .limit(1);
if (!provider) { if (!provider) {
@@ -79,7 +42,14 @@ export async function verifyApiKeyAiProviderAccess(
); );
} }
if (!req.apiKeyOrg) { if (apiKey.isRoot) {
req.aiProvider = provider;
return next();
}
const orgId = provider.orgId;
if (!req.apiKeyOrg || req.apiKeyOrg.orgId !== orgId) {
const apiKeyOrgRes = await db const apiKeyOrgRes = await db
.select() .select()
.from(apiKeyOrg) .from(apiKeyOrg)
@@ -1,6 +1,6 @@
import { Request, Response, NextFunction } from "express"; import { Request, Response, NextFunction } from "express";
import { db } from "@server/db"; import { db } from "@server/db";
import { resources, targets, apiKeyOrg } from "@server/db"; import { aiProviders, resources, targets, apiKeyOrg } from "@server/db";
import { and, eq } from "drizzle-orm"; import { and, eq } from "drizzle-orm";
import createHttpError from "http-errors"; import createHttpError from "http-errors";
import HttpCode from "@server/types/HttpCode"; import HttpCode from "@server/types/HttpCode";
@@ -43,43 +43,65 @@ export async function verifyApiKeyTargetAccess(
); );
} }
const resourceId = target.resourceId; const { resourceId, providerId } = target;
if (!resourceId) { if ((!resourceId && !providerId) || (resourceId && providerId)) {
return next( return next(
createHttpError( createHttpError(
HttpCode.INTERNAL_SERVER_ERROR, HttpCode.INTERNAL_SERVER_ERROR,
`Target with ID ${targetId} does not have a resource ID` `Target with ID ${targetId} has invalid ownership`
)
);
}
const [resource] = await db
.select()
.from(resources)
.where(eq(resources.resourceId, resourceId))
.limit(1);
if (!resource) {
return next(
createHttpError(
HttpCode.NOT_FOUND,
`Resource with ID ${resourceId} not found`
) )
); );
} }
if (apiKey.isRoot) { if (apiKey.isRoot) {
// Root keys can access any key in any org // Root keys can access any target
return next(); return next();
} }
if (!resource.orgId) { let orgId: string;
return next( if (resourceId) {
createHttpError( const [resource] = await db
HttpCode.INTERNAL_SERVER_ERROR, .select()
`Resource with ID ${resourceId} does not have an organization ID` .from(resources)
) .where(eq(resources.resourceId, resourceId))
); .limit(1);
if (!resource) {
return next(
createHttpError(
HttpCode.NOT_FOUND,
`Resource with ID ${resourceId} not found`
)
);
}
if (!resource.orgId) {
return next(
createHttpError(
HttpCode.INTERNAL_SERVER_ERROR,
`Resource with ID ${resourceId} does not have an organization ID`
)
);
}
orgId = resource.orgId;
} else {
const [provider] = await db
.select()
.from(aiProviders)
.where(eq(aiProviders.providerId, providerId!))
.limit(1);
if (!provider) {
return next(
createHttpError(
HttpCode.NOT_FOUND,
`AI provider with ID ${providerId} not found`
)
);
}
orgId = provider.orgId;
} }
if (!req.apiKeyOrg) { if (!req.apiKeyOrg) {
@@ -89,7 +111,7 @@ export async function verifyApiKeyTargetAccess(
.where( .where(
and( and(
eq(apiKeyOrg.apiKeyId, apiKey.apiKeyId), eq(apiKeyOrg.apiKeyId, apiKey.apiKeyId),
eq(apiKeyOrg.orgId, resource.orgId) eq(apiKeyOrg.orgId, orgId)
) )
) )
.limit(1); .limit(1);
@@ -98,7 +120,7 @@ export async function verifyApiKeyTargetAccess(
} }
} }
if (!req.apiKeyOrg) { if (!req.apiKeyOrg || req.apiKeyOrg.orgId !== orgId) {
return next( return next(
createHttpError( createHttpError(
HttpCode.FORBIDDEN, HttpCode.FORBIDDEN,
+5 -23
View File
@@ -16,9 +16,6 @@ export async function verifyAiModelAccess(
const userId = req.user!.userId; const userId = req.user!.userId;
const modelIdRaw = getFirstString(req.params.modelId); const modelIdRaw = getFirstString(req.params.modelId);
const modelId = Number.parseInt(modelIdRaw ?? "", 10); const modelId = Number.parseInt(modelIdRaw ?? "", 10);
const providerIdRaw = getFirstString(req.params.providerId);
const providerId = Number.parseInt(providerIdRaw ?? "", 10);
const orgId = getFirstString(req.params.orgId);
if (!userId) { if (!userId) {
return next( return next(
@@ -26,18 +23,6 @@ export async function verifyAiModelAccess(
); );
} }
if (!orgId) {
return next(
createHttpError(HttpCode.BAD_REQUEST, "Invalid organization ID")
);
}
if (Number.isNaN(providerId)) {
return next(
createHttpError(HttpCode.BAD_REQUEST, "Invalid provider ID")
);
}
if (Number.isNaN(modelId)) { if (Number.isNaN(modelId)) {
return next( return next(
createHttpError(HttpCode.BAD_REQUEST, "Invalid model ID") createHttpError(HttpCode.BAD_REQUEST, "Invalid model ID")
@@ -54,13 +39,7 @@ export async function verifyAiModelAccess(
aiProviders, aiProviders,
eq(aiModels.providerId, aiProviders.providerId) eq(aiModels.providerId, aiProviders.providerId)
) )
.where( .where(eq(aiModels.modelId, modelId))
and(
eq(aiModels.modelId, modelId),
eq(aiModels.providerId, providerId),
eq(aiProviders.orgId, orgId)
)
)
.limit(1); .limit(1);
if (!row) { if (!row) {
@@ -72,7 +51,9 @@ export async function verifyAiModelAccess(
); );
} }
if (!req.userOrg) { const orgId = row.provider.orgId;
if (!req.userOrg || req.userOrg.orgId !== orgId) {
const userOrgRole = await db const userOrgRole = await db
.select() .select()
.from(userOrgs) .from(userOrgs)
@@ -109,6 +90,7 @@ export async function verifyAiModelAccess(
} }
} }
req.userOrgId = orgId;
req.userOrgRoleIds = await getUserOrgRoleIds(req.userOrg.userId, orgId); req.userOrgRoleIds = await getUserOrgRoleIds(req.userOrg.userId, orgId);
req.aiProvider = row.provider; req.aiProvider = row.provider;
req.aiModel = row.model; req.aiModel = row.model;
+5 -14
View File
@@ -16,7 +16,6 @@ export async function verifyAiProviderAccess(
const userId = req.user!.userId; const userId = req.user!.userId;
const providerIdRaw = getFirstString(req.params.providerId); const providerIdRaw = getFirstString(req.params.providerId);
const providerId = Number.parseInt(providerIdRaw ?? "", 10); const providerId = Number.parseInt(providerIdRaw ?? "", 10);
const orgId = getFirstString(req.params.orgId);
if (!userId) { if (!userId) {
return next( return next(
@@ -24,12 +23,6 @@ export async function verifyAiProviderAccess(
); );
} }
if (!orgId) {
return next(
createHttpError(HttpCode.BAD_REQUEST, "Invalid organization ID")
);
}
if (Number.isNaN(providerId)) { if (Number.isNaN(providerId)) {
return next( return next(
createHttpError(HttpCode.BAD_REQUEST, "Invalid provider ID") createHttpError(HttpCode.BAD_REQUEST, "Invalid provider ID")
@@ -39,12 +32,7 @@ export async function verifyAiProviderAccess(
const [provider] = await db const [provider] = await db
.select() .select()
.from(aiProviders) .from(aiProviders)
.where( .where(eq(aiProviders.providerId, providerId))
and(
eq(aiProviders.providerId, providerId),
eq(aiProviders.orgId, orgId)
)
)
.limit(1); .limit(1);
if (!provider) { if (!provider) {
@@ -56,7 +44,9 @@ export async function verifyAiProviderAccess(
); );
} }
if (!req.userOrg) { const orgId = provider.orgId;
if (!req.userOrg || req.userOrg.orgId !== orgId) {
const userOrgRole = await db const userOrgRole = await db
.select() .select()
.from(userOrgs) .from(userOrgs)
@@ -93,6 +83,7 @@ export async function verifyAiProviderAccess(
} }
} }
req.userOrgId = orgId;
req.userOrgRoleIds = await getUserOrgRoleIds(req.userOrg.userId, orgId); req.userOrgRoleIds = await getUserOrgRoleIds(req.userOrg.userId, orgId);
req.aiProvider = provider; req.aiProvider = provider;
+71 -56
View File
@@ -1,6 +1,6 @@
import { Request, Response, NextFunction } from "express"; import { Request, Response, NextFunction } from "express";
import { db } from "@server/db"; import { db } from "@server/db";
import { resources, targets, userOrgs } from "@server/db"; import { aiProviders, resources, targets, userOrgs } from "@server/db";
import { and, eq } from "drizzle-orm"; import { and, eq } from "drizzle-orm";
import createHttpError from "http-errors"; import createHttpError from "http-errors";
import HttpCode from "@server/types/HttpCode"; import HttpCode from "@server/types/HttpCode";
@@ -25,9 +25,7 @@ export async function verifyTargetAccess(
} }
if (isNaN(targetId)) { if (isNaN(targetId)) {
return next( return next(createHttpError(HttpCode.BAD_REQUEST, "Invalid target ID"));
createHttpError(HttpCode.BAD_REQUEST, "Invalid organization ID")
);
} }
const target = await db const target = await db
@@ -45,73 +43,88 @@ export async function verifyTargetAccess(
); );
} }
const resourceId = target[0].resourceId; const { resourceId, providerId } = target[0];
if (!resourceId) { if ((!resourceId && !providerId) || (resourceId && providerId)) {
return next( return next(
createHttpError( createHttpError(
HttpCode.INTERNAL_SERVER_ERROR, HttpCode.INTERNAL_SERVER_ERROR,
`Target with ID ${targetId} does not have a resource ID` `Target with ID ${targetId} has invalid ownership`
) )
); );
} }
try { try {
const resource = await db let orgId: string;
.select()
.from(resources)
.where(eq(resources.resourceId, resourceId!))
.limit(1);
if (resource.length === 0) { if (resourceId) {
return next( const [resource] = await db
createHttpError( .select()
HttpCode.NOT_FOUND, .from(resources)
`Resource with ID ${resourceId} not found` .where(eq(resources.resourceId, resourceId))
) .limit(1);
);
}
if (!resource[0].orgId) { if (!resource) {
return next( return next(
createHttpError( createHttpError(
HttpCode.INTERNAL_SERVER_ERROR, HttpCode.NOT_FOUND,
`resource with ID ${resourceId} does not have an organization ID` `Resource with ID ${resourceId} not found`
) )
); );
}
if (!resource.orgId) {
return next(
createHttpError(
HttpCode.INTERNAL_SERVER_ERROR,
`Resource with ID ${resourceId} does not have an organization ID`
)
);
}
orgId = resource.orgId;
} else {
const [provider] = await db
.select()
.from(aiProviders)
.where(eq(aiProviders.providerId, providerId!))
.limit(1);
if (!provider) {
return next(
createHttpError(
HttpCode.NOT_FOUND,
`AI provider with ID ${providerId} not found`
)
);
}
orgId = provider.orgId;
} }
if (!req.userOrg) { if (!req.userOrg) {
const res = await db const userOrgResult = await db
.select() .select()
.from(userOrgs) .from(userOrgs)
.where( .where(
and( and(eq(userOrgs.userId, userId), eq(userOrgs.orgId, orgId))
eq(userOrgs.userId, userId),
eq(userOrgs.orgId, resource[0].orgId)
)
); );
req.userOrg = res[0]; req.userOrg = userOrgResult[0];
} }
if (!req.userOrg) { if (!req.userOrg || req.userOrg.orgId !== orgId) {
next( return next(
createHttpError( createHttpError(
HttpCode.FORBIDDEN, HttpCode.FORBIDDEN,
"User does not have access to this organization" "User does not have access to this organization"
) )
); );
} else {
req.userOrgRoleIds = await getUserOrgRoleIds(
req.userOrg.userId,
resource[0].orgId!
);
req.userOrgId = resource[0].orgId!;
} }
const orgId = req.userOrg.orgId; req.userOrgRoleIds = await getUserOrgRoleIds(req.userOrg.userId, orgId);
req.userOrgId = orgId;
if (req.orgPolicyAllowed === undefined && orgId) { if (req.orgPolicyAllowed === undefined) {
const policyCheck = await checkOrgAccessPolicy({ const policyCheck = await checkOrgAccessPolicy({
orgId, orgId,
userId, userId,
@@ -128,22 +141,24 @@ export async function verifyTargetAccess(
} }
} }
const resourceAllowed = await canUserAccessResource({ if (resourceId) {
userId, const resourceAllowed = await canUserAccessResource({
resourceId, userId,
roleIds: req.userOrgRoleIds ?? [] resourceId,
}); roleIds: req.userOrgRoleIds ?? []
});
if (!resourceAllowed) { if (!resourceAllowed) {
return next( return next(
createHttpError( createHttpError(
HttpCode.FORBIDDEN, HttpCode.FORBIDDEN,
"User does not have access to this resource" "User does not have access to this resource"
) )
); );
}
} }
next(); return next();
} catch (e) { } catch (e) {
return next( return next(
createHttpError( createHttpError(
+3 -9
View File
@@ -15,7 +15,6 @@ import {
} from "@server/routers/aiProvider/validation"; } from "@server/routers/aiProvider/validation";
const paramsSchema = z.strictObject({ const paramsSchema = z.strictObject({
orgId: z.string().nonempty(),
providerId: z.coerce.number().int().positive() providerId: z.coerce.number().int().positive()
}); });
@@ -33,7 +32,7 @@ const bodySchema = z
registry.registerPath({ registry.registerPath({
method: "put", method: "put",
path: "/org/{orgId}/ai-provider/{providerId}/model", path: "/ai-provider/{providerId}/model",
description: "Create an AI model under a provider.", description: "Create an AI model under a provider.",
tags: [OpenAPITags.AiModel], tags: [OpenAPITags.AiModel],
request: { request: {
@@ -79,7 +78,7 @@ export async function createAiModel(
); );
} }
const { orgId, providerId } = parsedParams.data; const { providerId } = parsedParams.data;
const { modelKey, name, budgetAmount, budgetUnit, enabled } = const { modelKey, name, budgetAmount, budgetUnit, enabled } =
parsedBody.data; parsedBody.data;
@@ -89,12 +88,7 @@ export async function createAiModel(
: await db : await db
.select() .select()
.from(aiProviders) .from(aiProviders)
.where( .where(eq(aiProviders.providerId, providerId))
and(
eq(aiProviders.providerId, providerId),
eq(aiProviders.orgId, orgId)
)
)
.limit(1); .limit(1);
if (!provider) { if (!provider) {
+10 -1
View File
@@ -15,6 +15,7 @@ import {
aiAuthTypeSchema, aiAuthTypeSchema,
aiBudgetUnitSchema, aiBudgetUnitSchema,
aiProviderTypeSchema, aiProviderTypeSchema,
aiRoutingModeSchema,
refineBudgetFields, refineBudgetFields,
refineProviderUpstreamFields refineProviderUpstreamFields
} from "@server/routers/aiProvider/validation"; } from "@server/routers/aiProvider/validation";
@@ -30,6 +31,7 @@ const bodySchema = z
upstreamUrl: z.url().optional().nullable(), upstreamUrl: z.url().optional().nullable(),
apiKey: z.string().optional(), apiKey: z.string().optional(),
authType: aiAuthTypeSchema.optional().nullable(), authType: aiAuthTypeSchema.optional().nullable(),
routingMode: aiRoutingModeSchema.optional(),
skipTlsVerification: z.boolean().optional(), skipTlsVerification: z.boolean().optional(),
budgetAmount: z.number().positive().optional().nullable(), budgetAmount: z.number().positive().optional().nullable(),
budgetUnit: aiBudgetUnitSchema.optional().nullable(), budgetUnit: aiBudgetUnitSchema.optional().nullable(),
@@ -95,6 +97,7 @@ export async function createAiProvider(
upstreamUrl, upstreamUrl,
apiKey, apiKey,
authType, authType,
routingMode,
skipTlsVerification, skipTlsVerification,
budgetAmount, budgetAmount,
budgetUnit, budgetUnit,
@@ -105,6 +108,8 @@ export async function createAiProvider(
const encryptedApiKey = apiKey ? encrypt(apiKey, key) : null; const encryptedApiKey = apiKey ? encrypt(apiKey, key) : null;
const apiKeyLastChars = apiKey ? apiKey.slice(-4) : null; const apiKeyLastChars = apiKey ? apiKey.slice(-4) : null;
const now = Date.now(); const now = Date.now();
const resolvedRoutingMode =
type === "custom" ? (routingMode ?? "url") : "url";
const [provider] = await db const [provider] = await db
.insert(aiProviders) .insert(aiProviders)
@@ -112,10 +117,14 @@ export async function createAiProvider(
orgId, orgId,
name, name,
type, type,
upstreamUrl: upstreamUrl ?? null, upstreamUrl:
resolvedRoutingMode === "target"
? null
: (upstreamUrl ?? null),
apiKey: encryptedApiKey, apiKey: encryptedApiKey,
apiKeyLastChars, apiKeyLastChars,
authType: authType ?? null, authType: authType ?? null,
routingMode: resolvedRoutingMode,
skipTlsVerification: skipTlsVerification ?? false, skipTlsVerification: skipTlsVerification ?? false,
budgetAmount: budgetAmount ?? null, budgetAmount: budgetAmount ?? null,
budgetUnit: budgetUnit ?? null, budgetUnit: budgetUnit ?? null,
+6 -25
View File
@@ -1,23 +1,21 @@
import { Request, Response, NextFunction } from "express"; import { Request, Response, NextFunction } from "express";
import { z } from "zod"; import { z } from "zod";
import { aiModels, aiProviders, db } from "@server/db"; import { aiModels, db } from "@server/db";
import response from "@server/lib/response"; import response from "@server/lib/response";
import HttpCode from "@server/types/HttpCode"; import HttpCode from "@server/types/HttpCode";
import createHttpError from "http-errors"; import createHttpError from "http-errors";
import logger from "@server/logger"; import logger from "@server/logger";
import { fromError } from "zod-validation-error"; import { fromError } from "zod-validation-error";
import { OpenAPITags, registry } from "@server/openApi"; import { OpenAPITags, registry } from "@server/openApi";
import { and, eq } from "drizzle-orm"; import { eq } from "drizzle-orm";
const paramsSchema = z.strictObject({ const paramsSchema = z.strictObject({
orgId: z.string().nonempty(),
providerId: z.coerce.number().int().positive(),
modelId: z.coerce.number().int().positive() modelId: z.coerce.number().int().positive()
}); });
registry.registerPath({ registry.registerPath({
method: "delete", method: "delete",
path: "/org/{orgId}/ai-provider/{providerId}/model/{modelId}", path: "/ai-model/{modelId}",
description: "Delete an AI model.", description: "Delete an AI model.",
tags: [OpenAPITags.AiModel], tags: [OpenAPITags.AiModel],
request: { request: {
@@ -46,22 +44,12 @@ export async function deleteAiModel(
); );
} }
const { orgId, providerId, modelId } = parsedParams.data; const { modelId } = parsedParams.data;
const [existing] = await db const [existing] = await db
.select({ modelId: aiModels.modelId }) .select({ modelId: aiModels.modelId })
.from(aiModels) .from(aiModels)
.innerJoin( .where(eq(aiModels.modelId, modelId))
aiProviders,
eq(aiModels.providerId, aiProviders.providerId)
)
.where(
and(
eq(aiModels.modelId, modelId),
eq(aiModels.providerId, providerId),
eq(aiProviders.orgId, orgId)
)
)
.limit(1); .limit(1);
if (!existing) { if (!existing) {
@@ -73,14 +61,7 @@ export async function deleteAiModel(
); );
} }
await db await db.delete(aiModels).where(eq(aiModels.modelId, modelId));
.delete(aiModels)
.where(
and(
eq(aiModels.modelId, modelId),
eq(aiModels.providerId, providerId)
)
);
return response(res, { return response(res, {
data: null, data: null,
+5 -16
View File
@@ -7,16 +7,15 @@ import createHttpError from "http-errors";
import logger from "@server/logger"; import logger from "@server/logger";
import { fromError } from "zod-validation-error"; import { fromError } from "zod-validation-error";
import { OpenAPITags, registry } from "@server/openApi"; import { OpenAPITags, registry } from "@server/openApi";
import { and, eq } from "drizzle-orm"; import { eq } from "drizzle-orm";
const paramsSchema = z.strictObject({ const paramsSchema = z.strictObject({
orgId: z.string().nonempty(),
providerId: z.coerce.number().int().positive() providerId: z.coerce.number().int().positive()
}); });
registry.registerPath({ registry.registerPath({
method: "delete", method: "delete",
path: "/org/{orgId}/ai-provider/{providerId}", path: "/ai-provider/{providerId}",
description: "Delete an AI provider.", description: "Delete an AI provider.",
tags: [OpenAPITags.AiProvider], tags: [OpenAPITags.AiProvider],
request: { request: {
@@ -45,17 +44,12 @@ export async function deleteAiProvider(
); );
} }
const { orgId, providerId } = parsedParams.data; const { providerId } = parsedParams.data;
const [existing] = await db const [existing] = await db
.select({ providerId: aiProviders.providerId }) .select({ providerId: aiProviders.providerId })
.from(aiProviders) .from(aiProviders)
.where( .where(eq(aiProviders.providerId, providerId))
and(
eq(aiProviders.providerId, providerId),
eq(aiProviders.orgId, orgId)
)
)
.limit(1); .limit(1);
if (!existing) { if (!existing) {
@@ -69,12 +63,7 @@ export async function deleteAiProvider(
await db await db
.delete(aiProviders) .delete(aiProviders)
.where( .where(eq(aiProviders.providerId, providerId));
and(
eq(aiProviders.providerId, providerId),
eq(aiProviders.orgId, orgId)
)
);
return response(res, { return response(res, {
data: null, data: null,
+7 -20
View File
@@ -1,24 +1,22 @@
import { Request, Response, NextFunction } from "express"; import { Request, Response, NextFunction } from "express";
import { z } from "zod"; import { z } from "zod";
import { aiModels, aiProviders, db } from "@server/db"; import { aiModels, db } from "@server/db";
import response from "@server/lib/response"; import response from "@server/lib/response";
import HttpCode from "@server/types/HttpCode"; import HttpCode from "@server/types/HttpCode";
import createHttpError from "http-errors"; import createHttpError from "http-errors";
import logger from "@server/logger"; import logger from "@server/logger";
import { fromError } from "zod-validation-error"; import { fromError } from "zod-validation-error";
import { OpenAPITags, registry } from "@server/openApi"; import { OpenAPITags, registry } from "@server/openApi";
import { and, eq } from "drizzle-orm"; import { eq } from "drizzle-orm";
import type { GetAiModelResponse } from "@server/routers/aiProvider/types"; import type { GetAiModelResponse } from "@server/routers/aiProvider/types";
const paramsSchema = z.strictObject({ const paramsSchema = z.strictObject({
orgId: z.string().nonempty(),
providerId: z.coerce.number().int().positive(),
modelId: z.coerce.number().int().positive() modelId: z.coerce.number().int().positive()
}); });
registry.registerPath({ registry.registerPath({
method: "get", method: "get",
path: "/org/{orgId}/ai-provider/{providerId}/model/{modelId}", path: "/ai-model/{modelId}",
description: "Get an AI model by ID.", description: "Get an AI model by ID.",
tags: [OpenAPITags.AiModel], tags: [OpenAPITags.AiModel],
request: { request: {
@@ -47,27 +45,16 @@ export async function getAiModel(
); );
} }
const { orgId, providerId, modelId } = parsedParams.data; const { modelId } = parsedParams.data;
const [model] = const [model] =
req.aiModel && req.aiModel.modelId === modelId req.aiModel && req.aiModel.modelId === modelId
? [req.aiModel] ? [req.aiModel]
: await db : await db
.select({ model: aiModels }) .select()
.from(aiModels) .from(aiModels)
.innerJoin( .where(eq(aiModels.modelId, modelId))
aiProviders, .limit(1);
eq(aiModels.providerId, aiProviders.providerId)
)
.where(
and(
eq(aiModels.modelId, modelId),
eq(aiModels.providerId, providerId),
eq(aiProviders.orgId, orgId)
)
)
.limit(1)
.then((rows) => rows.map((r) => r.model));
if (!model) { if (!model) {
return next( return next(
+4 -10
View File
@@ -7,18 +7,17 @@ import createHttpError from "http-errors";
import logger from "@server/logger"; import logger from "@server/logger";
import { fromError } from "zod-validation-error"; import { fromError } from "zod-validation-error";
import { OpenAPITags, registry } from "@server/openApi"; import { OpenAPITags, registry } from "@server/openApi";
import { and, eq } from "drizzle-orm"; import { eq } from "drizzle-orm";
import type { GetAiProviderResponse } from "@server/routers/aiProvider/types"; import type { GetAiProviderResponse } from "@server/routers/aiProvider/types";
import { toPublicAiProvider } from "@server/routers/aiProvider/types"; import { toPublicAiProvider } from "@server/routers/aiProvider/types";
const paramsSchema = z.strictObject({ const paramsSchema = z.strictObject({
orgId: z.string().nonempty(),
providerId: z.coerce.number().int().positive() providerId: z.coerce.number().int().positive()
}); });
registry.registerPath({ registry.registerPath({
method: "get", method: "get",
path: "/org/{orgId}/ai-provider/{providerId}", path: "/ai-provider/{providerId}",
description: "Get an AI provider by ID.", description: "Get an AI provider by ID.",
tags: [OpenAPITags.AiProvider], tags: [OpenAPITags.AiProvider],
request: { request: {
@@ -47,7 +46,7 @@ export async function getAiProvider(
); );
} }
const { orgId, providerId } = parsedParams.data; const { providerId } = parsedParams.data;
const [provider] = const [provider] =
req.aiProvider && req.aiProvider.providerId === providerId req.aiProvider && req.aiProvider.providerId === providerId
@@ -55,12 +54,7 @@ export async function getAiProvider(
: await db : await db
.select() .select()
.from(aiProviders) .from(aiProviders)
.where( .where(eq(aiProviders.providerId, providerId))
and(
eq(aiProviders.providerId, providerId),
eq(aiProviders.orgId, orgId)
)
)
.limit(1); .limit(1);
if (!provider) { if (!provider) {
+3 -9
View File
@@ -11,7 +11,6 @@ import { and, asc, eq, like, sql } from "drizzle-orm";
import type { ListAiModelsResponse } from "@server/routers/aiProvider/types"; import type { ListAiModelsResponse } from "@server/routers/aiProvider/types";
const paramsSchema = z.strictObject({ const paramsSchema = z.strictObject({
orgId: z.string().nonempty(),
providerId: z.coerce.number().int().positive() providerId: z.coerce.number().int().positive()
}); });
@@ -45,7 +44,7 @@ const listSchema = z.object({
registry.registerPath({ registry.registerPath({
method: "get", method: "get",
path: "/org/{orgId}/ai-provider/{providerId}/models", path: "/ai-provider/{providerId}/models",
description: "List AI models for a provider.", description: "List AI models for a provider.",
tags: [OpenAPITags.AiModel], tags: [OpenAPITags.AiModel],
request: { request: {
@@ -85,7 +84,7 @@ export async function listAiModels(
); );
} }
const { orgId, providerId } = parsedParams.data; const { providerId } = parsedParams.data;
const [provider] = const [provider] =
req.aiProvider && req.aiProvider.providerId === providerId req.aiProvider && req.aiProvider.providerId === providerId
@@ -93,12 +92,7 @@ export async function listAiModels(
: await db : await db
.select({ providerId: aiProviders.providerId }) .select({ providerId: aiProviders.providerId })
.from(aiProviders) .from(aiProviders)
.where( .where(eq(aiProviders.providerId, providerId))
and(
eq(aiProviders.providerId, providerId),
eq(aiProviders.orgId, orgId)
)
)
.limit(1); .limit(1);
if (!provider) { if (!provider) {
+3 -1
View File
@@ -3,6 +3,7 @@ import type { PaginatedResponse } from "@server/types/Pagination";
import { import {
resolveAiProviderConfig, resolveAiProviderConfig,
type AiProviderAuthType, type AiProviderAuthType,
type AiProviderRoutingMode,
type AiProviderType type AiProviderType
} from "@server/lib/aiProviderDefaults"; } from "@server/lib/aiProviderDefaults";
@@ -40,7 +41,8 @@ export function toPublicAiProvider(provider: AiProvider): AiProviderPublic {
const resolved = resolveAiProviderConfig({ const resolved = resolveAiProviderConfig({
type: provider.type as AiProviderType, type: provider.type as AiProviderType,
upstreamUrl: provider.upstreamUrl, upstreamUrl: provider.upstreamUrl,
authType: provider.authType as AiProviderAuthType | null authType: provider.authType as AiProviderAuthType | null,
routingMode: provider.routingMode as AiProviderRoutingMode | null
}); });
return { return {
+8 -26
View File
@@ -1,6 +1,6 @@
import { Request, Response, NextFunction } from "express"; import { Request, Response, NextFunction } from "express";
import { z } from "zod"; import { z } from "zod";
import { aiModels, aiProviders, db } from "@server/db"; import { aiModels, db } from "@server/db";
import response from "@server/lib/response"; import response from "@server/lib/response";
import HttpCode from "@server/types/HttpCode"; import HttpCode from "@server/types/HttpCode";
import createHttpError from "http-errors"; import createHttpError from "http-errors";
@@ -15,8 +15,6 @@ import {
} from "@server/routers/aiProvider/validation"; } from "@server/routers/aiProvider/validation";
const paramsSchema = z.strictObject({ const paramsSchema = z.strictObject({
orgId: z.string().nonempty(),
providerId: z.coerce.number().int().positive(),
modelId: z.coerce.number().int().positive() modelId: z.coerce.number().int().positive()
}); });
@@ -34,7 +32,7 @@ const bodySchema = z
registry.registerPath({ registry.registerPath({
method: "post", method: "post",
path: "/org/{orgId}/ai-provider/{providerId}/model/{modelId}", path: "/ai-model/{modelId}",
description: "Update an AI model.", description: "Update an AI model.",
tags: [OpenAPITags.AiModel], tags: [OpenAPITags.AiModel],
request: { request: {
@@ -80,28 +78,17 @@ export async function updateAiModel(
); );
} }
const { orgId, providerId, modelId } = parsedParams.data; const { modelId } = parsedParams.data;
const body = parsedBody.data; const body = parsedBody.data;
const [existing] = const [existing] =
req.aiModel && req.aiModel.modelId === modelId req.aiModel && req.aiModel.modelId === modelId
? [req.aiModel] ? [req.aiModel]
: await db : await db
.select({ model: aiModels }) .select()
.from(aiModels) .from(aiModels)
.innerJoin( .where(eq(aiModels.modelId, modelId))
aiProviders, .limit(1);
eq(aiModels.providerId, aiProviders.providerId)
)
.where(
and(
eq(aiModels.modelId, modelId),
eq(aiModels.providerId, providerId),
eq(aiProviders.orgId, orgId)
)
)
.limit(1)
.then((rows) => rows.map((r) => r.model));
if (!existing) { if (!existing) {
return next( return next(
@@ -121,7 +108,7 @@ export async function updateAiModel(
.from(aiModels) .from(aiModels)
.where( .where(
and( and(
eq(aiModels.providerId, providerId), eq(aiModels.providerId, existing.providerId),
eq(aiModels.modelKey, body.modelKey), eq(aiModels.modelKey, body.modelKey),
ne(aiModels.modelId, modelId) ne(aiModels.modelId, modelId)
) )
@@ -161,12 +148,7 @@ export async function updateAiModel(
const [model] = await db const [model] = await db
.update(aiModels) .update(aiModels)
.set(updateData) .set(updateData)
.where( .where(eq(aiModels.modelId, modelId))
and(
eq(aiModels.modelId, modelId),
eq(aiModels.providerId, providerId)
)
)
.returning(); .returning();
return response<CreateOrEditAiModelResponse>(res, { return response<CreateOrEditAiModelResponse>(res, {
+41 -48
View File
@@ -7,7 +7,7 @@ import createHttpError from "http-errors";
import logger from "@server/logger"; import logger from "@server/logger";
import { fromError } from "zod-validation-error"; import { fromError } from "zod-validation-error";
import { OpenAPITags, registry } from "@server/openApi"; import { OpenAPITags, registry } from "@server/openApi";
import { and, eq } from "drizzle-orm"; import { eq } from "drizzle-orm";
import { encrypt } from "@server/lib/crypto"; import { encrypt } from "@server/lib/crypto";
import config from "@server/lib/config"; import config from "@server/lib/config";
import type { CreateOrEditAiProviderResponse } from "@server/routers/aiProvider/types"; import type { CreateOrEditAiProviderResponse } from "@server/routers/aiProvider/types";
@@ -16,16 +16,16 @@ import {
aiAuthTypeSchema, aiAuthTypeSchema,
aiBudgetUnitSchema, aiBudgetUnitSchema,
aiProviderTypeSchema, aiProviderTypeSchema,
aiRoutingModeSchema,
refineBudgetFields, refineBudgetFields,
refineProviderUpstreamFields refineProviderUpstreamFields
} from "@server/routers/aiProvider/validation"; } from "@server/routers/aiProvider/validation";
import { import type {
providerRequiresUpstreamUrl, AiProviderRoutingMode,
type AiProviderType AiProviderType
} from "@server/lib/aiProviderDefaults"; } from "@server/lib/aiProviderDefaults";
const paramsSchema = z.strictObject({ const paramsSchema = z.strictObject({
orgId: z.string().nonempty(),
providerId: z.coerce.number().int().positive() providerId: z.coerce.number().int().positive()
}); });
@@ -35,6 +35,7 @@ const bodySchema = z
upstreamUrl: z.url().optional().nullable(), upstreamUrl: z.url().optional().nullable(),
apiKey: z.string().optional(), apiKey: z.string().optional(),
authType: aiAuthTypeSchema.optional().nullable(), authType: aiAuthTypeSchema.optional().nullable(),
routingMode: aiRoutingModeSchema.optional(),
skipTlsVerification: z.boolean().optional(), skipTlsVerification: z.boolean().optional(),
budgetAmount: z.number().positive().optional().nullable(), budgetAmount: z.number().positive().optional().nullable(),
budgetUnit: aiBudgetUnitSchema.optional().nullable(), budgetUnit: aiBudgetUnitSchema.optional().nullable(),
@@ -46,7 +47,7 @@ const bodySchema = z
registry.registerPath({ registry.registerPath({
method: "post", method: "post",
path: "/org/{orgId}/ai-provider/{providerId}", path: "/ai-provider/{providerId}",
description: "Update an AI provider.", description: "Update an AI provider.",
tags: [OpenAPITags.AiProvider], tags: [OpenAPITags.AiProvider],
request: { request: {
@@ -92,7 +93,7 @@ export async function updateAiProvider(
); );
} }
const { orgId, providerId } = parsedParams.data; const { providerId } = parsedParams.data;
const body = parsedBody.data; const body = parsedBody.data;
const [existing] = const [existing] =
@@ -101,12 +102,7 @@ export async function updateAiProvider(
: await db : await db
.select() .select()
.from(aiProviders) .from(aiProviders)
.where( .where(eq(aiProviders.providerId, providerId))
and(
eq(aiProviders.providerId, providerId),
eq(aiProviders.orgId, orgId)
)
)
.limit(1); .limit(1);
if (!existing) { if (!existing) {
@@ -119,6 +115,11 @@ export async function updateAiProvider(
} }
const providerType = existing.type as AiProviderType; const providerType = existing.type as AiProviderType;
const nextRoutingMode: AiProviderRoutingMode =
providerType === "custom"
? ((body.routingMode ??
existing.routingMode) as AiProviderRoutingMode)
: "url";
const nextUpstreamUrl = const nextUpstreamUrl =
body.upstreamUrl !== undefined body.upstreamUrl !== undefined
? body.upstreamUrl ? body.upstreamUrl
@@ -126,38 +127,33 @@ export async function updateAiProvider(
const nextAuthType = const nextAuthType =
body.authType !== undefined ? body.authType : existing.authType; body.authType !== undefined ? body.authType : existing.authType;
if ( const validation = z
providerRequiresUpstreamUrl(providerType) || .object({
body.upstreamUrl !== undefined || type: aiProviderTypeSchema,
body.authType !== undefined upstreamUrl: z.string().nullable().optional(),
) { authType: aiAuthTypeSchema.nullable().optional(),
const validation = z routingMode: aiRoutingModeSchema.optional()
.object({ })
type: aiProviderTypeSchema, .superRefine((data, ctx) => refineProviderUpstreamFields(data, ctx))
upstreamUrl: z.string().nullable().optional(), .safeParse({
authType: aiAuthTypeSchema.nullable().optional() type: providerType,
}) upstreamUrl: nextUpstreamUrl,
.superRefine((data, ctx) => authType: nextAuthType,
refineProviderUpstreamFields(data, ctx) routingMode: nextRoutingMode
) });
.safeParse({
type: providerType,
upstreamUrl: nextUpstreamUrl,
authType: nextAuthType
});
if (!validation.success) { if (!validation.success) {
return next( return next(
createHttpError( createHttpError(
HttpCode.BAD_REQUEST, HttpCode.BAD_REQUEST,
fromError(validation.error).toString() fromError(validation.error).toString()
) )
); );
}
} }
const updateData: Partial<typeof aiProviders.$inferInsert> = { const updateData: Partial<typeof aiProviders.$inferInsert> = {
updatedAt: Date.now() updatedAt: Date.now(),
routingMode: nextRoutingMode
}; };
if (body.name !== undefined) { if (body.name !== undefined) {
@@ -175,7 +171,9 @@ export async function updateAiProvider(
if (body.budgetUnit !== undefined) { if (body.budgetUnit !== undefined) {
updateData.budgetUnit = body.budgetUnit; updateData.budgetUnit = body.budgetUnit;
} }
if (body.upstreamUrl !== undefined) { if (nextRoutingMode === "target") {
updateData.upstreamUrl = null;
} else if (body.upstreamUrl !== undefined) {
updateData.upstreamUrl = body.upstreamUrl; updateData.upstreamUrl = body.upstreamUrl;
} }
if (body.authType !== undefined) { if (body.authType !== undefined) {
@@ -191,12 +189,7 @@ export async function updateAiProvider(
const [provider] = await db const [provider] = await db
.update(aiProviders) .update(aiProviders)
.set(updateData) .set(updateData)
.where( .where(eq(aiProviders.providerId, providerId))
and(
eq(aiProviders.providerId, providerId),
eq(aiProviders.orgId, orgId)
)
)
.returning(); .returning();
return response<CreateOrEditAiProviderResponse>(res, { return response<CreateOrEditAiProviderResponse>(res, {
+19 -2
View File
@@ -2,6 +2,7 @@ import { z } from "zod";
import { import {
providerRequiresUpstreamUrl, providerRequiresUpstreamUrl,
type AiBudgetUnit, type AiBudgetUnit,
type AiProviderRoutingMode,
type AiProviderType type AiProviderType
} from "@server/lib/aiProviderDefaults"; } from "@server/lib/aiProviderDefaults";
@@ -21,6 +22,8 @@ export const aiBudgetUnitSchema = z.enum(["usd", "tokens"]);
export const aiAuthTypeSchema = z.enum(["bearer"]); export const aiAuthTypeSchema = z.enum(["bearer"]);
export const aiRoutingModeSchema = z.enum(["url", "target"]);
export function refineBudgetFields( export function refineBudgetFields(
data: { data: {
budgetAmount?: number | null; budgetAmount?: number | null;
@@ -47,10 +50,24 @@ export function refineProviderUpstreamFields(
type: AiProviderType; type: AiProviderType;
upstreamUrl?: string | null; upstreamUrl?: string | null;
authType?: "bearer" | null; authType?: "bearer" | null;
routingMode?: AiProviderRoutingMode | null;
}, },
ctx: z.RefinementCtx ctx: z.RefinementCtx
) { ) {
if (providerRequiresUpstreamUrl(data.type) && !data.upstreamUrl) { const routingMode = data.routingMode ?? "url";
if (data.type !== "custom" && routingMode === "target") {
ctx.addIssue({
code: "custom",
message: "routingMode target is only allowed for custom providers",
path: ["routingMode"]
});
}
if (
providerRequiresUpstreamUrl(data.type, routingMode) &&
!data.upstreamUrl
) {
ctx.addIssue({ ctx.addIssue({
code: "custom", code: "custom",
message: `upstreamUrl is required for ${data.type} providers`, message: `upstreamUrl is required for ${data.type} providers`,
@@ -58,7 +75,7 @@ export function refineProviderUpstreamFields(
}); });
} }
if (data.type === "custom" && !data.authType) { if (data.type === "custom" && routingMode === "url" && !data.authType) {
ctx.addIssue({ ctx.addIssue({
code: "custom", code: "custom",
message: "authType is required for custom providers", message: "authType is required for custom providers",
+25 -16
View File
@@ -1385,16 +1385,31 @@ authenticated.get(
); );
authenticated.get( authenticated.get(
"/org/:orgId/ai-provider/:providerId", "/ai-provider/:providerId",
verifyOrgAccess,
verifyAiProviderAccess, verifyAiProviderAccess,
verifyUserHasAction(ActionsEnum.getAiProvider), verifyUserHasAction(ActionsEnum.getAiProvider),
aiProvider.getAiProvider aiProvider.getAiProvider
); );
authenticated.put(
"/ai-provider/:providerId/target",
verifyAiProviderAccess,
verifySiteAccess,
verifyLimits,
verifyUserHasAction(ActionsEnum.createTarget),
logActionAudit(ActionsEnum.createTarget),
target.createTarget
);
authenticated.get(
"/ai-provider/:providerId/targets",
verifyAiProviderAccess,
verifyUserHasAction(ActionsEnum.listTargets),
target.listTargets
);
authenticated.post( authenticated.post(
"/org/:orgId/ai-provider/:providerId", "/ai-provider/:providerId",
verifyOrgAccess,
verifyAiProviderAccess, verifyAiProviderAccess,
verifyUserHasAction(ActionsEnum.updateAiProvider), verifyUserHasAction(ActionsEnum.updateAiProvider),
logActionAudit(ActionsEnum.updateAiProvider), logActionAudit(ActionsEnum.updateAiProvider),
@@ -1402,8 +1417,7 @@ authenticated.post(
); );
authenticated.delete( authenticated.delete(
"/org/:orgId/ai-provider/:providerId", "/ai-provider/:providerId",
verifyOrgAccess,
verifyAiProviderAccess, verifyAiProviderAccess,
verifyUserHasAction(ActionsEnum.deleteAiProvider), verifyUserHasAction(ActionsEnum.deleteAiProvider),
logActionAudit(ActionsEnum.deleteAiProvider), logActionAudit(ActionsEnum.deleteAiProvider),
@@ -1411,8 +1425,7 @@ authenticated.delete(
); );
authenticated.put( authenticated.put(
"/org/:orgId/ai-provider/:providerId/model", "/ai-provider/:providerId/model",
verifyOrgAccess,
verifyAiProviderAccess, verifyAiProviderAccess,
verifyUserHasAction(ActionsEnum.createAiModel), verifyUserHasAction(ActionsEnum.createAiModel),
logActionAudit(ActionsEnum.createAiModel), logActionAudit(ActionsEnum.createAiModel),
@@ -1420,24 +1433,21 @@ authenticated.put(
); );
authenticated.get( authenticated.get(
"/org/:orgId/ai-provider/:providerId/models", "/ai-provider/:providerId/models",
verifyOrgAccess,
verifyAiProviderAccess, verifyAiProviderAccess,
verifyUserHasAction(ActionsEnum.listAiModels), verifyUserHasAction(ActionsEnum.listAiModels),
aiProvider.listAiModels aiProvider.listAiModels
); );
authenticated.get( authenticated.get(
"/org/:orgId/ai-provider/:providerId/model/:modelId", "/ai-model/:modelId",
verifyOrgAccess,
verifyAiModelAccess, verifyAiModelAccess,
verifyUserHasAction(ActionsEnum.getAiModel), verifyUserHasAction(ActionsEnum.getAiModel),
aiProvider.getAiModel aiProvider.getAiModel
); );
authenticated.post( authenticated.post(
"/org/:orgId/ai-provider/:providerId/model/:modelId", "/ai-model/:modelId",
verifyOrgAccess,
verifyAiModelAccess, verifyAiModelAccess,
verifyUserHasAction(ActionsEnum.updateAiModel), verifyUserHasAction(ActionsEnum.updateAiModel),
logActionAudit(ActionsEnum.updateAiModel), logActionAudit(ActionsEnum.updateAiModel),
@@ -1445,8 +1455,7 @@ authenticated.post(
); );
authenticated.delete( authenticated.delete(
"/org/:orgId/ai-provider/:providerId/model/:modelId", "/ai-model/:modelId",
verifyOrgAccess,
verifyAiModelAccess, verifyAiModelAccess,
verifyUserHasAction(ActionsEnum.deleteAiModel), verifyUserHasAction(ActionsEnum.deleteAiModel),
logActionAudit(ActionsEnum.deleteAiModel), logActionAudit(ActionsEnum.deleteAiModel),
+24 -16
View File
@@ -1386,16 +1386,30 @@ authenticated.get(
); );
authenticated.get( authenticated.get(
"/org/:orgId/ai-provider/:providerId", "/ai-provider/:providerId",
verifyApiKeyOrgAccess,
verifyApiKeyAiProviderAccess, verifyApiKeyAiProviderAccess,
verifyApiKeyHasAction(ActionsEnum.getAiProvider), verifyApiKeyHasAction(ActionsEnum.getAiProvider),
aiProvider.getAiProvider aiProvider.getAiProvider
); );
authenticated.put(
"/ai-provider/:providerId/target",
verifyApiKeyAiProviderAccess,
verifyLimits,
verifyApiKeyHasAction(ActionsEnum.createTarget),
logActionAudit(ActionsEnum.createTarget),
target.createTarget
);
authenticated.get(
"/ai-provider/:providerId/targets",
verifyApiKeyAiProviderAccess,
verifyApiKeyHasAction(ActionsEnum.listTargets),
target.listTargets
);
authenticated.post( authenticated.post(
"/org/:orgId/ai-provider/:providerId", "/ai-provider/:providerId",
verifyApiKeyOrgAccess,
verifyApiKeyAiProviderAccess, verifyApiKeyAiProviderAccess,
verifyApiKeyHasAction(ActionsEnum.updateAiProvider), verifyApiKeyHasAction(ActionsEnum.updateAiProvider),
logActionAudit(ActionsEnum.updateAiProvider), logActionAudit(ActionsEnum.updateAiProvider),
@@ -1403,8 +1417,7 @@ authenticated.post(
); );
authenticated.delete( authenticated.delete(
"/org/:orgId/ai-provider/:providerId", "/ai-provider/:providerId",
verifyApiKeyOrgAccess,
verifyApiKeyAiProviderAccess, verifyApiKeyAiProviderAccess,
verifyApiKeyHasAction(ActionsEnum.deleteAiProvider), verifyApiKeyHasAction(ActionsEnum.deleteAiProvider),
logActionAudit(ActionsEnum.deleteAiProvider), logActionAudit(ActionsEnum.deleteAiProvider),
@@ -1412,8 +1425,7 @@ authenticated.delete(
); );
authenticated.put( authenticated.put(
"/org/:orgId/ai-provider/:providerId/model", "/ai-provider/:providerId/model",
verifyApiKeyOrgAccess,
verifyApiKeyAiProviderAccess, verifyApiKeyAiProviderAccess,
verifyApiKeyHasAction(ActionsEnum.createAiModel), verifyApiKeyHasAction(ActionsEnum.createAiModel),
logActionAudit(ActionsEnum.createAiModel), logActionAudit(ActionsEnum.createAiModel),
@@ -1421,24 +1433,21 @@ authenticated.put(
); );
authenticated.get( authenticated.get(
"/org/:orgId/ai-provider/:providerId/models", "/ai-provider/:providerId/models",
verifyApiKeyOrgAccess,
verifyApiKeyAiProviderAccess, verifyApiKeyAiProviderAccess,
verifyApiKeyHasAction(ActionsEnum.listAiModels), verifyApiKeyHasAction(ActionsEnum.listAiModels),
aiProvider.listAiModels aiProvider.listAiModels
); );
authenticated.get( authenticated.get(
"/org/:orgId/ai-provider/:providerId/model/:modelId", "/ai-model/:modelId",
verifyApiKeyOrgAccess,
verifyApiKeyAiModelAccess, verifyApiKeyAiModelAccess,
verifyApiKeyHasAction(ActionsEnum.getAiModel), verifyApiKeyHasAction(ActionsEnum.getAiModel),
aiProvider.getAiModel aiProvider.getAiModel
); );
authenticated.post( authenticated.post(
"/org/:orgId/ai-provider/:providerId/model/:modelId", "/ai-model/:modelId",
verifyApiKeyOrgAccess,
verifyApiKeyAiModelAccess, verifyApiKeyAiModelAccess,
verifyApiKeyHasAction(ActionsEnum.updateAiModel), verifyApiKeyHasAction(ActionsEnum.updateAiModel),
logActionAudit(ActionsEnum.updateAiModel), logActionAudit(ActionsEnum.updateAiModel),
@@ -1446,8 +1455,7 @@ authenticated.post(
); );
authenticated.delete( authenticated.delete(
"/org/:orgId/ai-provider/:providerId/model/:modelId", "/ai-model/:modelId",
verifyApiKeyOrgAccess,
verifyApiKeyAiModelAccess, verifyApiKeyAiModelAccess,
verifyApiKeyHasAction(ActionsEnum.deleteAiModel), verifyApiKeyHasAction(ActionsEnum.deleteAiModel),
logActionAudit(ActionsEnum.deleteAiModel), logActionAudit(ActionsEnum.deleteAiModel),
+10 -6
View File
@@ -15,7 +15,7 @@ import {
} from "@server/db"; } from "@server/db";
import logger from "@server/logger"; import logger from "@server/logger";
import { initPeerAddHandshake, updatePeer } from "../olm/peers"; import { initPeerAddHandshake, updatePeer } from "../olm/peers";
import { eq, and, inArray } from "drizzle-orm"; import { eq, and, inArray, or, isNotNull, sql } from "drizzle-orm";
import config from "@server/lib/config"; import config from "@server/lib/config";
import { decrypt } from "@server/lib/crypto"; import { decrypt } from "@server/lib/crypto";
import { import {
@@ -211,7 +211,8 @@ export async function buildClientConfigurationForNewtClient(
// call rather than letting each resource fetch its own — with thousands // call rather than letting each resource fetch its own — with thousands
// of resources this avoids a concurrent DB/cache stampede for what is // of resources this avoids a concurrent DB/cache stampede for what is
// often the very same (e.g. wildcard) certificate. // often the very same (e.g. wildcard) certificate.
const certByDomain = await batchFetchCertsForSiteResources(allSiteResources); const certByDomain =
await batchFetchCertsForSiteResources(allSiteResources);
const resourceTargetsArr = await Promise.all( const resourceTargetsArr = await Promise.all(
allSiteResources.map((resource) => allSiteResources.map((resource) =>
@@ -240,7 +241,7 @@ export async function buildTargetConfigurationForNewtClient(
version?: string | null, version?: string | null,
remoteExitNodeId?: string remoteExitNodeId?: string
) { ) {
// Get all enabled targets with their resource mode information // Get enabled HTTP/TCP/UDP targets for resources and AI providers
const allTargets = await db const allTargets = await db
.select({ .select({
resourceId: targets.resourceId, resourceId: targets.resourceId,
@@ -250,15 +251,18 @@ export async function buildTargetConfigurationForNewtClient(
port: targets.port, port: targets.port,
internalPort: targets.internalPort, internalPort: targets.internalPort,
enabled: targets.enabled, enabled: targets.enabled,
mode: resources.mode mode: sql<string>`COALESCE(${resources.mode}, ${targets.mode})`.mapWith(
String
)
}) })
.from(targets) .from(targets)
.innerJoin(resources, eq(targets.resourceId, resources.resourceId)) .leftJoin(resources, eq(targets.resourceId, resources.resourceId))
.where( .where(
and( and(
eq(targets.siteId, siteId), eq(targets.siteId, siteId),
eq(targets.enabled, true), eq(targets.enabled, true),
inArray(targets.mode, ["http", "udp", "tcp"]) inArray(targets.mode, ["http", "udp", "tcp"]),
or(isNotNull(targets.resourceId), isNotNull(targets.providerId))
) )
); );
+163 -31
View File
@@ -6,7 +6,14 @@ import {
TargetHealthCheck, TargetHealthCheck,
targetHealthCheck targetHealthCheck
} from "@server/db"; } from "@server/db";
import { newts, resources, sites, Target, targets } from "@server/db"; import {
aiProviders,
newts,
resources,
sites,
Target,
targets
} from "@server/db";
import response from "@server/lib/response"; import response from "@server/lib/response";
import HttpCode from "@server/types/HttpCode"; import HttpCode from "@server/types/HttpCode";
import createHttpError from "http-errors"; import createHttpError from "http-errors";
@@ -29,10 +36,19 @@ import { generateId } from "@server/auth/sessions/app";
import config from "@server/lib/config"; import config from "@server/lib/config";
import { sendBrowserGatewayTargets } from "@server/routers/newt/targets"; import { sendBrowserGatewayTargets } from "@server/routers/newt/targets";
const createTargetParamsSchema = z.strictObject({ const resourceTargetParamsSchema = z.strictObject({
resourceId: z.coerce.number().int().positive() resourceId: z.coerce.number().int().positive()
}); });
const providerTargetParamsSchema = z.strictObject({
providerId: z.coerce.number().int().positive()
});
const createTargetParamsSchema = z.union([
resourceTargetParamsSchema,
providerTargetParamsSchema
]);
const createTargetSchema = z const createTargetSchema = z
.strictObject({ .strictObject({
siteId: z.int().positive(), siteId: z.int().positive(),
@@ -95,7 +111,7 @@ registry.registerPath({
description: "Create a target for a resource.", description: "Create a target for a resource.",
tags: [OpenAPITags.PublicResourceLegacy], tags: [OpenAPITags.PublicResourceLegacy],
request: { request: {
params: createTargetParamsSchema, params: resourceTargetParamsSchema,
body: { body: {
content: { content: {
"application/json": { "application/json": {
@@ -128,7 +144,40 @@ registry.registerPath({
description: "Create a target for a resource.", description: "Create a target for a resource.",
tags: [OpenAPITags.PublicResource, OpenAPITags.Target], tags: [OpenAPITags.PublicResource, OpenAPITags.Target],
request: { request: {
params: createTargetParamsSchema, params: resourceTargetParamsSchema,
body: {
content: {
"application/json": {
schema: createTargetSchema
}
}
}
},
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()
})
}
}
}
}
});
registry.registerPath({
method: "put",
path: "/ai-provider/{providerId}/target",
description: "Create a target for an AI provider.",
tags: [OpenAPITags.AiProvider],
request: {
params: providerTargetParamsSchema,
body: { body: {
content: { content: {
"application/json": { "application/json": {
@@ -183,21 +232,74 @@ export async function createTarget(
); );
} }
const { resourceId } = parsedParams.data; let resource: typeof resources.$inferSelect | undefined;
let provider: typeof aiProviders.$inferSelect | undefined;
// get the resource if ("providerId" in parsedParams.data) {
const [resource] = await db const { providerId } = parsedParams.data;
.select() [provider] =
.from(resources) req.aiProvider && req.aiProvider.providerId === providerId
.where(eq(resources.resourceId, resourceId)); ? [req.aiProvider]
: await db
.select()
.from(aiProviders)
.where(eq(aiProviders.providerId, providerId))
.limit(1);
if (!resource) { if (!provider) {
return next( return next(
createHttpError( createHttpError(
HttpCode.NOT_FOUND, HttpCode.NOT_FOUND,
`Resource with ID ${resourceId} not found` `AI provider with ID ${providerId} not found`
) )
); );
}
if (provider.routingMode !== "target") {
return next(
createHttpError(
HttpCode.BAD_REQUEST,
"AI provider must use target routing mode"
)
);
}
if (provider.type !== "custom") {
return next(
createHttpError(
HttpCode.BAD_REQUEST,
"Only custom AI providers support targets"
)
);
}
if (
targetData.method &&
!["http", "https"].includes(targetData.method.toLowerCase())
) {
return next(
createHttpError(
HttpCode.BAD_REQUEST,
"AI provider target method must be http or https"
)
);
}
} else {
const { resourceId } = parsedParams.data;
[resource] = await db
.select()
.from(resources)
.where(eq(resources.resourceId, resourceId))
.limit(1);
if (!resource) {
return next(
createHttpError(
HttpCode.NOT_FOUND,
`Resource with ID ${resourceId} not found`
)
);
}
} }
const siteId = targetData.siteId; const siteId = targetData.siteId;
@@ -217,6 +319,24 @@ export async function createTarget(
); );
} }
if (provider && site.orgId && site.orgId !== provider.orgId) {
return next(
createHttpError(
HttpCode.BAD_REQUEST,
"Site must belong to the AI provider organization"
)
);
}
const resourceId = resource?.resourceId ?? null;
const providerId = provider?.providerId ?? null;
const targetMode = provider
? "http"
: (targetData.mode ?? resource?.mode ?? "http");
const targetMethod = provider
? (targetData.method?.toLowerCase() ?? "https")
: targetData.method;
const plainToken = generateId(48); const plainToken = generateId(48);
const encryptedToken = encrypt( const encryptedToken = encrypt(
plainToken, plainToken,
@@ -230,20 +350,24 @@ export async function createTarget(
const existingTargets = await trx const existingTargets = await trx
.select() .select()
.from(targets) .from(targets)
.where(eq(targets.resourceId, resourceId)); .where(
providerId
? eq(targets.providerId, providerId)
: eq(targets.resourceId, resourceId!)
);
const existingTarget = existingTargets.find( const existingTarget = existingTargets.find(
(target) => (target) =>
target.ip === targetData.ip && target.ip === targetData.ip &&
target.port === targetData.port && target.port === targetData.port &&
target.method === targetData.method && target.method === targetMethod &&
target.siteId === targetData.siteId target.siteId === targetData.siteId
); );
if (existingTarget) { if (existingTarget) {
// log a warning // log a warning
logger.warn( logger.warn(
`Target with IP ${targetData.ip}, port ${targetData.port}, method ${targetData.method} already exists for resource ID ${resourceId}` `Target with IP ${targetData.ip}, port ${targetData.port}, method ${targetMethod} already exists for ${providerId ? `AI provider ID ${providerId}` : `resource ID ${resourceId}`}`
); );
} }
@@ -252,10 +376,10 @@ export async function createTarget(
.insert(targets) .insert(targets)
.values({ .values({
resourceId, resourceId,
providerId,
...targetData, ...targetData,
mode: (targetData.mode ?? mode: targetMode as Target["mode"],
resource.mode ?? method: targetMethod,
"http") as Target["mode"],
priority: targetData.priority || 100 priority: targetData.priority || 100
}) })
.returning(); .returning();
@@ -289,13 +413,12 @@ export async function createTarget(
.insert(targets) .insert(targets)
.values({ .values({
resourceId, resourceId,
providerId,
siteId: site.siteId, siteId: site.siteId,
ip: targetData.ip, ip: targetData.ip,
mode: (targetData.mode ?? mode: targetMode as Target["mode"],
resource.mode ??
"http") as Target["mode"],
authToken: encryptedToken, authToken: encryptedToken,
method: targetData.method, method: targetMethod,
port: targetData.port, port: targetData.port,
internalPort, internalPort,
enabled: targetData.enabled, enabled: targetData.enabled,
@@ -321,10 +444,12 @@ export async function createTarget(
healthCheck = await trx healthCheck = await trx
.insert(targetHealthCheck) .insert(targetHealthCheck)
.values({ .values({
orgId: resource.orgId, orgId: provider?.orgId ?? resource!.orgId,
targetId: newTarget[0].targetId, targetId: newTarget[0].targetId,
siteId: targetData.siteId, siteId: targetData.siteId,
name: `Resource ${resource.name} - ${targetData.ip}:${targetData.port}`, name: provider
? `AI Provider ${provider.name} - ${targetData.ip}:${targetData.port}`
: `Resource ${resource!.name} - ${targetData.ip}:${targetData.port}`,
hcEnabled: targetData.hcEnabled ?? false, hcEnabled: targetData.hcEnabled ?? false,
hcPath: targetData.hcPath ?? null, hcPath: targetData.hcPath ?? null,
hcScheme: targetData.hcScheme ?? null, hcScheme: targetData.hcScheme ?? null,
@@ -399,10 +524,17 @@ export async function createTarget(
newt.newtId, newt.newtId,
newTarget, newTarget,
healthCheck, healthCheck,
resource.mode === "udp" ? "udp" : "tcp", provider
? "tcp"
: (resource!.mode as string) === "udp"
? "udp"
: "tcp",
newt.version newt.version
); );
} else if (["ssh", "rdp", "vnc"].includes(newTarget[0].mode)) { } else if (
!provider &&
["ssh", "rdp", "vnc"].includes(newTarget[0].mode)
) {
await sendBrowserGatewayTargets( await sendBrowserGatewayTargets(
newt.newtId, newt.newtId,
newTarget, newTarget,
+53 -27
View File
@@ -79,38 +79,54 @@ export async function deleteTarget(
) )
); );
} }
// get the resource
const [resource] = await db
.select()
.from(resources)
.where(eq(resources.resourceId, deletedTarget.resourceId!));
if (!resource) { if (
(!deletedTarget.resourceId && !deletedTarget.providerId) ||
(deletedTarget.resourceId && deletedTarget.providerId)
) {
return next( return next(
createHttpError( createHttpError(
HttpCode.NOT_FOUND, HttpCode.INTERNAL_SERVER_ERROR,
`Resource with ID ${deletedTarget.resourceId} not found` `Target with ID ${targetId} has invalid ownership`
) )
); );
} }
// check if there are other targets on the resource let resource: typeof resources.$inferSelect | undefined;
const otherTargets = await db if (deletedTarget.resourceId) {
.select() [resource] = await db
.from(targets) .select()
.where( .from(resources)
and( .where(eq(resources.resourceId, deletedTarget.resourceId))
eq(targets.resourceId, resource.resourceId), .limit(1);
ne(targets.targetId, targetId)
)
);
if (otherTargets.length == 0) { if (!resource) {
// set the resource status return next(
await db createHttpError(
.update(resources) HttpCode.NOT_FOUND,
.set({ health: "unknown" }) `Resource with ID ${deletedTarget.resourceId} not found`
.where(eq(resources.resourceId, resource.resourceId)); )
);
}
// check if there are other targets on the resource
const otherTargets = await db
.select()
.from(targets)
.where(
and(
eq(targets.resourceId, resource.resourceId),
ne(targets.targetId, targetId)
)
);
if (otherTargets.length == 0) {
// set the resource status
await db
.update(resources)
.set({ health: "unknown" })
.where(eq(resources.resourceId, resource.resourceId));
}
} }
const [site] = await db const [site] = await db
@@ -137,16 +153,26 @@ export async function deleteTarget(
.where(eq(newts.siteId, site.siteId)) .where(eq(newts.siteId, site.siteId))
.limit(1); .limit(1);
if (["http", "tcp", "udp"].includes(deletedTarget.mode)) { if (
deletedTarget.providerId ||
["http", "tcp", "udp"].includes(deletedTarget.mode)
) {
await removeTargets( await removeTargets(
newt.newtId, newt.newtId,
// [deletedTarget], // [deletedTarget],
[], // deleting the target from newt causes issues because we cant unbind the port. this needs to be fixed in newt before we can do this [], // deleting the target from newt causes issues because we cant unbind the port. this needs to be fixed in newt before we can do this
[deletedHealthCheck], [deletedHealthCheck],
resource.mode === "udp" ? "udp" : "tcp", deletedTarget.providerId
? "tcp"
: (resource!.mode as string) === "udp"
? "udp"
: "tcp",
newt.version newt.version
); );
} else if (["ssh", "rdp", "vnc"].includes(deletedTarget.mode)) { } else if (
!deletedTarget.providerId &&
["ssh", "rdp", "vnc"].includes(deletedTarget.mode)
) {
await removeBrowserGatewayTarget( await removeBrowserGatewayTarget(
newt.newtId, newt.newtId,
deletedTarget.targetId, deletedTarget.targetId,
+53 -8
View File
@@ -10,10 +10,19 @@ import { fromError } from "zod-validation-error";
import logger from "@server/logger"; import logger from "@server/logger";
import { OpenAPITags, registry } from "@server/openApi"; import { OpenAPITags, registry } from "@server/openApi";
const listTargetsParamsSchema = z.strictObject({ const resourceTargetsParamsSchema = z.strictObject({
resourceId: z.coerce.number().int().positive() resourceId: z.coerce.number().int().positive()
}); });
const providerTargetsParamsSchema = z.strictObject({
providerId: z.coerce.number().int().positive()
});
const listTargetsParamsSchema = z.union([
resourceTargetsParamsSchema,
providerTargetsParamsSchema
]);
const listTargetsSchema = z.strictObject({ const listTargetsSchema = z.strictObject({
limit: z limit: z
.string() .string()
@@ -29,7 +38,7 @@ const listTargetsSchema = z.strictObject({
.pipe(z.int().nonnegative()) .pipe(z.int().nonnegative())
}); });
function queryTargets(resourceId: number) { function queryTargets(owner: { resourceId: number } | { providerId: number }) {
const baseQuery = db const baseQuery = db
.select({ .select({
targetId: targets.targetId, targetId: targets.targetId,
@@ -39,6 +48,7 @@ function queryTargets(resourceId: number) {
port: targets.port, port: targets.port,
enabled: targets.enabled, enabled: targets.enabled,
resourceId: targets.resourceId, resourceId: targets.resourceId,
providerId: targets.providerId,
siteId: targets.siteId, siteId: targets.siteId,
siteType: sites.type, siteType: sites.type,
siteName: sites.name, siteName: sites.name,
@@ -71,7 +81,11 @@ function queryTargets(resourceId: number) {
targetHealthCheck, targetHealthCheck,
eq(targetHealthCheck.targetId, targets.targetId) eq(targetHealthCheck.targetId, targets.targetId)
) )
.where(eq(targets.resourceId, resourceId)); .where(
"providerId" in owner
? eq(targets.providerId, owner.providerId)
: eq(targets.resourceId, owner.resourceId)
);
return baseQuery; return baseQuery;
} }
@@ -94,7 +108,7 @@ registry.registerPath({
description: "List targets for a resource.", description: "List targets for a resource.",
tags: [OpenAPITags.PublicResourceLegacy], tags: [OpenAPITags.PublicResourceLegacy],
request: { request: {
params: listTargetsParamsSchema, params: resourceTargetsParamsSchema,
query: listTargetsSchema query: listTargetsSchema
}, },
responses: { responses: {
@@ -121,7 +135,34 @@ registry.registerPath({
description: "List targets for a resource.", description: "List targets for a resource.",
tags: [OpenAPITags.PublicResource, OpenAPITags.Target], tags: [OpenAPITags.PublicResource, OpenAPITags.Target],
request: { request: {
params: listTargetsParamsSchema, params: resourceTargetsParamsSchema,
query: listTargetsSchema
},
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()
})
}
}
}
}
});
registry.registerPath({
method: "get",
path: "/ai-provider/{providerId}/targets",
description: "List targets for an AI provider.",
tags: [OpenAPITags.AiProvider],
request: {
params: providerTargetsParamsSchema,
query: listTargetsSchema query: listTargetsSchema
}, },
responses: { responses: {
@@ -168,14 +209,18 @@ export async function listTargets(
) )
); );
} }
const { resourceId } = parsedParams.data; const owner = parsedParams.data;
const ownerCondition =
"providerId" in owner
? eq(targets.providerId, owner.providerId)
: eq(targets.resourceId, owner.resourceId);
const baseQuery = queryTargets(resourceId); const baseQuery = queryTargets(owner);
const countQuery = db const countQuery = db
.select({ count: sql<number>`cast(count(*) as integer)` }) .select({ count: sql<number>`cast(count(*) as integer)` })
.from(targets) .from(targets)
.where(eq(targets.resourceId, resourceId)); .where(ownerCondition);
const targetsList = await baseQuery.limit(limit).offset(offset); const targetsList = await baseQuery.limit(limit).offset(offset);
const totalCountResult = await countQuery; const totalCountResult = await countQuery;
+90 -16
View File
@@ -1,7 +1,7 @@
import { Request, Response, NextFunction } from "express"; import { Request, Response, NextFunction } from "express";
import { z } from "zod"; import { z } from "zod";
import { db, targetHealthCheck } from "@server/db"; import { db, targetHealthCheck } from "@server/db";
import { newts, resources, sites, targets } from "@server/db"; import { aiProviders, newts, resources, sites, targets } from "@server/db";
import { eq } from "drizzle-orm"; import { eq } from "drizzle-orm";
import response from "@server/lib/response"; import response from "@server/lib/response";
import HttpCode from "@server/types/HttpCode"; import HttpCode from "@server/types/HttpCode";
@@ -147,21 +147,68 @@ export async function updateTarget(
); );
} }
// get the resource if (
const [resource] = await db (!target.resourceId && !target.providerId) ||
.select() (target.resourceId && target.providerId)
.from(resources) ) {
.where(eq(resources.resourceId, target.resourceId!));
if (!resource) {
return next( return next(
createHttpError( createHttpError(
HttpCode.NOT_FOUND, HttpCode.INTERNAL_SERVER_ERROR,
`Resource with ID ${target.resourceId} not found` `Target with ID ${targetId} has invalid ownership`
) )
); );
} }
let resource: typeof resources.$inferSelect | undefined;
let provider: typeof aiProviders.$inferSelect | undefined;
if (target.resourceId) {
[resource] = await db
.select()
.from(resources)
.where(eq(resources.resourceId, target.resourceId))
.limit(1);
if (!resource) {
return next(
createHttpError(
HttpCode.NOT_FOUND,
`Resource with ID ${target.resourceId} not found`
)
);
}
} else {
[provider] = await db
.select()
.from(aiProviders)
.where(eq(aiProviders.providerId, target.providerId!))
.limit(1);
if (!provider) {
return next(
createHttpError(
HttpCode.NOT_FOUND,
`AI provider with ID ${target.providerId} not found`
)
);
}
if (
parsedBody.data.method !== undefined &&
(!parsedBody.data.method ||
!["http", "https"].includes(
parsedBody.data.method.toLowerCase()
))
) {
return next(
createHttpError(
HttpCode.BAD_REQUEST,
"AI provider target method must be http or https"
)
);
}
}
const [site] = await db const [site] = await db
.select() .select()
.from(sites) .from(sites)
@@ -177,6 +224,15 @@ export async function updateTarget(
); );
} }
if (provider && site.orgId && site.orgId !== provider.orgId) {
return next(
createHttpError(
HttpCode.BAD_REQUEST,
"Site must belong to the AI provider organization"
)
);
}
const { internalPort, targetIps } = await pickPort(site.siteId!, db); const { internalPort, targetIps } = await pickPort(site.siteId!, db);
if (!internalPort) { if (!internalPort) {
@@ -221,8 +277,13 @@ export async function updateTarget(
} }
const pathMatchTypeRemoved = parsedBody.data.pathMatchType === null; const pathMatchTypeRemoved = parsedBody.data.pathMatchType === null;
const nextMode = const nextMode = provider
parsedBody.data.mode === null ? undefined : parsedBody.data.mode; ? parsedBody.data.mode !== undefined
? "http"
: undefined
: parsedBody.data.mode === null
? undefined
: parsedBody.data.mode;
let updatedTarget: any; let updatedTarget: any;
let updatedHc: any; let updatedHc: any;
@@ -233,7 +294,10 @@ export async function updateTarget(
siteId: parsedBody.data.siteId, siteId: parsedBody.data.siteId,
ip: parsedBody.data.ip, ip: parsedBody.data.ip,
mode: nextMode, mode: nextMode,
method: parsedBody.data.method, method:
provider && parsedBody.data.method
? parsedBody.data.method.toLowerCase()
: parsedBody.data.method,
port: parsedBody.data.port, port: parsedBody.data.port,
internalPort, internalPort,
enabled: parsedBody.data.enabled, enabled: parsedBody.data.enabled,
@@ -368,15 +432,25 @@ export async function updateTarget(
.where(eq(newts.siteId, site.siteId)) .where(eq(newts.siteId, site.siteId))
.limit(1); .limit(1);
if (["http", "tcp", "udp"].includes(updatedTarget.mode)) { if (
provider ||
["http", "tcp", "udp"].includes(updatedTarget.mode)
) {
await addTargets( await addTargets(
newt.newtId, newt.newtId,
[updatedTarget], [updatedTarget],
[updatedHc], [updatedHc],
resource.mode === "udp" ? "udp" : "tcp", provider
? "tcp"
: (resource!.mode as string) === "udp"
? "udp"
: "tcp",
newt.version newt.version
); );
} else if (["ssh", "rdp", "vnc"].includes(updatedTarget.mode)) { } else if (
!provider &&
["ssh", "rdp", "vnc"].includes(updatedTarget.mode)
) {
await sendBrowserGatewayTargets( await sendBrowserGatewayTargets(
newt.newtId, newt.newtId,
[updatedTarget], [updatedTarget],
@@ -596,6 +596,7 @@ export function ProxyResourceTargetsForm({
priority: 100, priority: 100,
enabled: true, enabled: true,
resourceId: resource?.resourceId ?? 0, resourceId: resource?.resourceId ?? 0,
providerId: null,
hcEnabled: false, hcEnabled: false,
hcPath: null, hcPath: null,
hcMethod: null, hcMethod: null,