diff --git a/server/db/pg/schema/schema.ts b/server/db/pg/schema/schema.ts index 4e0719284..0bbf88631 100644 --- a/server/db/pg/schema/schema.ts +++ b/server/db/pg/schema/schema.ts @@ -195,7 +195,10 @@ export const resources = pgTable( postAuthPath: text("postAuthPath"), health: varchar("health").default("unknown"), // "healthy", "unhealthy", "unknown" wildcard: boolean("wildcard").notNull().default(false), - mode: text("mode").default("http").notNull(), // rdp, ssh, http, vnc + mode: text("mode") + .default("http") + .$type<"rdp" | "ssh" | "http" | "vnc" | "inference">() + .notNull(), pamMode: varchar("pamMode", { length: 32 }) .$type<"passthrough" | "push">() .default("passthrough"), @@ -429,7 +432,7 @@ export const siteResources = pgTable( name: varchar("name").notNull(), ssl: boolean("ssl").notNull().default(false), mode: varchar("mode") - .$type<"host" | "cidr" | "http" | "ssh">() + .$type<"host" | "cidr" | "http" | "ssh" | "inference">() .notNull(), // "host" | "cidr" | "http" scheme: varchar("scheme").$type<"http" | "https">(), // only for when we are doing https or http mode proxyPort: integer("proxyPort"), // only for port mode diff --git a/server/db/sqlite/schema/schema.ts b/server/db/sqlite/schema/schema.ts index 1e2d822e7..f13afa2cf 100644 --- a/server/db/sqlite/schema/schema.ts +++ b/server/db/sqlite/schema/schema.ts @@ -203,7 +203,10 @@ export const resources = sqliteTable("resources", { postAuthPath: text("postAuthPath"), health: text("health").default("unknown"), // "healthy", "unhealthy", "unknown" wildcard: integer("wildcard", { mode: "boolean" }).notNull().default(false), - mode: text("mode").default("http").notNull(), // rdp, ssh, http, vnc + mode: text("mode") + .default("http") + .$type<"rdp" | "ssh" | "http" | "vnc" | "inference">() + .notNull(), // rdp, ssh, http, vnc, inference pamMode: text("pamMode") .$type<"passthrough" | "push">() .default("passthrough"), @@ -425,7 +428,9 @@ export const siteResources = sqliteTable("siteResources", { niceId: text("niceId").notNull(), name: text("name").notNull(), ssl: integer("ssl", { mode: "boolean" }).notNull().default(false), - mode: text("mode").$type<"host" | "cidr" | "http" | "ssh">().notNull(), // "host" | "cidr" | "http" + mode: text("mode") + .$type<"host" | "cidr" | "http" | "ssh" | "inference">() + .notNull(), // "host" | "cidr" | "http" scheme: text("scheme").$type<"http" | "https">(), // only for when we are doing https or http mode proxyPort: integer("proxyPort"), // only for port mode destinationPort: integer("destinationPort"), // only for port mode diff --git a/server/routers/olm/handleOlmRegisterMessage.ts b/server/routers/olm/handleOlmRegisterMessage.ts index f97c3328c..8d148380c 100644 --- a/server/routers/olm/handleOlmRegisterMessage.ts +++ b/server/routers/olm/handleOlmRegisterMessage.ts @@ -1,4 +1,4 @@ -import { db, orgs, primaryDb } from "@server/db"; +import { db, ExitNode, orgs, primaryDb } from "@server/db"; import { MessageHandler } from "@server/routers/ws"; import { clients, @@ -311,14 +311,15 @@ export const handleOlmRegisterMessage: MessageHandler = async (context) => { } let clientSubnet = client.exitNodeSubnet; + let exitNode: ExitNode | null = null; if ( exitNodeId && (client.exitNodeId !== exitNodeId || !client.exitNodeSubnet) ) { - const { exitNode, hasAccess } = await verifyExitNodeOrgAccess( - exitNodeId, - client.orgId - ); + const { exitNode: exitNodeResult, hasAccess } = + await verifyExitNodeOrgAccess(exitNodeId, client.orgId); + + exitNode = exitNodeResult; if (!exitNode) { logger.warn("[handleOlmRegisterMessage] Exit node not found", { @@ -466,6 +467,18 @@ export const handleOlmRegisterMessage: MessageHandler = async (context) => { sites: siteConfigurations, tunnelIP: client.subnet, utilitySubnet: org.utilitySubnet, + exitNode: + exitNode && client.exitNodeSubnet + ? { + endpoint: `${exitNode.endpoint}:${exitNode.listenPort}`, + relayPort: + config.getRawConfig().gerbil + .clients_start_port, + publicKey: exitNode.publicKey, + serverIP: exitNode.address.split("/")[0], + tunnelIP: client.exitNodeSubnet.split("/")[0] + } + : undefined, chainId: chainId } },