mirror of
https://github.com/fosrl/pangolin.git
synced 2026-09-22 13:59:05 +02:00
Merge upstream/dev into feature-response-headers
Resolve conflicts against upstream's refactors: - server/db/sqlite/schema/schema.ts: adopt upstream's reindented sqliteTable(name, cols, indexes) form for sites/resources, re-applying the headers -> requestHeaders/responseHeaders split. Kept in sync with the Postgres schema. - server/lib/traefik/headersMiddleware.ts: extend upstream's extracted buildCustomHeadersMiddleware helper to take requestHeaders and responseHeaders and emit both customRequestHeaders and customResponseHeaders. - server/lib/traefik/getTraefikConfig.ts and server/private/lib/traefik/getTraefikConfig.ts: keep upstream's helper extraction and appendPathMatch refactor, dropping the superseded inline blocks. Also carry the feature forward onto code that moved upstream: - The resource settings UI moved from resources/proxy/[niceId]/proxy to resources/public/[niceId]/http, which dropped this branch's changes in the previous merge. Re-add the request/response header inputs there and rename the vestigial headers field on the tcp page. - messages/da-DK.json is new upstream and still had the old customHeaders key; rename it in line with the other locales. Per the contributing docs, versioned migrations are intentionally omitted so maintainers can write them at release time. Co-Authored-By: Claude Opus 5 <noreply@anthropic.com>
This commit is contained in:
+865
-2
@@ -1,3 +1,866 @@
|
||||
import fs from "fs";
|
||||
import path from "path";
|
||||
import crypto from "crypto";
|
||||
import {
|
||||
certificates,
|
||||
clients,
|
||||
clientSiteResourcesAssociationsCache,
|
||||
db,
|
||||
domains,
|
||||
newts,
|
||||
siteNetworks,
|
||||
SiteResource,
|
||||
siteResources
|
||||
} from "@server/db";
|
||||
import { and, eq } from "drizzle-orm";
|
||||
import { encrypt, decrypt } from "@server/lib/crypto";
|
||||
import logger from "@server/logger";
|
||||
import config from "@server/lib/config";
|
||||
import {
|
||||
generateSubnetProxyTargetV2,
|
||||
SubnetProxyTargetV2
|
||||
} from "@server/lib/ip";
|
||||
import { updateTargets } from "@server/routers/client/targets";
|
||||
import cache from "#dynamic/lib/cache";
|
||||
import { build } from "@server/build";
|
||||
|
||||
interface AcmeCert {
|
||||
domain: { main: string; sans?: string[] };
|
||||
certificate: string;
|
||||
key: string;
|
||||
Store: string;
|
||||
}
|
||||
|
||||
interface AcmeJson {
|
||||
[resolver: string]: {
|
||||
Certificates: AcmeCert[];
|
||||
};
|
||||
}
|
||||
|
||||
export async function pushCertUpdateToAffectedNewts(
|
||||
domain: string,
|
||||
domainId: string | null,
|
||||
oldCertPem: string | null,
|
||||
oldKeyPem: string | null
|
||||
): Promise<void> {
|
||||
// Find all SSL-enabled HTTP site resources that use this cert's domain
|
||||
let affectedResources: SiteResource[] = [];
|
||||
|
||||
if (domainId) {
|
||||
affectedResources = await db
|
||||
.select()
|
||||
.from(siteResources)
|
||||
.where(
|
||||
and(
|
||||
eq(siteResources.domainId, domainId),
|
||||
eq(siteResources.ssl, true)
|
||||
)
|
||||
);
|
||||
} else {
|
||||
// Fallback: match by exact fullDomain when no domainId is available
|
||||
affectedResources = await db
|
||||
.select()
|
||||
.from(siteResources)
|
||||
.where(
|
||||
and(
|
||||
eq(siteResources.fullDomain, domain),
|
||||
eq(siteResources.ssl, true)
|
||||
)
|
||||
);
|
||||
}
|
||||
|
||||
if (affectedResources.length === 0) {
|
||||
logger.debug(
|
||||
`acmeCertSync: no affected site resources for cert domain "${domain}"`
|
||||
);
|
||||
return;
|
||||
}
|
||||
|
||||
logger.debug(
|
||||
`acmeCertSync: pushing cert update to ${affectedResources.length} affected site resource(s) for domain "${domain}"`
|
||||
);
|
||||
|
||||
for (const resource of affectedResources) {
|
||||
try {
|
||||
// Get all sites for this resource via siteNetworks
|
||||
const resourceSiteRows = resource.networkId
|
||||
? await db
|
||||
.select({ siteId: siteNetworks.siteId })
|
||||
.from(siteNetworks)
|
||||
.where(eq(siteNetworks.networkId, resource.networkId))
|
||||
: [];
|
||||
|
||||
if (resourceSiteRows.length === 0) {
|
||||
logger.debug(
|
||||
`acmeCertSync: no sites for resource ${resource.siteResourceId}, skipping`
|
||||
);
|
||||
continue;
|
||||
}
|
||||
|
||||
// Get all clients with access to this resource
|
||||
const resourceClients = await db
|
||||
.select({
|
||||
clientId: clients.clientId,
|
||||
pubKey: clients.pubKey,
|
||||
subnet: clients.subnet
|
||||
})
|
||||
.from(clients)
|
||||
.innerJoin(
|
||||
clientSiteResourcesAssociationsCache,
|
||||
eq(
|
||||
clients.clientId,
|
||||
clientSiteResourcesAssociationsCache.clientId
|
||||
)
|
||||
)
|
||||
.where(
|
||||
eq(
|
||||
clientSiteResourcesAssociationsCache.siteResourceId,
|
||||
resource.siteResourceId
|
||||
)
|
||||
);
|
||||
|
||||
if (resourceClients.length === 0) {
|
||||
logger.debug(
|
||||
`acmeCertSync: no clients for resource ${resource.siteResourceId}, skipping`
|
||||
);
|
||||
continue;
|
||||
}
|
||||
|
||||
// Invalidate the cert cache so generateSubnetProxyTargetV2 fetches fresh data
|
||||
if (resource.fullDomain) {
|
||||
await cache.del(`cert:${resource.fullDomain}`);
|
||||
}
|
||||
|
||||
// Generate target once - same cert applies to all sites for this resource
|
||||
const newTargets = await generateSubnetProxyTargetV2(
|
||||
resource,
|
||||
resourceClients
|
||||
);
|
||||
|
||||
if (!newTargets) {
|
||||
logger.debug(
|
||||
`acmeCertSync: could not generate target for resource ${resource.siteResourceId}, skipping`
|
||||
);
|
||||
continue;
|
||||
}
|
||||
|
||||
// Construct the old targets - same routing shape but with the previous cert/key.
|
||||
// The newt only uses destPrefix/sourcePrefixes for removal, but we keep the
|
||||
// semantics correct so the update message accurately reflects what changed.
|
||||
const oldTargets: SubnetProxyTargetV2[] = newTargets.map((t) => ({
|
||||
...t,
|
||||
tlsCert: oldCertPem ?? undefined,
|
||||
tlsKey: oldKeyPem ?? undefined
|
||||
}));
|
||||
|
||||
// Push update to each site's newt
|
||||
for (const { siteId } of resourceSiteRows) {
|
||||
const [newt] = await db
|
||||
.select()
|
||||
.from(newts)
|
||||
.where(eq(newts.siteId, siteId))
|
||||
.limit(1);
|
||||
|
||||
if (!newt) {
|
||||
logger.debug(
|
||||
`acmeCertSync: no newt found for site ${siteId}, skipping resource ${resource.siteResourceId}`
|
||||
);
|
||||
continue;
|
||||
}
|
||||
|
||||
await updateTargets(
|
||||
newt.newtId,
|
||||
{ oldTargets: oldTargets, newTargets: newTargets },
|
||||
newt.version
|
||||
);
|
||||
|
||||
logger.debug(
|
||||
`acmeCertSync: pushed cert update to newt for site ${siteId}, resource ${resource.siteResourceId}`
|
||||
);
|
||||
}
|
||||
} catch (err) {
|
||||
logger.error(
|
||||
`acmeCertSync: error pushing cert update for resource ${resource?.siteResourceId}: ${err}`
|
||||
);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
async function findDomainId(certDomain: string): Promise<string | null> {
|
||||
// Strip wildcard prefix before lookup (*.example.com -> example.com)
|
||||
const lookupDomain = certDomain.startsWith("*.")
|
||||
? certDomain.slice(2)
|
||||
: certDomain;
|
||||
|
||||
// 1. Exact baseDomain match (any domain type)
|
||||
const exactMatch = await db
|
||||
.select({ domainId: domains.domainId })
|
||||
.from(domains)
|
||||
.where(eq(domains.baseDomain, lookupDomain))
|
||||
.limit(1);
|
||||
|
||||
if (exactMatch.length > 0) {
|
||||
return exactMatch[0].domainId;
|
||||
}
|
||||
|
||||
// 2. Walk up the domain hierarchy looking for a wildcard-type domain whose
|
||||
// baseDomain is a suffix of the cert domain. e.g. cert "sub.example.com"
|
||||
// matches a wildcard domain with baseDomain "example.com".
|
||||
const parts = lookupDomain.split(".");
|
||||
for (let i = 1; i < parts.length; i++) {
|
||||
const candidate = parts.slice(i).join(".");
|
||||
if (!candidate) continue;
|
||||
|
||||
const wildcardMatch = await db
|
||||
.select({ domainId: domains.domainId })
|
||||
.from(domains)
|
||||
.where(
|
||||
and(
|
||||
eq(domains.baseDomain, candidate),
|
||||
eq(domains.type, "wildcard")
|
||||
)
|
||||
)
|
||||
.limit(1);
|
||||
|
||||
if (wildcardMatch.length > 0) {
|
||||
return wildcardMatch[0].domainId;
|
||||
}
|
||||
}
|
||||
|
||||
return null;
|
||||
}
|
||||
|
||||
function extractFirstCert(pemBundle: string): string | null {
|
||||
const match = pemBundle.match(
|
||||
/-----BEGIN CERTIFICATE-----[\s\S]+?-----END CERTIFICATE-----/
|
||||
);
|
||||
return match ? match[0] : null;
|
||||
}
|
||||
|
||||
/**
|
||||
* Determine whether an ACME cert entry represents a wildcard cert by checking
|
||||
* both the primary domain (`main`) and the SANs. Some ACME clients (notably
|
||||
* Traefik) store the bare apex in `main` and only put the wildcard form in
|
||||
* `sans` (e.g. main="access.example.com", sans=["*.access.example.com"]).
|
||||
*/
|
||||
function detectWildcard(
|
||||
main: string,
|
||||
sans: string[] | undefined
|
||||
): { wildcard: boolean; wildcardSan: string | null } {
|
||||
if (main.startsWith("*.")) {
|
||||
return { wildcard: true, wildcardSan: null };
|
||||
}
|
||||
if (Array.isArray(sans)) {
|
||||
for (const san of sans) {
|
||||
if (typeof san !== "string") continue;
|
||||
if (san === `*.${main}` || san.startsWith("*.")) {
|
||||
return { wildcard: true, wildcardSan: san };
|
||||
}
|
||||
}
|
||||
}
|
||||
return { wildcard: false, wildcardSan: null };
|
||||
}
|
||||
|
||||
interface HttpCert {
|
||||
wildcard: boolean;
|
||||
altName: string;
|
||||
certName: string;
|
||||
commonName: string;
|
||||
certFile: string;
|
||||
keyFile: string;
|
||||
}
|
||||
|
||||
async function syncAcmeCertsFromHttp(endpoint: string): Promise<void> {
|
||||
let response: Response;
|
||||
try {
|
||||
response = await fetch(endpoint);
|
||||
} catch (err) {
|
||||
logger.debug(
|
||||
`acmeCertSync: could not reach HTTP endpoint ${endpoint}: ${err}`
|
||||
);
|
||||
return;
|
||||
}
|
||||
|
||||
if (!response.ok) {
|
||||
logger.debug(
|
||||
`acmeCertSync: HTTP endpoint returned status ${response.status}`
|
||||
);
|
||||
return;
|
||||
}
|
||||
|
||||
let httpCerts: HttpCert[];
|
||||
try {
|
||||
httpCerts = await response.json();
|
||||
} catch (err) {
|
||||
logger.debug(
|
||||
`acmeCertSync: could not parse JSON from HTTP endpoint: ${err}`
|
||||
);
|
||||
return;
|
||||
}
|
||||
|
||||
if (!Array.isArray(httpCerts) || httpCerts.length === 0) {
|
||||
logger.debug(
|
||||
`acmeCertSync: no certificates returned from HTTP endpoint`
|
||||
);
|
||||
return;
|
||||
}
|
||||
|
||||
for (const cert of httpCerts) {
|
||||
const domain = cert?.certName;
|
||||
|
||||
if (!domain || typeof domain !== "string") {
|
||||
logger.debug(
|
||||
`acmeCertSync: skipping HTTP cert with missing certName`
|
||||
);
|
||||
continue;
|
||||
}
|
||||
|
||||
const certPem = cert.certFile;
|
||||
const keyPem = cert.keyFile;
|
||||
|
||||
if (!certPem?.trim() || !keyPem?.trim()) {
|
||||
logger.debug(
|
||||
`acmeCertSync: skipping HTTP cert for ${domain} - empty certFile or keyFile`
|
||||
);
|
||||
continue;
|
||||
}
|
||||
|
||||
const firstCertPemForValidation = extractFirstCert(certPem);
|
||||
if (!firstCertPemForValidation) {
|
||||
logger.debug(
|
||||
`acmeCertSync: skipping HTTP cert for ${domain} - no PEM certificate block found`
|
||||
);
|
||||
continue;
|
||||
}
|
||||
|
||||
let validatedX509: crypto.X509Certificate;
|
||||
try {
|
||||
validatedX509 = new crypto.X509Certificate(
|
||||
firstCertPemForValidation
|
||||
);
|
||||
} catch (err) {
|
||||
logger.debug(
|
||||
`acmeCertSync: skipping HTTP cert for ${domain} - invalid X.509 certificate: ${err}`
|
||||
);
|
||||
continue;
|
||||
}
|
||||
|
||||
try {
|
||||
crypto.createPrivateKey(keyPem);
|
||||
} catch (err) {
|
||||
logger.debug(
|
||||
`acmeCertSync: skipping HTTP cert for ${domain} - invalid private key: ${err}`
|
||||
);
|
||||
continue;
|
||||
}
|
||||
|
||||
const wildcard = cert.wildcard ?? false;
|
||||
|
||||
const existing = await db
|
||||
.select()
|
||||
.from(certificates)
|
||||
.where(eq(certificates.domain, domain))
|
||||
.limit(1);
|
||||
|
||||
let oldCertPem: string | null = null;
|
||||
let oldKeyPem: string | null = null;
|
||||
|
||||
if (existing.length > 0 && existing[0].certFile) {
|
||||
try {
|
||||
const storedCertPem = decrypt(
|
||||
existing[0].certFile,
|
||||
config.getRawConfig().server.secret!
|
||||
);
|
||||
const wildcardUnchanged = existing[0].wildcard === wildcard;
|
||||
if (storedCertPem === certPem && wildcardUnchanged) {
|
||||
continue;
|
||||
}
|
||||
oldCertPem = storedCertPem;
|
||||
if (existing[0].keyFile) {
|
||||
try {
|
||||
oldKeyPem = decrypt(
|
||||
existing[0].keyFile,
|
||||
config.getRawConfig().server.secret!
|
||||
);
|
||||
} catch (keyErr) {
|
||||
logger.debug(
|
||||
`acmeCertSync: could not decrypt stored key for ${domain}: ${keyErr}`
|
||||
);
|
||||
}
|
||||
}
|
||||
} catch (err) {
|
||||
logger.debug(
|
||||
`acmeCertSync: could not decrypt stored cert for ${domain}, will update: ${err}`
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
let expiresAt: number | null = null;
|
||||
try {
|
||||
expiresAt = Math.floor(
|
||||
new Date(validatedX509.validTo).getTime() / 1000
|
||||
);
|
||||
} catch (err) {
|
||||
logger.debug(
|
||||
`acmeCertSync: could not parse cert expiry for ${domain}: ${err}`
|
||||
);
|
||||
}
|
||||
|
||||
const encryptedCert = encrypt(
|
||||
certPem,
|
||||
config.getRawConfig().server.secret!
|
||||
);
|
||||
const encryptedKey = encrypt(
|
||||
keyPem,
|
||||
config.getRawConfig().server.secret!
|
||||
);
|
||||
const now = Math.floor(Date.now() / 1000);
|
||||
|
||||
const domainId = await findDomainId(domain);
|
||||
if (domainId) {
|
||||
logger.debug(
|
||||
`acmeCertSync: resolved domainId "${domainId}" for HTTP cert domain "${domain}"`
|
||||
);
|
||||
} else {
|
||||
logger.debug(
|
||||
`acmeCertSync: no matching domain record found for HTTP cert domain "${domain}"`
|
||||
);
|
||||
}
|
||||
|
||||
if (existing.length > 0) {
|
||||
logger.debug(
|
||||
`acmeCertSync: updating existing certificate (HTTP) for ${domain} (expires ${expiresAt ? new Date(expiresAt * 1000).toISOString() : "unknown"})`
|
||||
);
|
||||
await db
|
||||
.update(certificates)
|
||||
.set({
|
||||
certFile: encryptedCert,
|
||||
keyFile: encryptedKey,
|
||||
status: "valid",
|
||||
expiresAt,
|
||||
updatedAt: now,
|
||||
wildcard,
|
||||
...(domainId !== null && { domainId })
|
||||
})
|
||||
.where(eq(certificates.domain, domain));
|
||||
|
||||
await pushCertUpdateToAffectedNewts(
|
||||
domain,
|
||||
domainId,
|
||||
oldCertPem,
|
||||
oldKeyPem
|
||||
);
|
||||
} else {
|
||||
logger.debug(
|
||||
`acmeCertSync: inserting new certificate (HTTP) for ${domain} (expires ${expiresAt ? new Date(expiresAt * 1000).toISOString() : "unknown"})`
|
||||
);
|
||||
await db.insert(certificates).values({
|
||||
domain,
|
||||
domainId,
|
||||
certFile: encryptedCert,
|
||||
keyFile: encryptedKey,
|
||||
status: "valid",
|
||||
expiresAt,
|
||||
createdAt: now,
|
||||
updatedAt: now,
|
||||
wildcard
|
||||
});
|
||||
|
||||
await pushCertUpdateToAffectedNewts(domain, domainId, null, null);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
async function storeCertForDomain(
|
||||
domain: string,
|
||||
certPem: string,
|
||||
keyPem: string,
|
||||
validatedX509: crypto.X509Certificate
|
||||
): Promise<void> {
|
||||
const wildcard = domain.startsWith("*.");
|
||||
|
||||
const existing = await db
|
||||
.select()
|
||||
.from(certificates)
|
||||
.where(eq(certificates.domain, domain))
|
||||
.limit(1);
|
||||
|
||||
let oldCertPem: string | null = null;
|
||||
let oldKeyPem: string | null = null;
|
||||
|
||||
if (existing.length > 0 && existing[0].certFile) {
|
||||
try {
|
||||
const storedCertPem = decrypt(
|
||||
existing[0].certFile,
|
||||
config.getRawConfig().server.secret!
|
||||
);
|
||||
const wildcardUnchanged = existing[0].wildcard === wildcard;
|
||||
if (storedCertPem === certPem && wildcardUnchanged) {
|
||||
return;
|
||||
}
|
||||
oldCertPem = storedCertPem;
|
||||
if (existing[0].keyFile) {
|
||||
try {
|
||||
oldKeyPem = decrypt(
|
||||
existing[0].keyFile,
|
||||
config.getRawConfig().server.secret!
|
||||
);
|
||||
} catch (keyErr) {
|
||||
logger.debug(
|
||||
`acmeCertSync: could not decrypt stored key for ${domain}: ${keyErr}`
|
||||
);
|
||||
}
|
||||
}
|
||||
} catch (err) {
|
||||
logger.debug(
|
||||
`acmeCertSync: could not decrypt stored cert for ${domain}, will update: ${err}`
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
let expiresAt: number | null = null;
|
||||
try {
|
||||
expiresAt = Math.floor(
|
||||
new Date(validatedX509.validTo).getTime() / 1000
|
||||
);
|
||||
} catch (err) {
|
||||
logger.debug(
|
||||
`acmeCertSync: could not parse cert expiry for ${domain}: ${err}`
|
||||
);
|
||||
}
|
||||
|
||||
const encryptedCert = encrypt(
|
||||
certPem,
|
||||
config.getRawConfig().server.secret!
|
||||
);
|
||||
const encryptedKey = encrypt(keyPem, config.getRawConfig().server.secret!);
|
||||
const now = Math.floor(Date.now() / 1000);
|
||||
|
||||
const domainId = await findDomainId(domain);
|
||||
if (domainId) {
|
||||
logger.debug(
|
||||
`acmeCertSync: resolved domainId "${domainId}" for cert domain "${domain}"`
|
||||
);
|
||||
} else {
|
||||
logger.debug(
|
||||
`acmeCertSync: no matching domain record found for cert domain "${domain}"`
|
||||
);
|
||||
}
|
||||
|
||||
if (existing.length > 0) {
|
||||
logger.debug(
|
||||
`acmeCertSync: updating existing certificate for ${domain} (expires ${expiresAt ? new Date(expiresAt * 1000).toISOString() : "unknown"})`
|
||||
);
|
||||
await db
|
||||
.update(certificates)
|
||||
.set({
|
||||
certFile: encryptedCert,
|
||||
keyFile: encryptedKey,
|
||||
status: "valid",
|
||||
expiresAt,
|
||||
updatedAt: now,
|
||||
wildcard,
|
||||
...(domainId !== null && { domainId })
|
||||
})
|
||||
.where(eq(certificates.domain, domain));
|
||||
|
||||
logger.debug(
|
||||
`acmeCertSync: updated certificate for ${domain} (expires ${expiresAt ? new Date(expiresAt * 1000).toISOString() : "unknown"})`
|
||||
);
|
||||
|
||||
await pushCertUpdateToAffectedNewts(
|
||||
domain,
|
||||
domainId,
|
||||
oldCertPem,
|
||||
oldKeyPem
|
||||
);
|
||||
} else {
|
||||
logger.debug(
|
||||
`acmeCertSync: inserting new certificate for ${domain} (expires ${expiresAt ? new Date(expiresAt * 1000).toISOString() : "unknown"})`
|
||||
);
|
||||
await db.insert(certificates).values({
|
||||
domain,
|
||||
domainId,
|
||||
certFile: encryptedCert,
|
||||
keyFile: encryptedKey,
|
||||
status: "valid",
|
||||
expiresAt,
|
||||
createdAt: now,
|
||||
updatedAt: now,
|
||||
wildcard
|
||||
});
|
||||
|
||||
logger.debug(
|
||||
`acmeCertSync: inserted new certificate for ${domain} (expires ${expiresAt ? new Date(expiresAt * 1000).toISOString() : "unknown"})`
|
||||
);
|
||||
|
||||
await pushCertUpdateToAffectedNewts(domain, domainId, null, null);
|
||||
}
|
||||
}
|
||||
|
||||
function findAcmeJsonFiles(dirPath: string): string[] {
|
||||
const results: string[] = [];
|
||||
let entries: fs.Dirent[];
|
||||
try {
|
||||
entries = fs.readdirSync(dirPath, { withFileTypes: true });
|
||||
} catch (err) {
|
||||
logger.warn(
|
||||
`acmeCertSync: could not read directory "${dirPath}": ${err}`
|
||||
);
|
||||
return results;
|
||||
}
|
||||
for (const entry of entries) {
|
||||
const fullPath = path.join(dirPath, entry.name);
|
||||
if (entry.isDirectory()) {
|
||||
results.push(...findAcmeJsonFiles(fullPath));
|
||||
} else if (entry.isFile()) {
|
||||
// check if it is a json file
|
||||
if (entry.name.endsWith(".json")) {
|
||||
let raw: string;
|
||||
try {
|
||||
raw = fs.readFileSync(fullPath, "utf8");
|
||||
} catch (err) {
|
||||
logger.warn(
|
||||
`acmeCertSync: could not read file "${fullPath}": ${err}`
|
||||
);
|
||||
continue;
|
||||
}
|
||||
|
||||
let parsed: any;
|
||||
try {
|
||||
parsed = JSON.parse(raw);
|
||||
} catch (err) {
|
||||
logger.warn(
|
||||
`acmeCertSync: could not parse "${fullPath}" as JSON: ${err}`
|
||||
);
|
||||
continue;
|
||||
}
|
||||
}
|
||||
|
||||
results.push(fullPath);
|
||||
}
|
||||
}
|
||||
return results;
|
||||
}
|
||||
|
||||
async function syncAcmeCerts(acmeJsonPath: string): Promise<void> {
|
||||
let raw: string;
|
||||
try {
|
||||
raw = fs.readFileSync(acmeJsonPath, "utf8");
|
||||
} catch (err) {
|
||||
logger.warn(`acmeCertSync: could not read "${acmeJsonPath}": ${err}`);
|
||||
return;
|
||||
}
|
||||
|
||||
let acmeJson: AcmeJson;
|
||||
try {
|
||||
acmeJson = JSON.parse(raw);
|
||||
} catch (err) {
|
||||
logger.warn(
|
||||
`acmeCertSync: could not parse "${acmeJsonPath}" as JSON: ${err}`
|
||||
);
|
||||
return;
|
||||
}
|
||||
|
||||
const resolvers = Object.keys(acmeJson || {});
|
||||
if (resolvers.length === 0) {
|
||||
logger.debug(`acmeCertSync: no resolvers found in acme.json`);
|
||||
return;
|
||||
}
|
||||
|
||||
// Collect certificates from every resolver. If the same domain appears in
|
||||
// multiple resolvers, the last one wins (resolvers iterated in object order).
|
||||
const allCerts: AcmeCert[] = [];
|
||||
for (const resolver of resolvers) {
|
||||
const resolverData = acmeJson[resolver];
|
||||
if (!resolverData || !Array.isArray(resolverData.Certificates)) {
|
||||
logger.debug(
|
||||
`acmeCertSync: no certificates found for resolver "${resolver}"`
|
||||
);
|
||||
continue;
|
||||
}
|
||||
// logger.debug(
|
||||
// `acmeCertSync: found ${resolverData.Certificates.length} certificate(s) for resolver "${resolver}"`
|
||||
// );
|
||||
for (const cert of resolverData.Certificates) {
|
||||
allCerts.push(cert);
|
||||
}
|
||||
}
|
||||
|
||||
for (const cert of allCerts) {
|
||||
const mainDomain = cert?.domain?.main;
|
||||
|
||||
if (!mainDomain || typeof mainDomain !== "string") {
|
||||
logger.debug(`acmeCertSync: skipping cert with missing domain`);
|
||||
continue;
|
||||
}
|
||||
|
||||
if (!cert.certificate || !cert.key) {
|
||||
logger.debug(
|
||||
`acmeCertSync: skipping cert for ${mainDomain} - empty certificate or key field`
|
||||
);
|
||||
continue;
|
||||
}
|
||||
|
||||
let certPem: string;
|
||||
let keyPem: string;
|
||||
try {
|
||||
certPem = Buffer.from(cert.certificate, "base64").toString("utf8");
|
||||
keyPem = Buffer.from(cert.key, "base64").toString("utf8");
|
||||
} catch (err) {
|
||||
logger.debug(
|
||||
`acmeCertSync: skipping cert for ${mainDomain} - failed to base64-decode cert/key: ${err}`
|
||||
);
|
||||
continue;
|
||||
}
|
||||
|
||||
if (!certPem.trim() || !keyPem.trim()) {
|
||||
logger.debug(
|
||||
`acmeCertSync: skipping cert for ${mainDomain} - blank PEM after base64 decode`
|
||||
);
|
||||
continue;
|
||||
}
|
||||
|
||||
// Validate that the decoded data actually parses as a real X.509 cert
|
||||
// before we touch the database. This prevents importing partially-written
|
||||
// or corrupted entries from acme.json.
|
||||
const firstCertPemForValidation = extractFirstCert(certPem);
|
||||
if (!firstCertPemForValidation) {
|
||||
logger.debug(
|
||||
`acmeCertSync: skipping cert for ${mainDomain} - no PEM certificate block found`
|
||||
);
|
||||
continue;
|
||||
}
|
||||
|
||||
let validatedX509: crypto.X509Certificate;
|
||||
try {
|
||||
validatedX509 = new crypto.X509Certificate(
|
||||
firstCertPemForValidation
|
||||
);
|
||||
} catch (err) {
|
||||
logger.debug(
|
||||
`acmeCertSync: skipping cert for ${mainDomain} - invalid X.509 certificate: ${err}`
|
||||
);
|
||||
continue;
|
||||
}
|
||||
|
||||
// Sanity-check the private key parses too
|
||||
try {
|
||||
crypto.createPrivateKey(keyPem);
|
||||
} catch (err) {
|
||||
logger.debug(
|
||||
`acmeCertSync: skipping cert for ${mainDomain} - invalid private key: ${err}`
|
||||
);
|
||||
continue;
|
||||
}
|
||||
|
||||
// Collect all domains covered by this cert: main + every SAN.
|
||||
// Each domain gets its own row in the certificates table so that
|
||||
// lookups by any hostname on the cert succeed independently.
|
||||
const allDomains = new Set<string>([mainDomain]);
|
||||
if (Array.isArray(cert.domain?.sans)) {
|
||||
for (const san of cert.domain.sans) {
|
||||
if (typeof san === "string" && san.trim()) {
|
||||
allDomains.add(san.trim());
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// logger.debug(
|
||||
// `acmeCertSync: cert for ${mainDomain} covers ${allDomains.size} domain(s): ${[...allDomains].join(", ")}`
|
||||
// );
|
||||
|
||||
for (const domain of allDomains) {
|
||||
try {
|
||||
await storeCertForDomain(
|
||||
domain,
|
||||
certPem,
|
||||
keyPem,
|
||||
validatedX509
|
||||
);
|
||||
} catch (err) {
|
||||
logger.error(
|
||||
`acmeCertSync: error storing cert for domain "${domain}": ${err}`
|
||||
);
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
export function initAcmeCertSync(): void {
|
||||
// stub
|
||||
}
|
||||
if (build == "saas") {
|
||||
logger.debug(`acmeCertSync: skipping ACME cert sync in SaaS build`);
|
||||
return;
|
||||
}
|
||||
|
||||
const configData = config.getRawConfig();
|
||||
|
||||
if (!configData.flags?.enable_acme_cert_sync) {
|
||||
logger.debug(
|
||||
`acmeCertSync: ACME cert sync is disabled by config flag, skipping`
|
||||
);
|
||||
return;
|
||||
}
|
||||
|
||||
const acmeJsonPath =
|
||||
configData.acme?.acme_json_path ?? "config/letsencrypt/acme.json";
|
||||
const intervalMs = configData.acme?.sync_interval_ms ?? 5000;
|
||||
const httpEndpoint = configData.acme?.acme_http_endpoint;
|
||||
|
||||
logger.debug(
|
||||
`acmeCertSync: starting ACME cert sync from "${acmeJsonPath}" across all resolvers every ${intervalMs}ms`
|
||||
);
|
||||
if (httpEndpoint) {
|
||||
logger.debug(
|
||||
`acmeCertSync: also syncing from HTTP endpoint "${httpEndpoint}" every ${intervalMs}ms`
|
||||
);
|
||||
}
|
||||
|
||||
const runSync = () => {
|
||||
if (httpEndpoint) {
|
||||
syncAcmeCertsFromHttp(httpEndpoint).catch((err) => {
|
||||
logger.error(`acmeCertSync: error during HTTP sync: ${err}`);
|
||||
});
|
||||
} else {
|
||||
// only run the file-based sync if the HTTP endpoint is not configured, to avoid doubling up
|
||||
let stat: fs.Stats | null = null;
|
||||
try {
|
||||
stat = fs.statSync(acmeJsonPath);
|
||||
} catch (err) {
|
||||
logger.warn(
|
||||
`acmeCertSync: cannot stat path "${acmeJsonPath}": ${err}`
|
||||
);
|
||||
return;
|
||||
}
|
||||
|
||||
if (stat.isDirectory()) {
|
||||
const files = findAcmeJsonFiles(acmeJsonPath);
|
||||
if (files.length === 0) {
|
||||
logger.debug(
|
||||
`acmeCertSync: no acme.json files found in directory "${acmeJsonPath}"`
|
||||
);
|
||||
return;
|
||||
}
|
||||
// logger.debug(
|
||||
// `acmeCertSync: found ${files.length} acme.json file(s) in directory "${acmeJsonPath}"`
|
||||
// );
|
||||
for (const file of files) {
|
||||
syncAcmeCerts(file).catch((err) => {
|
||||
logger.error(
|
||||
`acmeCertSync: error during sync of "${file}": ${err}`
|
||||
);
|
||||
});
|
||||
}
|
||||
} else {
|
||||
syncAcmeCerts(acmeJsonPath).catch((err) => {
|
||||
logger.error(`acmeCertSync: error during sync: ${err}`);
|
||||
});
|
||||
}
|
||||
}
|
||||
};
|
||||
|
||||
// Run immediately on init, then on the configured interval
|
||||
runSync();
|
||||
|
||||
setInterval(runSync, intervalMs);
|
||||
}
|
||||
|
||||
@@ -0,0 +1,616 @@
|
||||
import {
|
||||
and,
|
||||
eq,
|
||||
gte,
|
||||
inArray,
|
||||
isNull,
|
||||
or,
|
||||
sql,
|
||||
SQL,
|
||||
type InferInsertModel
|
||||
} from "drizzle-orm";
|
||||
import {
|
||||
AiBudget,
|
||||
aiBudgetBreachEvents,
|
||||
aiBudgets,
|
||||
aiModels,
|
||||
aiUsageRecords,
|
||||
db,
|
||||
logsDb,
|
||||
userOrgRoles
|
||||
} from "@server/db";
|
||||
import { modelKeyMatches } from "@server/lib/aiModelKeyMatch";
|
||||
import type { AiUsage } from "@server/lib/aiUsageExtraction";
|
||||
import { regionalCache as cache } from "#dynamic/lib/cache";
|
||||
import logger from "@server/logger";
|
||||
|
||||
type BudgetPeriod = AiBudget["period"];
|
||||
|
||||
const PERIOD_DURATIONS_MS: Record<Exclude<BudgetPeriod, "lifetime">, number> = {
|
||||
hourly: 60 * 60 * 1000,
|
||||
daily: 24 * 60 * 60 * 1000,
|
||||
weekly: 7 * 24 * 60 * 60 * 1000,
|
||||
monthly: 30 * 24 * 60 * 60 * 1000,
|
||||
yearly: 365 * 24 * 60 * 60 * 1000
|
||||
};
|
||||
|
||||
// Budgets are cheap to be a little stale about (enforcement is already
|
||||
// check-then-act, not transactional). Re-derive each budget's usage sum
|
||||
// from aiUsageRecords at most this often; in between, completed requests
|
||||
// just add their own contribution onto the cached sum instead of
|
||||
// re-querying/re-aggregating from scratch.
|
||||
const BUDGET_CACHE_REFRESH_MS = 8_000;
|
||||
// Redis-level TTL is only a safety net for eviction if a budget stops
|
||||
// seeing traffic - the actual staleness check is the computedAt timestamp
|
||||
// stored in the cached value, compared against BUDGET_CACHE_REFRESH_MS.
|
||||
const BUDGET_CACHE_SAFETY_TTL_SEC = 60;
|
||||
|
||||
function applicableBudgetsCacheKey(ctx: BudgetScopeContext): string {
|
||||
const roleKey = [...ctx.roleIds].sort((a, b) => a - b).join(",");
|
||||
return [
|
||||
"aiBudget:applicable",
|
||||
ctx.orgId,
|
||||
ctx.providerId,
|
||||
ctx.requestedModel,
|
||||
ctx.resourceId ?? "",
|
||||
ctx.siteResourceId ?? "",
|
||||
roleKey,
|
||||
ctx.virtualApiKeyId ?? ""
|
||||
].join(":");
|
||||
}
|
||||
|
||||
function budgetUsageCacheKey(budgetId: number): string {
|
||||
return `aiBudget:usage:${budgetId}`;
|
||||
}
|
||||
|
||||
type CachedBudgetUsage = {
|
||||
sum: number;
|
||||
computedAt: number;
|
||||
};
|
||||
|
||||
// Budget periods are trailing windows from "now", not calendar-aligned
|
||||
// (e.g. "daily" = last 24h). "lifetime" has no lower bound.
|
||||
function windowStart(period: BudgetPeriod, now: number): number {
|
||||
if (period === "lifetime") {
|
||||
return 0;
|
||||
}
|
||||
return now - PERIOD_DURATIONS_MS[period];
|
||||
}
|
||||
|
||||
export type BudgetScopeContext = {
|
||||
orgId: string;
|
||||
providerId: number;
|
||||
requestedModel: string;
|
||||
resourceId: number | null;
|
||||
siteResourceId: number | null;
|
||||
roleIds: number[];
|
||||
requestUserId: string | null;
|
||||
virtualApiKeyId: string | null;
|
||||
};
|
||||
|
||||
/**
|
||||
* Every budget that could apply to this request: the provider itself, any
|
||||
* model on that provider whose (possibly wildcarded) modelKey matches the
|
||||
* requested model, the target resource/site-resource, and any role the
|
||||
* requesting user holds in the org. Cached for BUDGET_CACHE_REFRESH_MS since
|
||||
* budget/model config changes are rare and a request-scoped org/provider/
|
||||
* model/resource/role combination repeats constantly under real traffic.
|
||||
*/
|
||||
export async function resolveApplicableBudgets(
|
||||
ctx: BudgetScopeContext
|
||||
): Promise<AiBudget[]> {
|
||||
const cacheKey = applicableBudgetsCacheKey(ctx);
|
||||
const cached = await cache.get<AiBudget[]>(cacheKey);
|
||||
if (cached !== undefined) {
|
||||
return cached;
|
||||
}
|
||||
|
||||
const budgets = await fetchApplicableBudgets(ctx);
|
||||
await cache.set(cacheKey, budgets, BUDGET_CACHE_REFRESH_MS / 1000);
|
||||
return budgets;
|
||||
}
|
||||
|
||||
async function fetchApplicableBudgets(
|
||||
ctx: BudgetScopeContext
|
||||
): Promise<AiBudget[]> {
|
||||
const providerModels = await db
|
||||
.select({ modelId: aiModels.modelId, modelKey: aiModels.modelKey })
|
||||
.from(aiModels)
|
||||
.where(
|
||||
and(
|
||||
eq(aiModels.providerId, ctx.providerId),
|
||||
eq(aiModels.enabled, true)
|
||||
)
|
||||
);
|
||||
|
||||
const matchingModelIds = providerModels
|
||||
.filter((m) => modelKeyMatches(m.modelKey, ctx.requestedModel))
|
||||
.map((m) => m.modelId);
|
||||
|
||||
const scopeConditions: SQL[] = [
|
||||
and(
|
||||
eq(aiBudgets.providerId, ctx.providerId),
|
||||
isNull(aiBudgets.modelId)
|
||||
)!
|
||||
];
|
||||
if (matchingModelIds.length > 0) {
|
||||
scopeConditions.push(inArray(aiBudgets.modelId, matchingModelIds));
|
||||
}
|
||||
if (ctx.resourceId != null) {
|
||||
scopeConditions.push(eq(aiBudgets.resourceId, ctx.resourceId));
|
||||
}
|
||||
if (ctx.siteResourceId != null) {
|
||||
scopeConditions.push(eq(aiBudgets.siteResourceId, ctx.siteResourceId));
|
||||
}
|
||||
if (ctx.roleIds.length > 0) {
|
||||
scopeConditions.push(inArray(aiBudgets.roleId, ctx.roleIds));
|
||||
}
|
||||
if (ctx.virtualApiKeyId != null) {
|
||||
scopeConditions.push(
|
||||
eq(aiBudgets.virtualApiKeyId, ctx.virtualApiKeyId)
|
||||
);
|
||||
}
|
||||
|
||||
return db
|
||||
.select()
|
||||
.from(aiBudgets)
|
||||
.where(
|
||||
and(
|
||||
eq(aiBudgets.orgId, ctx.orgId),
|
||||
eq(aiBudgets.enabled, true),
|
||||
or(...scopeConditions)
|
||||
)
|
||||
);
|
||||
}
|
||||
|
||||
async function sumUsageAmount(
|
||||
where: SQL,
|
||||
unit: AiBudget["unit"]
|
||||
): Promise<number> {
|
||||
const column =
|
||||
unit === "usd" ? aiUsageRecords.costUsd : aiUsageRecords.totalTokens;
|
||||
const [row] = await logsDb
|
||||
.select({ total: sql<number>`coalesce(sum(${column}), 0)` })
|
||||
.from(aiUsageRecords)
|
||||
.where(where);
|
||||
return Number(row?.total ?? 0);
|
||||
}
|
||||
|
||||
/**
|
||||
* Sums recorded usage for a single budget's scope + rolling window. Model
|
||||
* budgets can't be pushed down to SQL because the model's key may itself be
|
||||
* a glob, so those rows are fetched for the provider+window and matched in
|
||||
* JS the same way access-control matching does.
|
||||
*/
|
||||
export async function sumUsageForBudget(
|
||||
budget: AiBudget,
|
||||
ctx: BudgetScopeContext,
|
||||
now: number
|
||||
): Promise<number> {
|
||||
const start = windowStart(budget.period, now);
|
||||
|
||||
if (budget.modelId != null) {
|
||||
const [model] = await db
|
||||
.select({
|
||||
providerId: aiModels.providerId,
|
||||
modelKey: aiModels.modelKey
|
||||
})
|
||||
.from(aiModels)
|
||||
.where(eq(aiModels.modelId, budget.modelId))
|
||||
.limit(1);
|
||||
if (!model) {
|
||||
return 0;
|
||||
}
|
||||
const rows = await logsDb
|
||||
.select({
|
||||
requestedModel: aiUsageRecords.requestedModel,
|
||||
costUsd: aiUsageRecords.costUsd,
|
||||
totalTokens: aiUsageRecords.totalTokens
|
||||
})
|
||||
.from(aiUsageRecords)
|
||||
.where(
|
||||
and(
|
||||
eq(aiUsageRecords.orgId, ctx.orgId),
|
||||
eq(aiUsageRecords.providerId, model.providerId),
|
||||
gte(aiUsageRecords.createdAt, start)
|
||||
)
|
||||
);
|
||||
return rows
|
||||
.filter((r) => modelKeyMatches(model.modelKey, r.requestedModel))
|
||||
.reduce(
|
||||
(sum, r) =>
|
||||
sum +
|
||||
(budget.unit === "usd" ? (r.costUsd ?? 0) : r.totalTokens),
|
||||
0
|
||||
);
|
||||
}
|
||||
|
||||
if (budget.providerId != null) {
|
||||
return sumUsageAmount(
|
||||
and(
|
||||
eq(aiUsageRecords.orgId, ctx.orgId),
|
||||
eq(aiUsageRecords.providerId, budget.providerId),
|
||||
gte(aiUsageRecords.createdAt, start)
|
||||
)!,
|
||||
budget.unit
|
||||
);
|
||||
}
|
||||
|
||||
if (budget.resourceId != null) {
|
||||
return sumUsageAmount(
|
||||
and(
|
||||
eq(aiUsageRecords.orgId, ctx.orgId),
|
||||
eq(aiUsageRecords.resourceId, budget.resourceId),
|
||||
gte(aiUsageRecords.createdAt, start)
|
||||
)!,
|
||||
budget.unit
|
||||
);
|
||||
}
|
||||
|
||||
if (budget.siteResourceId != null) {
|
||||
return sumUsageAmount(
|
||||
and(
|
||||
eq(aiUsageRecords.orgId, ctx.orgId),
|
||||
eq(aiUsageRecords.siteResourceId, budget.siteResourceId),
|
||||
gte(aiUsageRecords.createdAt, start)
|
||||
)!,
|
||||
budget.unit
|
||||
);
|
||||
}
|
||||
|
||||
if (budget.roleId != null) {
|
||||
const members = await db
|
||||
.select({ userId: userOrgRoles.userId })
|
||||
.from(userOrgRoles)
|
||||
.where(
|
||||
and(
|
||||
eq(userOrgRoles.roleId, budget.roleId),
|
||||
eq(userOrgRoles.orgId, ctx.orgId)
|
||||
)
|
||||
);
|
||||
const userIds = members.map((m) => m.userId);
|
||||
if (userIds.length === 0) {
|
||||
return 0;
|
||||
}
|
||||
return sumUsageAmount(
|
||||
and(
|
||||
eq(aiUsageRecords.orgId, ctx.orgId),
|
||||
inArray(aiUsageRecords.userId, userIds),
|
||||
gte(aiUsageRecords.createdAt, start)
|
||||
)!,
|
||||
budget.unit
|
||||
);
|
||||
}
|
||||
|
||||
if (budget.virtualApiKeyId != null) {
|
||||
return sumUsageAmount(
|
||||
and(
|
||||
eq(aiUsageRecords.orgId, ctx.orgId),
|
||||
eq(aiUsageRecords.virtualApiKeyId, budget.virtualApiKeyId),
|
||||
gte(aiUsageRecords.createdAt, start)
|
||||
)!,
|
||||
budget.unit
|
||||
);
|
||||
}
|
||||
|
||||
return 0;
|
||||
}
|
||||
|
||||
/**
|
||||
* Cached wrapper around sumUsageForBudget. Reuses a per-budget cached sum
|
||||
* for up to BUDGET_CACHE_REFRESH_MS, and otherwise falls through to the DB
|
||||
* aggregation and reseeds the cache. Completed requests within that window
|
||||
* top the cached sum up via applyUsageToBudgetCache below rather than
|
||||
* forcing a re-aggregation on every request.
|
||||
*/
|
||||
async function getBudgetUsage(
|
||||
budget: AiBudget,
|
||||
ctx: BudgetScopeContext,
|
||||
now: number
|
||||
): Promise<number> {
|
||||
const cacheKey = budgetUsageCacheKey(budget.budgetId);
|
||||
const cached = await cache.get<CachedBudgetUsage>(cacheKey);
|
||||
if (cached && now - cached.computedAt < BUDGET_CACHE_REFRESH_MS) {
|
||||
return cached.sum;
|
||||
}
|
||||
|
||||
const sum = await sumUsageForBudget(budget, ctx, now);
|
||||
await cache.set(
|
||||
cacheKey,
|
||||
{ sum, computedAt: now } satisfies CachedBudgetUsage,
|
||||
BUDGET_CACHE_SAFETY_TTL_SEC
|
||||
);
|
||||
return sum;
|
||||
}
|
||||
|
||||
/**
|
||||
* Called once a request's actual usage is known, for every budget that was
|
||||
* resolved as applicable to it (i.e. checkBudgets' returned `budgets`).
|
||||
* Adds this request's contribution directly onto each budget's cached sum
|
||||
* so the next request in the same refresh window doesn't need to re-query
|
||||
* or re-aggregate. If there's no warm cache entry, or it's already due for
|
||||
* a refresh, this is a no-op - the next reader re-derives from the DB,
|
||||
* which by then already includes this request's row via recordUsage.
|
||||
*/
|
||||
export async function applyUsageToBudgetCache(
|
||||
budgets: AiBudget[],
|
||||
usage: { usd: number; tokens: number }
|
||||
): Promise<void> {
|
||||
await Promise.all(
|
||||
budgets.map(async (budget) => {
|
||||
const delta = budget.unit === "usd" ? usage.usd : usage.tokens;
|
||||
if (!delta) {
|
||||
return;
|
||||
}
|
||||
|
||||
const cacheKey = budgetUsageCacheKey(budget.budgetId);
|
||||
const cached = await cache.get<CachedBudgetUsage>(cacheKey);
|
||||
if (
|
||||
!cached ||
|
||||
Date.now() - cached.computedAt >= BUDGET_CACHE_REFRESH_MS
|
||||
) {
|
||||
return;
|
||||
}
|
||||
|
||||
await cache.set(
|
||||
cacheKey,
|
||||
{
|
||||
sum: cached.sum + delta,
|
||||
computedAt: cached.computedAt
|
||||
} satisfies CachedBudgetUsage,
|
||||
BUDGET_CACHE_SAFETY_TTL_SEC
|
||||
);
|
||||
})
|
||||
);
|
||||
}
|
||||
|
||||
// Throttled to one durable event per budget per breach window, so a soft
|
||||
// budget being exceeded doesn't write a row on every subsequent request
|
||||
// while it stays over.
|
||||
async function recordBreachEventIfNew(
|
||||
budget: AiBudget,
|
||||
ctx: BudgetScopeContext,
|
||||
usageAmount: number,
|
||||
now: number
|
||||
): Promise<void> {
|
||||
try {
|
||||
const start = windowStart(budget.period, now);
|
||||
const [existing] = await db
|
||||
.select({ id: aiBudgetBreachEvents.id })
|
||||
.from(aiBudgetBreachEvents)
|
||||
.where(
|
||||
and(
|
||||
eq(aiBudgetBreachEvents.budgetId, budget.budgetId),
|
||||
gte(aiBudgetBreachEvents.createdAt, start)
|
||||
)
|
||||
)
|
||||
.limit(1);
|
||||
if (existing) {
|
||||
return;
|
||||
}
|
||||
|
||||
await db.insert(aiBudgetBreachEvents).values({
|
||||
orgId: ctx.orgId,
|
||||
budgetId: budget.budgetId,
|
||||
enforcement: budget.enforcement,
|
||||
unit: budget.unit,
|
||||
period: budget.period,
|
||||
amount: budget.amount,
|
||||
usageAmount,
|
||||
blocked: budget.enforcement === "hard",
|
||||
requestUserId: ctx.requestUserId,
|
||||
createdAt: now
|
||||
});
|
||||
} catch (error) {
|
||||
logger.error("Failed to record AI budget breach event", {
|
||||
error,
|
||||
budgetId: budget.budgetId
|
||||
});
|
||||
}
|
||||
}
|
||||
|
||||
export type BudgetCheckResult = {
|
||||
blocked: boolean;
|
||||
blockingBudget?: AiBudget;
|
||||
// Every budget resolved as applicable to this request, regardless of
|
||||
// whether it was breached - pass to applyUsageToBudgetCache once this
|
||||
// request's actual usage is known.
|
||||
budgets: AiBudget[];
|
||||
};
|
||||
|
||||
export async function checkBudgets(
|
||||
ctx: BudgetScopeContext
|
||||
): Promise<BudgetCheckResult> {
|
||||
const budgets = await resolveApplicableBudgets(ctx);
|
||||
if (budgets.length === 0) {
|
||||
return { blocked: false, budgets: [] };
|
||||
}
|
||||
|
||||
const now = Date.now();
|
||||
let blockingBudget: AiBudget | undefined;
|
||||
|
||||
for (const budget of budgets) {
|
||||
const usage = await getBudgetUsage(budget, ctx, now);
|
||||
if (usage < budget.amount) {
|
||||
continue;
|
||||
}
|
||||
|
||||
await recordBreachEventIfNew(budget, ctx, usage, now);
|
||||
|
||||
if (budget.enforcement === "hard" && !blockingBudget) {
|
||||
blockingBudget = budget;
|
||||
}
|
||||
}
|
||||
|
||||
return blockingBudget
|
||||
? { blocked: true, blockingBudget, budgets }
|
||||
: { blocked: false, budgets };
|
||||
}
|
||||
|
||||
export type UsageRecordInput = {
|
||||
orgId: string;
|
||||
providerId: number;
|
||||
resourceId: number | null;
|
||||
siteResourceId: number | null;
|
||||
userId: string | null;
|
||||
virtualApiKeyId: string | null;
|
||||
requestedModel: string;
|
||||
usage: AiUsage;
|
||||
costUsd: number | null;
|
||||
createdAt?: number;
|
||||
// Same id as the aiSessionLog row logged for this request, so the two
|
||||
// can be joined to show token/cost usage alongside the session
|
||||
// transcript. Undefined when the session wasn't logged (e.g. session
|
||||
// log retention disabled for the org).
|
||||
sessionId?: string;
|
||||
};
|
||||
|
||||
type AiUsageRecordInsert = InferInsertModel<typeof aiUsageRecords>;
|
||||
|
||||
// In-memory buffer for batching AI usage record inserts, mirroring the
|
||||
// approach in server/routers/badger/logRequestAudit.ts. Usage rows are read
|
||||
// back on every budget-cache miss (see getBudgetUsage above), which happens
|
||||
// at least every BUDGET_CACHE_REFRESH_MS, so this buffer is flushed much
|
||||
// more aggressively than the request audit log to keep the table from
|
||||
// lagging behind what budget enforcement needs. Unlike the audit log, there
|
||||
// is no retention/cleanup job for this table - usage history is kept
|
||||
// indefinitely for billing and historical reporting.
|
||||
const usageRecordBuffer: AiUsageRecordInsert[] = [];
|
||||
|
||||
const USAGE_BATCH_SIZE = 20; // Write to DB every 20 records
|
||||
const USAGE_BATCH_INTERVAL_MS = 1000; // Or every 1 second, whichever comes first
|
||||
const USAGE_MAX_BUFFER_SIZE = 5000; // Prevent unbounded memory growth
|
||||
let usageFlushTimer: NodeJS.Timeout | null = null;
|
||||
let isUsageFlushInProgress = false;
|
||||
|
||||
async function flushUsageRecords() {
|
||||
if (usageRecordBuffer.length === 0 || isUsageFlushInProgress) {
|
||||
return;
|
||||
}
|
||||
|
||||
isUsageFlushInProgress = true;
|
||||
|
||||
const recordsToWrite = usageRecordBuffer.splice(
|
||||
0,
|
||||
usageRecordBuffer.length
|
||||
);
|
||||
|
||||
try {
|
||||
// Use a transaction to ensure all inserts succeed or fail together
|
||||
await logsDb.transaction(async (tx) => {
|
||||
// Batch insert in groups to avoid overwhelming the database
|
||||
const DB_BATCH_SIZE = 25;
|
||||
for (let i = 0; i < recordsToWrite.length; i += DB_BATCH_SIZE) {
|
||||
const batch = recordsToWrite.slice(i, i + DB_BATCH_SIZE);
|
||||
await tx.insert(aiUsageRecords).values(batch);
|
||||
}
|
||||
});
|
||||
logger.debug(
|
||||
`Flushed ${recordsToWrite.length} AI usage records to database`
|
||||
);
|
||||
} catch (error) {
|
||||
logger.error("Error flushing AI usage records:", error);
|
||||
// On transaction error, put records back at the front of the buffer
|
||||
// to retry, but only if the buffer isn't too large
|
||||
if (
|
||||
usageRecordBuffer.length <
|
||||
USAGE_MAX_BUFFER_SIZE - recordsToWrite.length
|
||||
) {
|
||||
usageRecordBuffer.unshift(...recordsToWrite);
|
||||
logger.info(
|
||||
`Re-queued ${recordsToWrite.length} AI usage records for retry`
|
||||
);
|
||||
} else {
|
||||
logger.error(
|
||||
`Buffer full, dropped ${recordsToWrite.length} AI usage records`
|
||||
);
|
||||
}
|
||||
} finally {
|
||||
isUsageFlushInProgress = false;
|
||||
// If buffer filled up while we were flushing, flush again
|
||||
if (usageRecordBuffer.length >= USAGE_BATCH_SIZE) {
|
||||
flushUsageRecords().catch((err) =>
|
||||
logger.error("Error in follow-up AI usage flush:", err)
|
||||
);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
function scheduleUsageFlush() {
|
||||
if (usageFlushTimer === null) {
|
||||
usageFlushTimer = setTimeout(() => {
|
||||
usageFlushTimer = null;
|
||||
flushUsageRecords().catch((err) =>
|
||||
logger.error("Error in scheduled AI usage flush:", err)
|
||||
);
|
||||
}, USAGE_BATCH_INTERVAL_MS);
|
||||
}
|
||||
}
|
||||
|
||||
/**
|
||||
* Gracefully flush all pending AI usage records (call this on shutdown).
|
||||
*/
|
||||
export async function shutdownUsageRecorder() {
|
||||
if (usageFlushTimer) {
|
||||
clearTimeout(usageFlushTimer);
|
||||
usageFlushTimer = null;
|
||||
}
|
||||
// Force flush even if one is in progress by waiting and retrying
|
||||
while (isUsageFlushInProgress) {
|
||||
await new Promise((resolve) => setTimeout(resolve, 100));
|
||||
}
|
||||
await flushUsageRecords();
|
||||
}
|
||||
|
||||
export async function recordUsage(input: UsageRecordInput): Promise<void> {
|
||||
try {
|
||||
const { usage } = input;
|
||||
const totalTokens =
|
||||
usage.promptTokens +
|
||||
usage.cacheReadTokens +
|
||||
usage.cacheWriteTokens +
|
||||
usage.completionTokens +
|
||||
usage.reasoningTokens;
|
||||
|
||||
// Prevent unbounded buffer growth - drop oldest entries if buffer is too large
|
||||
if (usageRecordBuffer.length >= USAGE_MAX_BUFFER_SIZE) {
|
||||
const dropped = usageRecordBuffer.splice(0, USAGE_BATCH_SIZE);
|
||||
logger.warn(
|
||||
`AI usage record buffer exceeded max size (${USAGE_MAX_BUFFER_SIZE}), dropped ${dropped.length} oldest entries`
|
||||
);
|
||||
}
|
||||
|
||||
const timestamp = Math.floor(Date.now() / 1000);
|
||||
|
||||
usageRecordBuffer.push({
|
||||
orgId: input.orgId,
|
||||
providerId: input.providerId,
|
||||
resourceId: input.resourceId,
|
||||
siteResourceId: input.siteResourceId,
|
||||
userId: input.userId,
|
||||
virtualApiKeyId: input.virtualApiKeyId,
|
||||
sessionId: input.sessionId,
|
||||
requestedModel: input.requestedModel,
|
||||
promptTokens: usage.promptTokens,
|
||||
cacheReadTokens: usage.cacheReadTokens,
|
||||
cacheWriteTokens: usage.cacheWriteTokens,
|
||||
completionTokens: usage.completionTokens,
|
||||
reasoningTokens: usage.reasoningTokens,
|
||||
totalTokens,
|
||||
costUsd: input.costUsd,
|
||||
estimated: usage.estimated,
|
||||
createdAt: input.createdAt ?? timestamp
|
||||
});
|
||||
|
||||
// Flush immediately if buffer is full, otherwise schedule a flush
|
||||
if (usageRecordBuffer.length >= USAGE_BATCH_SIZE) {
|
||||
flushUsageRecords().catch((err) =>
|
||||
logger.error("Error flushing AI usage records:", err)
|
||||
);
|
||||
} else {
|
||||
scheduleUsageFlush();
|
||||
}
|
||||
} catch (error) {
|
||||
logger.error("Failed to record AI usage", { error });
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,354 @@
|
||||
import type { Request } from "express";
|
||||
import { AI_CAPABILITIES, type AiCapability } from "@app/lib/aiCapabilities";
|
||||
|
||||
export { AI_CAPABILITIES, type AiCapability };
|
||||
|
||||
export type AiCapabilityRoute = {
|
||||
method: "GET" | "POST";
|
||||
path: string;
|
||||
};
|
||||
|
||||
export type AiProtocolFamily = "openai" | "anthropic" | "google" | "bedrock";
|
||||
|
||||
export type AiCapabilityDefinition = {
|
||||
id: AiCapability;
|
||||
protocolFamily: AiProtocolFamily;
|
||||
routes: AiCapabilityRoute[];
|
||||
extractModel: (req: Request) => string | undefined;
|
||||
resolveUpstreamUrl: (
|
||||
baseUrl: string,
|
||||
req: Request,
|
||||
model: string
|
||||
) => string;
|
||||
isStreaming: (req: Request, contentType: string) => boolean;
|
||||
};
|
||||
|
||||
function bodyModel(req: Request): string | undefined {
|
||||
return typeof req.body?.model === "string" ? req.body.model : undefined;
|
||||
}
|
||||
|
||||
function paramModel(req: Request): string | undefined {
|
||||
const model = req.params?.model;
|
||||
return typeof model === "string" && model.length > 0 ? model : undefined;
|
||||
}
|
||||
|
||||
export function joinUpstreamUrl(baseUrl: string, path: string): string {
|
||||
const base = baseUrl.replace(/\/+$/, "");
|
||||
let suffix = path.startsWith("/") ? path : `/${path}`;
|
||||
|
||||
let basePathname = "/";
|
||||
try {
|
||||
basePathname = new URL(base).pathname.replace(/\/+$/, "") || "/";
|
||||
} catch {
|
||||
// Fall through with "/" non-absolute bases are not expected in
|
||||
// production, but keep joining usable for malformed input.
|
||||
}
|
||||
|
||||
if (basePathname !== "/") {
|
||||
const baseSegs = basePathname.split("/").filter(Boolean);
|
||||
const pathSegs = suffix.split("/").filter(Boolean);
|
||||
const max = Math.min(baseSegs.length, pathSegs.length);
|
||||
let overlap = 0;
|
||||
for (let n = max; n >= 1; n--) {
|
||||
const baseSuffix = baseSegs.slice(-n);
|
||||
const pathPrefix = pathSegs.slice(0, n);
|
||||
if (baseSuffix.every((seg, i) => seg === pathPrefix[i])) {
|
||||
overlap = n;
|
||||
break;
|
||||
}
|
||||
}
|
||||
if (overlap > 0) {
|
||||
const remaining = pathSegs.slice(overlap);
|
||||
suffix = remaining.length > 0 ? `/${remaining.join("/")}` : "/";
|
||||
}
|
||||
}
|
||||
|
||||
if (suffix === "/") {
|
||||
return base;
|
||||
}
|
||||
|
||||
return `${base}${suffix}`;
|
||||
}
|
||||
|
||||
function pathFromRequest(req: Request): string {
|
||||
const raw = req.originalUrl || req.url || req.path;
|
||||
return raw.startsWith("/") ? raw : `/${raw}`;
|
||||
}
|
||||
|
||||
function bodyRequestsStream(req: Request): boolean {
|
||||
return req.body?.stream === true;
|
||||
}
|
||||
|
||||
function contentTypeIsSse(contentType: string): boolean {
|
||||
return contentType.includes("text/event-stream");
|
||||
}
|
||||
|
||||
function contentTypeIsAmazonEventStream(contentType: string): boolean {
|
||||
return contentType.includes("application/vnd.amazon.eventstream");
|
||||
}
|
||||
|
||||
function pathIncludes(req: Request, fragment: string): boolean {
|
||||
return pathFromRequest(req).includes(fragment);
|
||||
}
|
||||
|
||||
function isBodyOrSseStreaming(req: Request, contentType: string): boolean {
|
||||
return bodyRequestsStream(req) || contentTypeIsSse(contentType);
|
||||
}
|
||||
|
||||
function isGeminiStyleStreaming(req: Request, contentType: string): boolean {
|
||||
return (
|
||||
pathIncludes(req, "streamGenerateContent") ||
|
||||
pathIncludes(req, "alt=sse") ||
|
||||
contentTypeIsSse(contentType)
|
||||
);
|
||||
}
|
||||
|
||||
export const AI_CAPABILITY_DEFS: Record<AiCapability, AiCapabilityDefinition> =
|
||||
{
|
||||
openai_chat: {
|
||||
id: "openai_chat",
|
||||
protocolFamily: "openai",
|
||||
routes: [
|
||||
{ method: "POST", path: "/v1/chat/completions" },
|
||||
{ method: "POST", path: "/chat/completions" }
|
||||
],
|
||||
extractModel: bodyModel,
|
||||
resolveUpstreamUrl: (base, req) =>
|
||||
joinUpstreamUrl(base, pathFromRequest(req)),
|
||||
isStreaming: isBodyOrSseStreaming
|
||||
},
|
||||
openai_responses: {
|
||||
id: "openai_responses",
|
||||
protocolFamily: "openai",
|
||||
routes: [{ method: "POST", path: "/v1/responses" }],
|
||||
extractModel: bodyModel,
|
||||
resolveUpstreamUrl: (base, req) =>
|
||||
joinUpstreamUrl(base, pathFromRequest(req)),
|
||||
isStreaming: isBodyOrSseStreaming
|
||||
},
|
||||
anthropic_messages: {
|
||||
id: "anthropic_messages",
|
||||
protocolFamily: "anthropic",
|
||||
routes: [{ method: "POST", path: "/v1/messages" }],
|
||||
extractModel: bodyModel,
|
||||
resolveUpstreamUrl: (base, req) =>
|
||||
joinUpstreamUrl(base, pathFromRequest(req)),
|
||||
isStreaming: isBodyOrSseStreaming
|
||||
},
|
||||
v1_models: {
|
||||
id: "v1_models",
|
||||
protocolFamily: "anthropic",
|
||||
routes: [
|
||||
{ method: "GET", path: "/v1/models" },
|
||||
{ method: "GET", path: "/v1/models/:model" }
|
||||
],
|
||||
extractModel: paramModel,
|
||||
resolveUpstreamUrl: (base, req) =>
|
||||
joinUpstreamUrl(base, pathFromRequest(req)),
|
||||
// Model listings are answered from the gateway's own view of the
|
||||
// provider allow/block lists rather than proxied upstream, so
|
||||
// there is never a stream to detect.
|
||||
isStreaming: () => false
|
||||
},
|
||||
gemini_generate_content: {
|
||||
id: "gemini_generate_content",
|
||||
protocolFamily: "google",
|
||||
routes: [
|
||||
{
|
||||
method: "POST",
|
||||
path: "/v1beta/models/:model\\:generateContent"
|
||||
},
|
||||
{
|
||||
method: "POST",
|
||||
path: "/v1beta/models/:model\\:streamGenerateContent"
|
||||
}
|
||||
],
|
||||
extractModel: paramModel,
|
||||
resolveUpstreamUrl: (base, req) =>
|
||||
joinUpstreamUrl(base, pathFromRequest(req)),
|
||||
isStreaming: isGeminiStyleStreaming
|
||||
},
|
||||
google_generate_content: {
|
||||
id: "google_generate_content",
|
||||
protocolFamily: "google",
|
||||
routes: [
|
||||
{
|
||||
method: "POST",
|
||||
// Vertex publisher model generateContent
|
||||
path: "/v1/projects/:project/locations/:location/publishers/:publisher/models/:model\\:generateContent"
|
||||
},
|
||||
{
|
||||
method: "POST",
|
||||
path: "/v1/projects/:project/locations/:location/publishers/:publisher/models/:model\\:streamGenerateContent"
|
||||
}
|
||||
],
|
||||
extractModel: paramModel,
|
||||
resolveUpstreamUrl: (base, req) =>
|
||||
joinUpstreamUrl(base, pathFromRequest(req)),
|
||||
isStreaming: isGeminiStyleStreaming
|
||||
},
|
||||
google_raw_predict: {
|
||||
id: "google_raw_predict",
|
||||
protocolFamily: "google",
|
||||
routes: [
|
||||
{
|
||||
method: "POST",
|
||||
path: "/v1/projects/:project/locations/:location/publishers/:publisher/models/:model\\:rawPredict"
|
||||
},
|
||||
{
|
||||
method: "POST",
|
||||
path: "/v1/projects/:project/locations/:location/publishers/:publisher/models/:model\\:streamRawPredict"
|
||||
}
|
||||
],
|
||||
extractModel: paramModel,
|
||||
resolveUpstreamUrl: (base, req) =>
|
||||
joinUpstreamUrl(base, pathFromRequest(req)),
|
||||
isStreaming: (req, contentType) =>
|
||||
pathIncludes(req, "streamRawPredict") ||
|
||||
pathIncludes(req, "alt=sse") ||
|
||||
contentTypeIsSse(contentType)
|
||||
},
|
||||
bedrock_model_invoke: {
|
||||
id: "bedrock_model_invoke",
|
||||
protocolFamily: "bedrock",
|
||||
routes: [
|
||||
{ method: "POST", path: "/model/:model/invoke" },
|
||||
{
|
||||
method: "POST",
|
||||
path: "/model/:model/invoke-with-response-stream"
|
||||
}
|
||||
],
|
||||
extractModel: paramModel,
|
||||
resolveUpstreamUrl: (base, req) =>
|
||||
joinUpstreamUrl(base, pathFromRequest(req)),
|
||||
isStreaming: (req, contentType) =>
|
||||
pathIncludes(req, "invoke-with-response-stream") ||
|
||||
contentTypeIsAmazonEventStream(contentType) ||
|
||||
contentTypeIsSse(contentType)
|
||||
},
|
||||
bedrock_converse: {
|
||||
id: "bedrock_converse",
|
||||
protocolFamily: "bedrock",
|
||||
routes: [
|
||||
{ method: "POST", path: "/model/:model/converse" },
|
||||
{ method: "POST", path: "/model/:model/converse-stream" }
|
||||
],
|
||||
extractModel: paramModel,
|
||||
resolveUpstreamUrl: (base, req) =>
|
||||
joinUpstreamUrl(base, pathFromRequest(req)),
|
||||
isStreaming: (req, contentType) =>
|
||||
pathIncludes(req, "converse-stream") ||
|
||||
contentTypeIsAmazonEventStream(contentType) ||
|
||||
contentTypeIsSse(contentType)
|
||||
}
|
||||
};
|
||||
|
||||
export function isAiCapability(value: unknown): value is AiCapability {
|
||||
return (
|
||||
typeof value === "string" &&
|
||||
(AI_CAPABILITIES as readonly string[]).includes(value)
|
||||
);
|
||||
}
|
||||
|
||||
/**
|
||||
* Convert an Express-style route path from AI_CAPABILITY_DEFS into a RegExp.
|
||||
* Handles `:param` segments and escaped literal colons (`\:`).
|
||||
*/
|
||||
export function routePatternToRegExp(routePath: string): RegExp {
|
||||
let pattern = "";
|
||||
for (let i = 0; i < routePath.length; i++) {
|
||||
const ch = routePath[i];
|
||||
if (
|
||||
ch === "\\" &&
|
||||
i + 1 < routePath.length &&
|
||||
routePath[i + 1] === ":"
|
||||
) {
|
||||
pattern += ":";
|
||||
i++;
|
||||
continue;
|
||||
}
|
||||
if (ch === ":") {
|
||||
// Named param: consume until next / or end
|
||||
i++;
|
||||
while (
|
||||
i < routePath.length &&
|
||||
routePath[i] !== "/" &&
|
||||
!(routePath[i] === "\\" && routePath[i + 1] === ":")
|
||||
) {
|
||||
i++;
|
||||
}
|
||||
i--; // loop will ++
|
||||
pattern += "[^/]+";
|
||||
continue;
|
||||
}
|
||||
// Escape regex special chars
|
||||
if (/[.*+?^${}()|[\]\\]/.test(ch)) {
|
||||
pattern += "\\" + ch;
|
||||
} else {
|
||||
pattern += ch;
|
||||
}
|
||||
}
|
||||
return new RegExp(`^${pattern}$`);
|
||||
}
|
||||
|
||||
export function resolveAiCapabilityFromPath(path: string): AiCapability | null {
|
||||
const pathname = path.split("?")[0] || "/";
|
||||
const normalized = pathname.startsWith("/") ? pathname : `/${pathname}`;
|
||||
|
||||
for (const def of Object.values(AI_CAPABILITY_DEFS)) {
|
||||
for (const route of def.routes) {
|
||||
if (routePatternToRegExp(route.path).test(normalized)) {
|
||||
return def.id;
|
||||
}
|
||||
}
|
||||
}
|
||||
return null;
|
||||
}
|
||||
|
||||
export function parseCapabilities(raw: unknown): AiCapability[] {
|
||||
if (raw == null) {
|
||||
return [];
|
||||
}
|
||||
|
||||
let parsed: unknown = raw;
|
||||
if (typeof raw === "string") {
|
||||
const trimmed = raw.trim();
|
||||
if (!trimmed) {
|
||||
return [];
|
||||
}
|
||||
try {
|
||||
parsed = JSON.parse(trimmed);
|
||||
} catch {
|
||||
return [];
|
||||
}
|
||||
}
|
||||
|
||||
if (!Array.isArray(parsed)) {
|
||||
return [];
|
||||
}
|
||||
|
||||
const out: AiCapability[] = [];
|
||||
const seen = new Set<AiCapability>();
|
||||
for (const item of parsed) {
|
||||
if (isAiCapability(item) && !seen.has(item)) {
|
||||
seen.add(item);
|
||||
out.push(item);
|
||||
}
|
||||
}
|
||||
return out;
|
||||
}
|
||||
|
||||
export function serializeCapabilities(capabilities: AiCapability[]): string {
|
||||
return JSON.stringify(capabilities);
|
||||
}
|
||||
|
||||
export function providerHasCapability(
|
||||
capabilities: AiCapability[] | string | null | undefined,
|
||||
capability: AiCapability
|
||||
): boolean {
|
||||
const list =
|
||||
typeof capabilities === "string" || capabilities == null
|
||||
? parseCapabilities(capabilities)
|
||||
: capabilities;
|
||||
return list.includes(capability);
|
||||
}
|
||||
@@ -0,0 +1,139 @@
|
||||
import {
|
||||
AI_CAPABILITY_DEFS,
|
||||
type AiCapability,
|
||||
type AiProtocolFamily
|
||||
} from "@server/lib/aiCapabilities";
|
||||
import HttpCode from "@server/types/HttpCode";
|
||||
|
||||
export type ClientErrorResponse = {
|
||||
statusCode?: number;
|
||||
contentType?: string;
|
||||
body: string;
|
||||
};
|
||||
|
||||
export type AiCapabilityErrorKind =
|
||||
| "authentication"
|
||||
| "invalid_request"
|
||||
| "not_found"
|
||||
| "permission"
|
||||
| "rate_limit"
|
||||
| "internal";
|
||||
|
||||
const AUTH_MESSAGE = "Invalid API key provided.";
|
||||
|
||||
type KindFields = {
|
||||
openaiType: string;
|
||||
openaiCode: string | null;
|
||||
anthropicType: string;
|
||||
googleStatus: string;
|
||||
};
|
||||
|
||||
const KIND_FIELDS: Record<AiCapabilityErrorKind, KindFields> = {
|
||||
authentication: {
|
||||
openaiType: "authentication_error",
|
||||
openaiCode: "invalid_api_key",
|
||||
anthropicType: "authentication_error",
|
||||
googleStatus: "UNAUTHENTICATED"
|
||||
},
|
||||
invalid_request: {
|
||||
openaiType: "invalid_request_error",
|
||||
openaiCode: null,
|
||||
anthropicType: "invalid_request_error",
|
||||
googleStatus: "INVALID_ARGUMENT"
|
||||
},
|
||||
not_found: {
|
||||
openaiType: "invalid_request_error",
|
||||
openaiCode: null,
|
||||
anthropicType: "not_found_error",
|
||||
googleStatus: "NOT_FOUND"
|
||||
},
|
||||
permission: {
|
||||
openaiType: "invalid_request_error",
|
||||
openaiCode: null,
|
||||
anthropicType: "permission_error",
|
||||
googleStatus: "PERMISSION_DENIED"
|
||||
},
|
||||
rate_limit: {
|
||||
openaiType: "rate_limit_error",
|
||||
openaiCode: "rate_limit_exceeded",
|
||||
anthropicType: "rate_limit_error",
|
||||
googleStatus: "RESOURCE_EXHAUSTED"
|
||||
},
|
||||
internal: {
|
||||
openaiType: "api_error",
|
||||
openaiCode: null,
|
||||
anthropicType: "api_error",
|
||||
googleStatus: "INTERNAL"
|
||||
}
|
||||
};
|
||||
|
||||
function resolveProtocolFamily(
|
||||
capability: AiCapability | null
|
||||
): AiProtocolFamily {
|
||||
if (capability == null) {
|
||||
return "openai";
|
||||
}
|
||||
return AI_CAPABILITY_DEFS[capability].protocolFamily;
|
||||
}
|
||||
|
||||
/**
|
||||
* Build a protocol-native error body for the given capability.
|
||||
* Message stays contextual; only the envelope/machine fields follow the
|
||||
* capability's native API shape.
|
||||
*/
|
||||
export function buildAiCapabilityErrorBody(
|
||||
capability: AiCapability | null,
|
||||
kind: AiCapabilityErrorKind,
|
||||
message: string,
|
||||
httpStatus?: number
|
||||
): Record<string, unknown> {
|
||||
const family = resolveProtocolFamily(capability);
|
||||
const fields = KIND_FIELDS[kind];
|
||||
|
||||
switch (family) {
|
||||
case "openai":
|
||||
return {
|
||||
error: {
|
||||
message,
|
||||
type: fields.openaiType,
|
||||
param: null,
|
||||
code: fields.openaiCode
|
||||
}
|
||||
};
|
||||
case "anthropic":
|
||||
return {
|
||||
type: "error",
|
||||
error: {
|
||||
type: fields.anthropicType,
|
||||
message
|
||||
}
|
||||
};
|
||||
case "google":
|
||||
return {
|
||||
error: {
|
||||
code: httpStatus ?? HttpCode.BAD_REQUEST,
|
||||
message,
|
||||
status: fields.googleStatus
|
||||
}
|
||||
};
|
||||
case "bedrock":
|
||||
return { message };
|
||||
}
|
||||
}
|
||||
|
||||
export function buildInferenceAuthClientError(
|
||||
capability: AiCapability | null
|
||||
): ClientErrorResponse {
|
||||
return {
|
||||
statusCode: HttpCode.UNAUTHORIZED,
|
||||
contentType: "application/json",
|
||||
body: JSON.stringify(
|
||||
buildAiCapabilityErrorBody(
|
||||
capability,
|
||||
"authentication",
|
||||
AUTH_MESSAGE,
|
||||
HttpCode.UNAUTHORIZED
|
||||
)
|
||||
)
|
||||
};
|
||||
}
|
||||
@@ -0,0 +1,69 @@
|
||||
import { createHash } from "crypto";
|
||||
import config from "@server/lib/config";
|
||||
|
||||
export const AI_GATEWAY_TRUST_HEADER = "X-Pangolin-Ai-Gateway-Auth";
|
||||
|
||||
// Injected by the same Traefik trust middleware as AI_GATEWAY_TRUST_HEADER,
|
||||
// but its value differs per router (public inference resource vs. private
|
||||
// siteResource) so the gateway can tell which kind of resource a trusted
|
||||
// request arrived on without re-deriving it from resourceId/siteResourceId.
|
||||
export const AI_GATEWAY_RESOURCE_TYPE_HEADER =
|
||||
"X-Pangolin-Ai-Gateway-Resource-Type";
|
||||
|
||||
export type AiGatewayResourceType = "resource" | "site-resource";
|
||||
|
||||
// Opt-in (server.enable_ai_gateway_client_ip_header): carries the client IP
|
||||
// that Badger resolved at the Traefik hop, so it survives an intermediary
|
||||
// proxy between Traefik and the AI gateway that overwrites
|
||||
// X-Forwarded-For/X-Real-Ip instead of appending to them. Set by a
|
||||
// disableForwardAuth Badger middleware instance (see getTraefikConfig.ts)
|
||||
// on the site-resource inference router only, since that's the sole path
|
||||
// that resolves request identity from the client IP.
|
||||
export const AI_GATEWAY_CLIENT_IP_HEADER = "X-Pangolin-Client-Ip";
|
||||
|
||||
/**
|
||||
* Derive a Traefik-injected trust token from the server secret.
|
||||
* Traefik overwrites this header on inference routes so the AI gateway can
|
||||
* trust Badger-injected Remote-* identity without re-validating credentials.
|
||||
*/
|
||||
export function deriveAiGatewayTrustToken(secret: string): string {
|
||||
return createHash("sha256")
|
||||
.update(`ai-gateway-trust:${secret}`)
|
||||
.digest("hex");
|
||||
}
|
||||
|
||||
export function getAiGatewayTrustToken(): string {
|
||||
const secret = config.getRawConfig().server.secret;
|
||||
if (!secret) {
|
||||
throw new Error("Server secret is required for AI gateway trust token");
|
||||
}
|
||||
return deriveAiGatewayTrustToken(secret);
|
||||
}
|
||||
|
||||
export function isAiGatewayTrustHeaderValid(
|
||||
headers: Record<string, string | string[] | undefined> | undefined,
|
||||
expectedToken?: string
|
||||
): boolean {
|
||||
if (!headers) {
|
||||
return false;
|
||||
}
|
||||
const expected = expectedToken ?? getAiGatewayTrustToken();
|
||||
const raw =
|
||||
headers[AI_GATEWAY_TRUST_HEADER] ??
|
||||
headers[AI_GATEWAY_TRUST_HEADER.toLowerCase()];
|
||||
const value = Array.isArray(raw) ? raw[0] : raw;
|
||||
return typeof value === "string" && value === expected;
|
||||
}
|
||||
|
||||
export function getAiGatewayResourceType(
|
||||
headers: Record<string, string | string[] | undefined> | undefined
|
||||
): AiGatewayResourceType | null {
|
||||
if (!headers) {
|
||||
return null;
|
||||
}
|
||||
const raw =
|
||||
headers[AI_GATEWAY_RESOURCE_TYPE_HEADER] ??
|
||||
headers[AI_GATEWAY_RESOURCE_TYPE_HEADER.toLowerCase()];
|
||||
const value = Array.isArray(raw) ? raw[0] : raw;
|
||||
return value === "resource" || value === "site-resource" ? value : null;
|
||||
}
|
||||
@@ -0,0 +1,82 @@
|
||||
import http from "node:http";
|
||||
import https from "node:https";
|
||||
import { Readable } from "node:stream";
|
||||
|
||||
type UpstreamFetchInit = {
|
||||
method: string;
|
||||
headers: Record<string, string>;
|
||||
body?: string;
|
||||
skipTlsVerification?: boolean;
|
||||
signal?: AbortSignal;
|
||||
};
|
||||
|
||||
const insecureHttpsAgent = new https.Agent({
|
||||
rejectUnauthorized: false,
|
||||
keepAlive: true
|
||||
});
|
||||
|
||||
export function aiGatewayUpstreamFetch(
|
||||
url: string,
|
||||
init: UpstreamFetchInit
|
||||
): Promise<Response> {
|
||||
const parsed = new URL(url);
|
||||
const isHttps = parsed.protocol === "https:";
|
||||
const lib = isHttps ? https : http;
|
||||
const agent =
|
||||
isHttps && init.skipTlsVerification ? insecureHttpsAgent : undefined;
|
||||
|
||||
return new Promise((resolve, reject) => {
|
||||
if (init.signal?.aborted) {
|
||||
reject(init.signal.reason ?? new Error("Request aborted"));
|
||||
return;
|
||||
}
|
||||
|
||||
const req = lib.request(
|
||||
url,
|
||||
{
|
||||
method: init.method,
|
||||
headers: init.headers,
|
||||
agent
|
||||
},
|
||||
(res) => {
|
||||
const headers = new Headers();
|
||||
for (const [key, value] of Object.entries(res.headers)) {
|
||||
if (value === undefined) {
|
||||
continue;
|
||||
}
|
||||
if (Array.isArray(value)) {
|
||||
for (const entry of value) {
|
||||
headers.append(key, entry);
|
||||
}
|
||||
} else {
|
||||
headers.set(key, value);
|
||||
}
|
||||
}
|
||||
|
||||
const body = Readable.toWeb(res) as ReadableStream<Uint8Array>;
|
||||
resolve(
|
||||
new Response(body, {
|
||||
status: res.statusCode ?? 502,
|
||||
statusText: res.statusMessage,
|
||||
headers
|
||||
})
|
||||
);
|
||||
}
|
||||
);
|
||||
|
||||
req.on("error", reject);
|
||||
|
||||
if (init.signal) {
|
||||
const onAbort = () => req.destroy(init.signal!.reason);
|
||||
init.signal.addEventListener("abort", onAbort, { once: true });
|
||||
req.on("close", () =>
|
||||
init.signal!.removeEventListener("abort", onAbort)
|
||||
);
|
||||
}
|
||||
|
||||
if (init.body !== undefined) {
|
||||
req.write(init.body);
|
||||
}
|
||||
req.end();
|
||||
});
|
||||
}
|
||||
@@ -0,0 +1,721 @@
|
||||
import { and, eq, inArray } from "drizzle-orm";
|
||||
import {
|
||||
aiModels,
|
||||
aiProviders,
|
||||
db,
|
||||
resourceAiModels,
|
||||
resourceAiProviders,
|
||||
siteResourceAiModels,
|
||||
siteResourceAiProviders,
|
||||
type Transaction
|
||||
} from "@server/db";
|
||||
import { z } from "zod";
|
||||
|
||||
type DbOrTrx = Transaction | typeof db;
|
||||
|
||||
export const modelListTypeSchema = z.enum(["allow", "block"]);
|
||||
|
||||
export type ModelListType = z.infer<typeof modelListTypeSchema>;
|
||||
|
||||
export const accessModeSchema = z.enum(["inherit", "select"]);
|
||||
|
||||
export type AccessMode = z.infer<typeof accessModeSchema>;
|
||||
|
||||
export const resourceAiProviderAttachmentSchema = z.strictObject({
|
||||
providerId: z.number().int().positive(),
|
||||
accessMode: accessModeSchema.optional().default("inherit"),
|
||||
enabled: z.boolean().optional().default(true)
|
||||
});
|
||||
|
||||
export type ResourceAiProviderInput = z.infer<
|
||||
typeof resourceAiProviderAttachmentSchema
|
||||
>;
|
||||
|
||||
export type ResourceAiProviderAttachment = {
|
||||
providerId: number;
|
||||
accessMode: AccessMode;
|
||||
enabled: boolean;
|
||||
};
|
||||
|
||||
export const resourceAiModelEntrySchema = z.strictObject({
|
||||
modelId: z.number().int().positive(),
|
||||
listType: modelListTypeSchema
|
||||
});
|
||||
|
||||
export type ResourceAiModelEntry = z.infer<typeof resourceAiModelEntrySchema>;
|
||||
|
||||
export type InferenceFieldsError = {
|
||||
error: string;
|
||||
};
|
||||
|
||||
export function isInferenceFieldsError(
|
||||
value: { error: string } | object
|
||||
): value is InferenceFieldsError {
|
||||
return "error" in value;
|
||||
}
|
||||
|
||||
/**
|
||||
* Resolve which allow/block patterns apply for an attachment.
|
||||
* inherit → provider lists; select → resource-selected lists (replace).
|
||||
*/
|
||||
export function resolveEffectiveLists(input: {
|
||||
accessMode: AccessMode;
|
||||
providerAllows: string[];
|
||||
providerBlocks: string[];
|
||||
resourceAllows: string[];
|
||||
resourceBlocks: string[];
|
||||
}): { allows: string[]; blocks: string[] } {
|
||||
if (input.accessMode === "select") {
|
||||
return {
|
||||
allows: input.resourceAllows,
|
||||
blocks: input.resourceBlocks
|
||||
};
|
||||
}
|
||||
return {
|
||||
allows: input.providerAllows,
|
||||
blocks: input.providerBlocks
|
||||
};
|
||||
}
|
||||
|
||||
function normalizeAttachments(
|
||||
inputs: ResourceAiProviderInput[]
|
||||
): ResourceAiProviderAttachment[] {
|
||||
const byProviderId = new Map<
|
||||
number,
|
||||
{ accessMode: AccessMode; enabled: boolean }
|
||||
>();
|
||||
for (const input of inputs) {
|
||||
byProviderId.set(input.providerId, {
|
||||
accessMode: input.accessMode ?? "inherit",
|
||||
enabled: input.enabled ?? true
|
||||
});
|
||||
}
|
||||
return [...byProviderId.entries()].map(
|
||||
([providerId, { accessMode, enabled }]) => ({
|
||||
providerId,
|
||||
accessMode,
|
||||
enabled
|
||||
})
|
||||
);
|
||||
}
|
||||
|
||||
/**
|
||||
* Validate provider attachments for an org.
|
||||
*/
|
||||
export async function resolveProviderAttachments(input: {
|
||||
orgId: string;
|
||||
attachments: ResourceAiProviderInput[];
|
||||
requireAtLeastOne: boolean;
|
||||
}): Promise<ResourceAiProviderAttachment[] | InferenceFieldsError> {
|
||||
const attachments = normalizeAttachments(input.attachments);
|
||||
|
||||
if (input.requireAtLeastOne && attachments.length === 0) {
|
||||
return {
|
||||
error: "At least one AI provider is required for inference-mode resources"
|
||||
};
|
||||
}
|
||||
|
||||
if (attachments.length === 0) {
|
||||
return [];
|
||||
}
|
||||
|
||||
const providerIds = attachments.map((a) => a.providerId);
|
||||
const providers = await db
|
||||
.select({
|
||||
providerId: aiProviders.providerId,
|
||||
orgId: aiProviders.orgId,
|
||||
enabled: aiProviders.enabled
|
||||
})
|
||||
.from(aiProviders)
|
||||
.where(
|
||||
and(
|
||||
inArray(aiProviders.providerId, providerIds),
|
||||
eq(aiProviders.orgId, input.orgId)
|
||||
)
|
||||
);
|
||||
|
||||
if (providers.length !== providerIds.length) {
|
||||
return {
|
||||
error: "One or more AI providers were not found in this organization"
|
||||
};
|
||||
}
|
||||
|
||||
const disabled = providers.find((p) => !p.enabled);
|
||||
if (disabled) {
|
||||
return {
|
||||
error: `AI provider with ID ${disabled.providerId} is disabled`
|
||||
};
|
||||
}
|
||||
|
||||
return attachments;
|
||||
}
|
||||
|
||||
export async function assertInferenceModeAllowsProviderFields(input: {
|
||||
mode: string;
|
||||
hasProviderAttachments: boolean;
|
||||
}): Promise<InferenceFieldsError | null> {
|
||||
if (input.mode === "inference") {
|
||||
return null;
|
||||
}
|
||||
if (input.hasProviderAttachments) {
|
||||
return {
|
||||
error: "AI providers can only be attached to inference-mode resources"
|
||||
};
|
||||
}
|
||||
return null;
|
||||
}
|
||||
|
||||
/**
|
||||
* Attach providers to a resource. Inherit attachments use the provider lists
|
||||
* as-is (resource model rows for those providers are pruned). Select
|
||||
* attachments keep resource-selected allow/block subsets.
|
||||
*/
|
||||
export async function setPublicResourceAiProviders(
|
||||
resourceId: number,
|
||||
attachments: ResourceAiProviderAttachment[],
|
||||
trx: DbOrTrx = db
|
||||
): Promise<void> {
|
||||
await trx
|
||||
.delete(resourceAiProviders)
|
||||
.where(eq(resourceAiProviders.resourceId, resourceId));
|
||||
|
||||
if (attachments.length > 0) {
|
||||
await trx.insert(resourceAiProviders).values(
|
||||
attachments.map((a) => ({
|
||||
resourceId,
|
||||
providerId: a.providerId,
|
||||
accessMode: a.accessMode,
|
||||
enabled: a.enabled
|
||||
}))
|
||||
);
|
||||
}
|
||||
|
||||
await prunePublicResourceModelsToSelectProviders(
|
||||
resourceId,
|
||||
attachments,
|
||||
trx
|
||||
);
|
||||
}
|
||||
|
||||
export async function setSiteResourceAiProviders(
|
||||
siteResourceId: number,
|
||||
attachments: ResourceAiProviderAttachment[],
|
||||
trx: DbOrTrx = db
|
||||
): Promise<void> {
|
||||
await trx
|
||||
.delete(siteResourceAiProviders)
|
||||
.where(eq(siteResourceAiProviders.siteResourceId, siteResourceId));
|
||||
|
||||
if (attachments.length > 0) {
|
||||
await trx.insert(siteResourceAiProviders).values(
|
||||
attachments.map((a) => ({
|
||||
siteResourceId,
|
||||
providerId: a.providerId,
|
||||
accessMode: a.accessMode,
|
||||
enabled: a.enabled
|
||||
}))
|
||||
);
|
||||
}
|
||||
|
||||
await pruneSiteResourceModelsToSelectProviders(
|
||||
siteResourceId,
|
||||
attachments,
|
||||
trx
|
||||
);
|
||||
}
|
||||
|
||||
/**
|
||||
* Keep resource model rows only for providers in select mode.
|
||||
*/
|
||||
async function prunePublicResourceModelsToSelectProviders(
|
||||
resourceId: number,
|
||||
attachments: ResourceAiProviderAttachment[],
|
||||
trx: DbOrTrx
|
||||
): Promise<void> {
|
||||
const selectProviderIds = attachments
|
||||
.filter((a) => a.accessMode === "select")
|
||||
.map((a) => a.providerId);
|
||||
|
||||
if (selectProviderIds.length === 0) {
|
||||
await trx
|
||||
.delete(resourceAiModels)
|
||||
.where(eq(resourceAiModels.resourceId, resourceId));
|
||||
return;
|
||||
}
|
||||
|
||||
const existing = await trx
|
||||
.select({
|
||||
modelId: resourceAiModels.modelId,
|
||||
providerId: aiModels.providerId
|
||||
})
|
||||
.from(resourceAiModels)
|
||||
.innerJoin(aiModels, eq(resourceAiModels.modelId, aiModels.modelId))
|
||||
.where(eq(resourceAiModels.resourceId, resourceId));
|
||||
|
||||
const allowed = new Set(selectProviderIds);
|
||||
const toRemove = existing
|
||||
.filter((row) => !allowed.has(row.providerId))
|
||||
.map((row) => row.modelId);
|
||||
|
||||
if (toRemove.length > 0) {
|
||||
await trx
|
||||
.delete(resourceAiModels)
|
||||
.where(
|
||||
and(
|
||||
eq(resourceAiModels.resourceId, resourceId),
|
||||
inArray(resourceAiModels.modelId, toRemove)
|
||||
)
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
async function pruneSiteResourceModelsToSelectProviders(
|
||||
siteResourceId: number,
|
||||
attachments: ResourceAiProviderAttachment[],
|
||||
trx: DbOrTrx
|
||||
): Promise<void> {
|
||||
const selectProviderIds = attachments
|
||||
.filter((a) => a.accessMode === "select")
|
||||
.map((a) => a.providerId);
|
||||
|
||||
if (selectProviderIds.length === 0) {
|
||||
await trx
|
||||
.delete(siteResourceAiModels)
|
||||
.where(eq(siteResourceAiModels.siteResourceId, siteResourceId));
|
||||
return;
|
||||
}
|
||||
|
||||
const existing = await trx
|
||||
.select({
|
||||
modelId: siteResourceAiModels.modelId,
|
||||
providerId: aiModels.providerId
|
||||
})
|
||||
.from(siteResourceAiModels)
|
||||
.innerJoin(aiModels, eq(siteResourceAiModels.modelId, aiModels.modelId))
|
||||
.where(eq(siteResourceAiModels.siteResourceId, siteResourceId));
|
||||
|
||||
const allowed = new Set(selectProviderIds);
|
||||
const toRemove = existing
|
||||
.filter((row) => !allowed.has(row.providerId))
|
||||
.map((row) => row.modelId);
|
||||
|
||||
if (toRemove.length > 0) {
|
||||
await trx
|
||||
.delete(siteResourceAiModels)
|
||||
.where(
|
||||
and(
|
||||
eq(siteResourceAiModels.siteResourceId, siteResourceId),
|
||||
inArray(siteResourceAiModels.modelId, toRemove)
|
||||
)
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
export async function clearPublicResourceAiConfig(
|
||||
resourceId: number,
|
||||
trx: DbOrTrx = db
|
||||
): Promise<void> {
|
||||
await trx
|
||||
.delete(resourceAiModels)
|
||||
.where(eq(resourceAiModels.resourceId, resourceId));
|
||||
await trx
|
||||
.delete(resourceAiProviders)
|
||||
.where(eq(resourceAiProviders.resourceId, resourceId));
|
||||
}
|
||||
|
||||
export async function clearSiteResourceAiConfig(
|
||||
siteResourceId: number,
|
||||
trx: DbOrTrx = db
|
||||
): Promise<void> {
|
||||
await trx
|
||||
.delete(siteResourceAiModels)
|
||||
.where(eq(siteResourceAiModels.siteResourceId, siteResourceId));
|
||||
await trx
|
||||
.delete(siteResourceAiProviders)
|
||||
.where(eq(siteResourceAiProviders.siteResourceId, siteResourceId));
|
||||
}
|
||||
|
||||
export async function listPublicResourceAiProviders(resourceId: number) {
|
||||
return db
|
||||
.select({
|
||||
providerId: resourceAiProviders.providerId,
|
||||
niceId: aiProviders.niceId,
|
||||
name: aiProviders.name,
|
||||
type: aiProviders.type,
|
||||
enabled: resourceAiProviders.enabled,
|
||||
providerEnabled: aiProviders.enabled,
|
||||
accessMode: resourceAiProviders.accessMode
|
||||
})
|
||||
.from(resourceAiProviders)
|
||||
.innerJoin(
|
||||
aiProviders,
|
||||
eq(resourceAiProviders.providerId, aiProviders.providerId)
|
||||
)
|
||||
.where(eq(resourceAiProviders.resourceId, resourceId));
|
||||
}
|
||||
|
||||
export async function listSiteResourceAiProviders(siteResourceId: number) {
|
||||
return db
|
||||
.select({
|
||||
providerId: siteResourceAiProviders.providerId,
|
||||
niceId: aiProviders.niceId,
|
||||
name: aiProviders.name,
|
||||
type: aiProviders.type,
|
||||
enabled: siteResourceAiProviders.enabled,
|
||||
providerEnabled: aiProviders.enabled,
|
||||
accessMode: siteResourceAiProviders.accessMode
|
||||
})
|
||||
.from(siteResourceAiProviders)
|
||||
.innerJoin(
|
||||
aiProviders,
|
||||
eq(siteResourceAiProviders.providerId, aiProviders.providerId)
|
||||
)
|
||||
.where(eq(siteResourceAiProviders.siteResourceId, siteResourceId));
|
||||
}
|
||||
|
||||
export type EffectiveAllowModel = {
|
||||
modelId: number;
|
||||
modelKey: string;
|
||||
name: string;
|
||||
providerId: number;
|
||||
providerName: string;
|
||||
};
|
||||
|
||||
export async function listEffectiveAllowModels(options: {
|
||||
resourceId?: number;
|
||||
siteResourceId?: number;
|
||||
}): Promise<EffectiveAllowModel[]> {
|
||||
if (
|
||||
options.resourceId === undefined &&
|
||||
options.siteResourceId === undefined
|
||||
) {
|
||||
return [];
|
||||
}
|
||||
|
||||
const attachments =
|
||||
options.resourceId !== undefined
|
||||
? await listPublicResourceAiProviders(options.resourceId)
|
||||
: await listSiteResourceAiProviders(options.siteResourceId!);
|
||||
|
||||
const activeAttachments = attachments.filter(
|
||||
(a) => a.enabled && a.providerEnabled
|
||||
);
|
||||
if (activeAttachments.length === 0) {
|
||||
return [];
|
||||
}
|
||||
|
||||
const inheritProviderIds = activeAttachments
|
||||
.filter((a) => a.accessMode === "inherit")
|
||||
.map((a) => a.providerId);
|
||||
const selectProviderIds = activeAttachments
|
||||
.filter((a) => a.accessMode === "select")
|
||||
.map((a) => a.providerId);
|
||||
|
||||
const providerNameById = new Map(
|
||||
activeAttachments.map((a) => [a.providerId, a.name] as const)
|
||||
);
|
||||
|
||||
const models: EffectiveAllowModel[] = [];
|
||||
|
||||
if (inheritProviderIds.length > 0) {
|
||||
const rows = await db
|
||||
.select({
|
||||
modelId: aiModels.modelId,
|
||||
modelKey: aiModels.modelKey,
|
||||
name: aiModels.name,
|
||||
providerId: aiModels.providerId
|
||||
})
|
||||
.from(aiModels)
|
||||
.where(
|
||||
and(
|
||||
inArray(aiModels.providerId, inheritProviderIds),
|
||||
eq(aiModels.enabled, true),
|
||||
eq(aiModels.listType, "allow")
|
||||
)
|
||||
);
|
||||
for (const row of rows) {
|
||||
models.push({
|
||||
...row,
|
||||
providerName: providerNameById.get(row.providerId) ?? ""
|
||||
});
|
||||
}
|
||||
}
|
||||
|
||||
if (selectProviderIds.length > 0) {
|
||||
if (options.resourceId !== undefined) {
|
||||
const rows = await db
|
||||
.select({
|
||||
modelId: aiModels.modelId,
|
||||
modelKey: aiModels.modelKey,
|
||||
name: aiModels.name,
|
||||
providerId: aiModels.providerId
|
||||
})
|
||||
.from(resourceAiModels)
|
||||
.innerJoin(
|
||||
aiModels,
|
||||
eq(resourceAiModels.modelId, aiModels.modelId)
|
||||
)
|
||||
.where(
|
||||
and(
|
||||
eq(resourceAiModels.resourceId, options.resourceId),
|
||||
inArray(aiModels.providerId, selectProviderIds),
|
||||
eq(resourceAiModels.listType, "allow"),
|
||||
eq(aiModels.enabled, true)
|
||||
)
|
||||
);
|
||||
for (const row of rows) {
|
||||
models.push({
|
||||
...row,
|
||||
providerName: providerNameById.get(row.providerId) ?? ""
|
||||
});
|
||||
}
|
||||
} else if (options.siteResourceId !== undefined) {
|
||||
const rows = await db
|
||||
.select({
|
||||
modelId: aiModels.modelId,
|
||||
modelKey: aiModels.modelKey,
|
||||
name: aiModels.name,
|
||||
providerId: aiModels.providerId
|
||||
})
|
||||
.from(siteResourceAiModels)
|
||||
.innerJoin(
|
||||
aiModels,
|
||||
eq(siteResourceAiModels.modelId, aiModels.modelId)
|
||||
)
|
||||
.where(
|
||||
and(
|
||||
eq(
|
||||
siteResourceAiModels.siteResourceId,
|
||||
options.siteResourceId
|
||||
),
|
||||
inArray(aiModels.providerId, selectProviderIds),
|
||||
eq(siteResourceAiModels.listType, "allow"),
|
||||
eq(aiModels.enabled, true)
|
||||
)
|
||||
);
|
||||
for (const row of rows) {
|
||||
models.push({
|
||||
...row,
|
||||
providerName: providerNameById.get(row.providerId) ?? ""
|
||||
});
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
models.sort((a, b) => {
|
||||
const byProvider = a.providerName.localeCompare(
|
||||
b.providerName,
|
||||
undefined,
|
||||
{
|
||||
sensitivity: "base"
|
||||
}
|
||||
);
|
||||
if (byProvider !== 0) {
|
||||
return byProvider;
|
||||
}
|
||||
return a.name.localeCompare(b.name, undefined, { sensitivity: "base" });
|
||||
});
|
||||
|
||||
return models;
|
||||
}
|
||||
|
||||
/**
|
||||
* Model list APIs require an inference resource with at least one select-mode
|
||||
* attached provider.
|
||||
*/
|
||||
export async function assertPublicModelListApiEligible(resource: {
|
||||
resourceId: number;
|
||||
mode: string;
|
||||
}): Promise<string | null> {
|
||||
if (resource.mode !== "inference") {
|
||||
return "AI model lists are only supported on inference-mode resources";
|
||||
}
|
||||
|
||||
const [row] = await db
|
||||
.select({ providerId: resourceAiProviders.providerId })
|
||||
.from(resourceAiProviders)
|
||||
.where(
|
||||
and(
|
||||
eq(resourceAiProviders.resourceId, resource.resourceId),
|
||||
eq(resourceAiProviders.accessMode, "select")
|
||||
)
|
||||
)
|
||||
.limit(1);
|
||||
|
||||
if (!row) {
|
||||
return "Set at least one attached AI provider to select mode before managing model lists";
|
||||
}
|
||||
return null;
|
||||
}
|
||||
|
||||
export async function assertSiteModelListApiEligible(siteResource: {
|
||||
siteResourceId: number;
|
||||
mode: string;
|
||||
}): Promise<string | null> {
|
||||
if (siteResource.mode !== "inference") {
|
||||
return "AI model lists are only supported on inference-mode resources";
|
||||
}
|
||||
|
||||
const [row] = await db
|
||||
.select({ providerId: siteResourceAiProviders.providerId })
|
||||
.from(siteResourceAiProviders)
|
||||
.where(
|
||||
and(
|
||||
eq(
|
||||
siteResourceAiProviders.siteResourceId,
|
||||
siteResource.siteResourceId
|
||||
),
|
||||
eq(siteResourceAiProviders.accessMode, "select")
|
||||
)
|
||||
)
|
||||
.limit(1);
|
||||
|
||||
if (!row) {
|
||||
return "Set at least one attached AI provider to select mode before managing model lists";
|
||||
}
|
||||
return null;
|
||||
}
|
||||
|
||||
/**
|
||||
* Resource model entries must belong to select-mode attached providers, and
|
||||
* listType must match the provider catalog entry (allow→allow, block→block).
|
||||
*/
|
||||
export async function assertPublicResourceModelEntriesValid(input: {
|
||||
orgId: string;
|
||||
resourceId: number;
|
||||
models: ResourceAiModelEntry[];
|
||||
}): Promise<string | null> {
|
||||
const uniqueModels = dedupeModelEntries(input.models);
|
||||
if (uniqueModels.length === 0) {
|
||||
return null;
|
||||
}
|
||||
|
||||
const attachments = await db
|
||||
.select({
|
||||
providerId: resourceAiProviders.providerId,
|
||||
accessMode: resourceAiProviders.accessMode,
|
||||
enabled: resourceAiProviders.enabled
|
||||
})
|
||||
.from(resourceAiProviders)
|
||||
.innerJoin(
|
||||
aiProviders,
|
||||
eq(resourceAiProviders.providerId, aiProviders.providerId)
|
||||
)
|
||||
.where(
|
||||
and(
|
||||
eq(resourceAiProviders.resourceId, input.resourceId),
|
||||
eq(aiProviders.orgId, input.orgId)
|
||||
)
|
||||
);
|
||||
|
||||
return assertModelEntriesValid({
|
||||
orgId: input.orgId,
|
||||
modelEntries: uniqueModels,
|
||||
attachments,
|
||||
resourceLabel: "resource"
|
||||
});
|
||||
}
|
||||
|
||||
export async function assertSiteResourceModelEntriesValid(input: {
|
||||
orgId: string;
|
||||
siteResourceId: number;
|
||||
models: ResourceAiModelEntry[];
|
||||
}): Promise<string | null> {
|
||||
const uniqueModels = dedupeModelEntries(input.models);
|
||||
if (uniqueModels.length === 0) {
|
||||
return null;
|
||||
}
|
||||
|
||||
const attachments = await db
|
||||
.select({
|
||||
providerId: siteResourceAiProviders.providerId,
|
||||
accessMode: siteResourceAiProviders.accessMode,
|
||||
enabled: siteResourceAiProviders.enabled
|
||||
})
|
||||
.from(siteResourceAiProviders)
|
||||
.innerJoin(
|
||||
aiProviders,
|
||||
eq(siteResourceAiProviders.providerId, aiProviders.providerId)
|
||||
)
|
||||
.where(
|
||||
and(
|
||||
eq(
|
||||
siteResourceAiProviders.siteResourceId,
|
||||
input.siteResourceId
|
||||
),
|
||||
eq(aiProviders.orgId, input.orgId)
|
||||
)
|
||||
);
|
||||
|
||||
return assertModelEntriesValid({
|
||||
orgId: input.orgId,
|
||||
modelEntries: uniqueModels,
|
||||
attachments,
|
||||
resourceLabel: "site resource"
|
||||
});
|
||||
}
|
||||
|
||||
function dedupeModelEntries(
|
||||
models: ResourceAiModelEntry[]
|
||||
): ResourceAiModelEntry[] {
|
||||
const byModelId = new Map(
|
||||
models.map((m) => [m.modelId, m.listType] as const)
|
||||
);
|
||||
return [...byModelId.entries()].map(([modelId, listType]) => ({
|
||||
modelId,
|
||||
listType
|
||||
}));
|
||||
}
|
||||
|
||||
async function assertModelEntriesValid(input: {
|
||||
orgId: string;
|
||||
modelEntries: ResourceAiModelEntry[];
|
||||
attachments: ResourceAiProviderAttachment[];
|
||||
resourceLabel: string;
|
||||
}): Promise<string | null> {
|
||||
const selectProviderIds = input.attachments
|
||||
.filter((a) => a.accessMode === "select")
|
||||
.map((a) => a.providerId);
|
||||
|
||||
if (selectProviderIds.length === 0) {
|
||||
return "Set at least one attached AI provider to select mode before managing model lists";
|
||||
}
|
||||
|
||||
const modelIds = input.modelEntries.map((m) => m.modelId);
|
||||
const catalogRows = await db
|
||||
.select({
|
||||
modelId: aiModels.modelId,
|
||||
listType: aiModels.listType,
|
||||
providerId: aiModels.providerId,
|
||||
enabled: aiModels.enabled
|
||||
})
|
||||
.from(aiModels)
|
||||
.innerJoin(aiProviders, eq(aiModels.providerId, aiProviders.providerId))
|
||||
.where(
|
||||
and(
|
||||
inArray(aiModels.modelId, modelIds),
|
||||
inArray(aiModels.providerId, selectProviderIds),
|
||||
eq(aiProviders.orgId, input.orgId)
|
||||
)
|
||||
);
|
||||
|
||||
if (catalogRows.length !== modelIds.length) {
|
||||
return `One or more model IDs do not exist or do not belong to a select-mode provider on this ${input.resourceLabel}`;
|
||||
}
|
||||
|
||||
const catalogById = new Map(catalogRows.map((row) => [row.modelId, row]));
|
||||
for (const entry of input.modelEntries) {
|
||||
const catalog = catalogById.get(entry.modelId);
|
||||
if (!catalog) {
|
||||
return `One or more model IDs do not exist or do not belong to a select-mode provider on this ${input.resourceLabel}`;
|
||||
}
|
||||
if (catalog.listType !== entry.listType) {
|
||||
return `Model ${entry.modelId} must use listType "${catalog.listType}" to match the provider catalog entry`;
|
||||
}
|
||||
if (!catalog.enabled) {
|
||||
return `Model ${entry.modelId} is disabled on its provider`;
|
||||
}
|
||||
}
|
||||
|
||||
return null;
|
||||
}
|
||||
@@ -0,0 +1,541 @@
|
||||
import type { AiCapability } from "@server/lib/aiCapabilities";
|
||||
import { sseDataFrames, tryParseJson } from "@server/lib/aiUsageExtraction";
|
||||
import logger from "@server/logger";
|
||||
|
||||
// Uniform, capability-agnostic representation of a chat message, used so
|
||||
// the AI session log can be searched/displayed the same way regardless of
|
||||
// which provider/capability produced it. Content is flattened to plain text
|
||||
// - non-text parts (images, tool calls/results) are rendered as readable
|
||||
// placeholders rather than preserved as structured data, which is enough for
|
||||
// a transcript-style replay view without a per-capability renderer.
|
||||
export type NormalizedRole = "system" | "user" | "assistant" | "tool";
|
||||
|
||||
export type NormalizedAiMessage = {
|
||||
role: NormalizedRole;
|
||||
content: string;
|
||||
};
|
||||
|
||||
function normalizeRole(role: unknown): NormalizedRole {
|
||||
if (
|
||||
role === "system" ||
|
||||
role === "user" ||
|
||||
role === "assistant" ||
|
||||
role === "tool"
|
||||
) {
|
||||
return role;
|
||||
}
|
||||
if (role === "model") return "assistant"; // Gemini
|
||||
if (role === "function") return "tool"; // OpenAI legacy function role
|
||||
return "user";
|
||||
}
|
||||
|
||||
function safeJsonStringify(value: unknown): string {
|
||||
try {
|
||||
return JSON.stringify(value ?? {});
|
||||
} catch {
|
||||
return "";
|
||||
}
|
||||
}
|
||||
|
||||
/**
|
||||
* Flattens one message "part"/"block" (OpenAI content parts, Anthropic
|
||||
* content blocks, Gemini parts, Bedrock converse content blocks - they all
|
||||
* follow the same rough shape) into readable text.
|
||||
*/
|
||||
function flattenContentPart(part: unknown): string {
|
||||
if (typeof part === "string") return part;
|
||||
if (part == null || typeof part !== "object") return "";
|
||||
const p = part as Record<string, unknown>;
|
||||
|
||||
if (typeof p.text === "string") return p.text;
|
||||
|
||||
if (
|
||||
p.type === "image_url" ||
|
||||
p.type === "image" ||
|
||||
p.type === "input_image" ||
|
||||
p.type === "output_image" ||
|
||||
"inlineData" in p
|
||||
) {
|
||||
return "[image]";
|
||||
}
|
||||
|
||||
// Anthropic-style tool_use / tool_result blocks
|
||||
if (p.type === "tool_use") {
|
||||
const name = typeof p.name === "string" ? p.name : "tool";
|
||||
return `[tool_call: ${name}(${safeJsonStringify(p.input)})]`;
|
||||
}
|
||||
if (p.type === "tool_result") {
|
||||
const content = p.content;
|
||||
const text =
|
||||
typeof content === "string"
|
||||
? content
|
||||
: Array.isArray(content)
|
||||
? flattenContentParts(content)
|
||||
: "";
|
||||
return `[tool_result: ${text}]`;
|
||||
}
|
||||
|
||||
// Gemini-style functionCall / functionResponse parts
|
||||
if (p.functionCall && typeof p.functionCall === "object") {
|
||||
const fc = p.functionCall as Record<string, unknown>;
|
||||
return `[tool_call: ${fc.name}(${safeJsonStringify(fc.args)})]`;
|
||||
}
|
||||
if (p.functionResponse && typeof p.functionResponse === "object") {
|
||||
const fr = p.functionResponse as Record<string, unknown>;
|
||||
return `[tool_result: ${fr.name}(${safeJsonStringify(fr.response)})]`;
|
||||
}
|
||||
|
||||
// Bedrock converse-style toolUse / toolResult content blocks
|
||||
if (p.toolUse && typeof p.toolUse === "object") {
|
||||
const tu = p.toolUse as Record<string, unknown>;
|
||||
return `[tool_call: ${tu.name}(${safeJsonStringify(tu.input)})]`;
|
||||
}
|
||||
if (p.toolResult && typeof p.toolResult === "object") {
|
||||
const tr = p.toolResult as Record<string, unknown>;
|
||||
const content = tr.content;
|
||||
const text = Array.isArray(content) ? flattenContentParts(content) : "";
|
||||
return `[tool_result: ${text}]`;
|
||||
}
|
||||
|
||||
return "";
|
||||
}
|
||||
|
||||
function flattenContentParts(parts: unknown[]): string {
|
||||
return parts.map(flattenContentPart).join("");
|
||||
}
|
||||
|
||||
function flattenContent(content: unknown): string {
|
||||
if (typeof content === "string") return content;
|
||||
if (Array.isArray(content)) return flattenContentParts(content);
|
||||
return "";
|
||||
}
|
||||
|
||||
/**
|
||||
* Best-effort scan for every `"text":"..."` JSON string value in raw text,
|
||||
* concatenated in order. Fallback for streaming formats we can't fully parse
|
||||
* as JSON/SSE (Gemini's array-JSON stream, Bedrock's binary event-stream
|
||||
* framing) - same spirit as aiUsageExtraction's scanNumericFields.
|
||||
*/
|
||||
function scanTextFragments(text: string): string {
|
||||
const out: string[] = [];
|
||||
const re = /"text"\s*:\s*"((?:[^"\\]|\\.)*)"/g;
|
||||
let match: RegExpExecArray | null;
|
||||
while ((match = re.exec(text)) !== null) {
|
||||
try {
|
||||
out.push(JSON.parse(`"${match[1]}"`));
|
||||
} catch {
|
||||
out.push(match[1]);
|
||||
}
|
||||
}
|
||||
return out.join("");
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// Request (input) normalizers - operate on the already-parsed outbound body.
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
function normalizeOpenAiChatRequest(body: any): NormalizedAiMessage[] {
|
||||
const messages = Array.isArray(body?.messages) ? body.messages : [];
|
||||
return messages.map((m: any) => ({
|
||||
role: normalizeRole(m?.role),
|
||||
content: flattenContent(m?.content)
|
||||
}));
|
||||
}
|
||||
|
||||
function normalizeOpenAiResponsesRequest(body: any): NormalizedAiMessage[] {
|
||||
const out: NormalizedAiMessage[] = [];
|
||||
if (typeof body?.instructions === "string" && body.instructions) {
|
||||
out.push({ role: "system", content: body.instructions });
|
||||
}
|
||||
const input = body?.input;
|
||||
if (typeof input === "string") {
|
||||
out.push({ role: "user", content: input });
|
||||
} else if (Array.isArray(input)) {
|
||||
for (const item of input) {
|
||||
if (item?.role) {
|
||||
out.push({
|
||||
role: normalizeRole(item.role),
|
||||
content: flattenContent(item.content)
|
||||
});
|
||||
} else if (typeof item?.type === "string") {
|
||||
out.push({ role: "tool", content: `[${item.type}]` });
|
||||
}
|
||||
}
|
||||
}
|
||||
return out;
|
||||
}
|
||||
|
||||
function normalizeAnthropicRequest(body: any): NormalizedAiMessage[] {
|
||||
const out: NormalizedAiMessage[] = [];
|
||||
if (body?.system) {
|
||||
const sys = flattenContent(body.system);
|
||||
if (sys) out.push({ role: "system", content: sys });
|
||||
}
|
||||
const messages = Array.isArray(body?.messages) ? body.messages : [];
|
||||
for (const m of messages) {
|
||||
out.push({
|
||||
role: normalizeRole(m?.role),
|
||||
content: flattenContent(m?.content)
|
||||
});
|
||||
}
|
||||
return out;
|
||||
}
|
||||
|
||||
function normalizeGeminiRequest(body: any): NormalizedAiMessage[] {
|
||||
const out: NormalizedAiMessage[] = [];
|
||||
const sysParts = body?.systemInstruction?.parts;
|
||||
if (Array.isArray(sysParts)) {
|
||||
const text = flattenContentParts(sysParts);
|
||||
if (text) out.push({ role: "system", content: text });
|
||||
}
|
||||
const contents = Array.isArray(body?.contents) ? body.contents : [];
|
||||
for (const c of contents) {
|
||||
out.push({
|
||||
role: normalizeRole(c?.role),
|
||||
content: Array.isArray(c?.parts) ? flattenContentParts(c.parts) : ""
|
||||
});
|
||||
}
|
||||
return out;
|
||||
}
|
||||
|
||||
function normalizeBedrockConverseRequest(body: any): NormalizedAiMessage[] {
|
||||
const out: NormalizedAiMessage[] = [];
|
||||
if (Array.isArray(body?.system)) {
|
||||
const text = flattenContentParts(body.system);
|
||||
if (text) out.push({ role: "system", content: text });
|
||||
}
|
||||
const messages = Array.isArray(body?.messages) ? body.messages : [];
|
||||
for (const m of messages) {
|
||||
out.push({
|
||||
role: normalizeRole(m?.role),
|
||||
content: Array.isArray(m?.content)
|
||||
? flattenContentParts(m.content)
|
||||
: ""
|
||||
});
|
||||
}
|
||||
return out;
|
||||
}
|
||||
|
||||
/**
|
||||
* bedrock_model_invoke and google_raw_predict are passthroughs - the body
|
||||
* shape depends entirely on the underlying model, not the capability. Try
|
||||
* the two shapes we're most likely to see (Anthropic Claude, then plain
|
||||
* OpenAI-style) and give up otherwise, same fallback spirit
|
||||
* aiUsageExtraction.ts uses for these two capabilities' usage extraction.
|
||||
*/
|
||||
function normalizeBestEffortRequest(body: any): NormalizedAiMessage[] | null {
|
||||
if (!Array.isArray(body?.messages)) return null;
|
||||
const looksAnthropicShaped = body.messages.some((m: any) =>
|
||||
Array.isArray(m?.content)
|
||||
);
|
||||
return looksAnthropicShaped
|
||||
? normalizeAnthropicRequest(body)
|
||||
: normalizeOpenAiChatRequest(body);
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// Response (output) normalizers - operate on the raw response text, which
|
||||
// may be a single JSON document (non-streaming) or provider-framed streaming
|
||||
// text (SSE `data:` frames, a JSON-array stream, or binary event-stream
|
||||
// framing with JSON payloads embedded in it).
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
function normalizeOpenAiChatResponse(
|
||||
text: string,
|
||||
isStream: boolean
|
||||
): NormalizedAiMessage[] | null {
|
||||
if (isStream) {
|
||||
let role: unknown = "assistant";
|
||||
let content = "";
|
||||
let found = false;
|
||||
for (const frame of sseDataFrames(text)) {
|
||||
const delta = tryParseJson(frame)?.choices?.[0]?.delta;
|
||||
if (!delta) continue;
|
||||
found = true;
|
||||
if (typeof delta.role === "string") role = delta.role;
|
||||
if (typeof delta.content === "string") content += delta.content;
|
||||
}
|
||||
return found ? [{ role: normalizeRole(role), content }] : null;
|
||||
}
|
||||
|
||||
const message = tryParseJson(text)?.choices?.[0]?.message;
|
||||
if (!message) return null;
|
||||
return [
|
||||
{
|
||||
role: normalizeRole(message.role),
|
||||
content: flattenContent(message.content)
|
||||
}
|
||||
];
|
||||
}
|
||||
|
||||
function extractOpenAiResponsesOutputText(response: any): string | null {
|
||||
if (typeof response?.output_text === "string") return response.output_text;
|
||||
const output = Array.isArray(response?.output) ? response.output : [];
|
||||
const pieces: string[] = [];
|
||||
for (const item of output) {
|
||||
if (item?.type === "message" && Array.isArray(item.content)) {
|
||||
pieces.push(flattenContentParts(item.content));
|
||||
}
|
||||
}
|
||||
return pieces.length > 0 ? pieces.join("") : null;
|
||||
}
|
||||
|
||||
function normalizeOpenAiResponsesResponse(
|
||||
text: string,
|
||||
isStream: boolean
|
||||
): NormalizedAiMessage[] | null {
|
||||
if (isStream) {
|
||||
let content = "";
|
||||
let found = false;
|
||||
for (const frame of sseDataFrames(text)) {
|
||||
const parsed = tryParseJson(frame);
|
||||
if (!parsed) continue;
|
||||
if (
|
||||
parsed.type === "response.output_text.delta" &&
|
||||
typeof parsed.delta === "string"
|
||||
) {
|
||||
content += parsed.delta;
|
||||
found = true;
|
||||
} else if (
|
||||
parsed.type === "response.completed" &&
|
||||
parsed.response
|
||||
) {
|
||||
const outputText = extractOpenAiResponsesOutputText(
|
||||
parsed.response
|
||||
);
|
||||
if (outputText != null) {
|
||||
content = outputText;
|
||||
found = true;
|
||||
}
|
||||
}
|
||||
}
|
||||
return found ? [{ role: "assistant", content }] : null;
|
||||
}
|
||||
|
||||
const parsed = tryParseJson(text);
|
||||
const outputText = extractOpenAiResponsesOutputText(
|
||||
parsed?.response ?? parsed
|
||||
);
|
||||
return outputText != null
|
||||
? [{ role: "assistant", content: outputText }]
|
||||
: null;
|
||||
}
|
||||
|
||||
function normalizeAnthropicResponse(
|
||||
text: string,
|
||||
isStream: boolean
|
||||
): NormalizedAiMessage[] | null {
|
||||
if (isStream) {
|
||||
let role: unknown = "assistant";
|
||||
let content = "";
|
||||
let found = false;
|
||||
for (const frame of sseDataFrames(text)) {
|
||||
const parsed = tryParseJson(frame);
|
||||
if (!parsed) continue;
|
||||
if (parsed.type === "message_start" && parsed.message?.role) {
|
||||
role = parsed.message.role;
|
||||
}
|
||||
if (
|
||||
parsed.type === "content_block_start" &&
|
||||
parsed.content_block?.type === "tool_use"
|
||||
) {
|
||||
const name = parsed.content_block.name ?? "tool";
|
||||
content += `[tool_call: ${name}]`;
|
||||
found = true;
|
||||
}
|
||||
if (
|
||||
parsed.type === "content_block_delta" &&
|
||||
typeof parsed.delta?.text === "string"
|
||||
) {
|
||||
content += parsed.delta.text;
|
||||
found = true;
|
||||
}
|
||||
}
|
||||
return found ? [{ role: normalizeRole(role), content }] : null;
|
||||
}
|
||||
|
||||
const parsed = tryParseJson(text);
|
||||
if (!parsed || !Array.isArray(parsed.content)) return null;
|
||||
return [
|
||||
{
|
||||
role: normalizeRole(parsed.role ?? "assistant"),
|
||||
content: flattenContentParts(parsed.content)
|
||||
}
|
||||
];
|
||||
}
|
||||
|
||||
function geminiCandidateParts(node: any): string {
|
||||
const parts = node?.candidates?.[0]?.content?.parts;
|
||||
return Array.isArray(parts) ? flattenContentParts(parts) : "";
|
||||
}
|
||||
|
||||
function normalizeGeminiResponse(
|
||||
text: string,
|
||||
_isStream: boolean
|
||||
): NormalizedAiMessage[] | null {
|
||||
const frames = sseDataFrames(text);
|
||||
let content = "";
|
||||
let role: unknown = "model";
|
||||
let found = false;
|
||||
|
||||
if (frames.length > 0) {
|
||||
for (const frame of frames) {
|
||||
const parsed = tryParseJson(frame);
|
||||
const piece = geminiCandidateParts(parsed);
|
||||
if (piece) {
|
||||
content += piece;
|
||||
found = true;
|
||||
}
|
||||
const r = parsed?.candidates?.[0]?.content?.role;
|
||||
if (r) role = r;
|
||||
}
|
||||
} else {
|
||||
const parsed = tryParseJson(text);
|
||||
if (Array.isArray(parsed)) {
|
||||
for (const chunk of parsed) {
|
||||
const piece = geminiCandidateParts(chunk);
|
||||
if (piece) {
|
||||
content += piece;
|
||||
found = true;
|
||||
}
|
||||
const r = chunk?.candidates?.[0]?.content?.role;
|
||||
if (r) role = r;
|
||||
}
|
||||
} else if (parsed) {
|
||||
const piece = geminiCandidateParts(parsed);
|
||||
if (piece) {
|
||||
content = piece;
|
||||
found = true;
|
||||
}
|
||||
const r = parsed?.candidates?.[0]?.content?.role;
|
||||
if (r) role = r;
|
||||
}
|
||||
}
|
||||
|
||||
if (!found) {
|
||||
const scanned = scanTextFragments(text);
|
||||
return scanned
|
||||
? [{ role: normalizeRole(role), content: scanned }]
|
||||
: null;
|
||||
}
|
||||
return [{ role: normalizeRole(role), content }];
|
||||
}
|
||||
|
||||
function normalizeBedrockConverseResponse(
|
||||
text: string,
|
||||
isStream: boolean
|
||||
): NormalizedAiMessage[] | null {
|
||||
if (!isStream) {
|
||||
const message = tryParseJson(text)?.output?.message;
|
||||
if (!message) return null;
|
||||
return [
|
||||
{
|
||||
role: normalizeRole(message.role ?? "assistant"),
|
||||
content: Array.isArray(message.content)
|
||||
? flattenContentParts(message.content)
|
||||
: ""
|
||||
}
|
||||
];
|
||||
}
|
||||
// converse-stream uses AWS's binary event-stream framing, but the JSON
|
||||
// payload of each event survives intact inside it (same assumption
|
||||
// aiUsageExtraction.ts makes for usage) - scan for the text pieces.
|
||||
const scanned = scanTextFragments(text);
|
||||
return scanned ? [{ role: "assistant", content: scanned }] : null;
|
||||
}
|
||||
|
||||
function normalizeBedrockModelInvokeResponse(
|
||||
text: string,
|
||||
isStream: boolean
|
||||
): NormalizedAiMessage[] | null {
|
||||
const anthropicStyle = normalizeAnthropicResponse(text, isStream);
|
||||
if (anthropicStyle) return anthropicStyle;
|
||||
const scanned = scanTextFragments(text);
|
||||
return scanned ? [{ role: "assistant", content: scanned }] : null;
|
||||
}
|
||||
|
||||
function normalizeGoogleRawPredictResponse(
|
||||
text: string,
|
||||
isStream: boolean
|
||||
): NormalizedAiMessage[] | null {
|
||||
const anthropicStyle = normalizeAnthropicResponse(text, isStream);
|
||||
if (anthropicStyle) return anthropicStyle;
|
||||
const scanned = scanTextFragments(text);
|
||||
return scanned ? [{ role: "assistant", content: scanned }] : null;
|
||||
}
|
||||
|
||||
const REQUEST_NORMALIZERS: Record<
|
||||
AiCapability,
|
||||
(body: any) => NormalizedAiMessage[] | null
|
||||
> = {
|
||||
openai_chat: normalizeOpenAiChatRequest,
|
||||
openai_responses: normalizeOpenAiResponsesRequest,
|
||||
anthropic_messages: normalizeAnthropicRequest,
|
||||
// Model discovery carries no transcript to normalize.
|
||||
v1_models: () => null,
|
||||
gemini_generate_content: normalizeGeminiRequest,
|
||||
google_generate_content: normalizeGeminiRequest,
|
||||
google_raw_predict: normalizeBestEffortRequest,
|
||||
bedrock_model_invoke: normalizeBestEffortRequest,
|
||||
bedrock_converse: normalizeBedrockConverseRequest
|
||||
};
|
||||
|
||||
const RESPONSE_NORMALIZERS: Record<
|
||||
AiCapability,
|
||||
(text: string, isStream: boolean) => NormalizedAiMessage[] | null
|
||||
> = {
|
||||
openai_chat: normalizeOpenAiChatResponse,
|
||||
openai_responses: normalizeOpenAiResponsesResponse,
|
||||
anthropic_messages: normalizeAnthropicResponse,
|
||||
v1_models: () => null,
|
||||
gemini_generate_content: normalizeGeminiResponse,
|
||||
google_generate_content: normalizeGeminiResponse,
|
||||
google_raw_predict: normalizeGoogleRawPredictResponse,
|
||||
bedrock_model_invoke: normalizeBedrockModelInvokeResponse,
|
||||
bedrock_converse: normalizeBedrockConverseResponse
|
||||
};
|
||||
|
||||
/**
|
||||
* Normalizes an outbound AI gateway request body into a uniform message
|
||||
* transcript, regardless of capability/provider. Returns null if the body
|
||||
* doesn't contain any recognizable messages (or parsing failed) - callers
|
||||
* should fall back to showing the raw request body.
|
||||
*/
|
||||
export function normalizeAiRequest(
|
||||
capability: AiCapability,
|
||||
body: unknown
|
||||
): NormalizedAiMessage[] | null {
|
||||
try {
|
||||
const result = REQUEST_NORMALIZERS[capability](body);
|
||||
return result && result.length > 0 ? result : null;
|
||||
} catch (error) {
|
||||
logger.debug("Failed to normalize AI request messages", {
|
||||
capability,
|
||||
error
|
||||
});
|
||||
return null;
|
||||
}
|
||||
}
|
||||
|
||||
/**
|
||||
* Normalizes a completed (non-streaming or fully-accumulated streaming) AI
|
||||
* gateway response into a uniform message transcript. Returns null if
|
||||
* nothing recognizable could be extracted - callers should fall back to
|
||||
* showing the raw response body.
|
||||
*/
|
||||
export function normalizeAiResponse(
|
||||
capability: AiCapability,
|
||||
responseText: string,
|
||||
isStream: boolean
|
||||
): NormalizedAiMessage[] | null {
|
||||
try {
|
||||
const result = RESPONSE_NORMALIZERS[capability](responseText, isStream);
|
||||
return result && result.length > 0 ? result : null;
|
||||
} catch (error) {
|
||||
logger.debug("Failed to normalize AI response messages", {
|
||||
capability,
|
||||
error
|
||||
});
|
||||
return null;
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,387 @@
|
||||
import fs from "node:fs";
|
||||
import axios from "axios";
|
||||
import { z } from "zod";
|
||||
import config from "@server/lib/config";
|
||||
import logger from "@server/logger";
|
||||
import type { AiProviderType } from "@server/lib/aiProviderDefaults";
|
||||
|
||||
export const CATALOG_PROVIDERS = [
|
||||
"openai",
|
||||
"anthropic",
|
||||
"gemini",
|
||||
"vertex",
|
||||
"azure",
|
||||
"bedrock"
|
||||
] as const;
|
||||
|
||||
export type CatalogProvider = (typeof CATALOG_PROVIDERS)[number];
|
||||
|
||||
const CATALOG_PROVIDER_SET = new Set<string>(CATALOG_PROVIDERS);
|
||||
|
||||
// Each of our provider types maps to at most one catalog provider. Provider
|
||||
// types that proxy arbitrary underlying models (openRouter, vercelAiGateway,
|
||||
// custom) have no mapping.
|
||||
const PROVIDER_CATALOG_MAP: Record<
|
||||
Exclude<AiProviderType, "custom">,
|
||||
CatalogProvider | null
|
||||
> = {
|
||||
openai: "openai",
|
||||
anthropic: "anthropic",
|
||||
googleGemini: "gemini",
|
||||
vertexAi: "vertex",
|
||||
bedrock: "bedrock",
|
||||
microsoftFoundry: "azure",
|
||||
openRouter: null,
|
||||
vercelAiGateway: null
|
||||
};
|
||||
|
||||
export function getCatalogProviderForType(
|
||||
type: AiProviderType
|
||||
): CatalogProvider | null {
|
||||
if (type === "custom") {
|
||||
return null;
|
||||
}
|
||||
return PROVIDER_CATALOG_MAP[type];
|
||||
}
|
||||
|
||||
/**
|
||||
* Per-model feature flags as reported upstream. `null` means the catalog has
|
||||
* no data for that model - deliberately distinct from `false`, so consumers
|
||||
* can tell "unsupported" apart from "unknown".
|
||||
*/
|
||||
export type AiModelCapabilityFlags = {
|
||||
functionCalling: boolean | null;
|
||||
vision: boolean | null;
|
||||
promptCaching: boolean | null;
|
||||
reasoning: boolean | null;
|
||||
responseSchema: boolean | null;
|
||||
webSearch: boolean | null;
|
||||
};
|
||||
|
||||
export type AiModelCatalogEntry = {
|
||||
provider: CatalogProvider;
|
||||
model: string;
|
||||
pricing: {
|
||||
in: number | null;
|
||||
out: number | null;
|
||||
cache: number | null;
|
||||
reasoning: number | null;
|
||||
};
|
||||
limits: {
|
||||
/** Context window. */
|
||||
input: number | null;
|
||||
/** Cap on the output/max_tokens request parameter. */
|
||||
output: number | null;
|
||||
};
|
||||
capabilities: AiModelCapabilityFlags;
|
||||
};
|
||||
|
||||
const flag = z.boolean().nullable().optional();
|
||||
|
||||
// limits/capabilities are optional so a catalog published before they were
|
||||
// added (or an operator's own merge_file) still parses - those entries just
|
||||
// report unknown metadata rather than failing the whole payload.
|
||||
const catalogEntrySchema = z.object({
|
||||
model: z.string(),
|
||||
provider: z.string(),
|
||||
pricing: z
|
||||
.object({
|
||||
in: z.number().nullable().optional(),
|
||||
out: z.number().nullable().optional(),
|
||||
cache: z.number().nullable().optional(),
|
||||
reasoning: z.number().nullable().optional()
|
||||
})
|
||||
.optional(),
|
||||
limits: z
|
||||
.object({
|
||||
input: z.number().nullable().optional(),
|
||||
output: z.number().nullable().optional()
|
||||
})
|
||||
.optional(),
|
||||
capabilities: z
|
||||
.object({
|
||||
functionCalling: flag,
|
||||
vision: flag,
|
||||
promptCaching: flag,
|
||||
reasoning: flag,
|
||||
responseSchema: flag,
|
||||
webSearch: flag
|
||||
})
|
||||
.optional()
|
||||
});
|
||||
|
||||
const catalogFileSchema = z.object({
|
||||
data: z.array(catalogEntrySchema).optional().default([])
|
||||
});
|
||||
|
||||
type RawCatalogEntry = z.infer<typeof catalogEntrySchema>;
|
||||
|
||||
function normalizeCatalogProvider(raw: string): CatalogProvider | null {
|
||||
if (CATALOG_PROVIDER_SET.has(raw)) {
|
||||
return raw as CatalogProvider;
|
||||
}
|
||||
if (raw.startsWith("bedrock")) {
|
||||
return "bedrock";
|
||||
}
|
||||
if (raw.startsWith("vertex")) {
|
||||
return "vertex";
|
||||
}
|
||||
if (raw.startsWith("azure")) {
|
||||
return "azure";
|
||||
}
|
||||
return null;
|
||||
}
|
||||
|
||||
function normalizeEntry(raw: RawCatalogEntry): AiModelCatalogEntry | null {
|
||||
const provider = normalizeCatalogProvider(raw.provider);
|
||||
if (!provider) {
|
||||
return null;
|
||||
}
|
||||
|
||||
if (!raw.model) {
|
||||
return null;
|
||||
}
|
||||
|
||||
return {
|
||||
provider,
|
||||
model: raw.model,
|
||||
pricing: {
|
||||
in: raw.pricing?.in ?? null,
|
||||
out: raw.pricing?.out ?? null,
|
||||
cache: raw.pricing?.cache ?? null,
|
||||
reasoning: raw.pricing?.reasoning ?? null
|
||||
},
|
||||
limits: {
|
||||
input: raw.limits?.input ?? null,
|
||||
output: raw.limits?.output ?? null
|
||||
},
|
||||
capabilities: {
|
||||
functionCalling: raw.capabilities?.functionCalling ?? null,
|
||||
vision: raw.capabilities?.vision ?? null,
|
||||
promptCaching: raw.capabilities?.promptCaching ?? null,
|
||||
reasoning: raw.capabilities?.reasoning ?? null,
|
||||
responseSchema: raw.capabilities?.responseSchema ?? null,
|
||||
webSearch: raw.capabilities?.webSearch ?? null
|
||||
}
|
||||
};
|
||||
}
|
||||
|
||||
function providerKey(provider: CatalogProvider, key: string): string {
|
||||
return `${provider}\0${key}`;
|
||||
}
|
||||
|
||||
export class AiModelCatalog {
|
||||
private entries: AiModelCatalogEntry[] = [];
|
||||
private byProvider = new Map<CatalogProvider, AiModelCatalogEntry[]>();
|
||||
private byProviderAndKey = new Map<string, AiModelCatalogEntry>();
|
||||
private byKey = new Map<string, AiModelCatalogEntry[]>();
|
||||
private refreshTimer: NodeJS.Timeout | null = null;
|
||||
|
||||
/**
|
||||
* Loads the catalog into memory and schedules periodic background refreshes.
|
||||
* Call once at server startup.
|
||||
*/
|
||||
async init(): Promise<void> {
|
||||
await this.refresh();
|
||||
this.scheduleNextRefresh();
|
||||
}
|
||||
|
||||
/** Exact lookup by catalog provider and model key. */
|
||||
get(
|
||||
provider: CatalogProvider,
|
||||
key: string
|
||||
): AiModelCatalogEntry | undefined {
|
||||
return this.byProviderAndKey.get(providerKey(provider, key));
|
||||
}
|
||||
|
||||
/** All models for a catalog provider. */
|
||||
list(provider: CatalogProvider): AiModelCatalogEntry[] {
|
||||
return this.byProvider.get(provider) ?? [];
|
||||
}
|
||||
|
||||
/** All catalog entries that share a model key, across providers. */
|
||||
listByKey(key: string): AiModelCatalogEntry[] {
|
||||
return this.byKey.get(key) ?? [];
|
||||
}
|
||||
|
||||
/** Full in-memory catalog. */
|
||||
getAll(): AiModelCatalogEntry[] {
|
||||
return this.entries;
|
||||
}
|
||||
|
||||
private setEntries(entries: AiModelCatalogEntry[]): void {
|
||||
const byProvider = new Map<CatalogProvider, AiModelCatalogEntry[]>();
|
||||
const byProviderAndKey = new Map<string, AiModelCatalogEntry>();
|
||||
const byKey = new Map<string, AiModelCatalogEntry[]>();
|
||||
|
||||
for (const entry of entries) {
|
||||
const list = byProvider.get(entry.provider) ?? [];
|
||||
list.push(entry);
|
||||
byProvider.set(entry.provider, list);
|
||||
|
||||
const mapKey = providerKey(entry.provider, entry.model);
|
||||
if (!byProviderAndKey.has(mapKey)) {
|
||||
byProviderAndKey.set(mapKey, entry);
|
||||
}
|
||||
|
||||
const keyList = byKey.get(entry.model) ?? [];
|
||||
keyList.push(entry);
|
||||
byKey.set(entry.model, keyList);
|
||||
}
|
||||
|
||||
this.entries = entries;
|
||||
this.byProvider = byProvider;
|
||||
this.byProviderAndKey = byProviderAndKey;
|
||||
this.byKey = byKey;
|
||||
}
|
||||
|
||||
private async fetchFromFile(
|
||||
filePath: string
|
||||
): Promise<AiModelCatalogEntry[] | null> {
|
||||
try {
|
||||
if (!fs.existsSync(filePath)) {
|
||||
logger.warn(
|
||||
`AI model catalog file not found at ${filePath}; cost calculation will fall back to unknown pricing`
|
||||
);
|
||||
return null;
|
||||
}
|
||||
const raw = fs.readFileSync(filePath, "utf-8");
|
||||
const result = catalogFileSchema.safeParse(JSON.parse(raw));
|
||||
if (!result.success) {
|
||||
logger.warn(
|
||||
`AI model catalog file at ${filePath} failed validation: ${result.error.message}`
|
||||
);
|
||||
return null;
|
||||
}
|
||||
return result.data.data
|
||||
.map(normalizeEntry)
|
||||
.filter((e): e is AiModelCatalogEntry => e != null);
|
||||
} catch (error) {
|
||||
logger.warn("Failed to read AI model catalog file", { error });
|
||||
return null;
|
||||
}
|
||||
}
|
||||
|
||||
private async fetchFromUpstream(
|
||||
upstreamUrl: string
|
||||
): Promise<AiModelCatalogEntry[] | null> {
|
||||
try {
|
||||
const res = await axios.get(upstreamUrl, { timeout: 15_000 });
|
||||
const result = catalogFileSchema.safeParse(res.data);
|
||||
if (!result.success) {
|
||||
logger.warn(
|
||||
`AI model catalog response from ${upstreamUrl} failed validation: ${result.error.message}`
|
||||
);
|
||||
return null;
|
||||
}
|
||||
return result.data.data
|
||||
.map(normalizeEntry)
|
||||
.filter((e): e is AiModelCatalogEntry => e != null);
|
||||
} catch (error: any) {
|
||||
logger.warn(
|
||||
`Failed to fetch AI model catalog from ${upstreamUrl}: ${error.message || error}`
|
||||
);
|
||||
return null;
|
||||
}
|
||||
}
|
||||
|
||||
private async refresh(): Promise<void> {
|
||||
const { file, merge_file, upstream_url } =
|
||||
config.getRawConfig().ai.model_catalog;
|
||||
|
||||
const fetched = file
|
||||
? await this.fetchFromFile(file)
|
||||
: await this.fetchFromUpstream(upstream_url);
|
||||
|
||||
if (!fetched) {
|
||||
logger.debug(
|
||||
"AI model catalog refresh failed; keeping previously loaded catalog in memory"
|
||||
);
|
||||
return;
|
||||
}
|
||||
|
||||
let merged = fetched;
|
||||
if (merge_file) {
|
||||
const mergeEntries = await this.fetchFromFile(merge_file);
|
||||
if (mergeEntries) {
|
||||
// Entries from the base catalog take precedence; the merge
|
||||
// file only adds models not already present.
|
||||
merged = [...fetched, ...mergeEntries];
|
||||
}
|
||||
}
|
||||
|
||||
this.setEntries(merged);
|
||||
logger.debug(
|
||||
`AI model catalog refreshed: ${this.entries.length} models loaded`
|
||||
);
|
||||
}
|
||||
|
||||
private scheduleNextRefresh(): void {
|
||||
const { refresh_interval_min_hours, refresh_interval_max_hours } =
|
||||
config.getRawConfig().ai.model_catalog;
|
||||
|
||||
// Jittered rather than fixed so that many self-hosted instances don't
|
||||
// all hit the upstream catalog endpoint at the same moment.
|
||||
const minMs = refresh_interval_min_hours * 60 * 60 * 1000;
|
||||
const maxMs = refresh_interval_max_hours * 60 * 60 * 1000;
|
||||
const delayMs = minMs + Math.random() * Math.max(0, maxMs - minMs);
|
||||
|
||||
if (this.refreshTimer) {
|
||||
clearTimeout(this.refreshTimer);
|
||||
}
|
||||
this.refreshTimer = setTimeout(async () => {
|
||||
await this.refresh();
|
||||
this.scheduleNextRefresh();
|
||||
}, delayMs);
|
||||
}
|
||||
}
|
||||
|
||||
export const aiModelCatalog = new AiModelCatalog();
|
||||
|
||||
/**
|
||||
* Full catalog entries for a provider type, deduplicated by model id and
|
||||
* sorted by id. Model discovery uses these to report real token limits and
|
||||
* capability flags; `listCatalogModelsForType` is the id-only view of the
|
||||
* same list.
|
||||
*/
|
||||
export function listCatalogEntriesForType(
|
||||
type: AiProviderType,
|
||||
query?: string
|
||||
): AiModelCatalogEntry[] {
|
||||
const catalogProvider = getCatalogProviderForType(type);
|
||||
|
||||
let entries = catalogProvider ? aiModelCatalog.list(catalogProvider) : [];
|
||||
|
||||
if (query) {
|
||||
const q = query.toLowerCase();
|
||||
entries = entries.filter((e) => e.model.toLowerCase().includes(q));
|
||||
}
|
||||
|
||||
const seen = new Set<string>();
|
||||
entries = entries.filter((e) => {
|
||||
if (seen.has(e.model)) {
|
||||
return false;
|
||||
}
|
||||
seen.add(e.model);
|
||||
return true;
|
||||
});
|
||||
|
||||
return [...entries].sort((a, b) => a.model.localeCompare(b.model));
|
||||
}
|
||||
|
||||
export function listCatalogModelsForType(
|
||||
type: AiProviderType,
|
||||
query?: string
|
||||
): { model: string }[] {
|
||||
return listCatalogEntriesForType(type, query).map((entry) => ({
|
||||
model: entry.model
|
||||
}));
|
||||
}
|
||||
|
||||
/**
|
||||
* Loads the AI model pricing catalog into memory and schedules periodic
|
||||
* background refreshes. Call once at server startup.
|
||||
*/
|
||||
export async function initAiModelCatalog(): Promise<void> {
|
||||
await aiModelCatalog.init();
|
||||
}
|
||||
@@ -0,0 +1,235 @@
|
||||
import {
|
||||
isAllowedByLists,
|
||||
isModelKeyPattern
|
||||
} from "@server/lib/aiModelKeyMatch";
|
||||
import type { AiModelCapabilityFlags } from "@server/lib/aiModelCatalog";
|
||||
|
||||
// Anthropic's Models API pagination: 20 per page by default, 1..1000.
|
||||
export const MODEL_PAGE_DEFAULT_LIMIT = 20;
|
||||
export const MODEL_PAGE_MAX_LIMIT = 1000;
|
||||
|
||||
// Release dates aren't something we can know for a wildcard allow pattern or a
|
||||
// catalog entry. The Models API explicitly permits an epoch value when the
|
||||
// release date is unknown.
|
||||
const UNKNOWN_CREATED_AT = new Date(0).toISOString();
|
||||
|
||||
/**
|
||||
* One entry of Anthropic's `GET /v1/models` response. Only the identity fields
|
||||
* can be filled in from a provider's model lists - token limits and
|
||||
* per-model capability flags aren't derivable from an allow/block list, and the
|
||||
* API schema declares all three nullable.
|
||||
*/
|
||||
export type AnthropicModelInfo = {
|
||||
type: "model";
|
||||
id: string;
|
||||
display_name: string;
|
||||
created_at: string;
|
||||
max_input_tokens: number | null;
|
||||
max_tokens: number | null;
|
||||
capabilities: Record<string, unknown> | null;
|
||||
};
|
||||
|
||||
/** A model row an administrator configured explicitly on a provider. */
|
||||
export type ConfiguredModel = { name: string; createdAt: number };
|
||||
|
||||
/** What the pricing catalog knows about a model beyond its id. */
|
||||
export type CatalogModelMetadata = {
|
||||
maxInputTokens: number | null;
|
||||
maxOutputTokens: number | null;
|
||||
capabilities: AiModelCapabilityFlags;
|
||||
};
|
||||
|
||||
/**
|
||||
* Translates the catalog's flat feature flags into the nested shape
|
||||
* Anthropic's Models API uses. Best-effort by nature: the catalog carries a
|
||||
* coarser set of flags than the Models API describes, so anything it reports
|
||||
* as unknown (`null`) is surfaced as unsupported rather than invented.
|
||||
*/
|
||||
export function capabilitiesFromCatalog(
|
||||
flags: AiModelCapabilityFlags
|
||||
): Record<string, unknown> {
|
||||
const supported = (value: boolean | null) => ({
|
||||
supported: value === true
|
||||
});
|
||||
// The catalog has a single `reasoning` flag and no way to distinguish
|
||||
// adaptive from budget_tokens-style thinking, so both variants follow it.
|
||||
const reasoning = flags.reasoning === true;
|
||||
|
||||
return {
|
||||
batch: supported(null),
|
||||
citations: supported(null),
|
||||
code_execution: supported(null),
|
||||
context_management: {
|
||||
supported: false,
|
||||
clear_thinking_20251015: null,
|
||||
clear_tool_uses_20250919: null,
|
||||
compact_20260112: null
|
||||
},
|
||||
effort: {
|
||||
supported: reasoning,
|
||||
low: supported(flags.reasoning),
|
||||
medium: supported(flags.reasoning),
|
||||
high: supported(flags.reasoning),
|
||||
max: supported(flags.reasoning),
|
||||
xhigh: null
|
||||
},
|
||||
image_input: supported(flags.vision),
|
||||
pdf_input: supported(null),
|
||||
structured_outputs: supported(flags.responseSchema),
|
||||
thinking: {
|
||||
supported: reasoning,
|
||||
types: {
|
||||
adaptive: { supported: reasoning },
|
||||
enabled: { supported: reasoning }
|
||||
}
|
||||
}
|
||||
};
|
||||
}
|
||||
|
||||
/**
|
||||
* One attached provider's contribution to a resource's model listing, with the
|
||||
* allow/block lists already resolved for the attachment's access mode.
|
||||
*/
|
||||
export type ModelDiscoveryProvider = {
|
||||
providerId: number;
|
||||
allows: string[];
|
||||
blocks: string[];
|
||||
/**
|
||||
* Concrete model ids the provider's type is known to serve, with whatever
|
||||
* the catalog knows about each. This is what lets a wildcard allow such as
|
||||
* `claude-*` enumerate into real ids; provider types with no catalog
|
||||
* (aggregators, custom) pass an empty map and surface only their exact
|
||||
* allow entries.
|
||||
*/
|
||||
catalog: Map<string, CatalogModelMetadata>;
|
||||
/** Keyed by model key, for display names and creation times. */
|
||||
configured: Map<string, ConfiguredModel>;
|
||||
};
|
||||
|
||||
export type ModelPage = {
|
||||
data: AnthropicModelInfo[];
|
||||
has_more: boolean;
|
||||
};
|
||||
|
||||
/**
|
||||
* Expands one provider's effective allow/block lists into concrete model ids.
|
||||
* Two sources feed the candidate set: exact (non-wildcard) allow entries, which
|
||||
* are already concrete ids, and the catalog for the provider's type, which is
|
||||
* what makes wildcard allows enumerable. Every candidate is then run back
|
||||
* through the same allow/block check the inference pipeline applies, so a block
|
||||
* pattern hides a model here exactly as it would reject it at request time.
|
||||
*/
|
||||
export function expandProviderModels(
|
||||
provider: ModelDiscoveryProvider
|
||||
): AnthropicModelInfo[] {
|
||||
const candidates = new Set<string>();
|
||||
|
||||
for (const allow of provider.allows) {
|
||||
if (!isModelKeyPattern(allow)) {
|
||||
candidates.add(allow);
|
||||
}
|
||||
}
|
||||
for (const modelId of provider.catalog.keys()) {
|
||||
candidates.add(modelId);
|
||||
}
|
||||
|
||||
const models: AnthropicModelInfo[] = [];
|
||||
for (const modelKey of candidates) {
|
||||
if (!isAllowedByLists(modelKey, provider.allows, provider.blocks)) {
|
||||
continue;
|
||||
}
|
||||
const configured = provider.configured.get(modelKey);
|
||||
const catalog = provider.catalog.get(modelKey);
|
||||
|
||||
models.push({
|
||||
type: "model",
|
||||
id: modelKey,
|
||||
display_name: configured?.name || modelKey,
|
||||
created_at: configured
|
||||
? new Date(configured.createdAt).toISOString()
|
||||
: UNKNOWN_CREATED_AT,
|
||||
max_input_tokens: catalog?.maxInputTokens ?? null,
|
||||
max_tokens: catalog?.maxOutputTokens ?? null,
|
||||
capabilities: catalog
|
||||
? capabilitiesFromCatalog(catalog.capabilities)
|
||||
: null
|
||||
});
|
||||
}
|
||||
|
||||
return models;
|
||||
}
|
||||
|
||||
/**
|
||||
* Aggregates the permitted models across every provider attached to a
|
||||
* resource. Unlike an inference request there is no requested model to
|
||||
* disambiguate on, so no provider selection happens - the listing is the union
|
||||
* of what each provider would accept, deduplicated by model id.
|
||||
*/
|
||||
export function listPermittedModels(
|
||||
providers: ModelDiscoveryProvider[]
|
||||
): AnthropicModelInfo[] {
|
||||
const byModelId = new Map<string, AnthropicModelInfo>();
|
||||
|
||||
// Sorted so a model offered by two providers always resolves to the same
|
||||
// entry, which keeps the cursor ordering stable across requests.
|
||||
const ordered = [...providers].sort((a, b) => a.providerId - b.providerId);
|
||||
|
||||
for (const provider of ordered) {
|
||||
for (const model of expandProviderModels(provider)) {
|
||||
if (!byModelId.has(model.id)) {
|
||||
byModelId.set(model.id, model);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// "More recently released models are listed first" per the Models API,
|
||||
// with the id as a tie-break so the ordering is total - cursor pagination
|
||||
// needs it to be stable between calls.
|
||||
return [...byModelId.values()].sort((a, b) => {
|
||||
const byCreated = b.created_at.localeCompare(a.created_at);
|
||||
return byCreated !== 0 ? byCreated : a.id.localeCompare(b.id);
|
||||
});
|
||||
}
|
||||
|
||||
/**
|
||||
* Applies Anthropic's cursor pagination to an ordered model list. `after_id`
|
||||
* returns the page immediately after that model, `before_id` the page
|
||||
* immediately before it. Returns an error message for a caller mistake
|
||||
* (both cursors, or a cursor naming a model that isn't in the list).
|
||||
*/
|
||||
export function paginateModels(
|
||||
models: AnthropicModelInfo[],
|
||||
limit: number,
|
||||
cursor: { afterId?: string; beforeId?: string }
|
||||
): ModelPage | { error: string } {
|
||||
if (cursor.afterId && cursor.beforeId) {
|
||||
return { error: "Only one of after_id and before_id may be provided" };
|
||||
}
|
||||
|
||||
const cursorId = cursor.afterId ?? cursor.beforeId;
|
||||
if (!cursorId) {
|
||||
return {
|
||||
data: models.slice(0, limit),
|
||||
has_more: models.length > limit
|
||||
};
|
||||
}
|
||||
|
||||
const index = models.findIndex((model) => model.id === cursorId);
|
||||
if (index === -1) {
|
||||
return { error: `Unknown cursor id "${cursorId}"` };
|
||||
}
|
||||
|
||||
if (cursor.afterId) {
|
||||
const start = index + 1;
|
||||
return {
|
||||
data: models.slice(start, start + limit),
|
||||
has_more: models.length > start + limit
|
||||
};
|
||||
}
|
||||
|
||||
const start = Math.max(0, index - limit);
|
||||
return {
|
||||
data: models.slice(start, index),
|
||||
has_more: start > 0
|
||||
};
|
||||
}
|
||||
@@ -0,0 +1,101 @@
|
||||
const modelKeyRegexCache = new Map<string, RegExp>();
|
||||
|
||||
export function isModelKeyPattern(key: string): boolean {
|
||||
return key.includes("*") || key.includes("?");
|
||||
}
|
||||
|
||||
function getModelKeyRegex(pattern: string): RegExp {
|
||||
let regex = modelKeyRegexCache.get(pattern);
|
||||
if (!regex) {
|
||||
const escaped = pattern.replace(/[.+^${}()|[\]\\]/g, "\\$&");
|
||||
regex = new RegExp(
|
||||
`^${escaped.replace(/\*/g, ".*").replace(/\?/g, ".")}$`
|
||||
);
|
||||
modelKeyRegexCache.set(pattern, regex);
|
||||
}
|
||||
return regex;
|
||||
}
|
||||
|
||||
export function modelKeyMatches(
|
||||
pattern: string,
|
||||
requestedModel: string
|
||||
): boolean {
|
||||
return getModelKeyRegex(pattern).test(requestedModel);
|
||||
}
|
||||
|
||||
function wildcardCharCount(key: string): number {
|
||||
let count = 0;
|
||||
for (const char of key) {
|
||||
if (char === "*" || char === "?") {
|
||||
count += 1;
|
||||
}
|
||||
}
|
||||
return count;
|
||||
}
|
||||
|
||||
function literalLength(key: string): number {
|
||||
return key.replace(/[*?]/g, "").length;
|
||||
}
|
||||
|
||||
/**
|
||||
* Sort comparator: more specific patterns sort before less specific ones
|
||||
* (negative when `a` is more specific than `b`).
|
||||
*
|
||||
* 1. Exact keys beat patterns
|
||||
* 2. Fewer wildcard characters win
|
||||
* 3. Longer literal length wins
|
||||
*/
|
||||
export function compareModelKeySpecificity(a: string, b: string): number {
|
||||
const aIsPattern = isModelKeyPattern(a);
|
||||
const bIsPattern = isModelKeyPattern(b);
|
||||
|
||||
if (aIsPattern !== bIsPattern) {
|
||||
return aIsPattern ? 1 : -1;
|
||||
}
|
||||
|
||||
const wildcardDiff = wildcardCharCount(a) - wildcardCharCount(b);
|
||||
if (wildcardDiff !== 0) {
|
||||
return wildcardDiff;
|
||||
}
|
||||
|
||||
return literalLength(b) - literalLength(a);
|
||||
}
|
||||
|
||||
/**
|
||||
* Provider-layer policy: empty allowlist denies all. Blocklist only applies
|
||||
* after an allow match.
|
||||
*/
|
||||
export function isAllowedByLists(
|
||||
requested: string,
|
||||
allows: string[],
|
||||
blocks: string[]
|
||||
): boolean {
|
||||
if (allows.length === 0) {
|
||||
return false;
|
||||
}
|
||||
if (!allows.some((pattern) => modelKeyMatches(pattern, requested))) {
|
||||
return false;
|
||||
}
|
||||
if (blocks.some((pattern) => modelKeyMatches(pattern, requested))) {
|
||||
return false;
|
||||
}
|
||||
return true;
|
||||
}
|
||||
|
||||
/**
|
||||
* Among allow patterns that match `requested`, return the most specific one,
|
||||
* or null if none match.
|
||||
*/
|
||||
export function mostSpecificMatchingAllow(
|
||||
requested: string,
|
||||
allows: string[]
|
||||
): string | null {
|
||||
const matching = allows.filter((pattern) =>
|
||||
modelKeyMatches(pattern, requested)
|
||||
);
|
||||
if (matching.length === 0) {
|
||||
return null;
|
||||
}
|
||||
matching.sort(compareModelKeySpecificity);
|
||||
return matching[0];
|
||||
}
|
||||
@@ -0,0 +1,148 @@
|
||||
import type { AiProviderType } from "@server/lib/aiProviderDefaults";
|
||||
import type { AiUsage } from "@server/lib/aiUsageExtraction";
|
||||
import {
|
||||
aiModelCatalog,
|
||||
getCatalogProviderForType,
|
||||
type AiModelCatalogEntry,
|
||||
type CatalogProvider
|
||||
} from "@server/lib/aiModelCatalog";
|
||||
|
||||
export type AiModelPricing = {
|
||||
inputCostPerToken: number | null;
|
||||
outputCostPerToken: number | null;
|
||||
cacheReadInputTokenCost: number | null;
|
||||
outputCostPerReasoningToken: number | null;
|
||||
// True when the match came from a different catalog provider than the
|
||||
// one mapped to this provider's type (e.g. an openRouter/custom model
|
||||
// id that only matched a global search across every provider). Costs
|
||||
// found this way are a best-effort approximation, not a guarantee the
|
||||
// upstream provider bills at the same rate.
|
||||
approximate: boolean;
|
||||
};
|
||||
|
||||
function stripVendorPrefix(modelId: string): string | null {
|
||||
const idx = modelId.indexOf("/");
|
||||
if (idx === -1 || idx === modelId.length - 1) {
|
||||
return null;
|
||||
}
|
||||
return modelId.slice(idx + 1);
|
||||
}
|
||||
|
||||
function toPricing(
|
||||
entry: AiModelCatalogEntry,
|
||||
approximate: boolean
|
||||
): AiModelPricing {
|
||||
return {
|
||||
inputCostPerToken: entry.pricing.in,
|
||||
outputCostPerToken: entry.pricing.out,
|
||||
cacheReadInputTokenCost: entry.pricing.cache,
|
||||
outputCostPerReasoningToken: entry.pricing.reasoning,
|
||||
approximate
|
||||
};
|
||||
}
|
||||
|
||||
function findEntry(
|
||||
modelId: string,
|
||||
provider: CatalogProvider | null
|
||||
): AiModelCatalogEntry | null {
|
||||
const candidates = [modelId, stripVendorPrefix(modelId)].filter(
|
||||
(v): v is string => v != null
|
||||
);
|
||||
|
||||
for (const key of candidates) {
|
||||
if (provider) {
|
||||
const match = aiModelCatalog.get(provider, key);
|
||||
if (match) {
|
||||
return match;
|
||||
}
|
||||
continue;
|
||||
}
|
||||
|
||||
const match = aiModelCatalog.listByKey(key)[0];
|
||||
if (match) {
|
||||
return match;
|
||||
}
|
||||
}
|
||||
return null;
|
||||
}
|
||||
|
||||
/**
|
||||
* Looks up per-token pricing for a model, scoped first to the catalog
|
||||
* provider that corresponds to our provider type, then falling back to a
|
||||
* global search across every provider (marked `approximate`) for provider
|
||||
* types that proxy arbitrary underlying models.
|
||||
*/
|
||||
export function getModelPricing(
|
||||
providerType: AiProviderType,
|
||||
modelId: string | undefined
|
||||
): AiModelPricing | null {
|
||||
if (!modelId) {
|
||||
return null;
|
||||
}
|
||||
|
||||
const catalogProvider = getCatalogProviderForType(providerType);
|
||||
|
||||
if (catalogProvider) {
|
||||
const scoped = findEntry(modelId, catalogProvider);
|
||||
if (scoped) {
|
||||
return toPricing(scoped, false);
|
||||
}
|
||||
}
|
||||
|
||||
const fallback = findEntry(modelId, null);
|
||||
if (fallback) {
|
||||
return toPricing(fallback, true);
|
||||
}
|
||||
|
||||
return null;
|
||||
}
|
||||
|
||||
export type AiCostBreakdown = {
|
||||
promptCost: number;
|
||||
cacheReadCost: number;
|
||||
cacheWriteCost: number;
|
||||
completionCost: number;
|
||||
reasoningCost: number;
|
||||
totalCost: number;
|
||||
};
|
||||
|
||||
/**
|
||||
* Computes a $ cost breakdown for a usage record given a model's pricing.
|
||||
* Cache writes and reasoning tokens fall back to the normal input/output
|
||||
* rate respectively when the catalog has no dedicated rate for them (the
|
||||
* catalog has no cache-write field at all, and only some models report a
|
||||
* distinct reasoning rate).
|
||||
*/
|
||||
export function calculateAiCost(
|
||||
pricing: AiModelPricing | null,
|
||||
usage: AiUsage
|
||||
): AiCostBreakdown | null {
|
||||
if (!pricing) {
|
||||
return null;
|
||||
}
|
||||
|
||||
const inputRate = pricing.inputCostPerToken ?? 0;
|
||||
const outputRate = pricing.outputCostPerToken ?? 0;
|
||||
const cacheReadRate = pricing.cacheReadInputTokenCost ?? inputRate;
|
||||
const reasoningRate = pricing.outputCostPerReasoningToken ?? outputRate;
|
||||
|
||||
const promptCost = usage.promptTokens * inputRate;
|
||||
const cacheReadCost = usage.cacheReadTokens * cacheReadRate;
|
||||
const cacheWriteCost = usage.cacheWriteTokens * inputRate;
|
||||
const completionCost = usage.completionTokens * outputRate;
|
||||
const reasoningCost = usage.reasoningTokens * reasoningRate;
|
||||
|
||||
return {
|
||||
promptCost,
|
||||
cacheReadCost,
|
||||
cacheWriteCost,
|
||||
completionCost,
|
||||
reasoningCost,
|
||||
totalCost:
|
||||
promptCost +
|
||||
cacheReadCost +
|
||||
cacheWriteCost +
|
||||
completionCost +
|
||||
reasoningCost
|
||||
};
|
||||
}
|
||||
@@ -0,0 +1,186 @@
|
||||
import { decrypt, encrypt } from "@server/lib/crypto";
|
||||
import {
|
||||
parseCapabilities,
|
||||
type AiCapability
|
||||
} from "@server/lib/aiCapabilities";
|
||||
import { stripVirtualApiKeyAuthHeaders } from "@app/lib/virtualApiKeyFormat";
|
||||
import {
|
||||
AI_PROVIDER_AUTH_TYPES,
|
||||
AI_PROVIDER_DEFAULTS,
|
||||
authTypeRequiresApiKey,
|
||||
defaultsForProviderType,
|
||||
providerRequiresUpstreamUrl,
|
||||
type AiBudgetUnit,
|
||||
type AiProviderAuthType,
|
||||
type AiProviderRoutingMode,
|
||||
type AiProviderType
|
||||
} from "@app/lib/aiProviderDefaults";
|
||||
|
||||
export {
|
||||
AI_PROVIDER_AUTH_TYPES,
|
||||
AI_PROVIDER_DEFAULTS,
|
||||
authTypeRequiresApiKey,
|
||||
defaultsForProviderType,
|
||||
providerRequiresUpstreamUrl,
|
||||
type AiBudgetUnit,
|
||||
type AiProviderAuthType,
|
||||
type AiProviderRoutingMode,
|
||||
type AiProviderType
|
||||
};
|
||||
|
||||
const CONFLICTING_AUTH_HEADERS = [
|
||||
"authorization",
|
||||
"x-api-key",
|
||||
"x-goog-api-key",
|
||||
"cf-aig-authorization"
|
||||
] as const;
|
||||
|
||||
export function resolveAiProviderCreateFields(input: {
|
||||
type: AiProviderType;
|
||||
upstreamUrl?: string | null;
|
||||
authType?: AiProviderAuthType | null;
|
||||
routingMode?: AiProviderRoutingMode | null;
|
||||
}): {
|
||||
upstreamUrl: string | null;
|
||||
authType: AiProviderAuthType;
|
||||
routingMode: AiProviderRoutingMode;
|
||||
} {
|
||||
const routingMode =
|
||||
input.type === "custom" ? (input.routingMode ?? "url") : "url";
|
||||
|
||||
if (routingMode === "target") {
|
||||
return {
|
||||
upstreamUrl: null,
|
||||
authType: input.authType ?? "bearer",
|
||||
routingMode
|
||||
};
|
||||
}
|
||||
|
||||
if (input.type === "custom") {
|
||||
return {
|
||||
upstreamUrl: input.upstreamUrl ?? null,
|
||||
authType: input.authType ?? "bearer",
|
||||
routingMode
|
||||
};
|
||||
}
|
||||
|
||||
const defaults = AI_PROVIDER_DEFAULTS[input.type];
|
||||
return {
|
||||
upstreamUrl: input.upstreamUrl ?? defaults.upstreamUrl,
|
||||
authType: input.authType ?? defaults.authType,
|
||||
routingMode
|
||||
};
|
||||
}
|
||||
|
||||
export type AiProviderHeader = { name: string; value: string };
|
||||
|
||||
export function serializeAiProviderHeaders(
|
||||
headers: AiProviderHeader[] | null | undefined,
|
||||
secret: string
|
||||
): string | null {
|
||||
if (!headers || headers.length === 0) {
|
||||
return null;
|
||||
}
|
||||
return encrypt(JSON.stringify(headers), secret);
|
||||
}
|
||||
|
||||
export function parseAiProviderHeaders(
|
||||
raw: string | null | undefined,
|
||||
secret: string
|
||||
): AiProviderHeader[] {
|
||||
if (!raw) {
|
||||
return [];
|
||||
}
|
||||
try {
|
||||
const decrypted = decrypt(raw, secret);
|
||||
const parsed = JSON.parse(decrypted);
|
||||
if (!Array.isArray(parsed)) {
|
||||
return [];
|
||||
}
|
||||
return parsed.filter(
|
||||
(h): h is AiProviderHeader =>
|
||||
h != null &&
|
||||
typeof h === "object" &&
|
||||
typeof h.name === "string" &&
|
||||
typeof h.value === "string"
|
||||
);
|
||||
} catch {
|
||||
return [];
|
||||
}
|
||||
}
|
||||
|
||||
export function applyAiProviderCustomHeaders(
|
||||
headers: Record<string, string>,
|
||||
raw: string | null | undefined,
|
||||
secret: string
|
||||
): void {
|
||||
for (const { name, value } of parseAiProviderHeaders(raw, secret)) {
|
||||
headers[name] = value;
|
||||
}
|
||||
}
|
||||
|
||||
/**
|
||||
* Apply provider auth to upstream headers.
|
||||
* - Always strips Pangolin virtual API key credentials from client auth headers.
|
||||
* - Injected modes: strip conflicting client auth headers, then set the provider key.
|
||||
* - none: strip conflicting client auth headers, send no auth.
|
||||
* - passthrough: leave remaining client auth headers as-is (after VAK strip).
|
||||
*/
|
||||
export function applyAiProviderAuthHeaders(
|
||||
headers: Record<string, string>,
|
||||
authType: AiProviderAuthType,
|
||||
apiKey: string | null
|
||||
): void {
|
||||
stripVirtualApiKeyAuthHeaders(headers);
|
||||
|
||||
if (authType === "passthrough") {
|
||||
return;
|
||||
}
|
||||
|
||||
for (const name of CONFLICTING_AUTH_HEADERS) {
|
||||
for (const key of Object.keys(headers)) {
|
||||
if (key.toLowerCase() === name) {
|
||||
delete headers[key];
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
if (authType === "none") {
|
||||
return;
|
||||
}
|
||||
|
||||
if (!apiKey) {
|
||||
throw new Error(`API key required for authType ${authType}`);
|
||||
}
|
||||
|
||||
switch (authType) {
|
||||
case "bearer":
|
||||
headers["Authorization"] = `Bearer ${apiKey}`;
|
||||
break;
|
||||
case "x-api-key":
|
||||
headers["x-api-key"] = apiKey;
|
||||
break;
|
||||
case "x-goog-api-key":
|
||||
headers["x-goog-api-key"] = apiKey;
|
||||
break;
|
||||
case "hec":
|
||||
headers["Authorization"] = `Splunk ${apiKey}`;
|
||||
break;
|
||||
case "cf-aig-authorization":
|
||||
headers["cf-aig-authorization"] = `Bearer ${apiKey}`;
|
||||
break;
|
||||
}
|
||||
}
|
||||
|
||||
export function resolveCapabilitiesForCreate(input: {
|
||||
type: AiProviderType;
|
||||
capabilities?: AiCapability[] | null;
|
||||
}): AiCapability[] {
|
||||
if (input.capabilities != null) {
|
||||
return parseCapabilities(input.capabilities);
|
||||
}
|
||||
if (input.type === "custom") {
|
||||
return [];
|
||||
}
|
||||
return [...AI_PROVIDER_DEFAULTS[input.type].capabilities];
|
||||
}
|
||||
@@ -0,0 +1,97 @@
|
||||
import {
|
||||
aiModelCatalog,
|
||||
getCatalogProviderForType,
|
||||
type CatalogProvider
|
||||
} from "@server/lib/aiModelCatalog";
|
||||
import type { AiProviderType } from "@server/lib/aiProviderDefaults";
|
||||
|
||||
function stripVendorPrefix(modelId: string): string | null {
|
||||
const idx = modelId.indexOf("/");
|
||||
if (idx === -1 || idx === modelId.length - 1) {
|
||||
return null;
|
||||
}
|
||||
return modelId.slice(idx + 1);
|
||||
}
|
||||
|
||||
function modelKeysToTry(modelId: string): string[] {
|
||||
const keys = [modelId];
|
||||
const stripped = stripVendorPrefix(modelId);
|
||||
if (stripped) {
|
||||
keys.push(stripped);
|
||||
}
|
||||
return keys;
|
||||
}
|
||||
|
||||
function catalogOwnsModel(
|
||||
catalogProvider: CatalogProvider,
|
||||
modelId: string
|
||||
): boolean {
|
||||
for (const key of modelKeysToTry(modelId)) {
|
||||
if (aiModelCatalog.get(catalogProvider, key)) {
|
||||
return true;
|
||||
}
|
||||
}
|
||||
return false;
|
||||
}
|
||||
|
||||
function modelKnownInAnyCatalog(modelId: string): boolean {
|
||||
for (const key of modelKeysToTry(modelId)) {
|
||||
if (aiModelCatalog.listByKey(key).length > 0) {
|
||||
return true;
|
||||
}
|
||||
}
|
||||
return false;
|
||||
}
|
||||
|
||||
/**
|
||||
* How strongly a provider "owns" a requested model id via the known catalog.
|
||||
*
|
||||
* 2 - Typed provider whose catalog contains the model
|
||||
* 1 - Aggregator/custom that can proxy a catalog-known model
|
||||
* 0 - No ownership signal (typed miss, or unknown model on aggregator/custom)
|
||||
*/
|
||||
export function catalogOwnershipScore(
|
||||
type: AiProviderType,
|
||||
modelId: string
|
||||
): number {
|
||||
const catalogProvider = getCatalogProviderForType(type);
|
||||
if (catalogProvider != null) {
|
||||
return catalogOwnsModel(catalogProvider, modelId) ? 2 : 0;
|
||||
}
|
||||
return modelKnownInAnyCatalog(modelId) ? 1 : 0;
|
||||
}
|
||||
|
||||
/**
|
||||
* Prefer native vendor providers over aggregators over custom when catalog
|
||||
* ownership is tied.
|
||||
*
|
||||
* 2 - Native typed provider (openai, anthropic, gemini, ...)
|
||||
* 1 - Aggregator gateway (openRouter, vercelAiGateway)
|
||||
* 0 - Custom
|
||||
*/
|
||||
export function providerClassRank(type: AiProviderType): number {
|
||||
if (type === "custom") {
|
||||
return 0;
|
||||
}
|
||||
if (type === "openRouter" || type === "vercelAiGateway") {
|
||||
return 1;
|
||||
}
|
||||
return 2;
|
||||
}
|
||||
|
||||
export function keepBestScored<T>(
|
||||
items: T[],
|
||||
scoreFn: (item: T) => number
|
||||
): T[] {
|
||||
if (items.length <= 1) {
|
||||
return items;
|
||||
}
|
||||
let best = Number.NEGATIVE_INFINITY;
|
||||
for (const item of items) {
|
||||
const score = scoreFn(item);
|
||||
if (score > best) {
|
||||
best = score;
|
||||
}
|
||||
}
|
||||
return items.filter((item) => scoreFn(item) === best);
|
||||
}
|
||||
@@ -0,0 +1,484 @@
|
||||
import { encode } from "gpt-tokenizer";
|
||||
import type { AiCapability } from "@server/lib/aiCapabilities";
|
||||
import logger from "@server/logger";
|
||||
|
||||
export type AiUsage = {
|
||||
// Input tokens billed at the normal input rate (i.e. NOT already
|
||||
// covered by cacheReadTokens/cacheWriteTokens below).
|
||||
promptTokens: number;
|
||||
cacheReadTokens: number;
|
||||
cacheWriteTokens: number;
|
||||
// Output tokens billed at the normal output rate (i.e. NOT already
|
||||
// covered by reasoningTokens below).
|
||||
completionTokens: number;
|
||||
reasoningTokens: number;
|
||||
// True when these numbers are our own best-guess estimate (the upstream
|
||||
// response didn't report usage), rather than provider-reported figures.
|
||||
estimated: boolean;
|
||||
};
|
||||
|
||||
export function emptyUsage(): AiUsage {
|
||||
return {
|
||||
promptTokens: 0,
|
||||
cacheReadTokens: 0,
|
||||
cacheWriteTokens: 0,
|
||||
completionTokens: 0,
|
||||
reasoningTokens: 0,
|
||||
estimated: false
|
||||
};
|
||||
}
|
||||
|
||||
/**
|
||||
* Scans raw (possibly binary-framed, e.g. Bedrock's vnd.amazon.eventstream)
|
||||
* text for `"fieldName":123` occurrences and returns the last value seen for
|
||||
* each field. Used as a best-effort fallback for response shapes we can't
|
||||
* fully parse as JSON/SSE (streaming Bedrock, raw predict passthroughs).
|
||||
*/
|
||||
function scanNumericFields(
|
||||
text: string,
|
||||
fields: string[]
|
||||
): Record<string, number> {
|
||||
const out: Record<string, number> = {};
|
||||
for (const field of fields) {
|
||||
const re = new RegExp(`"${field}"\\s*:\\s*(\\d+)`, "g");
|
||||
let match: RegExpExecArray | null;
|
||||
while ((match = re.exec(text)) !== null) {
|
||||
out[field] = Number(match[1]);
|
||||
}
|
||||
}
|
||||
return out;
|
||||
}
|
||||
|
||||
// Exported for reuse by server/lib/aiMessageNormalization.ts, which needs
|
||||
// the same SSE-frame/JSON-parsing groundwork to extract message content
|
||||
// instead of usage numbers.
|
||||
export function sseDataFrames(text: string): string[] {
|
||||
const frames: string[] = [];
|
||||
for (const rawFrame of text.split(/\r?\n\r?\n/)) {
|
||||
for (const line of rawFrame.split(/\r?\n/)) {
|
||||
if (!line.startsWith("data:")) continue;
|
||||
const data = line.slice("data:".length).trim();
|
||||
if (data && data !== "[DONE]") {
|
||||
frames.push(data);
|
||||
}
|
||||
}
|
||||
}
|
||||
return frames;
|
||||
}
|
||||
|
||||
export function tryParseJson(text: string): any | null {
|
||||
try {
|
||||
return JSON.parse(text);
|
||||
} catch {
|
||||
return null;
|
||||
}
|
||||
}
|
||||
|
||||
function extractOpenAiChat(text: string, isStream: boolean): AiUsage | null {
|
||||
let usage: any = null;
|
||||
|
||||
if (isStream) {
|
||||
for (const frame of sseDataFrames(text)) {
|
||||
const parsed = tryParseJson(frame);
|
||||
if (parsed?.usage) {
|
||||
usage = parsed.usage;
|
||||
}
|
||||
}
|
||||
} else {
|
||||
usage = tryParseJson(text)?.usage ?? null;
|
||||
}
|
||||
|
||||
if (!usage) {
|
||||
return null;
|
||||
}
|
||||
|
||||
const cacheReadTokens = usage.prompt_tokens_details?.cached_tokens ?? 0;
|
||||
const reasoningTokens =
|
||||
usage.completion_tokens_details?.reasoning_tokens ?? 0;
|
||||
|
||||
return {
|
||||
promptTokens: Math.max(0, (usage.prompt_tokens ?? 0) - cacheReadTokens),
|
||||
cacheReadTokens,
|
||||
cacheWriteTokens: 0,
|
||||
completionTokens: Math.max(
|
||||
0,
|
||||
(usage.completion_tokens ?? 0) - reasoningTokens
|
||||
),
|
||||
reasoningTokens,
|
||||
estimated: false
|
||||
};
|
||||
}
|
||||
|
||||
function extractOpenAiResponses(
|
||||
text: string,
|
||||
isStream: boolean
|
||||
): AiUsage | null {
|
||||
let usage: any = null;
|
||||
|
||||
if (isStream) {
|
||||
for (const frame of sseDataFrames(text)) {
|
||||
const parsed = tryParseJson(frame);
|
||||
if (
|
||||
parsed?.type === "response.completed" &&
|
||||
parsed?.response?.usage
|
||||
) {
|
||||
usage = parsed.response.usage;
|
||||
} else if (parsed?.usage) {
|
||||
usage = parsed.usage;
|
||||
}
|
||||
}
|
||||
} else {
|
||||
const parsed = tryParseJson(text);
|
||||
usage = parsed?.usage ?? parsed?.response?.usage ?? null;
|
||||
}
|
||||
|
||||
if (!usage) {
|
||||
return null;
|
||||
}
|
||||
|
||||
const cacheReadTokens = usage.input_tokens_details?.cached_tokens ?? 0;
|
||||
const reasoningTokens = usage.output_tokens_details?.reasoning_tokens ?? 0;
|
||||
|
||||
return {
|
||||
promptTokens: Math.max(0, (usage.input_tokens ?? 0) - cacheReadTokens),
|
||||
cacheReadTokens,
|
||||
cacheWriteTokens: 0,
|
||||
completionTokens: Math.max(
|
||||
0,
|
||||
(usage.output_tokens ?? 0) - reasoningTokens
|
||||
),
|
||||
reasoningTokens,
|
||||
estimated: false
|
||||
};
|
||||
}
|
||||
|
||||
function extractAnthropicMessages(
|
||||
text: string,
|
||||
isStream: boolean
|
||||
): AiUsage | null {
|
||||
let inputTokens = 0;
|
||||
let cacheReadTokens = 0;
|
||||
let cacheWriteTokens = 0;
|
||||
let outputTokens = 0;
|
||||
let found = false;
|
||||
|
||||
const applyUsage = (usage: any) => {
|
||||
if (!usage) return;
|
||||
found = true;
|
||||
if (typeof usage.input_tokens === "number") {
|
||||
inputTokens = usage.input_tokens;
|
||||
}
|
||||
if (typeof usage.cache_read_input_tokens === "number") {
|
||||
cacheReadTokens = usage.cache_read_input_tokens;
|
||||
}
|
||||
if (typeof usage.cache_creation_input_tokens === "number") {
|
||||
cacheWriteTokens = usage.cache_creation_input_tokens;
|
||||
}
|
||||
if (typeof usage.output_tokens === "number") {
|
||||
outputTokens = usage.output_tokens;
|
||||
}
|
||||
};
|
||||
|
||||
if (isStream) {
|
||||
for (const frame of sseDataFrames(text)) {
|
||||
const parsed = tryParseJson(frame);
|
||||
if (!parsed) continue;
|
||||
applyUsage(parsed.message?.usage);
|
||||
applyUsage(parsed.usage);
|
||||
}
|
||||
} else {
|
||||
applyUsage(tryParseJson(text)?.usage);
|
||||
}
|
||||
|
||||
if (!found) {
|
||||
return null;
|
||||
}
|
||||
|
||||
return {
|
||||
promptTokens: inputTokens,
|
||||
cacheReadTokens,
|
||||
cacheWriteTokens,
|
||||
completionTokens: outputTokens,
|
||||
// Anthropic bills extended-thinking output at the normal output
|
||||
// rate, so there's no separate reasoning bucket to report.
|
||||
reasoningTokens: 0,
|
||||
estimated: false
|
||||
};
|
||||
}
|
||||
|
||||
function extractGoogleGenerateContent(
|
||||
text: string,
|
||||
_isStream: boolean
|
||||
): AiUsage | null {
|
||||
// Both the plain-JSON-array stream format and the SSE (?alt=sse) format
|
||||
// repeat a cumulative `usageMetadata` object per chunk; the regex scan
|
||||
// below naturally picks up the last (most complete) one either way.
|
||||
const fields = scanNumericFields(text, [
|
||||
"promptTokenCount",
|
||||
"candidatesTokenCount",
|
||||
"cachedContentTokenCount",
|
||||
"thoughtsTokenCount"
|
||||
]);
|
||||
|
||||
if (fields.promptTokenCount === undefined) {
|
||||
return null;
|
||||
}
|
||||
|
||||
const cacheReadTokens = fields.cachedContentTokenCount ?? 0;
|
||||
const reasoningTokens = fields.thoughtsTokenCount ?? 0;
|
||||
|
||||
return {
|
||||
promptTokens: Math.max(0, fields.promptTokenCount - cacheReadTokens),
|
||||
cacheReadTokens,
|
||||
cacheWriteTokens: 0,
|
||||
completionTokens: fields.candidatesTokenCount ?? 0,
|
||||
reasoningTokens,
|
||||
estimated: false
|
||||
};
|
||||
}
|
||||
|
||||
function extractBedrockConverse(
|
||||
text: string,
|
||||
_isStream: boolean
|
||||
): AiUsage | null {
|
||||
// Non-streaming responses are plain JSON; converse-stream frames the
|
||||
// final `metadata` event's usage object inside binary event-stream
|
||||
// framing, but the JSON text survives intact inside that binary
|
||||
// envelope, so the same field scan works for both.
|
||||
const parsed = tryParseJson(text);
|
||||
const usage = parsed?.usage;
|
||||
if (usage) {
|
||||
const cacheReadTokens = usage.cacheReadInputTokens ?? 0;
|
||||
return {
|
||||
promptTokens: Math.max(
|
||||
0,
|
||||
(usage.inputTokens ?? 0) - cacheReadTokens
|
||||
),
|
||||
cacheReadTokens,
|
||||
cacheWriteTokens: usage.cacheWriteInputTokens ?? 0,
|
||||
completionTokens: usage.outputTokens ?? 0,
|
||||
reasoningTokens: 0,
|
||||
estimated: false
|
||||
};
|
||||
}
|
||||
|
||||
const fields = scanNumericFields(text, [
|
||||
"inputTokens",
|
||||
"outputTokens",
|
||||
"cacheReadInputTokens",
|
||||
"cacheWriteInputTokens"
|
||||
]);
|
||||
if (fields.inputTokens === undefined) {
|
||||
return null;
|
||||
}
|
||||
const cacheReadTokens = fields.cacheReadInputTokens ?? 0;
|
||||
return {
|
||||
promptTokens: Math.max(0, fields.inputTokens - cacheReadTokens),
|
||||
cacheReadTokens,
|
||||
cacheWriteTokens: fields.cacheWriteInputTokens ?? 0,
|
||||
completionTokens: fields.outputTokens ?? 0,
|
||||
reasoningTokens: 0,
|
||||
estimated: false
|
||||
};
|
||||
}
|
||||
|
||||
function extractBedrockModelInvoke(
|
||||
text: string,
|
||||
_isStream: boolean,
|
||||
headers: Headers
|
||||
): AiUsage | null {
|
||||
// Non-streaming invoke reports counts via response headers regardless
|
||||
// of the underlying model's payload format.
|
||||
const headerInput = headers.get("x-amzn-bedrock-input-token-count");
|
||||
const headerOutput = headers.get("x-amzn-bedrock-output-token-count");
|
||||
if (headerInput !== null || headerOutput !== null) {
|
||||
return {
|
||||
promptTokens: Number(headerInput ?? 0),
|
||||
cacheReadTokens: 0,
|
||||
cacheWriteTokens: 0,
|
||||
completionTokens: Number(headerOutput ?? 0),
|
||||
reasoningTokens: 0,
|
||||
estimated: false
|
||||
};
|
||||
}
|
||||
|
||||
// invoke-with-response-stream has no equivalent headers; the model's
|
||||
// own usage shape (frequently Anthropic-style on Bedrock) is embedded
|
||||
// inside binary event-stream framing, so fall back to a couple of
|
||||
// known field-name shapes via regex.
|
||||
const anthropicStyle = extractAnthropicMessages(text, true);
|
||||
if (anthropicStyle) {
|
||||
return anthropicStyle;
|
||||
}
|
||||
|
||||
const fields = scanNumericFields(text, [
|
||||
"inputTokenCount",
|
||||
"outputTokenCount"
|
||||
]);
|
||||
if (fields.inputTokenCount === undefined) {
|
||||
return null;
|
||||
}
|
||||
return {
|
||||
promptTokens: fields.inputTokenCount,
|
||||
cacheReadTokens: 0,
|
||||
cacheWriteTokens: 0,
|
||||
completionTokens: fields.outputTokenCount ?? 0,
|
||||
reasoningTokens: 0,
|
||||
estimated: false
|
||||
};
|
||||
}
|
||||
|
||||
const EXTRACTORS: Record<
|
||||
AiCapability,
|
||||
(text: string, isStream: boolean, headers: Headers) => AiUsage | null
|
||||
> = {
|
||||
openai_chat: extractOpenAiChat,
|
||||
openai_responses: extractOpenAiResponses,
|
||||
anthropic_messages: extractAnthropicMessages,
|
||||
// Model discovery never runs a model, so there are no tokens to bill.
|
||||
v1_models: () => null,
|
||||
gemini_generate_content: extractGoogleGenerateContent,
|
||||
google_generate_content: extractGoogleGenerateContent,
|
||||
// rawPredict is a passthrough to whatever the underlying publisher
|
||||
// model speaks (often Anthropic-shaped on Vertex); try that, then give
|
||||
// up to the token-count estimate.
|
||||
google_raw_predict: (text, isStream) =>
|
||||
extractAnthropicMessages(text, isStream),
|
||||
bedrock_model_invoke: extractBedrockModelInvoke,
|
||||
bedrock_converse: extractBedrockConverse
|
||||
};
|
||||
|
||||
/**
|
||||
* Attempts to pull provider-reported token usage out of an upstream AI
|
||||
* gateway response. Returns null if the response didn't contain (or we
|
||||
* couldn't find) usage data, in which case callers should fall back to
|
||||
* `estimateUsage`.
|
||||
*/
|
||||
export function extractUsage(
|
||||
capability: AiCapability,
|
||||
responseText: string,
|
||||
isStream: boolean,
|
||||
headers: Headers
|
||||
): AiUsage | null {
|
||||
try {
|
||||
return EXTRACTORS[capability](responseText, isStream, headers);
|
||||
} catch (error) {
|
||||
logger.debug("Failed to extract AI usage from response", {
|
||||
capability,
|
||||
error
|
||||
});
|
||||
return null;
|
||||
}
|
||||
}
|
||||
|
||||
/**
|
||||
* Best-guess token estimate for when the provider doesn't report usage.
|
||||
* Uses OpenAI's BPE tokenizer as a stand-in for whatever tokenizer the
|
||||
* actual model uses - close enough for an approximate cost figure, not
|
||||
* exact for non-OpenAI models.
|
||||
*/
|
||||
export function estimateUsage(
|
||||
promptText: string,
|
||||
completionText: string
|
||||
): AiUsage {
|
||||
const usage = emptyUsage();
|
||||
usage.estimated = true;
|
||||
try {
|
||||
usage.promptTokens = promptText ? encode(promptText).length : 0;
|
||||
} catch (error) {
|
||||
logger.debug("Failed to estimate prompt tokens", { error });
|
||||
}
|
||||
try {
|
||||
usage.completionTokens = completionText
|
||||
? encode(completionText).length
|
||||
: 0;
|
||||
} catch (error) {
|
||||
logger.debug("Failed to estimate completion tokens", { error });
|
||||
}
|
||||
return usage;
|
||||
}
|
||||
|
||||
/**
|
||||
* OpenAI's Chat Completions API only includes a `usage` field in a
|
||||
* streaming response when the request opts in via `stream_options:
|
||||
* {include_usage: true}` - unlike the Responses API, Anthropic, Gemini and
|
||||
* Bedrock, which report usage in a streaming response by default. Returns
|
||||
* whether we need to inject that option ourselves to be able to track cost.
|
||||
*/
|
||||
export function needsStreamUsageInjection(
|
||||
capability: AiCapability,
|
||||
body: any
|
||||
): boolean {
|
||||
return (
|
||||
capability === "openai_chat" &&
|
||||
body?.stream === true &&
|
||||
body?.stream_options?.include_usage !== true
|
||||
);
|
||||
}
|
||||
|
||||
/**
|
||||
* Returns a shallow-cloned body with `stream_options.include_usage`
|
||||
* injected, for capabilities/requests where `needsStreamUsageInjection`
|
||||
* is true. Leaves the original body untouched.
|
||||
*/
|
||||
export function withStreamUsageOption(body: any): any {
|
||||
return {
|
||||
...body,
|
||||
stream_options: { ...body.stream_options, include_usage: true }
|
||||
};
|
||||
}
|
||||
|
||||
/**
|
||||
* When we injected stream_options.include_usage ourselves (the caller
|
||||
* didn't ask for it), OpenAI appends an extra terminal SSE frame with an
|
||||
* empty `choices: []` array carrying only the usage data. Callers that
|
||||
* don't expect that shape (most minimal SSE parsers assume a non-empty
|
||||
* choices array) shouldn't see it, so it's stripped back out of the bytes
|
||||
* forwarded to the client.
|
||||
*/
|
||||
export function stripInjectedUsageFrame(sseText: string): string {
|
||||
const parts = sseText.split(/(\r?\n\r?\n)/);
|
||||
let out = "";
|
||||
for (let i = 0; i < parts.length; i += 2) {
|
||||
const frame = parts[i];
|
||||
const separator = parts[i + 1] ?? "";
|
||||
const dataLine = frame
|
||||
.split(/\r?\n/)
|
||||
.find((line) => line.startsWith("data:"));
|
||||
if (dataLine) {
|
||||
const data = dataLine.slice("data:".length).trim();
|
||||
const parsed = data !== "[DONE]" ? tryParseJson(data) : null;
|
||||
if (
|
||||
parsed &&
|
||||
Array.isArray(parsed.choices) &&
|
||||
parsed.choices.length === 0 &&
|
||||
parsed.usage
|
||||
) {
|
||||
continue;
|
||||
}
|
||||
}
|
||||
out += frame + separator;
|
||||
}
|
||||
return out;
|
||||
}
|
||||
|
||||
/**
|
||||
* Best-effort extraction of the model the upstream provider actually
|
||||
* served, which some gateways/routers echo back and which may differ from
|
||||
* the model the caller requested (e.g. an alias resolving to a dated
|
||||
* snapshot). Falls back to the caller's requested model when absent.
|
||||
*/
|
||||
export function extractResponseModel(responseText: string): string | null {
|
||||
const match = responseText.match(/"model"\s*:\s*"([^"]+)"/);
|
||||
return match ? match[1] : null;
|
||||
}
|
||||
|
||||
export function isUsageEmpty(usage: AiUsage): boolean {
|
||||
return (
|
||||
usage.promptTokens === 0 &&
|
||||
usage.cacheReadTokens === 0 &&
|
||||
usage.cacheWriteTokens === 0 &&
|
||||
usage.completionTokens === 0 &&
|
||||
usage.reasoningTokens === 0
|
||||
);
|
||||
}
|
||||
@@ -202,6 +202,10 @@ async function handleResource(
|
||||
return;
|
||||
}
|
||||
|
||||
if (!target.resourceId) {
|
||||
return;
|
||||
}
|
||||
|
||||
const [resource] = await trx
|
||||
.select()
|
||||
.from(resources)
|
||||
@@ -227,9 +231,7 @@ async function handleResource(
|
||||
|
||||
let health = "healthy";
|
||||
const allUnknown = monitoredTargets.length === 0;
|
||||
const allHealthy = monitoredTargets.every(
|
||||
(t) => t.hcHealth === "healthy"
|
||||
);
|
||||
const allHealthy = monitoredTargets.every((t) => t.hcHealth === "healthy");
|
||||
const allUnhealthy = monitoredTargets.every(
|
||||
(t) => t.hcHealth === "unhealthy"
|
||||
);
|
||||
|
||||
@@ -9,8 +9,9 @@ export enum TierFeature {
|
||||
AccessLogs = "accessLogs", // set the retention period to none on downgrade
|
||||
ActionLogs = "actionLogs", // set the retention period to none on downgrade
|
||||
ConnectionLogs = "connectionLogs",
|
||||
AISessionLogs = "aiSessionLogs",
|
||||
RotateCredentials = "rotateCredentials",
|
||||
MaintencePage = "maintencePage", // handle downgrade
|
||||
MaintenancePage = "maintenancePage", // handle downgrade
|
||||
DevicePosture = "devicePosture",
|
||||
TwoFactorEnforcement = "twoFactorEnforcement", // handle downgrade by setting to optional
|
||||
SessionDurationPolicies = "sessionDurationPolicies", // handle downgrade by setting to default duration
|
||||
@@ -23,15 +24,12 @@ export enum TierFeature {
|
||||
StandaloneHealthChecks = "standaloneHealthChecks",
|
||||
AlertingRules = "alertingRules",
|
||||
WildcardSubdomain = "wildcardSubdomain",
|
||||
Labels = "labels",
|
||||
NewtAutoUpdate = "newtAutoUpdate",
|
||||
ResourcePolicies = "resourcePolicies",
|
||||
AdvancedPublicResources = "advancedPublicResources",
|
||||
AdvancedPrivateResources = "advancedPrivateResources"
|
||||
RoleBasedSSHControls = "roleBasedSSHControls"
|
||||
}
|
||||
|
||||
export const tierMatrix: Record<TierFeature, Tier[]> = {
|
||||
[TierFeature.Labels]: ["tier1", "tier2", "tier3", "enterprise"],
|
||||
[TierFeature.OrgOidc]: ["tier1", "tier2", "tier3", "enterprise"],
|
||||
[TierFeature.LoginPageDomain]: ["tier1", "tier2", "tier3", "enterprise"],
|
||||
[TierFeature.DeviceApprovals]: ["tier1", "tier3", "enterprise"],
|
||||
@@ -40,8 +38,9 @@ export const tierMatrix: Record<TierFeature, Tier[]> = {
|
||||
[TierFeature.AccessLogs]: ["tier2", "tier3", "enterprise"],
|
||||
[TierFeature.ActionLogs]: ["tier2", "tier3", "enterprise"],
|
||||
[TierFeature.ConnectionLogs]: ["tier2", "tier3", "enterprise"],
|
||||
[TierFeature.AISessionLogs]: ["tier2", "tier3", "enterprise"],
|
||||
[TierFeature.RotateCredentials]: ["tier1", "tier2", "tier3", "enterprise"],
|
||||
[TierFeature.MaintencePage]: ["tier1", "tier2", "tier3", "enterprise"],
|
||||
[TierFeature.MaintenancePage]: ["tier1", "tier2", "tier3", "enterprise"],
|
||||
[TierFeature.DevicePosture]: ["tier2", "tier3", "enterprise"],
|
||||
[TierFeature.TwoFactorEnforcement]: [
|
||||
"tier1",
|
||||
@@ -71,6 +70,5 @@ export const tierMatrix: Record<TierFeature, Tier[]> = {
|
||||
[TierFeature.WildcardSubdomain]: ["tier1", "tier2", "tier3", "enterprise"],
|
||||
[TierFeature.NewtAutoUpdate]: ["tier1", "tier2", "tier3", "enterprise"],
|
||||
[TierFeature.ResourcePolicies]: ["tier3", "enterprise"],
|
||||
[TierFeature.AdvancedPublicResources]: ["tier3", "enterprise"],
|
||||
[TierFeature.AdvancedPrivateResources]: ["tier3", "enterprise"]
|
||||
[TierFeature.RoleBasedSSHControls]: ["tier3", "enterprise"]
|
||||
};
|
||||
|
||||
@@ -0,0 +1,89 @@
|
||||
import { eq } from "drizzle-orm";
|
||||
import { aiBudgets, Transaction } from "@server/db";
|
||||
|
||||
export type BlueprintAiBudgetInput = {
|
||||
amount: number;
|
||||
unit: "usd" | "tokens";
|
||||
period:
|
||||
| "monthly"
|
||||
| "yearly"
|
||||
| "lifetime"
|
||||
| "daily"
|
||||
| "hourly"
|
||||
| "weekly";
|
||||
enforcement: "hard" | "soft";
|
||||
enabled: boolean;
|
||||
};
|
||||
|
||||
type SyncAiBudgetsInput = {
|
||||
orgId: string;
|
||||
trx: Transaction;
|
||||
budgets: BlueprintAiBudgetInput[];
|
||||
} & (
|
||||
| { scope: "public"; resourceId: number }
|
||||
| { scope: "site"; siteResourceId: number }
|
||||
);
|
||||
|
||||
/**
|
||||
* Fully declarative: makes the resource's/site resource's AI budgets match
|
||||
* exactly what the blueprint declares (omitted unit/period budgets are removed).
|
||||
*/
|
||||
export async function syncAiBudgets(input: SyncAiBudgetsInput): Promise<void> {
|
||||
const { orgId, trx, budgets } = input;
|
||||
|
||||
const existing = await trx
|
||||
.select()
|
||||
.from(aiBudgets)
|
||||
.where(
|
||||
input.scope === "public"
|
||||
? eq(aiBudgets.resourceId, input.resourceId)
|
||||
: eq(aiBudgets.siteResourceId, input.siteResourceId)
|
||||
);
|
||||
|
||||
const existingByKey = new Map(
|
||||
existing.map((b) => [`${b.unit}::${b.period}`, b])
|
||||
);
|
||||
|
||||
const seenKeys = new Set<string>();
|
||||
const now = Date.now();
|
||||
|
||||
for (const budget of budgets) {
|
||||
const key = `${budget.unit}::${budget.period}`;
|
||||
seenKeys.add(key);
|
||||
const existingBudget = existingByKey.get(key);
|
||||
|
||||
if (existingBudget) {
|
||||
await trx
|
||||
.update(aiBudgets)
|
||||
.set({
|
||||
amount: budget.amount,
|
||||
enforcement: budget.enforcement,
|
||||
enabled: budget.enabled,
|
||||
updatedAt: now
|
||||
})
|
||||
.where(eq(aiBudgets.budgetId, existingBudget.budgetId));
|
||||
} else {
|
||||
await trx.insert(aiBudgets).values({
|
||||
orgId,
|
||||
resourceId: input.scope === "public" ? input.resourceId : null,
|
||||
siteResourceId:
|
||||
input.scope === "site" ? input.siteResourceId : null,
|
||||
amount: budget.amount,
|
||||
unit: budget.unit,
|
||||
period: budget.period,
|
||||
enforcement: budget.enforcement,
|
||||
enabled: budget.enabled,
|
||||
createdAt: now,
|
||||
updatedAt: now
|
||||
});
|
||||
}
|
||||
}
|
||||
|
||||
for (const [key, existingBudget] of existingByKey) {
|
||||
if (!seenKeys.has(key)) {
|
||||
await trx
|
||||
.delete(aiBudgets)
|
||||
.where(eq(aiBudgets.budgetId, existingBudget.budgetId));
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,280 @@
|
||||
import { and, eq, inArray } from "drizzle-orm";
|
||||
import {
|
||||
aiModels,
|
||||
aiProviders,
|
||||
resourceAiModels,
|
||||
siteResourceAiModels,
|
||||
Transaction
|
||||
} from "@server/db";
|
||||
import {
|
||||
AccessMode,
|
||||
ModelListType,
|
||||
clearPublicResourceAiConfig,
|
||||
clearSiteResourceAiConfig,
|
||||
isInferenceFieldsError,
|
||||
resolveProviderAttachments,
|
||||
setPublicResourceAiProviders,
|
||||
setSiteResourceAiProviders
|
||||
} from "@server/lib/aiInferenceResource";
|
||||
|
||||
export type BlueprintAiModelInput = {
|
||||
model: string;
|
||||
listType: ModelListType;
|
||||
};
|
||||
|
||||
export type BlueprintAiProviderInput = {
|
||||
provider: string;
|
||||
accessMode: AccessMode;
|
||||
enabled: boolean;
|
||||
models: string[];
|
||||
};
|
||||
|
||||
async function resolveProviderNiceIds(
|
||||
orgId: string,
|
||||
niceIds: string[],
|
||||
trx: Transaction
|
||||
): Promise<Map<string, number>> {
|
||||
const unique = [...new Set(niceIds)];
|
||||
if (unique.length === 0) {
|
||||
return new Map();
|
||||
}
|
||||
|
||||
const rows = await trx
|
||||
.select({
|
||||
providerId: aiProviders.providerId,
|
||||
niceId: aiProviders.niceId
|
||||
})
|
||||
.from(aiProviders)
|
||||
.where(
|
||||
and(
|
||||
eq(aiProviders.orgId, orgId),
|
||||
inArray(aiProviders.niceId, unique)
|
||||
)
|
||||
);
|
||||
|
||||
const byNiceId = new Map(rows.map((r) => [r.niceId, r.providerId]));
|
||||
const missing = unique.filter((id) => !byNiceId.has(id));
|
||||
if (missing.length > 0) {
|
||||
throw new Error(
|
||||
`AI provider(s) not found in this org: ${missing.join(", ")}`
|
||||
);
|
||||
}
|
||||
|
||||
return byNiceId;
|
||||
}
|
||||
|
||||
async function resolveModelKeys(
|
||||
providers: BlueprintAiProviderInput[],
|
||||
providerIdByNiceId: Map<string, number>,
|
||||
trx: Transaction
|
||||
): Promise<Map<string, number>> {
|
||||
const providerIds = [
|
||||
...new Set(
|
||||
providers
|
||||
.filter((p) => p.models.length > 0)
|
||||
.map((p) => providerIdByNiceId.get(p.provider)!)
|
||||
)
|
||||
];
|
||||
|
||||
if (providerIds.length === 0) {
|
||||
return new Map();
|
||||
}
|
||||
|
||||
const rows = await trx
|
||||
.select({
|
||||
modelId: aiModels.modelId,
|
||||
modelKey: aiModels.modelKey,
|
||||
providerId: aiModels.providerId
|
||||
})
|
||||
.from(aiModels)
|
||||
.where(inArray(aiModels.providerId, providerIds));
|
||||
|
||||
const byProviderAndKey = new Map<string, number>();
|
||||
for (const row of rows) {
|
||||
byProviderAndKey.set(`${row.providerId}::${row.modelKey}`, row.modelId);
|
||||
}
|
||||
|
||||
const modelIdByEntryKey = new Map<string, number>();
|
||||
const missing: string[] = [];
|
||||
for (const provider of providers) {
|
||||
const providerId = providerIdByNiceId.get(provider.provider)!;
|
||||
for (const m of provider.models) {
|
||||
const modelId = byProviderAndKey.get(`${providerId}::${m}`);
|
||||
if (modelId === undefined) {
|
||||
missing.push(`${provider.provider}/${m}`);
|
||||
continue;
|
||||
}
|
||||
modelIdByEntryKey.set(`${provider.provider}::${m}`, modelId);
|
||||
}
|
||||
}
|
||||
|
||||
if (missing.length > 0) {
|
||||
throw new Error(`AI model(s) not found: ${missing.join(", ")}`);
|
||||
}
|
||||
|
||||
return modelIdByEntryKey;
|
||||
}
|
||||
|
||||
async function validateModelEntries(input: {
|
||||
orgId: string;
|
||||
entries: { modelId: number }[];
|
||||
selectProviderIds: number[];
|
||||
trx: Transaction;
|
||||
}): Promise<void> {
|
||||
if (input.entries.length === 0) {
|
||||
return;
|
||||
}
|
||||
|
||||
if (input.selectProviderIds.length === 0) {
|
||||
throw new Error(
|
||||
"Set at least one attached AI provider to access-mode 'select' before declaring models"
|
||||
);
|
||||
}
|
||||
|
||||
const modelIds = input.entries.map((e) => e.modelId);
|
||||
const catalogRows = await input.trx
|
||||
.select({
|
||||
modelId: aiModels.modelId,
|
||||
listType: aiModels.listType,
|
||||
providerId: aiModels.providerId,
|
||||
enabled: aiModels.enabled
|
||||
})
|
||||
.from(aiModels)
|
||||
.innerJoin(aiProviders, eq(aiModels.providerId, aiProviders.providerId))
|
||||
.where(
|
||||
and(
|
||||
inArray(aiModels.modelId, modelIds),
|
||||
inArray(aiModels.providerId, input.selectProviderIds),
|
||||
eq(aiProviders.orgId, input.orgId)
|
||||
)
|
||||
);
|
||||
|
||||
const catalogById = new Map(catalogRows.map((row) => [row.modelId, row]));
|
||||
for (const entry of input.entries) {
|
||||
const catalog = catalogById.get(entry.modelId);
|
||||
if (!catalog) {
|
||||
throw new Error(
|
||||
`Model ${entry.modelId} does not exist or does not belong to a select-mode attached provider`
|
||||
);
|
||||
}
|
||||
if (!catalog.enabled) {
|
||||
throw new Error(
|
||||
`Model ${entry.modelId} is disabled on its provider`
|
||||
);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
type SyncInferenceAiConfigInput = {
|
||||
orgId: string;
|
||||
trx: Transaction;
|
||||
mode: string;
|
||||
providers: BlueprintAiProviderInput[];
|
||||
} & (
|
||||
| { scope: "public"; resourceId: number }
|
||||
| { scope: "site"; siteResourceId: number }
|
||||
);
|
||||
|
||||
/**
|
||||
* Fully declarative: makes the resource's attached AI providers/models match
|
||||
* exactly what the blueprint declares (omitted providers/models are removed).
|
||||
* Non-inference resources have any leftover AI config cleared.
|
||||
*/
|
||||
export async function syncInferenceAiConfig(
|
||||
input: SyncInferenceAiConfigInput
|
||||
): Promise<void> {
|
||||
const { orgId, trx, mode } = input;
|
||||
|
||||
if (mode !== "inference") {
|
||||
if (input.scope === "public") {
|
||||
await clearPublicResourceAiConfig(input.resourceId, trx);
|
||||
} else {
|
||||
await clearSiteResourceAiConfig(input.siteResourceId, trx);
|
||||
}
|
||||
return;
|
||||
}
|
||||
|
||||
const providerIdByNiceId = await resolveProviderNiceIds(
|
||||
orgId,
|
||||
input.providers.map((p) => p.provider),
|
||||
trx
|
||||
);
|
||||
|
||||
const resolvedAttachments = await resolveProviderAttachments({
|
||||
orgId,
|
||||
attachments: input.providers.map((p) => ({
|
||||
providerId: providerIdByNiceId.get(p.provider)!,
|
||||
accessMode: p.accessMode,
|
||||
enabled: p.enabled
|
||||
})),
|
||||
requireAtLeastOne: false
|
||||
});
|
||||
if (isInferenceFieldsError(resolvedAttachments)) {
|
||||
throw new Error(resolvedAttachments.error);
|
||||
}
|
||||
|
||||
if (input.scope === "public") {
|
||||
await setPublicResourceAiProviders(
|
||||
input.resourceId,
|
||||
resolvedAttachments,
|
||||
trx
|
||||
);
|
||||
} else {
|
||||
await setSiteResourceAiProviders(
|
||||
input.siteResourceId,
|
||||
resolvedAttachments,
|
||||
trx
|
||||
);
|
||||
}
|
||||
|
||||
const modelIdByEntryKey = await resolveModelKeys(
|
||||
input.providers,
|
||||
providerIdByNiceId,
|
||||
trx
|
||||
);
|
||||
|
||||
const modelEntries = input.providers.flatMap((p) =>
|
||||
p.models.map((m) => ({
|
||||
modelId: modelIdByEntryKey.get(`${p.provider}::${m}`)!
|
||||
}))
|
||||
);
|
||||
|
||||
const selectProviderIds = resolvedAttachments
|
||||
.filter((a) => a.accessMode === "select")
|
||||
.map((a) => a.providerId);
|
||||
|
||||
await validateModelEntries({
|
||||
orgId,
|
||||
entries: modelEntries,
|
||||
selectProviderIds,
|
||||
trx
|
||||
});
|
||||
|
||||
if (input.scope === "public") {
|
||||
await trx
|
||||
.delete(resourceAiModels)
|
||||
.where(eq(resourceAiModels.resourceId, input.resourceId));
|
||||
if (modelEntries.length > 0) {
|
||||
await trx.insert(resourceAiModels).values(
|
||||
modelEntries.map((m) => ({
|
||||
resourceId: input.resourceId,
|
||||
modelId: m.modelId
|
||||
}))
|
||||
);
|
||||
}
|
||||
} else {
|
||||
await trx
|
||||
.delete(siteResourceAiModels)
|
||||
.where(
|
||||
eq(siteResourceAiModels.siteResourceId, input.siteResourceId)
|
||||
);
|
||||
if (modelEntries.length > 0) {
|
||||
await trx.insert(siteResourceAiModels).values(
|
||||
modelEntries.map((m) => ({
|
||||
siteResourceId: input.siteResourceId,
|
||||
modelId: m.modelId
|
||||
}))
|
||||
);
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -1,5 +1,6 @@
|
||||
import {
|
||||
db,
|
||||
primaryDb,
|
||||
newts,
|
||||
blueprints,
|
||||
Blueprint,
|
||||
@@ -34,12 +35,6 @@ import {
|
||||
rebuildClientAssociationsFromSiteResource,
|
||||
waitForSiteResourceRebuildIdle
|
||||
} from "../rebuildClientAssociations";
|
||||
import { build } from "@server/build";
|
||||
import HttpCode from "@server/types/HttpCode";
|
||||
import createHttpError from "http-errors";
|
||||
import next from "next";
|
||||
import { LimitId } from "../billing";
|
||||
import { usageService } from "../billing/usageService";
|
||||
|
||||
type ApplyBlueprintArgs = {
|
||||
orgId: string;
|
||||
@@ -86,93 +81,103 @@ export async function applyBlueprint({
|
||||
trx,
|
||||
siteId
|
||||
);
|
||||
});
|
||||
|
||||
// We need to update the targets on the newts from the successfully updated information
|
||||
for (const result of publicResourcesResults) {
|
||||
for (const target of result.targetsToUpdate) {
|
||||
const [site] = await trx
|
||||
.select()
|
||||
.from(sites)
|
||||
.innerJoin(newts, eq(sites.siteId, newts.siteId))
|
||||
.where(
|
||||
and(
|
||||
eq(sites.siteId, target.siteId),
|
||||
eq(sites.orgId, orgId),
|
||||
eq(sites.type, "newt"),
|
||||
isNotNull(sites.pubKey)
|
||||
)
|
||||
// Push updates to newts/clients only after the transaction has
|
||||
// committed. Doing this while the transaction is still open can
|
||||
// race with the writes (e.g. newts requesting config before the
|
||||
// new targets/resources are actually visible), leaving them out
|
||||
// of sync until manually toggled.
|
||||
|
||||
// We need to update the targets on the newts from the successfully updated information
|
||||
for (const result of publicResourcesResults) {
|
||||
for (const target of result.targetsToUpdate) {
|
||||
// read from the primary: this determines whether/how we push
|
||||
// the just-created target to the newt, so a lagging replica
|
||||
// returning stale or missing data here would silently skip
|
||||
// the push
|
||||
const [site] = await primaryDb
|
||||
.select()
|
||||
.from(sites)
|
||||
.innerJoin(newts, eq(sites.siteId, newts.siteId))
|
||||
.where(
|
||||
and(
|
||||
eq(sites.siteId, target.siteId),
|
||||
eq(sites.orgId, orgId),
|
||||
eq(sites.type, "newt"),
|
||||
isNotNull(sites.pubKey)
|
||||
)
|
||||
.limit(1);
|
||||
)
|
||||
.limit(1);
|
||||
|
||||
if (site) {
|
||||
logger.debug(
|
||||
`Updating target ${target.targetId} on site ${site.sites.siteId}`
|
||||
if (site) {
|
||||
logger.debug(
|
||||
`Updating target ${target.targetId} on site ${site.sites.siteId}`
|
||||
);
|
||||
|
||||
// see if you can find a matching target health check from the healthchecksToUpdate array
|
||||
const matchingHealthcheck =
|
||||
result.healthchecksToUpdate.find(
|
||||
(hc) => hc.targetId === target.targetId
|
||||
);
|
||||
|
||||
// see if you can find a matching target health check from the healthchecksToUpdate array
|
||||
const matchingHealthcheck =
|
||||
result.healthchecksToUpdate.find(
|
||||
(hc) => hc.targetId === target.targetId
|
||||
);
|
||||
|
||||
if (["http", "tcp", "udp"].includes(target.mode)) {
|
||||
await addProxyTargets(
|
||||
site.newt.newtId,
|
||||
[target],
|
||||
matchingHealthcheck
|
||||
? [matchingHealthcheck]
|
||||
: [],
|
||||
result.proxyResource.mode === "udp"
|
||||
? "udp"
|
||||
: "tcp",
|
||||
site.newt.version
|
||||
);
|
||||
} else if (
|
||||
["ssh", "rdp", "vnc"].includes(target.mode)
|
||||
) {
|
||||
await sendBrowserGatewayTargets(
|
||||
site.newt.newtId,
|
||||
[target],
|
||||
site.newt.version
|
||||
);
|
||||
}
|
||||
if (["http", "tcp", "udp"].includes(target.mode)) {
|
||||
await addProxyTargets(
|
||||
site.newt.newtId,
|
||||
[target],
|
||||
matchingHealthcheck
|
||||
? [matchingHealthcheck]
|
||||
: [],
|
||||
result.proxyResource.mode === "udp"
|
||||
? "udp"
|
||||
: "tcp",
|
||||
site.newt.version
|
||||
);
|
||||
} else if (
|
||||
["ssh", "rdp", "vnc"].includes(target.mode)
|
||||
) {
|
||||
await sendBrowserGatewayTargets(
|
||||
site.newt.newtId,
|
||||
[target],
|
||||
site.newt.version
|
||||
);
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
logger.debug(
|
||||
`Successfully updated public resources for org ${orgId}: ${JSON.stringify(publicResourcesResults)}`
|
||||
);
|
||||
logger.debug(
|
||||
`Successfully updated public resources for org ${orgId}: ${JSON.stringify(publicResourcesResults)}`
|
||||
);
|
||||
|
||||
// We need to update the targets on the newts from the successfully updated information
|
||||
for (const result of privateResourcesResults) {
|
||||
rebuildClientAssociationsFromSiteResource(
|
||||
result.newSiteResource
|
||||
// We need to update the targets on the newts from the successfully updated information
|
||||
for (const result of privateResourcesResults) {
|
||||
rebuildClientAssociationsFromSiteResource(
|
||||
result.newSiteResource
|
||||
)
|
||||
.then(() =>
|
||||
waitForSiteResourceRebuildIdle(
|
||||
result.newSiteResource.siteResourceId
|
||||
)
|
||||
)
|
||||
.then(() =>
|
||||
waitForSiteResourceRebuildIdle(
|
||||
result.newSiteResource.siteResourceId
|
||||
)
|
||||
.then(() =>
|
||||
handleMessagingForUpdatedSiteResource(
|
||||
result.oldSiteResource,
|
||||
result.newSiteResource,
|
||||
result.oldSites.map((s) => s.siteId),
|
||||
result.newSites.map((s) => s.siteId)
|
||||
)
|
||||
.then(() =>
|
||||
handleMessagingForUpdatedSiteResource(
|
||||
result.oldSiteResource,
|
||||
result.newSiteResource,
|
||||
result.oldSites.map((s) => s.siteId),
|
||||
result.newSites.map((s) => s.siteId)
|
||||
)
|
||||
)
|
||||
.catch((e) => {
|
||||
logger.error(
|
||||
`Failed to rebuild and handle messaging for site resource ${result.newSiteResource.siteResourceId}. Error: ${e}`
|
||||
);
|
||||
});
|
||||
}
|
||||
)
|
||||
.catch((e) => {
|
||||
logger.error(
|
||||
`Failed to rebuild and handle messaging for site resource ${result.newSiteResource.siteResourceId}. Error: ${e}`
|
||||
);
|
||||
});
|
||||
}
|
||||
|
||||
logger.debug(
|
||||
`Successfully updated private resources for org ${orgId}: ${JSON.stringify(privateResourcesResults)}`
|
||||
);
|
||||
});
|
||||
logger.debug(
|
||||
`Successfully updated private resources for org ${orgId}: ${JSON.stringify(privateResourcesResults)}`
|
||||
);
|
||||
|
||||
blueprintSucceeded = true;
|
||||
blueprintMessage = "Blueprint applied successfully";
|
||||
|
||||
@@ -0,0 +1,86 @@
|
||||
import {
|
||||
labels,
|
||||
resourceLabels,
|
||||
siteResourceLabels,
|
||||
Transaction
|
||||
} from "@server/db";
|
||||
import logger from "@server/logger";
|
||||
import { and, eq, sql } from "drizzle-orm";
|
||||
|
||||
// Matches the "gray" swatch in the label color palette used by the UI
|
||||
// (src/components/labels-selector.tsx), used as the default for labels
|
||||
// auto-created from a blueprint where no color is specified.
|
||||
const DEFAULT_LABEL_COLOR = "#b4b4b4";
|
||||
|
||||
/**
|
||||
* Looks up labels by name (case-insensitive) within an org, auto-creating
|
||||
* any that don't already exist. Returns the resolved, de-duplicated labelIds.
|
||||
*/
|
||||
export async function getOrCreateLabelIds(
|
||||
orgId: string,
|
||||
labelNames: string[],
|
||||
trx: Transaction
|
||||
): Promise<number[]> {
|
||||
const labelIds = new Set<number>();
|
||||
|
||||
for (const name of labelNames) {
|
||||
let [label] = await trx
|
||||
.select({ labelId: labels.labelId })
|
||||
.from(labels)
|
||||
.where(
|
||||
and(
|
||||
eq(labels.orgId, orgId),
|
||||
sql`LOWER(${labels.name}) = ${name.toLowerCase()}`
|
||||
)
|
||||
)
|
||||
.limit(1);
|
||||
|
||||
if (!label) {
|
||||
[label] = await trx
|
||||
.insert(labels)
|
||||
.values({ name, color: DEFAULT_LABEL_COLOR, orgId })
|
||||
.returning({ labelId: labels.labelId });
|
||||
logger.info(
|
||||
`Auto-created label "${name}" in org ${orgId} from blueprint`
|
||||
);
|
||||
}
|
||||
|
||||
labelIds.add(label.labelId);
|
||||
}
|
||||
|
||||
return Array.from(labelIds);
|
||||
}
|
||||
|
||||
export async function syncResourceLabels(
|
||||
resourceId: number,
|
||||
labelIds: number[],
|
||||
trx: Transaction
|
||||
) {
|
||||
await trx
|
||||
.delete(resourceLabels)
|
||||
.where(eq(resourceLabels.resourceId, resourceId));
|
||||
|
||||
if (labelIds.length > 0) {
|
||||
await trx
|
||||
.insert(resourceLabels)
|
||||
.values(labelIds.map((labelId) => ({ resourceId, labelId })));
|
||||
}
|
||||
}
|
||||
|
||||
export async function syncSiteResourceLabels(
|
||||
siteResourceId: number,
|
||||
labelIds: number[],
|
||||
trx: Transaction
|
||||
) {
|
||||
await trx
|
||||
.delete(siteResourceLabels)
|
||||
.where(eq(siteResourceLabels.siteResourceId, siteResourceId));
|
||||
|
||||
if (labelIds.length > 0) {
|
||||
await trx
|
||||
.insert(siteResourceLabels)
|
||||
.values(
|
||||
labelIds.map((labelId) => ({ siteResourceId, labelId }))
|
||||
);
|
||||
}
|
||||
}
|
||||
@@ -19,23 +19,22 @@ import {
|
||||
import { sites } from "@server/db";
|
||||
import { eq, and, ne, inArray, or, isNotNull } from "drizzle-orm";
|
||||
import { Config } from "./types";
|
||||
import { getOrCreateLabelIds, syncSiteResourceLabels } from "./labels";
|
||||
import logger from "@server/logger";
|
||||
import { defaultRoleAllowedActions } from "@server/routers/role/createRole";
|
||||
import { getNextAvailableAliasAddress } from "../ip";
|
||||
import { createCertificate } from "#dynamic/routers/certificates/createCertificate";
|
||||
import { isLicensedOrSubscribed } from "#dynamic/lib/isLicencedOrSubscribed";
|
||||
import { tierMatrix } from "../billing/tierMatrix";
|
||||
import { createCertificate } from "@server/routers/certificates/createCertificate";
|
||||
import { build } from "@server/build";
|
||||
import HttpCode from "@server/types/HttpCode";
|
||||
import createHttpError from "http-errors";
|
||||
import next from "next";
|
||||
import { LimitId } from "../billing";
|
||||
import { usageService } from "../billing/usageService";
|
||||
import { syncInferenceAiConfig } from "./aiProviders";
|
||||
import { syncAiBudgets } from "./aiBudgets";
|
||||
|
||||
async function getDomainForSiteResource(
|
||||
siteResourceId: number | undefined,
|
||||
fullDomain: string,
|
||||
orgId: string,
|
||||
isInference: boolean,
|
||||
trx: Transaction
|
||||
): Promise<{ subdomain: string | null; domainId: string }> {
|
||||
const [fullDomainExists] = await trx
|
||||
@@ -45,6 +44,11 @@ async function getDomainForSiteResource(
|
||||
and(
|
||||
eq(siteResources.fullDomain, fullDomain),
|
||||
eq(siteResources.orgId, orgId),
|
||||
// exclude looking at the ones on exit nodes if this is an inference resource,
|
||||
// and vice versa, so inference and non-inference resources can share a full-domain
|
||||
isInference
|
||||
? ne(siteResources.mode, "inference")
|
||||
: eq(siteResources.mode, "inference"),
|
||||
siteResourceId
|
||||
? ne(siteResources.siteResourceId, siteResourceId)
|
||||
: isNotNull(siteResources.siteResourceId)
|
||||
@@ -122,30 +126,6 @@ export async function updatePrivateResources(
|
||||
for (const [resourceNiceId, resourceData] of Object.entries(
|
||||
config["client-resources"]
|
||||
)) {
|
||||
if (resourceData.mode === "http") {
|
||||
const hasHttpFeature = await isLicensedOrSubscribed(
|
||||
orgId,
|
||||
tierMatrix.advancedPrivateResources
|
||||
);
|
||||
if (!hasHttpFeature) {
|
||||
throw new Error(
|
||||
"HTTP private resources are not included in your current plan. Please upgrade."
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
if (resourceData.mode === "ssh") {
|
||||
const hasSshFeature = await isLicensedOrSubscribed(
|
||||
orgId,
|
||||
tierMatrix.advancedPrivateResources
|
||||
);
|
||||
if (!hasSshFeature) {
|
||||
throw new Error(
|
||||
"SSH private resources are not included in your current plan. Please upgrade."
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
const [existingResource] = await trx
|
||||
.select()
|
||||
.from(siteResources)
|
||||
@@ -201,60 +181,109 @@ export async function updatePrivateResources(
|
||||
}
|
||||
}
|
||||
|
||||
let resourceStatusFromSite: "approved" | "pending" = "approved";
|
||||
if (siteId && allSites.length === 0) {
|
||||
// only add if there are not provided sites
|
||||
// Use the provided siteId directly, but verify it belongs to the org
|
||||
const [siteSingle] = await trx
|
||||
.select({ siteId: sites.siteId })
|
||||
.select({ siteId: sites.siteId, status: sites.status })
|
||||
.from(sites)
|
||||
.where(and(eq(sites.siteId, siteId), eq(sites.orgId, orgId)))
|
||||
.limit(1);
|
||||
if (siteSingle) {
|
||||
allSites.push(siteSingle);
|
||||
}
|
||||
resourceStatusFromSite = siteSingle.status ?? "approved";
|
||||
}
|
||||
|
||||
if (allSites.length === 0) {
|
||||
if (resourceData.mode !== "inference" && allSites.length === 0) {
|
||||
throw new Error(
|
||||
`No valid sites found for private private resource ${resourceNiceId} in org ${orgId}`
|
||||
);
|
||||
}
|
||||
|
||||
const resourceEnabled =
|
||||
resourceData.enabled == undefined || resourceData.enabled == null
|
||||
? true
|
||||
: resourceStatusFromSite === "pending"
|
||||
? false
|
||||
: resourceData.enabled;
|
||||
|
||||
const resourceSsl =
|
||||
resourceData.mode === "inference" || resourceData.mode === "http"
|
||||
? resourceData.ssl == undefined || resourceData.ssl == null
|
||||
? true
|
||||
: resourceData.ssl
|
||||
: resourceData.ssl;
|
||||
|
||||
if (existingResource) {
|
||||
let domainInfo:
|
||||
| { subdomain: string | null; domainId: string }
|
||||
| undefined;
|
||||
if (resourceData["full-domain"] && resourceData.mode === "http") {
|
||||
if (
|
||||
resourceData["full-domain"] &&
|
||||
(resourceData.mode === "http" ||
|
||||
resourceData.mode === "inference")
|
||||
) {
|
||||
domainInfo = await getDomainForSiteResource(
|
||||
existingResource.siteResourceId,
|
||||
resourceData["full-domain"],
|
||||
orgId,
|
||||
resourceData.mode === "inference",
|
||||
trx
|
||||
);
|
||||
}
|
||||
|
||||
if (resourceData.alias) {
|
||||
const [aliasConflict] = await trx
|
||||
.select({
|
||||
siteResourceId: siteResources.siteResourceId
|
||||
})
|
||||
.from(siteResources)
|
||||
.where(
|
||||
and(
|
||||
eq(siteResources.orgId, orgId),
|
||||
eq(siteResources.alias, resourceData.alias),
|
||||
ne(
|
||||
siteResources.siteResourceId,
|
||||
existingResource.siteResourceId
|
||||
)
|
||||
)
|
||||
)
|
||||
.limit(1);
|
||||
|
||||
if (aliasConflict) {
|
||||
throw new Error(
|
||||
`Alias ${resourceData.alias} already in use by another site resource in org ${orgId}`
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
const isInference = resourceData.mode === "inference";
|
||||
|
||||
// Update existing resource
|
||||
const [updatedResource] = await trx
|
||||
.update(siteResources)
|
||||
.set({
|
||||
name: resourceData.name || resourceNiceId,
|
||||
mode: resourceData.mode,
|
||||
ssl: resourceData.ssl,
|
||||
ssl: resourceSsl,
|
||||
scheme: resourceData.scheme,
|
||||
destination: resourceData.destination,
|
||||
destinationPort: resourceData["destination-port"],
|
||||
enabled: true, // hardcoded for now
|
||||
// enabled: resourceData.enabled ?? true,
|
||||
enabled: resourceEnabled,
|
||||
alias: resourceData.alias || null,
|
||||
disableIcmp:
|
||||
resourceData["disable-icmp"] ||
|
||||
(resourceData.mode == "http" ? true : false), // default to true for http resources, otherwise false
|
||||
(resourceData.mode == "http" || isInference
|
||||
? true
|
||||
: false), // default to true for http/inference resources, otherwise false
|
||||
tcpPortRangeString:
|
||||
resourceData.mode == "http"
|
||||
resourceData.mode == "http" || isInference
|
||||
? "443,80"
|
||||
: resourceData["tcp-ports"],
|
||||
udpPortRangeString:
|
||||
resourceData.mode == "http"
|
||||
resourceData.mode == "http" || isInference
|
||||
? ""
|
||||
: resourceData["udp-ports"],
|
||||
fullDomain: resourceData["full-domain"] || null,
|
||||
@@ -263,7 +292,10 @@ export async function updatePrivateResources(
|
||||
pamMode: resourceData["auth-daemon"]?.pam || "passthrough",
|
||||
authDaemonMode:
|
||||
resourceData["auth-daemon"]?.mode || "native",
|
||||
authDaemonPort: resourceData["auth-daemon"]?.port || 22123
|
||||
authDaemonPort: resourceData["auth-daemon"]?.port || 22123,
|
||||
status: resourceStatusFromSite,
|
||||
networkId: isInference ? null : undefined,
|
||||
requiresExitNodeConnection: isInference
|
||||
})
|
||||
.where(
|
||||
eq(
|
||||
@@ -275,7 +307,19 @@ export async function updatePrivateResources(
|
||||
|
||||
const siteResourceId = existingResource.siteResourceId;
|
||||
|
||||
if (updatedResource.networkId) {
|
||||
if (isInference) {
|
||||
// inference resources are not attached to any site network
|
||||
if (existingResource.networkId) {
|
||||
await trx
|
||||
.delete(siteNetworks)
|
||||
.where(
|
||||
eq(
|
||||
siteNetworks.networkId,
|
||||
existingResource.networkId
|
||||
)
|
||||
);
|
||||
}
|
||||
} else if (updatedResource.networkId) {
|
||||
await trx
|
||||
.delete(siteNetworks)
|
||||
.where(
|
||||
@@ -290,6 +334,28 @@ export async function updatePrivateResources(
|
||||
}
|
||||
}
|
||||
|
||||
await syncInferenceAiConfig({
|
||||
orgId,
|
||||
trx,
|
||||
mode: resourceData.mode,
|
||||
scope: "site",
|
||||
siteResourceId,
|
||||
providers: resourceData["ai-providers"].map((p) => ({
|
||||
provider: p.provider,
|
||||
accessMode: p["access-mode"],
|
||||
enabled: p.enabled,
|
||||
models: p.models
|
||||
}))
|
||||
});
|
||||
|
||||
await syncAiBudgets({
|
||||
orgId,
|
||||
trx,
|
||||
scope: "site",
|
||||
siteResourceId,
|
||||
budgets: resourceData["ai-budget"]
|
||||
});
|
||||
|
||||
await trx
|
||||
.delete(clientSiteResources)
|
||||
.where(eq(clientSiteResources.siteResourceId, siteResourceId));
|
||||
@@ -412,6 +478,13 @@ export async function updatePrivateResources(
|
||||
);
|
||||
}
|
||||
|
||||
const labelIds = await getOrCreateLabelIds(
|
||||
orgId,
|
||||
resourceData.labels,
|
||||
trx
|
||||
);
|
||||
await syncSiteResourceLabels(siteResourceId, labelIds, trx);
|
||||
|
||||
results.push({
|
||||
newSiteResource: updatedResource,
|
||||
oldSiteResource: existingResource,
|
||||
@@ -462,25 +535,55 @@ export async function updatePrivateResources(
|
||||
releaseAliasLock = release;
|
||||
}
|
||||
|
||||
const isInference = resourceData.mode === "inference";
|
||||
|
||||
let domainInfo:
|
||||
| { subdomain: string | null; domainId: string }
|
||||
| undefined;
|
||||
if (resourceData["full-domain"] && resourceData.mode === "http") {
|
||||
if (
|
||||
resourceData["full-domain"] &&
|
||||
(resourceData.mode === "http" || isInference)
|
||||
) {
|
||||
domainInfo = await getDomainForSiteResource(
|
||||
undefined,
|
||||
resourceData["full-domain"],
|
||||
orgId,
|
||||
isInference,
|
||||
trx
|
||||
);
|
||||
}
|
||||
|
||||
const [network] = await trx
|
||||
.insert(networks)
|
||||
.values({
|
||||
scope: "resource",
|
||||
orgId: orgId
|
||||
})
|
||||
.returning();
|
||||
if (resourceData.alias) {
|
||||
const [aliasConflict] = await trx
|
||||
.select({
|
||||
siteResourceId: siteResources.siteResourceId
|
||||
})
|
||||
.from(siteResources)
|
||||
.where(
|
||||
and(
|
||||
eq(siteResources.orgId, orgId),
|
||||
eq(siteResources.alias, resourceData.alias)
|
||||
)
|
||||
)
|
||||
.limit(1);
|
||||
|
||||
if (aliasConflict) {
|
||||
throw new Error(
|
||||
`Alias ${resourceData.alias} already in use by another site resource in org ${orgId}`
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
let network: typeof networks.$inferSelect | undefined;
|
||||
if (!isInference) {
|
||||
[network] = await trx
|
||||
.insert(networks)
|
||||
.values({
|
||||
scope: "resource",
|
||||
orgId: orgId
|
||||
})
|
||||
.returning();
|
||||
}
|
||||
|
||||
// Create new resource
|
||||
const [newResource] = await trx
|
||||
@@ -488,27 +591,28 @@ export async function updatePrivateResources(
|
||||
.values({
|
||||
orgId: orgId,
|
||||
niceId: resourceNiceId,
|
||||
networkId: network.networkId,
|
||||
defaultNetworkId: network.networkId,
|
||||
networkId: network ? network.networkId : null,
|
||||
defaultNetworkId: network ? network.networkId : null,
|
||||
name: resourceData.name || resourceNiceId,
|
||||
mode: resourceData.mode,
|
||||
ssl: resourceData.ssl,
|
||||
ssl: resourceSsl,
|
||||
scheme: resourceData.scheme,
|
||||
destination: resourceData.destination,
|
||||
destinationPort: resourceData["destination-port"],
|
||||
enabled: true, // hardcoded for now
|
||||
// enabled: resourceData.enabled ?? true,
|
||||
enabled: resourceEnabled,
|
||||
alias: resourceData.alias || null,
|
||||
aliasAddress: aliasAddress,
|
||||
disableIcmp:
|
||||
resourceData["disable-icmp"] ||
|
||||
(resourceData.mode == "http" ? true : false), // default to true for http resources, otherwise false
|
||||
(resourceData.mode == "http" || isInference
|
||||
? true
|
||||
: false), // default to true for http/inference resources, otherwise false
|
||||
tcpPortRangeString:
|
||||
resourceData.mode == "http"
|
||||
resourceData.mode == "http" || isInference
|
||||
? "443,80"
|
||||
: resourceData["tcp-ports"],
|
||||
udpPortRangeString:
|
||||
resourceData.mode == "http"
|
||||
resourceData.mode == "http" || isInference
|
||||
? ""
|
||||
: resourceData["udp-ports"],
|
||||
fullDomain: resourceData["full-domain"] || null,
|
||||
@@ -517,7 +621,9 @@ export async function updatePrivateResources(
|
||||
pamMode: resourceData["auth-daemon"]?.pam || "passthrough",
|
||||
authDaemonMode:
|
||||
resourceData["auth-daemon"]?.mode || "native",
|
||||
authDaemonPort: resourceData["auth-daemon"]?.port || 22123
|
||||
authDaemonPort: resourceData["auth-daemon"]?.port || 22123,
|
||||
status: resourceStatusFromSite,
|
||||
requiresExitNodeConnection: isInference
|
||||
})
|
||||
.returning();
|
||||
|
||||
@@ -525,13 +631,37 @@ export async function updatePrivateResources(
|
||||
|
||||
const siteResourceId = newResource.siteResourceId;
|
||||
|
||||
for (const site of allSites) {
|
||||
await trx.insert(siteNetworks).values({
|
||||
siteId: site.siteId,
|
||||
networkId: network.networkId
|
||||
});
|
||||
if (network) {
|
||||
for (const site of allSites) {
|
||||
await trx.insert(siteNetworks).values({
|
||||
siteId: site.siteId,
|
||||
networkId: network.networkId
|
||||
});
|
||||
}
|
||||
}
|
||||
|
||||
await syncInferenceAiConfig({
|
||||
orgId,
|
||||
trx,
|
||||
mode: resourceData.mode,
|
||||
scope: "site",
|
||||
siteResourceId,
|
||||
providers: resourceData["ai-providers"].map((p) => ({
|
||||
provider: p.provider,
|
||||
accessMode: p["access-mode"],
|
||||
enabled: p.enabled,
|
||||
models: p.models
|
||||
}))
|
||||
});
|
||||
|
||||
await syncAiBudgets({
|
||||
orgId,
|
||||
trx,
|
||||
scope: "site",
|
||||
siteResourceId,
|
||||
budgets: resourceData["ai-budget"]
|
||||
});
|
||||
|
||||
const [adminRole] = await trx
|
||||
.select()
|
||||
.from(roles)
|
||||
@@ -645,6 +775,13 @@ export async function updatePrivateResources(
|
||||
|
||||
await usageService.add(orgId, LimitId.PRIVATE_RESOURCES, 1, trx);
|
||||
|
||||
const labelIds = await getOrCreateLabelIds(
|
||||
orgId,
|
||||
resourceData.labels,
|
||||
trx
|
||||
);
|
||||
await syncSiteResourceLabels(siteResourceId, labelIds, trx);
|
||||
|
||||
results.push({
|
||||
newSiteResource: newResource,
|
||||
newSites: allSites,
|
||||
|
||||
@@ -1,61 +1,60 @@
|
||||
import { isLicensedOrSubscribed } from "#dynamic/lib/isLicencedOrSubscribed";
|
||||
import { createCertificate } from "@server/routers/certificates/createCertificate";
|
||||
import { hashPassword } from "@server/auth/password";
|
||||
import { generateId } from "@server/auth/sessions/app";
|
||||
import { build } from "@server/build";
|
||||
import {
|
||||
domains,
|
||||
domainNamespaces,
|
||||
domains,
|
||||
orgDomains,
|
||||
Resource,
|
||||
resourceHeaderAuth,
|
||||
resourceHeaderAuthExtendedCompatibility,
|
||||
resourcePassword,
|
||||
resourcePincode,
|
||||
resourcePolicies,
|
||||
resourcePolicyHeaderAuth,
|
||||
resourcePolicyPassword,
|
||||
resourcePolicyPincode,
|
||||
resourcePolicyRules,
|
||||
resourcePolicyWhiteList,
|
||||
resourceRules,
|
||||
resources,
|
||||
resourceWhitelist,
|
||||
roleActions,
|
||||
rolePolicies,
|
||||
roleResources,
|
||||
roles,
|
||||
Site,
|
||||
sites,
|
||||
Target,
|
||||
TargetHealthCheck,
|
||||
targetHealthCheck,
|
||||
targets,
|
||||
Transaction,
|
||||
userOrgs,
|
||||
userPolicies,
|
||||
userResources,
|
||||
users,
|
||||
resourcePolicies,
|
||||
resourcePolicyPassword,
|
||||
resourcePolicyPincode,
|
||||
resourcePolicyHeaderAuth,
|
||||
resourcePolicyRules,
|
||||
resourcePolicyWhiteList,
|
||||
rolePolicies,
|
||||
userPolicies
|
||||
type ResourceRule
|
||||
} from "@server/db";
|
||||
import { resources, targets, sites } from "@server/db";
|
||||
import { eq, and, asc, or, ne, count, isNotNull } from "drizzle-orm";
|
||||
import {
|
||||
Config,
|
||||
ConfigSchema,
|
||||
isTargetsOnlyResource,
|
||||
TargetData
|
||||
} from "./types";
|
||||
import logger from "@server/logger";
|
||||
import { createCertificate } from "#dynamic/routers/certificates/createCertificate";
|
||||
import { pickPort } from "@server/routers/target/helpers";
|
||||
import { resourcePassword } from "@server/db";
|
||||
import { getUniqueResourcePolicyName } from "@server/db/names";
|
||||
import { hashPassword } from "@server/auth/password";
|
||||
import { isValidCIDR, isValidIP, isValidUrlGlobPattern } from "../validators";
|
||||
import { isValidRegionId } from "@server/db/regions";
|
||||
import { isLicensedOrSubscribed } from "#dynamic/lib/isLicencedOrSubscribed";
|
||||
import { fireHealthCheckUnknownAlert } from "@server/lib/alerts";
|
||||
import { tierMatrix } from "../billing/tierMatrix";
|
||||
import { defaultRoleAllowedActions } from "@server/routers/role/createRole";
|
||||
import { build } from "@server/build";
|
||||
import { encrypt } from "@server/lib/crypto";
|
||||
import { generateId } from "@server/auth/sessions/app";
|
||||
import serverConfig from "@server/lib/config";
|
||||
import HttpCode from "@server/types/HttpCode";
|
||||
import createHttpError from "http-errors";
|
||||
import next from "next";
|
||||
import { encrypt } from "@server/lib/crypto";
|
||||
import logger from "@server/logger";
|
||||
import { defaultRoleAllowedActions } from "@server/routers/role/createRole";
|
||||
import { pickPort } from "@server/routers/target/helpers";
|
||||
import { and, asc, eq, isNotNull, ne, or } from "drizzle-orm";
|
||||
import { tierMatrix } from "../billing/tierMatrix";
|
||||
import { isValidCIDR, isValidIP, isValidUrlGlobPattern } from "../validators";
|
||||
import { Config, isTargetsOnlyResource, TargetData } from "./types";
|
||||
import { getOrCreateLabelIds, syncResourceLabels } from "./labels";
|
||||
import { LimitId } from "../billing";
|
||||
import { usageService } from "../billing/usageService";
|
||||
import { syncInferenceAiConfig } from "./aiProviders";
|
||||
import { syncAiBudgets } from "./aiBudgets";
|
||||
|
||||
export type PublicResourcesResults = {
|
||||
proxyResource: Resource;
|
||||
@@ -76,19 +75,40 @@ export async function updatePublicResources(
|
||||
)) {
|
||||
const targetsToUpdate: Target[] = [];
|
||||
const healthchecksToUpdate: TargetHealthCheck[] = [];
|
||||
|
||||
let resource: Resource;
|
||||
let resourceStatusFromSite: "approved" | "pending" = "approved";
|
||||
let providedSite: Partial<Site> | undefined;
|
||||
if (siteId) {
|
||||
// Use the provided siteId directly, but verify it belongs to the org
|
||||
[providedSite] = await trx
|
||||
.select({
|
||||
siteId: sites.siteId,
|
||||
type: sites.type,
|
||||
status: sites.status
|
||||
})
|
||||
.from(sites)
|
||||
.where(and(eq(sites.siteId, siteId), eq(sites.orgId, orgId)))
|
||||
.limit(1);
|
||||
|
||||
resourceStatusFromSite = providedSite?.status ?? "approved";
|
||||
}
|
||||
|
||||
async function createTarget( // reusable function to create a target
|
||||
resourceId: number,
|
||||
targetData: TargetData
|
||||
) {
|
||||
const targetSiteId = targetData.site;
|
||||
let site;
|
||||
let site: Partial<Site> | undefined;
|
||||
|
||||
if (targetSiteId) {
|
||||
// Look up site by niceId
|
||||
[site] = await trx
|
||||
.select({ siteId: sites.siteId, type: sites.type })
|
||||
.select({
|
||||
siteId: sites.siteId,
|
||||
type: sites.type,
|
||||
status: sites.status
|
||||
})
|
||||
.from(sites)
|
||||
.where(
|
||||
and(
|
||||
@@ -97,15 +117,9 @@ export async function updatePublicResources(
|
||||
)
|
||||
)
|
||||
.limit(1);
|
||||
} else if (siteId) {
|
||||
} else if (siteId && providedSite) {
|
||||
// Use the provided siteId directly, but verify it belongs to the org
|
||||
[site] = await trx
|
||||
.select({ siteId: sites.siteId, type: sites.type })
|
||||
.from(sites)
|
||||
.where(
|
||||
and(eq(sites.siteId, siteId), eq(sites.orgId, orgId))
|
||||
)
|
||||
.limit(1);
|
||||
site = providedSite;
|
||||
} else {
|
||||
throw new Error(`Target site is required`);
|
||||
}
|
||||
@@ -141,7 +155,7 @@ export async function updatePublicResources(
|
||||
.insert(targets)
|
||||
.values({
|
||||
resourceId: resourceId,
|
||||
siteId: site.siteId,
|
||||
siteId: site.siteId!,
|
||||
ip: targetData.hostname,
|
||||
mode: resourceData.mode as Target["mode"],
|
||||
method: targetData.method,
|
||||
@@ -174,7 +188,7 @@ export async function updatePublicResources(
|
||||
.insert(targetHealthCheck)
|
||||
.values({
|
||||
name: `${targetData.hostname}:${targetData.port}`,
|
||||
siteId: site.siteId,
|
||||
siteId: site.siteId!,
|
||||
targetId: newTarget.targetId,
|
||||
orgId: orgId,
|
||||
hcEnabled: healthcheckData?.enabled || false,
|
||||
@@ -232,7 +246,10 @@ export async function updatePublicResources(
|
||||
const resourceEnabled =
|
||||
resourceData.enabled == undefined || resourceData.enabled == null
|
||||
? true
|
||||
: resourceData.enabled;
|
||||
: resourceStatusFromSite === "pending"
|
||||
? false
|
||||
: resourceData.enabled;
|
||||
|
||||
const resourceSsl =
|
||||
resourceData.ssl == undefined || resourceData.ssl == null
|
||||
? true
|
||||
@@ -250,18 +267,6 @@ export async function updatePublicResources(
|
||||
? JSON.stringify(resourceData.responseHeaders)
|
||||
: null;
|
||||
|
||||
if (["ssh", "rdp", "vnc"].includes(resourceData.mode || "")) {
|
||||
const isLicensed = await isLicensedOrSubscribed(
|
||||
orgId,
|
||||
tierMatrix.advancedPublicResources
|
||||
);
|
||||
if (!isLicensed) {
|
||||
throw new Error(
|
||||
"Your current subscription does not support browser gateway resources. Please upgrade to access this feature."
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
if (resourceData.policy) {
|
||||
const isLicensed = await isLicensedOrSubscribed(
|
||||
orgId,
|
||||
@@ -277,7 +282,9 @@ export async function updatePublicResources(
|
||||
if (existingResource) {
|
||||
let domain;
|
||||
if (
|
||||
["http", "ssh", "rdp", "vnc"].includes(resourceData.mode || "")
|
||||
["http", "ssh", "rdp", "vnc", "inference"].includes(
|
||||
resourceData.mode || ""
|
||||
)
|
||||
) {
|
||||
if (resourceData["full-domain"]?.startsWith("*.")) {
|
||||
const isLicensed = await isLicensedOrSubscribed(
|
||||
@@ -295,6 +302,7 @@ export async function updatePublicResources(
|
||||
existingResource.resourceId,
|
||||
resourceData["full-domain"]!,
|
||||
orgId,
|
||||
resourceData.mode === "inference",
|
||||
trx
|
||||
);
|
||||
|
||||
@@ -316,7 +324,7 @@ export async function updatePublicResources(
|
||||
|
||||
const isLicensed = await isLicensedOrSubscribed(
|
||||
orgId,
|
||||
tierMatrix.maintencePage
|
||||
tierMatrix.maintenancePage
|
||||
);
|
||||
if (!isLicensed) {
|
||||
resourceData.maintenance = undefined;
|
||||
@@ -361,14 +369,22 @@ export async function updatePublicResources(
|
||||
name: resourceData.name || "Unnamed Resource",
|
||||
|
||||
mode: resourceData.mode,
|
||||
proxyPort: ["http", "ssh", "rdp", "vnc"].includes(
|
||||
resourceData.mode || ""
|
||||
)
|
||||
proxyPort: [
|
||||
"http",
|
||||
"ssh",
|
||||
"rdp",
|
||||
"vnc",
|
||||
"inference"
|
||||
].includes(resourceData.mode || "")
|
||||
? null
|
||||
: resourceData["proxy-port"],
|
||||
fullDomain: ["http", "ssh", "rdp", "vnc"].includes(
|
||||
resourceData.mode || ""
|
||||
)
|
||||
fullDomain: [
|
||||
"http",
|
||||
"ssh",
|
||||
"rdp",
|
||||
"vnc",
|
||||
"inference"
|
||||
].includes(resourceData.mode || "")
|
||||
? resourceData["full-domain"]
|
||||
: null,
|
||||
subdomain: domain ? domain.subdomain : null,
|
||||
@@ -417,7 +433,8 @@ export async function updatePublicResources(
|
||||
? (resourceData["proxy-protocol-version"] ??
|
||||
1)
|
||||
: 1,
|
||||
resourcePolicyId: sharedPolicy.resourcePolicyId
|
||||
resourcePolicyId: sharedPolicy.resourcePolicyId,
|
||||
status: resourceStatusFromSite
|
||||
})
|
||||
.where(
|
||||
eq(
|
||||
@@ -557,14 +574,23 @@ export async function updatePublicResources(
|
||||
.update(resources)
|
||||
.set({
|
||||
name: resourceData.name || "Unnamed Resource",
|
||||
proxyPort: ["http", "ssh", "rdp", "vnc"].includes(
|
||||
resourceData.mode || ""
|
||||
)
|
||||
mode: resourceData.mode,
|
||||
proxyPort: [
|
||||
"http",
|
||||
"ssh",
|
||||
"rdp",
|
||||
"vnc",
|
||||
"inference"
|
||||
].includes(resourceData.mode || "")
|
||||
? null
|
||||
: resourceData["proxy-port"],
|
||||
fullDomain: ["http", "ssh", "rdp", "vnc"].includes(
|
||||
resourceData.mode || ""
|
||||
)
|
||||
fullDomain: [
|
||||
"http",
|
||||
"ssh",
|
||||
"rdp",
|
||||
"vnc",
|
||||
"inference"
|
||||
].includes(resourceData.mode || "")
|
||||
? resourceData["full-domain"]
|
||||
: null,
|
||||
subdomain: domain ? domain.subdomain : null,
|
||||
@@ -602,7 +628,8 @@ export async function updatePublicResources(
|
||||
authDaemonPort:
|
||||
resourceData["auth-daemon"]?.port || 22123,
|
||||
resourcePolicyId: null,
|
||||
defaultResourcePolicyId: inlinePolicyId
|
||||
defaultResourcePolicyId: inlinePolicyId,
|
||||
status: resourceStatusFromSite
|
||||
})
|
||||
.where(
|
||||
eq(
|
||||
@@ -664,6 +691,30 @@ export async function updatePublicResources(
|
||||
trx
|
||||
);
|
||||
}
|
||||
|
||||
await syncInferenceAiConfig({
|
||||
orgId,
|
||||
trx,
|
||||
mode: resourceData.mode || "",
|
||||
scope: "public",
|
||||
resourceId: existingResource.resourceId,
|
||||
providers: (resourceData["ai-providers"] || []).map(
|
||||
(p) => ({
|
||||
provider: p.provider,
|
||||
accessMode: p["access-mode"],
|
||||
enabled: p.enabled,
|
||||
models: p.models
|
||||
})
|
||||
)
|
||||
});
|
||||
|
||||
await syncAiBudgets({
|
||||
orgId,
|
||||
trx,
|
||||
scope: "public",
|
||||
resourceId: existingResource.resourceId,
|
||||
budgets: resourceData["ai-budget"] || []
|
||||
});
|
||||
}
|
||||
|
||||
const existingResourceTargets = await trx
|
||||
@@ -744,7 +795,7 @@ export async function updatePublicResources(
|
||||
: undefined),
|
||||
rewritePathType: targetData["rewrite-match"],
|
||||
priority: targetData.priority,
|
||||
mode: resourceData.mode
|
||||
mode: resourceData.mode as Target["mode"]
|
||||
})
|
||||
.where(eq(targets.targetId, existingTarget.targetId))
|
||||
.returning();
|
||||
@@ -919,7 +970,7 @@ export async function updatePublicResources(
|
||||
.update(resourceRules)
|
||||
.set({
|
||||
action: getRuleAction(rule.action),
|
||||
match: rule.match.toUpperCase(),
|
||||
match: rule.match.toUpperCase() as ResourceRule["match"],
|
||||
value: getRuleValue(
|
||||
rule.match.toUpperCase(),
|
||||
rule.value
|
||||
@@ -938,7 +989,7 @@ export async function updatePublicResources(
|
||||
await trx.insert(resourceRules).values({
|
||||
resourceId: existingResource.resourceId,
|
||||
action: getRuleAction(rule.action),
|
||||
match: rule.match.toUpperCase(),
|
||||
match: rule.match.toUpperCase() as ResourceRule["match"],
|
||||
value: getRuleValue(
|
||||
rule.match.toUpperCase(),
|
||||
rule.value
|
||||
@@ -1021,7 +1072,7 @@ export async function updatePublicResources(
|
||||
} else {
|
||||
// create a brand new resource
|
||||
|
||||
if (build == "saas") {
|
||||
if (build === "saas") {
|
||||
const usage = await usageService.getUsage(
|
||||
orgId,
|
||||
LimitId.PUBLIC_RESOURCES
|
||||
@@ -1049,7 +1100,9 @@ export async function updatePublicResources(
|
||||
|
||||
let domain;
|
||||
if (
|
||||
["http", "ssh", "rdp", "vnc"].includes(resourceData.mode || "")
|
||||
["http", "ssh", "rdp", "vnc", "inference"].includes(
|
||||
resourceData.mode || ""
|
||||
)
|
||||
) {
|
||||
if (resourceData["full-domain"]?.startsWith("*.")) {
|
||||
const isLicensed = await isLicensedOrSubscribed(
|
||||
@@ -1067,6 +1120,7 @@ export async function updatePublicResources(
|
||||
undefined,
|
||||
resourceData["full-domain"]!,
|
||||
orgId,
|
||||
resourceData.mode === "inference",
|
||||
trx
|
||||
);
|
||||
|
||||
@@ -1079,7 +1133,7 @@ export async function updatePublicResources(
|
||||
|
||||
const isLicensed = await isLicensedOrSubscribed(
|
||||
orgId,
|
||||
tierMatrix.maintencePage
|
||||
tierMatrix.maintenancePage
|
||||
);
|
||||
if (!isLicensed) {
|
||||
resourceData.maintenance = undefined;
|
||||
@@ -1143,16 +1197,25 @@ export async function updatePublicResources(
|
||||
.values({
|
||||
orgId,
|
||||
niceId: resourceNiceId,
|
||||
status: resourceStatusFromSite,
|
||||
name: resourceData.name || "Unnamed Resource",
|
||||
mode: resourceData.mode,
|
||||
proxyPort: ["http", "ssh", "rdp", "vnc"].includes(
|
||||
resourceData.mode || ""
|
||||
)
|
||||
proxyPort: [
|
||||
"http",
|
||||
"ssh",
|
||||
"rdp",
|
||||
"vnc",
|
||||
"inference"
|
||||
].includes(resourceData.mode || "")
|
||||
? null
|
||||
: resourceData["proxy-port"],
|
||||
fullDomain: ["http", "ssh", "rdp", "vnc"].includes(
|
||||
resourceData.mode || ""
|
||||
)
|
||||
fullDomain: [
|
||||
"http",
|
||||
"ssh",
|
||||
"rdp",
|
||||
"vnc",
|
||||
"inference"
|
||||
].includes(resourceData.mode || "")
|
||||
? resourceData["full-domain"]
|
||||
: null,
|
||||
subdomain: domain ? domain.subdomain : null,
|
||||
@@ -1208,6 +1271,28 @@ export async function updatePublicResources(
|
||||
|
||||
resource = newResource;
|
||||
|
||||
await syncInferenceAiConfig({
|
||||
orgId,
|
||||
trx,
|
||||
mode: resourceData.mode || "",
|
||||
scope: "public",
|
||||
resourceId: newResource.resourceId,
|
||||
providers: (resourceData["ai-providers"] || []).map((p) => ({
|
||||
provider: p.provider,
|
||||
accessMode: p["access-mode"],
|
||||
enabled: p.enabled,
|
||||
models: p.models
|
||||
}))
|
||||
});
|
||||
|
||||
await syncAiBudgets({
|
||||
orgId,
|
||||
trx,
|
||||
scope: "public",
|
||||
resourceId: newResource.resourceId,
|
||||
budgets: resourceData["ai-budget"] || []
|
||||
});
|
||||
|
||||
await trx.insert(roleResources).values({
|
||||
roleId: adminRole.roleId,
|
||||
resourceId: newResource.resourceId
|
||||
@@ -1303,7 +1388,7 @@ export async function updatePublicResources(
|
||||
await trx.insert(resourceRules).values({
|
||||
resourceId: newResource.resourceId,
|
||||
action: getRuleAction(rule.action),
|
||||
match: rule.match.toUpperCase(),
|
||||
match: rule.match.toUpperCase() as ResourceRule["match"],
|
||||
value: getRuleValue(
|
||||
rule.match.toUpperCase(),
|
||||
rule.value
|
||||
@@ -1342,6 +1427,15 @@ export async function updatePublicResources(
|
||||
logger.debug(`Created resource ${newResource.resourceId}`);
|
||||
}
|
||||
|
||||
if (!isTargetsOnlyResource(resourceData)) {
|
||||
const labelIds = await getOrCreateLabelIds(
|
||||
orgId,
|
||||
resourceData.labels || [],
|
||||
trx
|
||||
);
|
||||
await syncResourceLabels(resource.resourceId, labelIds, trx);
|
||||
}
|
||||
|
||||
results.push({
|
||||
proxyResource: resource,
|
||||
targetsToUpdate,
|
||||
@@ -1366,7 +1460,7 @@ function getRuleAction(input: string) {
|
||||
|
||||
function getRuleValue(match: string, value: string) {
|
||||
// if the match is a country, uppercase the value
|
||||
if (match == "COUNTRY") {
|
||||
if (match === "COUNTRY" || match === "COUNTRY_IS_NOT") {
|
||||
return value.toUpperCase();
|
||||
}
|
||||
return value;
|
||||
@@ -2056,6 +2150,7 @@ export async function getDomain(
|
||||
resourceId: number | undefined,
|
||||
fullDomain: string,
|
||||
orgId: string,
|
||||
isInference: boolean,
|
||||
trx: Transaction
|
||||
) {
|
||||
const [fullDomainExists] = await trx
|
||||
@@ -2065,6 +2160,14 @@ export async function getDomain(
|
||||
and(
|
||||
eq(resources.fullDomain, fullDomain),
|
||||
eq(resources.orgId, orgId),
|
||||
// Inference resources route through the central AI gateway
|
||||
// rather than normal target-based proxying, so they're
|
||||
// allowed to share a full-domain with a non-inference
|
||||
// resource (and vice versa) - only conflicts within the
|
||||
// same routing category are rejected.
|
||||
isInference
|
||||
? ne(resources.mode, "inference")
|
||||
: eq(resources.mode, "inference"),
|
||||
resourceId
|
||||
? ne(resources.resourceId, resourceId)
|
||||
: isNotNull(resources.resourceId)
|
||||
|
||||
@@ -1,7 +1,5 @@
|
||||
import {
|
||||
db,
|
||||
idp,
|
||||
idpOrg,
|
||||
resourcePolicies,
|
||||
resourcePolicyHeaderAuth,
|
||||
resourcePolicyPassword,
|
||||
@@ -20,6 +18,7 @@ import { Config, ResourcePolicyData } from "./types";
|
||||
import logger from "@server/logger";
|
||||
import { getUniqueResourcePolicyName } from "@server/db/names";
|
||||
import { hashPassword } from "@server/auth/password";
|
||||
import { idpExistsForOrg } from "@server/lib/idp/idpExistsForOrg";
|
||||
import { isValidCIDR, isValidIP, isValidUrlGlobPattern } from "../validators";
|
||||
import { isLicensedOrSubscribed } from "#dynamic/lib/isLicencedOrSubscribed";
|
||||
import { tierMatrix } from "../billing/tierMatrix";
|
||||
@@ -71,19 +70,13 @@ export async function updateResourcePolicies(
|
||||
|
||||
// Validate auto-login-idp if provided
|
||||
if (policyData["auto-login-idp"]) {
|
||||
const [provider] = await trx
|
||||
.select()
|
||||
.from(idp)
|
||||
.innerJoin(idpOrg, eq(idpOrg.idpId, idp.idpId))
|
||||
.where(
|
||||
and(
|
||||
eq(idp.idpId, policyData["auto-login-idp"]),
|
||||
eq(idpOrg.orgId, orgId)
|
||||
)
|
||||
)
|
||||
.limit(1);
|
||||
const providerExists = await idpExistsForOrg(
|
||||
policyData["auto-login-idp"],
|
||||
orgId,
|
||||
trx
|
||||
);
|
||||
|
||||
if (!provider) {
|
||||
if (!providerExists) {
|
||||
throw new Error(
|
||||
`Identity provider not found for policy '${policyNiceId}' in this organization`
|
||||
);
|
||||
@@ -347,12 +340,13 @@ function getRuleAction(input: string): "ACCEPT" | "DROP" | "PASS" {
|
||||
|
||||
function getRuleMatch(
|
||||
input: string
|
||||
): "CIDR" | "IP" | "PATH" | "COUNTRY" | "ASN" | "REGION" {
|
||||
): "CIDR" | "IP" | "PATH" | "COUNTRY" | "COUNTRY_IS_NOT" | "ASN" | "REGION" {
|
||||
return input.toUpperCase() as
|
||||
| "CIDR"
|
||||
| "IP"
|
||||
| "PATH"
|
||||
| "COUNTRY"
|
||||
| "COUNTRY_IS_NOT"
|
||||
| "ASN"
|
||||
| "REGION";
|
||||
}
|
||||
|
||||
+135
-15
@@ -5,6 +5,11 @@ import { MaintenanceSchema } from "#dynamic/lib/blueprints/MaintenanceSchema";
|
||||
import { isValidRegionId } from "@server/db/regions";
|
||||
import { wildcardSubdomainSchema } from "@server/lib/schemas";
|
||||
import config from "@server/lib/config";
|
||||
import {
|
||||
aiBudgetEnforcementSchema,
|
||||
aiBudgetPeriodSchema,
|
||||
aiBudgetUnitSchema
|
||||
} from "@server/routers/aiBudget/validation";
|
||||
|
||||
const maxmindDbPath = config.getRawConfig().server.maxmind_db_path;
|
||||
const maxmindAsnPath = config.getRawConfig().server.maxmind_asn_path;
|
||||
@@ -28,7 +33,7 @@ export const TargetHealthCheckSchema = z.object({
|
||||
hostname: z.string(),
|
||||
port: z.int().min(1).max(65535),
|
||||
enabled: z.boolean().optional().default(true),
|
||||
path: z.string().optional(),
|
||||
path: z.string().optional().default("/"),
|
||||
scheme: z.string().optional(),
|
||||
mode: z.string().default("http"),
|
||||
interval: z.int().default(30),
|
||||
@@ -96,7 +101,7 @@ export const AuthSchema = z.object({
|
||||
export const RuleSchema = z
|
||||
.object({
|
||||
action: z.enum(["allow", "deny", "pass"]),
|
||||
match: z.enum(["cidr", "path", "ip", "country", "asn", "region"]),
|
||||
match: z.enum(["cidr", "path", "ip", "country", "country_is_not", "asn", "region"]),
|
||||
value: z.coerce.string(),
|
||||
priority: z.int().optional(),
|
||||
enabled: z.boolean().optional().default(true)
|
||||
@@ -131,7 +136,7 @@ export const RuleSchema = z
|
||||
)
|
||||
.refine(
|
||||
(rule) => {
|
||||
if (rule.match === "country") {
|
||||
if (rule.match === "country" || rule.match === "country_is_not") {
|
||||
if (!hasMaxmindCountryDb) {
|
||||
return false;
|
||||
}
|
||||
@@ -183,6 +188,56 @@ export const HeaderSchema = z.object({
|
||||
value: z.string().min(1)
|
||||
});
|
||||
|
||||
export const AiProviderAttachmentSchema = z
|
||||
.object({
|
||||
provider: z.string().min(1),
|
||||
"access-mode": z
|
||||
.enum(["inherit", "select"])
|
||||
.optional()
|
||||
.default("inherit"),
|
||||
enabled: z.boolean().optional().default(true),
|
||||
models: z.array(z.string()).optional().default([])
|
||||
})
|
||||
.refine(
|
||||
(provider) => {
|
||||
if (provider.models.length === 0) {
|
||||
return true;
|
||||
}
|
||||
return provider["access-mode"] === "select";
|
||||
},
|
||||
{
|
||||
path: ["models"],
|
||||
error: "'models' can only be set on a provider with access-mode 'select'"
|
||||
}
|
||||
);
|
||||
|
||||
export const AiBudgetSchema = z.object({
|
||||
amount: z.number().positive(),
|
||||
unit: aiBudgetUnitSchema,
|
||||
period: aiBudgetPeriodSchema.optional().default("monthly"),
|
||||
enforcement: aiBudgetEnforcementSchema.optional().default("hard"),
|
||||
enabled: z.boolean().optional().default(true)
|
||||
});
|
||||
|
||||
const aiBudgetArraySchema = z.array(AiBudgetSchema).refine(
|
||||
(budgets) => {
|
||||
const keys = budgets.map((b) => `${b.unit}::${b.period}`);
|
||||
return keys.length === new Set(keys).size;
|
||||
},
|
||||
{
|
||||
message:
|
||||
"'ai-budget' entries must not overlap: only one budget per unit/period combination is allowed"
|
||||
}
|
||||
);
|
||||
|
||||
// No default here: an object with only 'targets' set must remain
|
||||
// recognized as a targets-only resource by isTargetsOnlyResource().
|
||||
export const AiBudgetListSchema = aiBudgetArraySchema.optional();
|
||||
|
||||
export const AiBudgetListSchemaWithDefault = aiBudgetArraySchema
|
||||
.optional()
|
||||
.default([]);
|
||||
|
||||
export const AuthDaemonSchema = z
|
||||
.object({
|
||||
pam: z.enum(["passthrough", "push"]).optional().default("passthrough"),
|
||||
@@ -209,7 +264,9 @@ export const PublicResourceSchema = z
|
||||
protocol: z
|
||||
.enum(["http", "tcp", "udp", "ssh", "rdp", "vnc"])
|
||||
.optional(), // this was the old one and is now DEPRECATED in favor of the mode
|
||||
mode: z.enum(["http", "tcp", "udp", "ssh", "rdp", "vnc"]).optional(),
|
||||
mode: z
|
||||
.enum(["http", "tcp", "udp", "ssh", "rdp", "vnc", "inference"])
|
||||
.optional(),
|
||||
policy: z.string().optional(),
|
||||
ssl: z.boolean().optional(),
|
||||
scheme: z.enum(["http", "https"]).optional(),
|
||||
@@ -227,7 +284,10 @@ export const PublicResourceSchema = z
|
||||
maintenance: MaintenanceSchema.optional(),
|
||||
"auth-daemon": AuthDaemonSchema.optional(),
|
||||
"proxy-protocol": z.boolean().optional(),
|
||||
"proxy-protocol-version": z.int().min(1).optional()
|
||||
"proxy-protocol-version": z.int().min(1).optional(),
|
||||
labels: z.array(z.string().min(1)).optional(),
|
||||
"ai-providers": z.array(AiProviderAttachmentSchema).optional(),
|
||||
"ai-budget": AiBudgetListSchema
|
||||
})
|
||||
.refine(
|
||||
(resource) => {
|
||||
@@ -316,11 +376,13 @@ export const PublicResourceSchema = z
|
||||
return true;
|
||||
}
|
||||
|
||||
// If protocol/mode is http, ssh, rdp, or vnc, it must have a full-domain
|
||||
// If protocol/mode is http, ssh, rdp, vnc, or inference, it must have a full-domain
|
||||
const effectiveProtocol = resource.mode ?? resource.protocol;
|
||||
if (
|
||||
effectiveProtocol !== undefined &&
|
||||
["http", "ssh", "rdp", "vnc"].includes(effectiveProtocol)
|
||||
["http", "ssh", "rdp", "vnc", "inference"].includes(
|
||||
effectiveProtocol
|
||||
)
|
||||
) {
|
||||
return (
|
||||
resource["full-domain"] !== undefined &&
|
||||
@@ -331,7 +393,43 @@ export const PublicResourceSchema = z
|
||||
},
|
||||
{
|
||||
path: ["full-domain"],
|
||||
error: "When protocol is 'http', 'ssh', 'rdp', or 'vnc', a 'full-domain' must be provided"
|
||||
error: "When protocol is 'http', 'ssh', 'rdp', 'vnc', or 'inference', a 'full-domain' must be provided"
|
||||
}
|
||||
)
|
||||
.refine(
|
||||
(resource) => {
|
||||
if (isTargetsOnlyResource(resource)) {
|
||||
return true;
|
||||
}
|
||||
|
||||
const effectiveMode = resource.mode ?? resource.protocol;
|
||||
if (effectiveMode !== "inference") {
|
||||
return true;
|
||||
}
|
||||
|
||||
return resource.targets.every((target) => target == null);
|
||||
},
|
||||
{
|
||||
path: ["targets"],
|
||||
error: "When mode is 'inference', 'targets' must not be provided"
|
||||
}
|
||||
)
|
||||
.refine(
|
||||
(resource) => {
|
||||
if (isTargetsOnlyResource(resource)) {
|
||||
return true;
|
||||
}
|
||||
|
||||
const effectiveMode = resource.mode ?? resource.protocol;
|
||||
if (effectiveMode === "inference") {
|
||||
return true;
|
||||
}
|
||||
|
||||
return (resource["ai-providers"]?.length ?? 0) === 0;
|
||||
},
|
||||
{
|
||||
path: ["ai-providers"],
|
||||
error: "'ai-providers' can only be set when mode is 'inference'"
|
||||
}
|
||||
)
|
||||
.refine(
|
||||
@@ -465,14 +563,14 @@ export function isTargetsOnlyResource(resource: any): boolean {
|
||||
export const PrivateResourceSchema = z
|
||||
.object({
|
||||
name: z.string().min(1).max(255),
|
||||
mode: z.enum(["host", "cidr", "http", "ssh"]),
|
||||
mode: z.enum(["host", "cidr", "http", "ssh", "inference"]),
|
||||
site: z.string().optional(), // DEPRECATED IN FAVOR OF sites
|
||||
sites: z.array(z.string()).optional().default([]),
|
||||
// protocol: z.enum(["tcp", "udp"]).optional(),
|
||||
// proxyPort: z.int().positive().optional(),
|
||||
"destination-port": z.int().positive().optional(),
|
||||
destination: z.string().min(1).optional(),
|
||||
// enabled: z.boolean().default(true),
|
||||
enabled: z.boolean().default(true),
|
||||
"tcp-ports": portRangeStringSchema.optional().default("*"),
|
||||
"udp-ports": portRangeStringSchema.optional().default("*"),
|
||||
"disable-icmp": z.boolean().optional().default(false),
|
||||
@@ -495,16 +593,26 @@ export const PrivateResourceSchema = z
|
||||
}),
|
||||
users: z.array(z.string()).optional().default([]),
|
||||
machines: z.array(z.string()).optional().default([]),
|
||||
"auth-daemon": AuthDaemonSchema.optional()
|
||||
labels: z.array(z.string().min(1)).optional().default([]),
|
||||
"auth-daemon": AuthDaemonSchema.optional(),
|
||||
"ai-providers": z
|
||||
.array(AiProviderAttachmentSchema)
|
||||
.optional()
|
||||
.default([]),
|
||||
"ai-budget": AiBudgetListSchemaWithDefault
|
||||
})
|
||||
.refine(
|
||||
(data) => {
|
||||
// destination is optional only for ssh+native; required for everything else
|
||||
// destination is optional only for ssh+native or inference; required for everything else
|
||||
const isNativeSSH =
|
||||
data.mode === "ssh" &&
|
||||
(data["auth-daemon"] === undefined ||
|
||||
data["auth-daemon"].mode === "native");
|
||||
if (!isNativeSSH && !data.destination) {
|
||||
if (
|
||||
data.mode !== "inference" &&
|
||||
!isNativeSSH &&
|
||||
!data.destination
|
||||
) {
|
||||
return false;
|
||||
}
|
||||
return true;
|
||||
@@ -512,7 +620,19 @@ export const PrivateResourceSchema = z
|
||||
{
|
||||
path: ["destination"],
|
||||
message:
|
||||
"destination is required unless mode is 'ssh' with auth-daemon mode 'native'"
|
||||
"destination is required unless mode is 'ssh' with auth-daemon mode 'native', or mode is 'inference'"
|
||||
}
|
||||
)
|
||||
.refine(
|
||||
(data) => {
|
||||
if (data.mode === "inference") {
|
||||
return true;
|
||||
}
|
||||
return (data["ai-providers"]?.length ?? 0) === 0;
|
||||
},
|
||||
{
|
||||
path: ["ai-providers"],
|
||||
error: "'ai-providers' can only be set when mode is 'inference'"
|
||||
}
|
||||
)
|
||||
.refine(
|
||||
@@ -634,7 +754,6 @@ export const ResourcePolicySchema = z.object({
|
||||
})
|
||||
)
|
||||
)
|
||||
.max(50)
|
||||
.transform((v) => v.map((e) => e.toLowerCase()))
|
||||
.optional()
|
||||
.default([]),
|
||||
@@ -840,3 +959,4 @@ export type Target = z.infer<typeof TargetSchema>;
|
||||
export type Resource = z.infer<typeof PublicResourceSchema>;
|
||||
export type Config = z.infer<typeof ConfigSchema>;
|
||||
export type BlueprintResourcePolicy = z.infer<typeof ResourcePolicySchema>;
|
||||
export type BlueprintAiBudget = z.infer<typeof AiBudgetSchema>;
|
||||
|
||||
+15
-15
@@ -10,12 +10,12 @@ export const localCache = new NodeCache({
|
||||
});
|
||||
|
||||
// Log cache statistics periodically for monitoring
|
||||
setInterval(() => {
|
||||
const stats = localCache.getStats();
|
||||
logger.debug(
|
||||
`Local cache stats - Keys: ${stats.keys}, Hits: ${stats.hits}, Misses: ${stats.misses}, Hit rate: ${stats.hits > 0 ? ((stats.hits / (stats.hits + stats.misses)) * 100).toFixed(2) : 0}%`
|
||||
);
|
||||
}, 300000); // Every 5 minutes
|
||||
// setInterval(() => {
|
||||
// const stats = localCache.getStats();
|
||||
// logger.debug(
|
||||
// `Local cache stats - Keys: ${stats.keys}, Hits: ${stats.hits}, Misses: ${stats.misses}, Hit rate: ${stats.hits > 0 ? ((stats.hits / (stats.hits + stats.misses)) * 100).toFixed(2) : 0}%`
|
||||
// );
|
||||
// }, 300000); // Every 5 minutes
|
||||
|
||||
/**
|
||||
* Adaptive cache that uses Redis when available in multi-node environments,
|
||||
@@ -34,9 +34,9 @@ class AdaptiveCache {
|
||||
|
||||
// Use local cache as fallback or primary
|
||||
const success = localCache.set(key, value, effectiveTtl || 0);
|
||||
if (success) {
|
||||
logger.debug(`Set key in local cache: ${key}`);
|
||||
}
|
||||
// if (success) {
|
||||
// logger.debug(`Set key in local cache: ${key}`);
|
||||
// }
|
||||
return success;
|
||||
}
|
||||
|
||||
@@ -48,11 +48,11 @@ class AdaptiveCache {
|
||||
async get<T = any>(key: string): Promise<T | undefined> {
|
||||
// Use local cache as fallback or primary
|
||||
const value = localCache.get<T>(key);
|
||||
if (value !== undefined) {
|
||||
logger.debug(`Cache hit in local cache: ${key}`);
|
||||
} else {
|
||||
logger.debug(`Cache miss in local cache: ${key}`);
|
||||
}
|
||||
// if (value !== undefined) {
|
||||
// logger.debug(`Cache hit in local cache: ${key}`);
|
||||
// } else {
|
||||
// logger.debug(`Cache miss in local cache: ${key}`);
|
||||
// }
|
||||
return value;
|
||||
}
|
||||
|
||||
@@ -168,5 +168,5 @@ class AdaptiveCache {
|
||||
|
||||
// Export singleton instance
|
||||
export const cache = new AdaptiveCache();
|
||||
export const regionalCache = cache; // Alias for compatability with the private version
|
||||
export const regionalCache = cache; // Alias for compatibility with the private version
|
||||
export default cache;
|
||||
|
||||
@@ -339,19 +339,6 @@ export async function calculateUserClientsForOrgs(
|
||||
continue;
|
||||
}
|
||||
|
||||
// Get exit nodes for this org
|
||||
const exitNodesList = await getExitNodes(orgId);
|
||||
|
||||
if (exitNodesList.length === 0) {
|
||||
logger.warn(
|
||||
`Skipping org ${orgId} for OLM ${olm.olmId} (user ${userId}): no exit nodes found`
|
||||
);
|
||||
continue;
|
||||
}
|
||||
|
||||
const randomExitNode =
|
||||
exitNodesList[Math.floor(Math.random() * exitNodesList.length)];
|
||||
|
||||
// Get next available subnet
|
||||
const { value: newSubnet, release: releaseSubnetLock } =
|
||||
await getNextAvailableClientSubnet(orgId, trx);
|
||||
@@ -370,7 +357,6 @@ export async function calculateUserClientsForOrgs(
|
||||
const newClientData: InferInsertModel<typeof clients> = {
|
||||
userId,
|
||||
orgId: userOrg.orgId,
|
||||
exitNodeId: randomExitNode.exitNodeId,
|
||||
name: olm.name || "User Client",
|
||||
subnet: updatedSubnet,
|
||||
olmId: olm.olmId,
|
||||
|
||||
+222
-12
@@ -1,16 +1,226 @@
|
||||
import config from "@server/lib/config";
|
||||
import { certificates, db } from "@server/db";
|
||||
import { and, eq, isNotNull, or, inArray, sql } from "drizzle-orm";
|
||||
import { decrypt } from "@server/lib/crypto";
|
||||
import logger from "@server/logger";
|
||||
import { regionalCache as cache } from "#dynamic/lib/cache";
|
||||
import { build } from "@server/build";
|
||||
|
||||
// Define the return type for clarity and type safety
|
||||
export type CertificateResult = {
|
||||
id: number;
|
||||
domain: string;
|
||||
queriedDomain: string; // The domain that was originally requested (may differ for wildcards)
|
||||
wildcard: boolean | null;
|
||||
certFile: string | null;
|
||||
keyFile: string | null;
|
||||
expiresAt: number | null;
|
||||
updatedAt?: number | null;
|
||||
};
|
||||
|
||||
export async function getValidCertificatesForDomains(
|
||||
domains: Set<string>,
|
||||
useCache: boolean = true
|
||||
): Promise<
|
||||
Array<{
|
||||
id: number;
|
||||
domain: string;
|
||||
wildcard: boolean | null;
|
||||
certFile: string | null;
|
||||
keyFile: string | null;
|
||||
expiresAt: number | null;
|
||||
updatedAt?: number | null;
|
||||
}>
|
||||
> {
|
||||
return []; // stub
|
||||
): Promise<Array<CertificateResult>> {
|
||||
const finalResults: CertificateResult[] = [];
|
||||
const domainsToQuery = new Set<string>();
|
||||
|
||||
// 1. Check cache first if enabled
|
||||
if (useCache) {
|
||||
for (const domain of domains) {
|
||||
const cacheKey = `cert:${domain}`;
|
||||
const cachedCert = await cache.get<CertificateResult>(cacheKey);
|
||||
if (cachedCert) {
|
||||
finalResults.push(cachedCert); // Valid cache hit
|
||||
} else {
|
||||
// Also check for a wildcard cache entry covering this domain's parent
|
||||
const parts = domain.split(".");
|
||||
let wildcardHit = false;
|
||||
if (parts.length > 1) {
|
||||
const parentDomain = parts.slice(1).join(".");
|
||||
const wildcardCacheKey = `cert:*.${parentDomain}`;
|
||||
const cachedWildcard =
|
||||
await cache.get<CertificateResult>(wildcardCacheKey);
|
||||
if (cachedWildcard) {
|
||||
// Re-stamp queriedDomain so callers see the originally requested domain
|
||||
finalResults.push({
|
||||
...cachedWildcard,
|
||||
queriedDomain: domain
|
||||
});
|
||||
wildcardHit = true;
|
||||
}
|
||||
}
|
||||
if (!wildcardHit) {
|
||||
domainsToQuery.add(domain); // Cache miss or expired
|
||||
}
|
||||
}
|
||||
}
|
||||
} else {
|
||||
// If caching is disabled, add all domains to the query set
|
||||
domains.forEach((d) => domainsToQuery.add(d));
|
||||
}
|
||||
|
||||
// 2. If all domains were resolved from the cache, return early
|
||||
if (domainsToQuery.size === 0) {
|
||||
const decryptedResults = decryptFinalResults(
|
||||
finalResults,
|
||||
config.getRawConfig().server.secret!
|
||||
);
|
||||
return decryptedResults;
|
||||
}
|
||||
|
||||
// 3. Prepare domains for the database query
|
||||
const domainsToQueryArray = Array.from(domainsToQuery);
|
||||
const parentDomainsToQuery = new Set<string>();
|
||||
|
||||
domainsToQueryArray.forEach((domain) => {
|
||||
const parts = domain.split(".");
|
||||
// A wildcard can only match a domain with at least two parts (e.g., example.com)
|
||||
if (parts.length > 1) {
|
||||
parentDomainsToQuery.add(parts.slice(1).join("."));
|
||||
}
|
||||
});
|
||||
|
||||
const parentDomainsArray = Array.from(parentDomainsToQuery);
|
||||
|
||||
// Build wildcard variants: for each parent domain "example.com", also query "*.example.com"
|
||||
const wildcardPrefixedArray =
|
||||
build != "saas" ? parentDomainsArray.map((d) => `*.${d}`) : [];
|
||||
|
||||
// 4. Build and execute a single, efficient Drizzle query
|
||||
// This query fetches all potential exact and wildcard matches in one database round-trip.
|
||||
const potentialCerts = await db
|
||||
.select()
|
||||
.from(certificates)
|
||||
.where(
|
||||
and(
|
||||
eq(certificates.status, "valid"),
|
||||
isNotNull(certificates.certFile),
|
||||
isNotNull(certificates.keyFile),
|
||||
or(
|
||||
// Condition for exact matches on the requested domains
|
||||
inArray(certificates.domain, domainsToQueryArray),
|
||||
// Condition for wildcard matches on the parent domains (stored as "example.com" or "*.example.com")
|
||||
parentDomainsArray.length > 0
|
||||
? and(
|
||||
inArray(certificates.domain, [
|
||||
...parentDomainsArray,
|
||||
...wildcardPrefixedArray
|
||||
]),
|
||||
eq(certificates.wildcard, true)
|
||||
)
|
||||
: // If there are no possible parent domains, this condition is false
|
||||
sql`false`
|
||||
)
|
||||
)
|
||||
);
|
||||
|
||||
// Helper to normalize a wildcard cert's domain to its bare parent domain (strips leading "*.")
|
||||
const normalizeWildcardDomain = (domain: string): string =>
|
||||
domain.startsWith("*.") ? domain.slice(2) : domain;
|
||||
|
||||
// 5. Process the database results, prioritizing exact matches over wildcards
|
||||
const exactMatches = new Map<string, (typeof potentialCerts)[0]>();
|
||||
const wildcardMatches = new Map<string, (typeof potentialCerts)[0]>();
|
||||
|
||||
for (const cert of potentialCerts) {
|
||||
if (cert.wildcard) {
|
||||
// Normalize to bare parent domain so lookups are consistent regardless of storage format
|
||||
wildcardMatches.set(normalizeWildcardDomain(cert.domain), cert);
|
||||
} else {
|
||||
exactMatches.set(cert.domain, cert);
|
||||
}
|
||||
}
|
||||
|
||||
for (const domain of domainsToQuery) {
|
||||
let foundCert: (typeof potentialCerts)[0] | undefined = undefined;
|
||||
|
||||
// Priority 1: Check for an exact match (non-wildcard)
|
||||
if (exactMatches.has(domain)) {
|
||||
foundCert = exactMatches.get(domain);
|
||||
}
|
||||
// Priority 2: Check for a wildcard certificate whose normalized domain equals the queried domain
|
||||
else {
|
||||
const normalizedDomain = normalizeWildcardDomain(domain);
|
||||
if (wildcardMatches.has(normalizedDomain)) {
|
||||
foundCert = wildcardMatches.get(normalizedDomain);
|
||||
}
|
||||
// Priority 3: Check for a wildcard match on the parent domain
|
||||
else {
|
||||
const parts = normalizedDomain.split(".");
|
||||
if (parts.length > 1) {
|
||||
const parentDomain = parts.slice(1).join(".");
|
||||
if (wildcardMatches.has(parentDomain)) {
|
||||
foundCert = wildcardMatches.get(parentDomain);
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// If a certificate was found, format it, add to results, and cache it
|
||||
if (foundCert) {
|
||||
logger.debug(
|
||||
`Creating result cert for ${domain} using cert from ${foundCert.domain}`
|
||||
);
|
||||
const resultCert: CertificateResult = {
|
||||
id: foundCert.certId,
|
||||
domain: foundCert.domain, // The actual domain of the cert record
|
||||
queriedDomain: domain, // The domain that was originally requested
|
||||
wildcard: foundCert.wildcard,
|
||||
certFile: foundCert.certFile,
|
||||
keyFile: foundCert.keyFile,
|
||||
expiresAt: foundCert.expiresAt,
|
||||
updatedAt: foundCert.updatedAt
|
||||
};
|
||||
|
||||
finalResults.push(resultCert);
|
||||
|
||||
// Add to cache for future requests, using the *requested domain* as the key
|
||||
if (useCache) {
|
||||
const cacheKey = `cert:${domain}`;
|
||||
await cache.set(cacheKey, resultCert, 180);
|
||||
|
||||
// Also cache wildcard certs under a pattern key so other subdomains
|
||||
// can find them without a DB round-trip
|
||||
if (resultCert.wildcard) {
|
||||
const normalizedCertDomain = normalizeWildcardDomain(
|
||||
resultCert.domain
|
||||
);
|
||||
const wildcardCacheKey = `cert:*.${normalizedCertDomain}`;
|
||||
await cache.set(wildcardCacheKey, resultCert, 180);
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
const decryptedResults = decryptFinalResults(
|
||||
finalResults,
|
||||
config.getRawConfig().server.secret!
|
||||
);
|
||||
return decryptedResults;
|
||||
}
|
||||
|
||||
function decryptFinalResults(
|
||||
finalResults: CertificateResult[],
|
||||
secret: string
|
||||
): CertificateResult[] {
|
||||
const validCertsDecrypted = finalResults.map((cert) => {
|
||||
// Decrypt and save certificate file
|
||||
const decryptedCert = decrypt(
|
||||
cert.certFile!, // is not null from query
|
||||
secret
|
||||
);
|
||||
|
||||
// Decrypt and save key file
|
||||
const decryptedKey = decrypt(cert.keyFile!, secret);
|
||||
|
||||
// Return only the certificate data without org information
|
||||
return {
|
||||
...cert,
|
||||
certFile: decryptedCert,
|
||||
keyFile: decryptedKey
|
||||
};
|
||||
});
|
||||
|
||||
return validCertsDecrypted;
|
||||
}
|
||||
|
||||
@@ -3,12 +3,14 @@ import { cleanUpOldLogs as cleanUpOldAccessLogs } from "#dynamic/lib/logAccessAu
|
||||
import { cleanUpOldLogs as cleanUpOldActionLogs } from "#dynamic/middlewares/logActionAudit";
|
||||
import { cleanUpOldLogs as cleanUpOldRequestLogs } from "@server/routers/badger/logRequestAudit";
|
||||
import { cleanUpOldLogs as cleanUpOldConnectionLogs } from "#dynamic/routers/newt";
|
||||
import { cleanUpOldLogs as cleanUpOldAiSessionLogs } from "@server/routers/aiGateway/logAiSession";
|
||||
import { gt, or } from "drizzle-orm";
|
||||
import { cleanUpOldFingerprintSnapshots } from "@server/routers/olm/fingerprintingUtils";
|
||||
import { build } from "@server/build";
|
||||
|
||||
export function initLogCleanupInterval() {
|
||||
if (build == "saas") { // skip log cleanup for saas builds
|
||||
if (build == "saas") {
|
||||
// skip log cleanup for saas builds
|
||||
return null;
|
||||
}
|
||||
return setInterval(
|
||||
@@ -23,7 +25,9 @@ export function initLogCleanupInterval() {
|
||||
settingsLogRetentionDaysRequest:
|
||||
orgs.settingsLogRetentionDaysRequest,
|
||||
settingsLogRetentionDaysConnection:
|
||||
orgs.settingsLogRetentionDaysConnection
|
||||
orgs.settingsLogRetentionDaysConnection,
|
||||
settingsLogRetentionDaysAISessions:
|
||||
orgs.settingsLogRetentionDaysAISessions
|
||||
})
|
||||
.from(orgs)
|
||||
.where(
|
||||
@@ -31,7 +35,8 @@ export function initLogCleanupInterval() {
|
||||
gt(orgs.settingsLogRetentionDaysAction, 0),
|
||||
gt(orgs.settingsLogRetentionDaysAccess, 0),
|
||||
gt(orgs.settingsLogRetentionDaysRequest, 0),
|
||||
gt(orgs.settingsLogRetentionDaysConnection, 0)
|
||||
gt(orgs.settingsLogRetentionDaysConnection, 0),
|
||||
gt(orgs.settingsLogRetentionDaysAISessions, 0)
|
||||
)
|
||||
);
|
||||
|
||||
@@ -42,7 +47,8 @@ export function initLogCleanupInterval() {
|
||||
settingsLogRetentionDaysAction,
|
||||
settingsLogRetentionDaysAccess,
|
||||
settingsLogRetentionDaysRequest,
|
||||
settingsLogRetentionDaysConnection
|
||||
settingsLogRetentionDaysConnection,
|
||||
settingsLogRetentionDaysAISessions
|
||||
} = org;
|
||||
|
||||
if (settingsLogRetentionDaysAction > 0) {
|
||||
@@ -72,6 +78,13 @@ export function initLogCleanupInterval() {
|
||||
settingsLogRetentionDaysConnection
|
||||
);
|
||||
}
|
||||
|
||||
if (settingsLogRetentionDaysAISessions > 0) {
|
||||
await cleanUpOldAiSessionLogs(
|
||||
orgId,
|
||||
settingsLogRetentionDaysAISessions
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
await cleanUpOldFingerprintSnapshots(365);
|
||||
|
||||
@@ -18,3 +18,19 @@ export function canCompress(
|
||||
return false;
|
||||
}
|
||||
}
|
||||
|
||||
// Whether this newt client understands `tlsCertId` references into the
|
||||
// sync message's `certs` array, instead of requiring each target to carry
|
||||
// its own inline `tlsCert`/`tlsKey` PEM data. Bump the version floor here to
|
||||
// match whatever release first ships the newt-side support.
|
||||
export function supportsCertReferences(
|
||||
clientVersion: string | null | undefined
|
||||
): boolean {
|
||||
try {
|
||||
if (!clientVersion) return false;
|
||||
if (!semver.valid(clientVersion)) return false;
|
||||
return semver.gte(clientVersion, "1.16.0");
|
||||
} catch {
|
||||
return false;
|
||||
}
|
||||
}
|
||||
|
||||
@@ -2,7 +2,7 @@ import path from "path";
|
||||
import { fileURLToPath } from "url";
|
||||
|
||||
// This is a placeholder value replaced by the build process
|
||||
export const APP_VERSION = "1.19.0";
|
||||
export const APP_VERSION = "1.22.0";
|
||||
|
||||
export const __FILENAME = fileURLToPath(import.meta.url);
|
||||
export const __DIRNAME = path.dirname(__FILENAME);
|
||||
|
||||
@@ -93,6 +93,9 @@ export async function deleteOrgById(
|
||||
await trx.delete(sites).where(eq(sites.siteId, site.siteId));
|
||||
}
|
||||
for (const client of orgClients) {
|
||||
if (client.exitNodeId && client.pubKey) {
|
||||
await deletePeer(client.exitNodeId, client.pubKey);
|
||||
}
|
||||
const [olm] = await trx
|
||||
.select()
|
||||
.from(olms)
|
||||
|
||||
@@ -14,8 +14,6 @@ import {
|
||||
} from "@server/db";
|
||||
import logger from "@server/logger";
|
||||
import { removeTargets } from "@server/routers/newt/targets";
|
||||
import createHttpError from "http-errors";
|
||||
import HttpCode from "@server/types/HttpCode";
|
||||
|
||||
export type DeleteResourceResult = {
|
||||
deletedResource: Resource;
|
||||
@@ -66,13 +64,20 @@ export async function performDeleteResources(
|
||||
|
||||
const targetsByResourceId = new Map<number, Target[]>();
|
||||
for (const target of targetsToBeRemoved) {
|
||||
if (target.resourceId == null) {
|
||||
continue;
|
||||
}
|
||||
const existing = targetsByResourceId.get(target.resourceId) ?? [];
|
||||
existing.push(target);
|
||||
targetsByResourceId.set(target.resourceId, existing);
|
||||
}
|
||||
|
||||
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[]>();
|
||||
@@ -117,10 +122,10 @@ export async function runResourceDeleteSideEffects(
|
||||
.limit(1);
|
||||
|
||||
if (!site) {
|
||||
throw createHttpError(
|
||||
HttpCode.NOT_FOUND,
|
||||
`Site with ID ${target.siteId} not found`
|
||||
logger.debug(
|
||||
`Site with ID ${target.siteId} not found during resource delete side effects; skipping target removal`
|
||||
);
|
||||
continue;
|
||||
}
|
||||
|
||||
if (site.pubKey && site.type === "newt") {
|
||||
|
||||
@@ -1,6 +1,7 @@
|
||||
import { and, eq, sql } from "drizzle-orm";
|
||||
import { and, eq, inArray, isNotNull, sql } from "drizzle-orm";
|
||||
import {
|
||||
db,
|
||||
resources,
|
||||
siteNetworks,
|
||||
siteResources,
|
||||
targets,
|
||||
@@ -32,9 +33,11 @@ export async function getResourceIdsForSite(
|
||||
const rows = await trx
|
||||
.selectDistinct({ resourceId: targets.resourceId })
|
||||
.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(
|
||||
@@ -97,6 +100,64 @@ export function exceedsSiteAssociatedResourceDeleteLimit(
|
||||
return resourceCount > MAX_SITE_ASSOCIATED_RESOURCES_FOR_BULK_DELETE;
|
||||
}
|
||||
|
||||
export async function getPendingResourceIdsForSite(
|
||||
siteId: number,
|
||||
trx: Transaction | typeof db = db
|
||||
): Promise<number[]> {
|
||||
const resourceIds = await getResourceIdsForSite(siteId, trx);
|
||||
if (resourceIds.length === 0) {
|
||||
return [];
|
||||
}
|
||||
|
||||
const rows = await trx
|
||||
.select({ resourceId: resources.resourceId })
|
||||
.from(resources)
|
||||
.where(
|
||||
and(
|
||||
inArray(resources.resourceId, resourceIds),
|
||||
eq(resources.status, "pending")
|
||||
)
|
||||
);
|
||||
|
||||
return rows.map((row) => row.resourceId);
|
||||
}
|
||||
|
||||
export async function getPendingSiteResourceIdsForSite(
|
||||
siteId: number,
|
||||
orgId: string,
|
||||
trx: Transaction | typeof db = db
|
||||
): Promise<number[]> {
|
||||
const siteResourceIds = await getSiteResourceIdsForSite(siteId, orgId, trx);
|
||||
if (siteResourceIds.length === 0) {
|
||||
return [];
|
||||
}
|
||||
|
||||
const rows = await trx
|
||||
.select({ siteResourceId: siteResources.siteResourceId })
|
||||
.from(siteResources)
|
||||
.where(
|
||||
and(
|
||||
inArray(siteResources.siteResourceId, siteResourceIds),
|
||||
eq(siteResources.status, "pending")
|
||||
)
|
||||
);
|
||||
|
||||
return rows.map((row) => row.siteResourceId);
|
||||
}
|
||||
|
||||
export async function getPendingAssociatedResourceCountForSite(
|
||||
siteId: number,
|
||||
orgId: string,
|
||||
trx: Transaction | typeof db = db
|
||||
): Promise<number> {
|
||||
const [resourceIds, siteResourceIds] = await Promise.all([
|
||||
getPendingResourceIdsForSite(siteId, trx),
|
||||
getPendingSiteResourceIdsForSite(siteId, orgId, trx)
|
||||
]);
|
||||
|
||||
return resourceIds.length + siteResourceIds.length;
|
||||
}
|
||||
|
||||
export async function deleteAssociatedResourcesForSite(
|
||||
siteId: number,
|
||||
orgId: string,
|
||||
@@ -105,12 +166,32 @@ export async function deleteAssociatedResourcesForSite(
|
||||
const resourceIds = await getResourceIdsForSite(siteId, trx);
|
||||
const siteResourceIds = await getSiteResourceIdsForSite(siteId, orgId, trx);
|
||||
|
||||
const [resources, siteResourcesDeleted] = await Promise.all([
|
||||
const [deletedResources, siteResourcesDeleted] = await Promise.all([
|
||||
performDeleteResources(resourceIds, trx),
|
||||
performDeleteSiteResources(siteResourceIds, trx)
|
||||
]);
|
||||
|
||||
return { resources, siteResources: siteResourcesDeleted };
|
||||
return { resources: deletedResources, siteResources: siteResourcesDeleted };
|
||||
}
|
||||
|
||||
export async function deletePendingAssociatedResourcesForSite(
|
||||
siteId: number,
|
||||
orgId: string,
|
||||
trx: Transaction | typeof db = db
|
||||
): Promise<DeleteSiteAssociatedResourcesSideEffects> {
|
||||
const resourceIds = await getPendingResourceIdsForSite(siteId, trx);
|
||||
const siteResourceIds = await getPendingSiteResourceIdsForSite(
|
||||
siteId,
|
||||
orgId,
|
||||
trx
|
||||
);
|
||||
|
||||
const [deletedResources, siteResourcesDeleted] = await Promise.all([
|
||||
performDeleteResources(resourceIds, trx),
|
||||
performDeleteSiteResources(siteResourceIds, trx)
|
||||
]);
|
||||
|
||||
return { resources: deletedResources, siteResources: siteResourcesDeleted };
|
||||
}
|
||||
|
||||
export async function runDeleteSiteAssociatedResourcesSideEffects(
|
||||
|
||||
+10
-16
@@ -31,7 +31,6 @@ export async function validateAndConstructDomain(
|
||||
subdomain?: string | null
|
||||
): Promise<DomainValidationResult> {
|
||||
try {
|
||||
// Query domain with organization access check
|
||||
const [domainRes] = await db
|
||||
.select()
|
||||
.from(domains)
|
||||
@@ -42,6 +41,10 @@ export async function validateAndConstructDomain(
|
||||
eq(orgDomains.orgId, orgId),
|
||||
eq(orgDomains.domainId, domainId)
|
||||
)
|
||||
)
|
||||
.leftJoin(
|
||||
domainNamespaces,
|
||||
eq(domainNamespaces.domainId, domainId)
|
||||
);
|
||||
|
||||
// Check if domain exists
|
||||
@@ -52,8 +55,7 @@ export async function validateAndConstructDomain(
|
||||
};
|
||||
}
|
||||
|
||||
// Check if organization has access to domain
|
||||
if (domainRes.orgDomains && domainRes.orgDomains.orgId !== orgId) {
|
||||
if (!domainRes.orgDomains && !domainRes.domainNamespaces) {
|
||||
return {
|
||||
success: false,
|
||||
error: `Organization does not have access to domain with ID ${domainId}`
|
||||
@@ -84,19 +86,11 @@ export async function validateAndConstructDomain(
|
||||
}
|
||||
|
||||
// Wildcard subdomains are not allowed on namespace (provided/free) domains
|
||||
if (isWildcard) {
|
||||
const [namespaceDomain] = await db
|
||||
.select()
|
||||
.from(domainNamespaces)
|
||||
.where(eq(domainNamespaces.domainId, domainId))
|
||||
.limit(1);
|
||||
|
||||
if (namespaceDomain) {
|
||||
return {
|
||||
success: false,
|
||||
error: "Wildcard subdomains are not supported for provided or free domains. Use a specific subdomain instead."
|
||||
};
|
||||
}
|
||||
if (isWildcard && domainRes.domainNamespaces) {
|
||||
return {
|
||||
success: false,
|
||||
error: "Wildcard subdomains are not supported for provided or free domains. Use a specific subdomain instead."
|
||||
};
|
||||
}
|
||||
|
||||
if (
|
||||
|
||||
@@ -0,0 +1,40 @@
|
||||
import { db, sites, clients } from "@server/db";
|
||||
import { and, eq, count } from "drizzle-orm";
|
||||
|
||||
// (MAX_CONNECTIONS - current_connections) / MAX_CONNECTIONS)
|
||||
// higher = more desirable
|
||||
// like saying, this node has x% of its capacity left
|
||||
export async function calculateExitNodeWeight(
|
||||
exitNodeId: number,
|
||||
maxConnections: number | null | undefined
|
||||
): Promise<number | null> {
|
||||
if (maxConnections === null || maxConnections === undefined) {
|
||||
return 1;
|
||||
}
|
||||
|
||||
const [[siteConnections], [clientConnections]] = await Promise.all([
|
||||
db
|
||||
.select({ count: count() })
|
||||
.from(sites)
|
||||
.where(
|
||||
and(eq(sites.exitNodeId, exitNodeId), eq(sites.online, true))
|
||||
),
|
||||
db
|
||||
.select({ count: count() })
|
||||
.from(clients)
|
||||
.where(
|
||||
and(
|
||||
eq(clients.exitNodeId, exitNodeId),
|
||||
eq(clients.online, true)
|
||||
)
|
||||
)
|
||||
]);
|
||||
|
||||
const currentConnections = siteConnections.count + clientConnections.count;
|
||||
|
||||
if (currentConnections >= maxConnections) {
|
||||
return null;
|
||||
}
|
||||
|
||||
return (maxConnections - currentConnections) / maxConnections;
|
||||
}
|
||||
@@ -1,6 +1,5 @@
|
||||
import { db, exitNodes, Transaction } from "@server/db";
|
||||
import logger from "@server/logger";
|
||||
import { ExitNodePingResult } from "@server/routers/newt";
|
||||
import { eq } from "drizzle-orm";
|
||||
|
||||
export async function verifyExitNodeOrgAccess(
|
||||
@@ -23,7 +22,10 @@ export async function listExitNodes(
|
||||
// Accepted for parity with the enterprise implementation (used there for
|
||||
// site-label filtering of remote exit nodes). The OSS build has no remote
|
||||
// exit nodes, so it is unused here.
|
||||
siteId?: number
|
||||
siteId?: number,
|
||||
// Same as above: accepted for parity, unused since the OSS build has no
|
||||
// remote exit nodes to exclude.
|
||||
noRemote = false
|
||||
) {
|
||||
// TODO: pick which nodes to send and ping better than just all of them that are not remote
|
||||
const allExitNodes = await db
|
||||
@@ -52,6 +54,16 @@ export async function listExitNodes(
|
||||
return allExitNodes;
|
||||
}
|
||||
|
||||
export type ExitNodePingResult = {
|
||||
exitNodeId: number;
|
||||
latencyMs: number;
|
||||
weight: number;
|
||||
error?: string;
|
||||
exitNodeName: string;
|
||||
endpoint: string;
|
||||
wasPreviouslyConnected: boolean;
|
||||
};
|
||||
|
||||
export function selectBestExitNode(
|
||||
pingResults: ExitNodePingResult[]
|
||||
): ExitNodePingResult | null {
|
||||
|
||||
@@ -0,0 +1,41 @@
|
||||
import { db, ExitNode, Transaction, sites, clients } from "@server/db";
|
||||
import { eq } from "drizzle-orm";
|
||||
import config from "@server/lib/config";
|
||||
import { findNextAvailableCidr } from "@server/lib/ip";
|
||||
import { lockManager } from "#dynamic/lib/lock";
|
||||
|
||||
export async function getUniqueSubnetForExitNode(
|
||||
exitNode: ExitNode,
|
||||
trx: Transaction | typeof db = db
|
||||
): Promise<string | null> {
|
||||
const lockKey = `subnet-allocation:${exitNode.exitNodeId}`;
|
||||
|
||||
return await lockManager.withLock(
|
||||
lockKey,
|
||||
async () => {
|
||||
const [sitesQuery, clientsQuery] = await Promise.all([
|
||||
trx
|
||||
.select({ subnet: sites.exitNodeSubnet })
|
||||
.from(sites)
|
||||
.where(eq(sites.exitNodeId, exitNode.exitNodeId)),
|
||||
trx
|
||||
.select({ subnet: clients.exitNodeSubnet })
|
||||
.from(clients)
|
||||
.where(eq(clients.exitNodeId, exitNode.exitNodeId))
|
||||
]);
|
||||
|
||||
const blockSize = config.getRawConfig().gerbil.site_block_size;
|
||||
const subnets = [...sitesQuery, ...clientsQuery]
|
||||
.map((row) => row.subnet)
|
||||
.filter(
|
||||
(subnet): subnet is string =>
|
||||
!!subnet &&
|
||||
/^(\d{1,3}\.){3}\d{1,3}\/\d{1,2}$/.test(subnet)
|
||||
);
|
||||
subnets.push(exitNode.address.replace(/\/\d+$/, `/${blockSize}`));
|
||||
|
||||
return findNextAvailableCidr(subnets, blockSize, exitNode.address);
|
||||
},
|
||||
5000 // 5 second lock TTL - subnet allocation should be quick
|
||||
);
|
||||
}
|
||||
@@ -2,3 +2,5 @@ export * from "./exitNodes";
|
||||
export * from "./exitNodeComms";
|
||||
export * from "./subnet";
|
||||
export * from "./getCurrentExitNodeId";
|
||||
export * from "./calculateExitNodeWeight";
|
||||
export * from "./getUniqueSubnetForExitNode";
|
||||
|
||||
@@ -1,20 +1,26 @@
|
||||
import { db, exitNodes, Transaction } from "@server/db";
|
||||
import { db, exitNodes, exitNodeOrgs, Transaction } from "@server/db";
|
||||
import config from "@server/lib/config";
|
||||
import { findNextAvailableCidr } from "@server/lib/ip";
|
||||
import { lockManager } from "#dynamic/lib/lock";
|
||||
import { eq } from "drizzle-orm";
|
||||
|
||||
/**
|
||||
* Reserves the next available exit node subnet.
|
||||
*
|
||||
* Exit node subnets must never overlap with one another - regardless of
|
||||
* which org(s) they belong to - since HA exit nodes can end up routing for
|
||||
* the same org. This acquires a lock that the caller MUST release (via the
|
||||
* returned `release`) only after the chosen address has been durably
|
||||
* persisted (e.g. after the enclosing transaction commits), otherwise
|
||||
* concurrent callers can race and pick the same subnet.
|
||||
* There isn't enough address space to give every exit node in every org a
|
||||
* globally unique subnet, so we only guarantee uniqueness among exit nodes
|
||||
* that already belong to the same org - that's all that actually matters,
|
||||
* since HA only routes multiple exit nodes for a single org. Pass `orgId` to
|
||||
* scope the search to that org's existing exit nodes; without it, the search
|
||||
* considers every exit node (used by flows with no org context, e.g. the
|
||||
* initial gerbil exit node bootstrap). This acquires a lock that the caller
|
||||
* MUST release (via the returned `release`) only after the chosen address
|
||||
* has been durably persisted (e.g. after the enclosing transaction commits),
|
||||
* otherwise concurrent callers can race and pick the same subnet.
|
||||
*/
|
||||
export async function getNextAvailableSubnet(
|
||||
trx: Transaction | typeof db = db
|
||||
trx: Transaction | typeof db = db,
|
||||
orgId?: string
|
||||
): Promise<{ value: string; release: () => Promise<void> }> {
|
||||
const lockKey = "exit-node-subnet-allocation";
|
||||
const acquired = await lockManager.acquireLockWithRetry(lockKey, 6000);
|
||||
@@ -24,12 +30,19 @@ export async function getNextAvailableSubnet(
|
||||
const release = () => lockManager.releaseLock(lockKey, acquired);
|
||||
|
||||
try {
|
||||
// Get all existing subnets from routes table
|
||||
const existingAddresses = await trx
|
||||
.select({
|
||||
address: exitNodes.address
|
||||
})
|
||||
.from(exitNodes);
|
||||
// Get existing subnets, scoped to this org's exit nodes when known
|
||||
const existingAddresses = orgId
|
||||
? await trx
|
||||
.select({ address: exitNodes.address })
|
||||
.from(exitNodes)
|
||||
.innerJoin(
|
||||
exitNodeOrgs,
|
||||
eq(exitNodeOrgs.exitNodeId, exitNodes.exitNodeId)
|
||||
)
|
||||
.where(eq(exitNodeOrgs.orgId, orgId))
|
||||
: await trx
|
||||
.select({ address: exitNodes.address })
|
||||
.from(exitNodes);
|
||||
|
||||
const addresses = existingAddresses.map((a) => a.address);
|
||||
let subnet = findNextAvailableCidr(
|
||||
|
||||
@@ -0,0 +1,37 @@
|
||||
import { db, idp, idpOrg, Transaction } from "@server/db";
|
||||
import { and, eq } from "drizzle-orm";
|
||||
import { build } from "@server/build";
|
||||
|
||||
export function isOrgIdentityProviderMode(): boolean {
|
||||
return build === "saas" || process.env.IDENTITY_PROVIDER_MODE === "org";
|
||||
}
|
||||
|
||||
/**
|
||||
* Checks whether an identity provider can be used for the given org.
|
||||
* In org IdP mode, the provider must be linked via idpOrg.
|
||||
* In global IdP mode, the provider only needs to exist.
|
||||
*/
|
||||
export async function idpExistsForOrg(
|
||||
idpId: number,
|
||||
orgId: string,
|
||||
dbOrTrx: typeof db | Transaction = db
|
||||
): Promise<boolean> {
|
||||
if (isOrgIdentityProviderMode()) {
|
||||
const [provider] = await dbOrTrx
|
||||
.select({ idpId: idp.idpId })
|
||||
.from(idp)
|
||||
.innerJoin(idpOrg, eq(idpOrg.idpId, idp.idpId))
|
||||
.where(and(eq(idp.idpId, idpId), eq(idpOrg.orgId, orgId)))
|
||||
.limit(1);
|
||||
|
||||
return !!provider;
|
||||
}
|
||||
|
||||
const [provider] = await dbOrTrx
|
||||
.select({ idpId: idp.idpId })
|
||||
.from(idp)
|
||||
.where(eq(idp.idpId, idpId))
|
||||
.limit(1);
|
||||
|
||||
return !!provider;
|
||||
}
|
||||
+136
-16
@@ -5,7 +5,8 @@ import config from "@server/lib/config";
|
||||
import z from "zod";
|
||||
import logger from "@server/logger";
|
||||
import semver from "semver";
|
||||
import { getValidCertificatesForDomains } from "#dynamic/lib/certificates";
|
||||
import { createHash } from "crypto";
|
||||
import { getValidCertificatesForDomains } from "@server/lib/certificates";
|
||||
import { lockManager } from "#dynamic/lib/lock";
|
||||
|
||||
interface IPRange {
|
||||
@@ -496,6 +497,7 @@ export function generateRemoteSubnets(
|
||||
): string[] {
|
||||
const remoteSubnets = allSiteResources
|
||||
.filter((sr) => {
|
||||
if (!sr.enabled) return false;
|
||||
if (!sr.destination) return false;
|
||||
|
||||
if (sr.mode === "cidr") {
|
||||
@@ -526,17 +528,21 @@ export function generateRemoteSubnets(
|
||||
|
||||
export type Alias = { alias: string | null; aliasAddress: string | null };
|
||||
|
||||
export function generateAliasConfig(allSiteResources: SiteResource[]): Alias[] {
|
||||
export function generateAliasConfig(
|
||||
allSiteResources: SiteResource[],
|
||||
overrideIp?: string
|
||||
): Alias[] {
|
||||
return allSiteResources
|
||||
.filter(
|
||||
(sr) =>
|
||||
sr.enabled &&
|
||||
sr.aliasAddress &&
|
||||
((sr.alias && (sr.mode == "host" || sr.mode == "ssh")) ||
|
||||
(sr.fullDomain && sr.mode == "http"))
|
||||
)
|
||||
.map((sr) => ({
|
||||
alias: sr.alias || sr.fullDomain,
|
||||
aliasAddress: sr.aliasAddress
|
||||
aliasAddress: overrideIp || sr.aliasAddress
|
||||
}));
|
||||
}
|
||||
|
||||
@@ -646,22 +652,115 @@ export type SubnetProxyTargetV2 = {
|
||||
httpTargets?: HTTPTarget[];
|
||||
tlsCert?: string;
|
||||
tlsKey?: string;
|
||||
tlsCertId?: string; // references an entry in the sync message's top-level `certs` array instead of inlining tlsCert/tlsKey
|
||||
};
|
||||
|
||||
export type CertRef = { id: string; cert: string; key: string };
|
||||
|
||||
/**
|
||||
* Replaces each target's inline tlsCert/tlsKey with a tlsCertId reference
|
||||
* into a deduplicated certs array, so that many targets sharing the same
|
||||
* certificate (e.g. a wildcard cert used by thousands of site resources)
|
||||
* only need that certificate sent once per sync message.
|
||||
*/
|
||||
export function dedupeCertsForTargets(targetsV2: SubnetProxyTargetV2[]): {
|
||||
targets: SubnetProxyTargetV2[];
|
||||
certs: CertRef[];
|
||||
} {
|
||||
const idByContent = new Map<string, string>();
|
||||
const certs: CertRef[] = [];
|
||||
|
||||
const targets = targetsV2.map((target) => {
|
||||
if (!target.tlsCert || !target.tlsKey) {
|
||||
return target;
|
||||
}
|
||||
|
||||
const contentKey = `${target.tlsCert}|${target.tlsKey}`;
|
||||
let id = idByContent.get(contentKey);
|
||||
if (!id) {
|
||||
id = createHash("sha1")
|
||||
.update(contentKey)
|
||||
.digest("hex")
|
||||
.slice(0, 16);
|
||||
idByContent.set(contentKey, id);
|
||||
certs.push({ id, cert: target.tlsCert, key: target.tlsKey });
|
||||
}
|
||||
|
||||
const { tlsCert, tlsKey, ...rest } = target;
|
||||
return { ...rest, tlsCertId: id };
|
||||
});
|
||||
|
||||
return { targets, certs };
|
||||
}
|
||||
|
||||
export type HTTPTarget = {
|
||||
destAddr: string; // must be an IP or hostname
|
||||
destPort: number;
|
||||
scheme: "http" | "https";
|
||||
};
|
||||
|
||||
export type CertByDomain = Map<string, { certFile: string; keyFile: string }>;
|
||||
|
||||
/**
|
||||
* Fetches the TLS certificates for every enabled, SSL-enabled HTTP site
|
||||
* resource's fullDomain in a single batched call, instead of one call per
|
||||
* resource. Many resources commonly resolve to the very same certificate
|
||||
* (e.g. a wildcard covering the org's domain), so batching turns what would
|
||||
* be N concurrent DB/cache round-trips into one, and a lookup failure fails
|
||||
* loudly for the whole batch rather than silently dropping the cert on a
|
||||
* random subset of otherwise-identical resources under load.
|
||||
*/
|
||||
export async function batchFetchCertsForSiteResources(
|
||||
allSiteResources: SiteResource[]
|
||||
): Promise<CertByDomain> {
|
||||
const domains = new Set(
|
||||
allSiteResources
|
||||
.filter(
|
||||
(r) => r.enabled && r.mode === "http" && r.ssl && r.fullDomain
|
||||
)
|
||||
.map((r) => r.fullDomain as string)
|
||||
);
|
||||
|
||||
const certByDomain: CertByDomain = new Map();
|
||||
if (domains.size === 0) {
|
||||
return certByDomain;
|
||||
}
|
||||
|
||||
try {
|
||||
const certResults = await getValidCertificatesForDomains(domains, true);
|
||||
for (const cert of certResults) {
|
||||
if (cert.certFile && cert.keyFile) {
|
||||
certByDomain.set(cert.queriedDomain, {
|
||||
certFile: cert.certFile,
|
||||
keyFile: cert.keyFile
|
||||
});
|
||||
}
|
||||
}
|
||||
} catch (err) {
|
||||
logger.error(
|
||||
`Failed to batch-retrieve certificates for ${domains.size} domain(s): ${err}`
|
||||
);
|
||||
}
|
||||
|
||||
return certByDomain;
|
||||
}
|
||||
|
||||
export async function generateSubnetProxyTargetV2(
|
||||
siteResource: SiteResource,
|
||||
clients: {
|
||||
clientId: number;
|
||||
pubKey: string | null;
|
||||
subnet: string | null;
|
||||
}[]
|
||||
}[],
|
||||
certByDomain?: CertByDomain
|
||||
): Promise<SubnetProxyTargetV2[] | undefined> {
|
||||
if (!siteResource.enabled) {
|
||||
logger.debug(
|
||||
`Site resource ${siteResource.siteResourceId} is disabled, skipping target generation.`
|
||||
);
|
||||
return;
|
||||
}
|
||||
|
||||
if (clients.length === 0) {
|
||||
logger.debug(
|
||||
`No clients have access to site resource ${siteResource.siteResourceId}, skipping target generation.`
|
||||
@@ -741,23 +840,44 @@ export async function generateSubnetProxyTargetV2(
|
||||
let tlsKey: string | undefined;
|
||||
|
||||
if (siteResource.ssl && siteResource.fullDomain) {
|
||||
try {
|
||||
const certs = await getValidCertificatesForDomains(
|
||||
new Set([siteResource.fullDomain]),
|
||||
true
|
||||
);
|
||||
if (certs.length > 0 && certs[0].certFile && certs[0].keyFile) {
|
||||
tlsCert = certs[0].certFile;
|
||||
tlsKey = certs[0].keyFile;
|
||||
if (certByDomain) {
|
||||
// Caller batch-fetched certs for all resources up front (the
|
||||
// common, high-scale path) — just look up this resource's
|
||||
// domain rather than issuing its own DB/cache round-trip.
|
||||
const cert = certByDomain.get(siteResource.fullDomain);
|
||||
if (cert) {
|
||||
tlsCert = cert.certFile;
|
||||
tlsKey = cert.keyFile;
|
||||
} else {
|
||||
logger.warn(
|
||||
`No valid certificate found for SSL site resource ${siteResource.siteResourceId} with domain ${siteResource.fullDomain}`
|
||||
);
|
||||
}
|
||||
} catch (err) {
|
||||
logger.error(
|
||||
`Failed to retrieve certificate for site resource ${siteResource.siteResourceId} domain ${siteResource.fullDomain}: ${err}`
|
||||
);
|
||||
} else {
|
||||
// No batched map supplied by the caller — fall back to a
|
||||
// single-domain lookup for this resource alone.
|
||||
try {
|
||||
const certs = await getValidCertificatesForDomains(
|
||||
new Set([siteResource.fullDomain]),
|
||||
true
|
||||
);
|
||||
if (
|
||||
certs.length > 0 &&
|
||||
certs[0].certFile &&
|
||||
certs[0].keyFile
|
||||
) {
|
||||
tlsCert = certs[0].certFile;
|
||||
tlsKey = certs[0].keyFile;
|
||||
} else {
|
||||
logger.warn(
|
||||
`No valid certificate found for SSL site resource ${siteResource.siteResourceId} with domain ${siteResource.fullDomain}`
|
||||
);
|
||||
}
|
||||
} catch (err) {
|
||||
logger.error(
|
||||
`Failed to retrieve certificate for site resource ${siteResource.siteResourceId} domain ${siteResource.fullDomain}: ${err}`
|
||||
);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
@@ -7,6 +7,7 @@ export async function logAccessAudit(data: {
|
||||
type: string;
|
||||
orgId: string;
|
||||
resourceId?: number;
|
||||
siteResourceId?: number;
|
||||
user?: { username: string; userId: string };
|
||||
apiKey?: { name: string | null; apiKeyId: string };
|
||||
metadata?: any;
|
||||
|
||||
+34
-3
@@ -15,10 +15,41 @@ function getSegmentRegex(patternPart: string): RegExp {
|
||||
return regex;
|
||||
}
|
||||
|
||||
// Decodes percent-encoding (so an encoded slash like `%2F` is treated as a
|
||||
// real path separator, matching what most backends will do) and then
|
||||
// resolves `.` / `..` segments, so a request like `/public%2F..%2Fadmin/`
|
||||
// or `/public/../admin/` is matched as `/admin/`, not as a literal segment
|
||||
// or a wildcard-swallowed sequence under `/public/*`.
|
||||
function decodeAndResolvePath(p: string): string[] {
|
||||
const rawParts = p.split("/").filter(Boolean);
|
||||
|
||||
const resolved: string[] = [];
|
||||
for (const rawPart of rawParts) {
|
||||
let part: string;
|
||||
try {
|
||||
part = decodeURIComponent(rawPart);
|
||||
} catch {
|
||||
part = rawPart;
|
||||
}
|
||||
|
||||
// an encoded slash can turn one raw segment into several real ones
|
||||
for (const segment of part.split("/").filter(Boolean)) {
|
||||
if (segment === ".") {
|
||||
continue;
|
||||
} else if (segment === "..") {
|
||||
resolved.pop();
|
||||
} else {
|
||||
resolved.push(segment);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
return resolved;
|
||||
}
|
||||
|
||||
export function isPathAllowed(pattern: string, path: string): boolean {
|
||||
const normalize = (p: string) => p.split("/").filter(Boolean);
|
||||
const patternParts = normalize(pattern);
|
||||
const pathParts = normalize(path);
|
||||
const patternParts = pattern.split("/").filter(Boolean);
|
||||
const pathParts = decodeAndResolvePath(path);
|
||||
|
||||
function matchSegments(
|
||||
patternIndex: number,
|
||||
|
||||
@@ -79,7 +79,13 @@ export const configSchema = z
|
||||
.default(3001)
|
||||
.transform(stoi)
|
||||
.pipe(portSchema),
|
||||
ai_gateway_port: portSchema
|
||||
.optional()
|
||||
.default(3005)
|
||||
.transform(stoi)
|
||||
.pipe(portSchema),
|
||||
badger_override: z.string().optional(),
|
||||
ai_gateway_override: z.string().optional(),
|
||||
next_port: portSchema
|
||||
.optional()
|
||||
.default(3002)
|
||||
@@ -105,6 +111,23 @@ export const configSchema = z
|
||||
})
|
||||
.optional()
|
||||
.prefault({}),
|
||||
remote_headers: z
|
||||
.object({
|
||||
user_id: z
|
||||
.string()
|
||||
.optional()
|
||||
.default("Remote-User-Id"),
|
||||
virtual_api_key_id: z
|
||||
.string()
|
||||
.optional()
|
||||
.default("Remote-Virtual-Api-Key-Id"),
|
||||
user: z.string().optional().default("Remote-User"),
|
||||
email: z.string().optional().default("Remote-Email"),
|
||||
name: z.string().optional().default("Remote-Name"),
|
||||
role: z.string().optional().default("Remote-Role")
|
||||
})
|
||||
.optional()
|
||||
.prefault({}),
|
||||
resource_session_request_param: z
|
||||
.string()
|
||||
.optional()
|
||||
@@ -130,6 +153,24 @@ export const configSchema = z
|
||||
})
|
||||
.optional(),
|
||||
trust_proxy: z.int().gte(0).optional().default(1),
|
||||
// Opt-in: have Traefik/Badger stamp the resolved client IP
|
||||
// into a dedicated header (X-Pangolin-Client-Ip) on the
|
||||
// site-resource AI gateway route, so it survives an
|
||||
// intermediary proxy between Traefik and the gateway that
|
||||
// overwrites X-Forwarded-For/X-Real-Ip instead of appending
|
||||
// to them. Off by default since it requires a Badger
|
||||
// version that supports realIpHeader.
|
||||
enable_ai_gateway_client_ip_header: z
|
||||
.boolean()
|
||||
.optional()
|
||||
.default(false)
|
||||
.transform((val) =>
|
||||
process.env.ENABLE_AI_GATEWAY_CLIENT_IP_HEADER !==
|
||||
undefined
|
||||
? process.env.ENABLE_AI_GATEWAY_CLIENT_IP_HEADER ===
|
||||
"true"
|
||||
: val
|
||||
),
|
||||
secret: z.string().pipe(z.string().min(8)).optional(),
|
||||
maxmind_db_path: z.string().optional(),
|
||||
maxmind_asn_path: z.string().optional()
|
||||
@@ -139,6 +180,7 @@ export const configSchema = z
|
||||
integration_port: 3003,
|
||||
external_port: 3000,
|
||||
internal_port: 3001,
|
||||
ai_gateway_port: 3005,
|
||||
next_port: 3002,
|
||||
internal_hostname: "pangolin",
|
||||
session_cookie_name: "p_session_token",
|
||||
@@ -147,11 +189,20 @@ export const configSchema = z
|
||||
id: "P-Access-Token-Id",
|
||||
token: "P-Access-Token"
|
||||
},
|
||||
remote_headers: {
|
||||
user_id: "Remote-User-Id",
|
||||
virtual_api_key_id: "Remote-Virtual-Api-Key-Id",
|
||||
user: "Remote-User",
|
||||
email: "Remote-Email",
|
||||
name: "Remote-Name",
|
||||
role: "Remote-Role"
|
||||
},
|
||||
resource_session_request_param:
|
||||
"resource_session_request_param",
|
||||
dashboard_session_length_hours: 720,
|
||||
resource_session_length_hours: 720,
|
||||
trust_proxy: 1
|
||||
trust_proxy: 1,
|
||||
enable_ai_gateway_client_ip_header: false
|
||||
}),
|
||||
postgres: z
|
||||
.object({
|
||||
@@ -184,7 +235,8 @@ export const configSchema = z
|
||||
.number()
|
||||
.positive()
|
||||
.optional()
|
||||
.default(5000)
|
||||
.default(5000),
|
||||
jit_mode: z.boolean().default(true)
|
||||
})
|
||||
.optional()
|
||||
.prefault({})
|
||||
@@ -257,7 +309,24 @@ export const configSchema = z
|
||||
pp_transport_prefix: z
|
||||
.string()
|
||||
.optional()
|
||||
.default("pp-transport-v")
|
||||
.default("pp-transport-v"),
|
||||
rate_limit: z
|
||||
.object({
|
||||
average: z
|
||||
.number()
|
||||
.positive()
|
||||
.gt(0)
|
||||
.optional()
|
||||
.default(30),
|
||||
burst: z
|
||||
.number()
|
||||
.positive()
|
||||
.gt(0)
|
||||
.optional()
|
||||
.default(50)
|
||||
})
|
||||
.optional()
|
||||
.prefault({})
|
||||
})
|
||||
.optional()
|
||||
.prefault({}),
|
||||
@@ -279,7 +348,6 @@ export const configSchema = z
|
||||
.optional()
|
||||
.pipe(z.string())
|
||||
.transform((url) => url.toLowerCase()),
|
||||
use_subdomain: z.boolean().optional().default(false),
|
||||
subnet_group: z.string().optional().default("100.89.137.0/20"),
|
||||
block_size: z.number().positive().gt(0).optional().default(24),
|
||||
site_block_size: z
|
||||
@@ -373,9 +441,58 @@ export const configSchema = z
|
||||
disable_basic_wireguard_sites: z.boolean().optional(),
|
||||
disable_config_managed_domains: z.boolean().optional(),
|
||||
disable_product_help_banners: z.boolean().optional(),
|
||||
disable_enterprise_features: z.boolean().optional()
|
||||
disable_enterprise_features: z.boolean().optional(),
|
||||
enable_acme_cert_sync: z.boolean().optional().default(true),
|
||||
disable_private_http_placeholder: z
|
||||
.boolean()
|
||||
.optional()
|
||||
.default(false)
|
||||
})
|
||||
.optional(),
|
||||
acme: z
|
||||
.object({
|
||||
acme_json_path: z
|
||||
.string()
|
||||
.optional()
|
||||
.default("config/letsencrypt/acme.json"),
|
||||
acme_http_endpoint: z.string().optional(),
|
||||
sync_interval_ms: z.number().optional().default(5000)
|
||||
})
|
||||
.optional(),
|
||||
ai: z
|
||||
.object({
|
||||
model_catalog: z
|
||||
.object({
|
||||
upstream_url: z
|
||||
.url()
|
||||
.optional()
|
||||
.default("https://api.fossorial.io/api/v1/models"),
|
||||
// No default - only used when an operator wants to
|
||||
// pin the catalog to a local file instead of
|
||||
// fetching it from upstream_url.
|
||||
file: z.string().optional(),
|
||||
// No default - only used when an operator wants to
|
||||
// merge the content of the json file with the upstream catalog. This is useful for adding
|
||||
// custom models to the catalog without having to maintain a separate fork of the upstream catalog.
|
||||
merge_file: z.string().optional(),
|
||||
refresh_interval_min_hours: z
|
||||
.number()
|
||||
.positive()
|
||||
.gt(0)
|
||||
.optional()
|
||||
.default(6),
|
||||
refresh_interval_max_hours: z
|
||||
.number()
|
||||
.positive()
|
||||
.gt(0)
|
||||
.optional()
|
||||
.default(12)
|
||||
})
|
||||
.optional()
|
||||
.prefault({})
|
||||
})
|
||||
.optional()
|
||||
.prefault({}),
|
||||
dns: z
|
||||
.object({
|
||||
nameservers: z
|
||||
|
||||
@@ -19,7 +19,7 @@ import {
|
||||
userOrgRoles,
|
||||
userSiteResources
|
||||
} from "@server/db";
|
||||
import { and, count, eq, inArray, ne } from "drizzle-orm";
|
||||
import { and, count, eq, inArray, isNotNull, ne } from "drizzle-orm";
|
||||
|
||||
import { deletePeersBatch as newtDeletePeersBatch } from "@server/routers/newt/peers";
|
||||
import {
|
||||
@@ -27,6 +27,9 @@ import {
|
||||
deletePeersBatch as olmDeletePeersBatch
|
||||
} from "@server/routers/olm/peers";
|
||||
import { sendToExitNode } from "#dynamic/lib/exitNodes";
|
||||
import { sendToClientsBatch } from "#dynamic/routers/ws";
|
||||
import { canCompress } from "@server/lib/clientVersionChecks";
|
||||
import config from "@server/lib/config";
|
||||
import logger from "@server/logger";
|
||||
import {
|
||||
generateAliasConfig,
|
||||
@@ -187,7 +190,12 @@ export async function getClientSiteResourceAccess(
|
||||
`rebuildClientAssociations: [getClientSiteResourceAccess] siteResourceId=${siteResource.siteResourceId} networkId=${siteResource.networkId} siteCount=${sitesList.length} siteIds=[${sitesList.map((s) => s.siteId).join(", ")}]`
|
||||
);
|
||||
|
||||
if (sitesList.length === 0) {
|
||||
if (sitesList.length === 0 && siteResource.networkId !== null) {
|
||||
// A site resource with a networkId is expected to have at least one
|
||||
// site attached via siteNetworks. Resources with no networkId (e.g.
|
||||
// inference-mode resources, which connect clients directly to the
|
||||
// exit node instead of any site) are expected to have no sites, so
|
||||
// don't warn for those.
|
||||
logger.warn(
|
||||
`No sites found for siteResource ${siteResource.siteResourceId} with networkId ${siteResource.networkId}`
|
||||
);
|
||||
@@ -687,6 +695,22 @@ async function rebuildClientAssociationsFromSiteResourceImpl(
|
||||
clientSiteResourcesToRemove,
|
||||
trx
|
||||
);
|
||||
|
||||
// If this resource requires clients to be connected to the exit node
|
||||
// (e.g. an inference resource), re-sync the connect/disconnect state for
|
||||
// every client whose access to it may have changed - both those who
|
||||
// currently have access and those who just lost it.
|
||||
if (siteResource.requiresExitNodeConnection) {
|
||||
await syncClientExitNodeConnections(
|
||||
Array.from(
|
||||
new Set([
|
||||
...mergedAllClientIds,
|
||||
...existingClientSiteResourceIds
|
||||
])
|
||||
),
|
||||
trx
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
async function handleMessagesForSiteClients(
|
||||
@@ -966,7 +990,7 @@ export async function updateClientSiteDestinations(
|
||||
.where(eq(clientSitesAssociationsCache.clientId, client.clientId));
|
||||
|
||||
for (const site of sitesData) {
|
||||
if (!site.sites.subnet) {
|
||||
if (!site.sites.exitNodeSubnet) {
|
||||
logger.debug(`Site ${site.sites.siteId} has no subnet, skipping`);
|
||||
continue;
|
||||
}
|
||||
@@ -1002,7 +1026,7 @@ export async function updateClientSiteDestinations(
|
||||
sourcePort: parsedEndpoint.port,
|
||||
destinations: [
|
||||
{
|
||||
destinationIP: site.sites.subnet.split("/")[0],
|
||||
destinationIP: site.sites.exitNodeSubnet.split("/")[0],
|
||||
destinationPort: site.sites.listenPort || 1 // this satisfies gerbil for now but should be reevaluated
|
||||
}
|
||||
]
|
||||
@@ -1010,7 +1034,7 @@ export async function updateClientSiteDestinations(
|
||||
} else {
|
||||
// add to the existing destinations
|
||||
destinations.destinations.push({
|
||||
destinationIP: site.sites.subnet.split("/")[0],
|
||||
destinationIP: site.sites.exitNodeSubnet.split("/")[0],
|
||||
destinationPort: site.sites.listenPort || 1 // this satisfies gerbil for now but should be reevaluated
|
||||
});
|
||||
}
|
||||
@@ -1052,6 +1076,265 @@ export async function updateClientSiteDestinations(
|
||||
}
|
||||
}
|
||||
|
||||
// Determines, for each of the given clients, whether they currently have
|
||||
// access to any enabled site resource with requiresExitNodeConnection set
|
||||
// (e.g. an inference-mode resource) and tells the client's olm to connect to
|
||||
// or disconnect from its assigned exit node accordingly. Site resources with
|
||||
// requiresExitNodeConnection don't belong to any site/network, so this can't
|
||||
// be derived from the per-site peer logic above - it has to be recomputed
|
||||
// from the client's full current resource access every time that access
|
||||
// changes.
|
||||
async function syncClientExitNodeConnections(
|
||||
clientIds: number[],
|
||||
trx: Transaction | typeof db = db
|
||||
): Promise<void> {
|
||||
const uniqueClientIds = Array.from(new Set(clientIds));
|
||||
if (uniqueClientIds.length === 0) {
|
||||
return;
|
||||
}
|
||||
|
||||
// Only clients with an exit node assigned can be told to connect/disconnect.
|
||||
const clientsData = await trx
|
||||
.select({
|
||||
clientId: clients.clientId,
|
||||
exitNodeId: clients.exitNodeId,
|
||||
exitNodeSubnet: clients.exitNodeSubnet
|
||||
})
|
||||
.from(clients)
|
||||
.where(
|
||||
and(
|
||||
inArray(clients.clientId, uniqueClientIds),
|
||||
isNotNull(clients.exitNodeId)
|
||||
)
|
||||
);
|
||||
|
||||
if (clientsData.length === 0) {
|
||||
return;
|
||||
}
|
||||
|
||||
const clientIdsWithExitNode = clientsData.map((c) => c.clientId);
|
||||
|
||||
const requiresExitNodeRows = await trx
|
||||
.select({
|
||||
clientId: clientSiteResourcesAssociationsCache.clientId,
|
||||
alias: siteResources.alias,
|
||||
fullDomain: siteResources.fullDomain
|
||||
})
|
||||
.from(clientSiteResourcesAssociationsCache)
|
||||
.innerJoin(
|
||||
siteResources,
|
||||
eq(
|
||||
clientSiteResourcesAssociationsCache.siteResourceId,
|
||||
siteResources.siteResourceId
|
||||
)
|
||||
)
|
||||
.where(
|
||||
and(
|
||||
inArray(
|
||||
clientSiteResourcesAssociationsCache.clientId,
|
||||
clientIdsWithExitNode
|
||||
),
|
||||
eq(siteResources.enabled, true),
|
||||
eq(siteResources.requiresExitNodeConnection, true)
|
||||
)
|
||||
);
|
||||
|
||||
const needsConnectSet = new Set(
|
||||
requiresExitNodeRows.map((r) => r.clientId)
|
||||
);
|
||||
|
||||
// Aliases for every exit-node-backed resource this client can reach, so
|
||||
// the live connect push carries the same alias list the register/reconnect
|
||||
// path (buildSiteConfigurationForOlmClient) would compute.
|
||||
const exitNodeAliasesByClientId = new Map<number, (string | null)[]>();
|
||||
for (const row of requiresExitNodeRows) {
|
||||
if (row.alias == null && row.fullDomain == null) continue;
|
||||
const existing = exitNodeAliasesByClientId.get(row.clientId);
|
||||
if (existing) {
|
||||
existing.push(row.fullDomain || row.alias); // accept both for now in case we have other resource types that dont use the full domain
|
||||
} else {
|
||||
exitNodeAliasesByClientId.set(row.clientId, [
|
||||
row.fullDomain || row.alias
|
||||
]);
|
||||
}
|
||||
}
|
||||
|
||||
const exitNodeIds = Array.from(
|
||||
new Set(
|
||||
clientsData
|
||||
.map((c) => c.exitNodeId)
|
||||
.filter((id): id is number => id !== null)
|
||||
)
|
||||
);
|
||||
|
||||
const exitNodeRows =
|
||||
exitNodeIds.length > 0
|
||||
? await trx
|
||||
.select()
|
||||
.from(exitNodes)
|
||||
.where(inArray(exitNodes.exitNodeId, exitNodeIds))
|
||||
: [];
|
||||
const exitNodeById = new Map(exitNodeRows.map((n) => [n.exitNodeId, n]));
|
||||
|
||||
const olmRows = await trx
|
||||
.select({
|
||||
clientId: olms.clientId,
|
||||
olmId: olms.olmId,
|
||||
version: olms.version
|
||||
})
|
||||
.from(olms)
|
||||
.where(inArray(olms.clientId, clientIdsWithExitNode));
|
||||
const olmByClientId = new Map(
|
||||
olmRows
|
||||
.filter((r) => r.clientId !== null)
|
||||
.map((r) => [r.clientId as number, r])
|
||||
);
|
||||
|
||||
const relayPort = config.getRawConfig().gerbil.clients_start_port;
|
||||
|
||||
const connectPayloads: {
|
||||
clientId: string;
|
||||
message: { type: string; data: any };
|
||||
options: { compress: boolean; incrementConfigVersion: boolean };
|
||||
}[] = [];
|
||||
const disconnectPayloads: {
|
||||
clientId: string;
|
||||
message: { type: string; data: any };
|
||||
options: { compress: boolean; incrementConfigVersion: boolean };
|
||||
}[] = [];
|
||||
|
||||
for (const client of clientsData) {
|
||||
const olm = olmByClientId.get(client.clientId);
|
||||
if (!olm) {
|
||||
// No olm registered for this client yet/anymore, nothing to send.
|
||||
continue;
|
||||
}
|
||||
|
||||
const needsConnect = needsConnectSet.has(client.clientId);
|
||||
|
||||
if (needsConnect) {
|
||||
const exitNode = client.exitNodeId
|
||||
? exitNodeById.get(client.exitNodeId)
|
||||
: undefined;
|
||||
if (!exitNode || !client.exitNodeSubnet) {
|
||||
logger.warn(
|
||||
`rebuildClientAssociations: [syncClientExitNodeConnections] client ${client.clientId} needs an exit node connection but has no exit node or subnet assigned`
|
||||
);
|
||||
continue;
|
||||
}
|
||||
|
||||
connectPayloads.push({
|
||||
clientId: olm.olmId,
|
||||
message: {
|
||||
type: "olm/wg/exitnode/connect",
|
||||
data: {
|
||||
connect: true,
|
||||
endpoint: `${exitNode.endpoint}:${exitNode.listenPort}`,
|
||||
relayPort,
|
||||
publicKey: exitNode.publicKey,
|
||||
serverIP: exitNode.address.split("/")[0],
|
||||
tunnelIP: client.exitNodeSubnet.split("/")[0],
|
||||
aliases:
|
||||
exitNodeAliasesByClientId.get(client.clientId) ?? []
|
||||
}
|
||||
},
|
||||
options: {
|
||||
compress: canCompress(olm.version, "olm"),
|
||||
incrementConfigVersion: true
|
||||
}
|
||||
});
|
||||
} else {
|
||||
disconnectPayloads.push({
|
||||
clientId: olm.olmId,
|
||||
message: {
|
||||
type: "olm/wg/exitnode/disconnect",
|
||||
data: {}
|
||||
},
|
||||
options: {
|
||||
compress: canCompress(olm.version, "olm"),
|
||||
incrementConfigVersion: true
|
||||
}
|
||||
});
|
||||
}
|
||||
}
|
||||
|
||||
if (connectPayloads.length > 0) {
|
||||
await sendToClientsBatch(connectPayloads).catch((error) => {
|
||||
logger.error(
|
||||
`rebuildClientAssociations: Error sending exit node connect messages:`,
|
||||
error
|
||||
);
|
||||
});
|
||||
}
|
||||
|
||||
if (disconnectPayloads.length > 0) {
|
||||
await sendToClientsBatch(disconnectPayloads).catch((error) => {
|
||||
logger.error(
|
||||
`rebuildClientAssociations: Error sending exit node disconnect messages:`,
|
||||
error
|
||||
);
|
||||
});
|
||||
}
|
||||
}
|
||||
|
||||
// Notifies the olms of every given client that the alias of the site resource
|
||||
// they're using an exit node connection for has changed, via the dedicated
|
||||
// exit node data-update message. Unlike syncClientExitNodeConnections, this
|
||||
// doesn't touch connect/disconnect state - it's purely a rename for clients
|
||||
// that are (and remain) connected to the exit node for this resource.
|
||||
async function syncClientExitNodeAliasUpdate(
|
||||
clientIds: number[],
|
||||
oldAlias: string | null,
|
||||
newAlias: string | null,
|
||||
trx: Transaction | typeof db = db
|
||||
): Promise<void> {
|
||||
const uniqueClientIds = Array.from(new Set(clientIds));
|
||||
if (uniqueClientIds.length === 0) {
|
||||
return;
|
||||
}
|
||||
|
||||
const oldAliases = oldAlias ? [oldAlias] : [];
|
||||
const newAliases = newAlias ? [newAlias] : [];
|
||||
if (oldAliases.length === 0 && newAliases.length === 0) {
|
||||
return;
|
||||
}
|
||||
|
||||
const olmRows = await trx
|
||||
.select({
|
||||
clientId: olms.clientId,
|
||||
olmId: olms.olmId,
|
||||
version: olms.version
|
||||
})
|
||||
.from(olms)
|
||||
.where(inArray(olms.clientId, uniqueClientIds));
|
||||
|
||||
const updatePayloads = olmRows
|
||||
.filter((r) => r.clientId !== null)
|
||||
.map((olm) => ({
|
||||
clientId: olm.olmId,
|
||||
message: {
|
||||
type: "olm/wg/exitnode/data/update",
|
||||
data: {
|
||||
oldAliases,
|
||||
newAliases
|
||||
}
|
||||
},
|
||||
options: {
|
||||
compress: canCompress(olm.version, "olm"),
|
||||
incrementConfigVersion: true // this is important information we would need to sync
|
||||
}
|
||||
}));
|
||||
|
||||
if (updatePayloads.length > 0) {
|
||||
await sendToClientsBatch(updatePayloads).catch((error) => {
|
||||
logger.error(
|
||||
`rebuildClientAssociations: Error sending exit node alias update messages:`,
|
||||
error
|
||||
);
|
||||
});
|
||||
}
|
||||
}
|
||||
|
||||
async function handleSubnetProxyTargetUpdates(
|
||||
siteResource: SiteResource,
|
||||
sitesList: Site[],
|
||||
@@ -1282,7 +1565,7 @@ export async function handleMessagingForUpdatedSiteResource(
|
||||
`handleMessagingForUpdatedSiteResource: fetched newts for ${newtsForSites.length}/${allSiteIds.length} site(s)`
|
||||
);
|
||||
|
||||
// WARNING: THIS RELIES ON THE CACHE TABLES BEING UP TO DATE, SO CALL THIS AFTER THE ASSOCIATION CACHE IS UPDATED
|
||||
// !!!!!!!!!!!!!!!!!! WARNING: THIS RELIES ON THE CACHE TABLES BEING UP TO DATE, SO CALL THIS AFTER THE ASSOCIATION CACHE IS UPDATED !!!!!!!!!!!!!!!!!!
|
||||
const mergedAllClients = await trx
|
||||
.select({
|
||||
clientId: clientSiteResourcesAssociationsCache.clientId,
|
||||
@@ -1561,9 +1844,19 @@ export async function handleMessagingForUpdatedSiteResource(
|
||||
updatedSiteResource.udpPortRangeString ||
|
||||
existingSiteResource.disableIcmp !==
|
||||
updatedSiteResource.disableIcmp);
|
||||
// Toggling enabled on/off doesn't change any of the fields above, but it
|
||||
// does change whether targets/peer data should exist at all, so it needs
|
||||
// to drive the same old->new diff machinery: going enabled->disabled
|
||||
// diffs "real data" against "nothing" (a remove), and disabled->enabled
|
||||
// diffs "nothing" against "real data" (an add). generateSubnetProxyTargetV2/
|
||||
// generateRemoteSubnets/generateAliasConfig already return nothing for a
|
||||
// disabled resource, so no other changes are needed here.
|
||||
const enabledChanged =
|
||||
existingSiteResource &&
|
||||
existingSiteResource.enabled !== updatedSiteResource.enabled;
|
||||
|
||||
logger.debug(
|
||||
`handleMessagingForUpdatedSiteResource: change flags destinationChanged=${Boolean(destinationChanged)} destinationPortChanged=${Boolean(destinationPortChanged)} aliasChanged=${Boolean(aliasChanged)} fullDomainChanged=${Boolean(fullDomainChanged)} sslChanged=${Boolean(sslChanged)} portRangesChanged=${Boolean(portRangesChanged)}`
|
||||
`handleMessagingForUpdatedSiteResource: change flags destinationChanged=${Boolean(destinationChanged)} destinationPortChanged=${Boolean(destinationPortChanged)} aliasChanged=${Boolean(aliasChanged)} fullDomainChanged=${Boolean(fullDomainChanged)} sslChanged=${Boolean(sslChanged)} portRangesChanged=${Boolean(portRangesChanged)} enabledChanged=${Boolean(enabledChanged)}`
|
||||
);
|
||||
|
||||
// if the existingSiteResource is undefined (new resource) we don't need to do anything here, the rebuild above handled it all
|
||||
@@ -1574,14 +1867,16 @@ export async function handleMessagingForUpdatedSiteResource(
|
||||
fullDomainChanged ||
|
||||
sslChanged ||
|
||||
portRangesChanged ||
|
||||
destinationPortChanged
|
||||
destinationPortChanged ||
|
||||
enabledChanged
|
||||
) {
|
||||
const shouldUpdateTargets =
|
||||
destinationChanged ||
|
||||
sslChanged ||
|
||||
portRangesChanged ||
|
||||
fullDomainChanged ||
|
||||
destinationPortChanged;
|
||||
destinationPortChanged ||
|
||||
enabledChanged;
|
||||
|
||||
logger.debug(
|
||||
`handleMessagingForUpdatedSiteResource: entering unchanged-site update path shouldUpdateTargets=${shouldUpdateTargets}`
|
||||
@@ -1657,20 +1952,22 @@ export async function handleMessagingForUpdatedSiteResource(
|
||||
peerDataUpdateBatch.push({
|
||||
clientId: client.clientId,
|
||||
siteId,
|
||||
remoteSubnets: destinationChanged
|
||||
? {
|
||||
oldRemoteSubnets: !oldDestinationStillInUseBySite
|
||||
? generateRemoteSubnets([
|
||||
existingSiteResource
|
||||
])
|
||||
: [],
|
||||
newRemoteSubnets: generateRemoteSubnets([
|
||||
updatedSiteResource
|
||||
])
|
||||
}
|
||||
: undefined,
|
||||
remoteSubnets:
|
||||
destinationChanged || enabledChanged
|
||||
? {
|
||||
oldRemoteSubnets:
|
||||
!oldDestinationStillInUseBySite
|
||||
? generateRemoteSubnets([
|
||||
existingSiteResource
|
||||
])
|
||||
: [],
|
||||
newRemoteSubnets: generateRemoteSubnets([
|
||||
updatedSiteResource
|
||||
])
|
||||
}
|
||||
: undefined,
|
||||
aliases:
|
||||
aliasChanged || fullDomainChanged // the full domain is sent down as an alias
|
||||
aliasChanged || fullDomainChanged || enabledChanged // the full domain is sent down as an alias
|
||||
? {
|
||||
oldAliases: generateAliasConfig([
|
||||
existingSiteResource
|
||||
@@ -1695,6 +1992,38 @@ export async function handleMessagingForUpdatedSiteResource(
|
||||
);
|
||||
}
|
||||
|
||||
// For a resource that stays on an exit node connection across the update,
|
||||
// the alias is the only field that affects already-connected clients (the
|
||||
// exit node itself, its endpoint, etc. are not per-resource). Tell those
|
||||
// clients' olms about the rename directly via the exit node data-update
|
||||
// message rather than a full connect/disconnect cycle.
|
||||
if (
|
||||
existingSiteResource?.requiresExitNodeConnection &&
|
||||
updatedSiteResource.requiresExitNodeConnection &&
|
||||
aliasChanged
|
||||
) {
|
||||
await syncClientExitNodeAliasUpdate(
|
||||
mergedAllClients.map((c) => c.clientId),
|
||||
existingSiteResource.alias,
|
||||
updatedSiteResource.alias,
|
||||
trx
|
||||
);
|
||||
}
|
||||
|
||||
// If this resource requires (or required) clients to be connected to the
|
||||
// exit node (e.g. an inference resource), re-sync connect/disconnect
|
||||
// state for every client currently associated with it - covers toggling
|
||||
// requiresExitNodeConnection on update as well as enabling/disabling it.
|
||||
if (
|
||||
updatedSiteResource.requiresExitNodeConnection ||
|
||||
existingSiteResource?.requiresExitNodeConnection
|
||||
) {
|
||||
await syncClientExitNodeConnections(
|
||||
mergedAllClients.map((c) => c.clientId),
|
||||
trx
|
||||
);
|
||||
}
|
||||
|
||||
logger.debug(
|
||||
`handleMessagingForUpdatedSiteResource: DONE siteResourceId=${updatedSiteResource.siteResourceId}`
|
||||
);
|
||||
@@ -1976,6 +2305,10 @@ async function rebuildClientAssociationsFromClientImpl(
|
||||
resourcesToRemove,
|
||||
trx
|
||||
);
|
||||
|
||||
// Re-sync exit node connect/disconnect state based on this client's
|
||||
// current full set of resource access (e.g. inference resources).
|
||||
await syncClientExitNodeConnections([client.clientId], trx);
|
||||
}
|
||||
|
||||
async function handleMessagesForClientSites(
|
||||
|
||||
@@ -0,0 +1,200 @@
|
||||
import { db, resources, users, virtualApiKeyResources } from "@server/db";
|
||||
import { and, asc, eq } from "drizzle-orm";
|
||||
import config from "@server/lib/config";
|
||||
import { sendEmail } from "@server/emails";
|
||||
import IdentityApiKeyGenerated from "@server/emails/templates/IdentityApiKeyGenerated";
|
||||
import VirtualApiKeyGenerated from "@server/emails/templates/VirtualApiKeyGenerated";
|
||||
import { formatVirtualApiKeyCredential } from "@server/lib/virtualApiKey";
|
||||
|
||||
const EMAIL_GATEWAY_URL_LIMIT = 5;
|
||||
const VIRTUAL_API_KEY_EMAIL_BATCH_SIZE = 50;
|
||||
|
||||
export async function mapInBatches<T>(
|
||||
items: T[],
|
||||
fn: (item: T) => Promise<void>,
|
||||
batchSize = VIRTUAL_API_KEY_EMAIL_BATCH_SIZE
|
||||
): Promise<void> {
|
||||
for (let i = 0; i < items.length; i += batchSize) {
|
||||
const batch = items.slice(i, i + batchSize);
|
||||
await Promise.all(batch.map(fn));
|
||||
}
|
||||
}
|
||||
|
||||
async function listVirtualApiKeyGatewayUrls(params: {
|
||||
orgId: string;
|
||||
allResources: boolean;
|
||||
virtualApiKeyId: string;
|
||||
}): Promise<{ urls: string[]; hasMore: boolean }> {
|
||||
const rows = params.allResources
|
||||
? await db
|
||||
.select({
|
||||
fullDomain: resources.fullDomain,
|
||||
ssl: resources.ssl
|
||||
})
|
||||
.from(resources)
|
||||
.where(
|
||||
and(
|
||||
eq(resources.orgId, params.orgId),
|
||||
eq(resources.mode, "inference")
|
||||
)
|
||||
)
|
||||
.orderBy(asc(resources.name))
|
||||
.limit(EMAIL_GATEWAY_URL_LIMIT + 1)
|
||||
: await db
|
||||
.select({
|
||||
fullDomain: resources.fullDomain,
|
||||
ssl: resources.ssl
|
||||
})
|
||||
.from(virtualApiKeyResources)
|
||||
.innerJoin(
|
||||
resources,
|
||||
eq(virtualApiKeyResources.resourceId, resources.resourceId)
|
||||
)
|
||||
.where(
|
||||
eq(
|
||||
virtualApiKeyResources.virtualApiKeyId,
|
||||
params.virtualApiKeyId
|
||||
)
|
||||
)
|
||||
.orderBy(asc(resources.name))
|
||||
.limit(EMAIL_GATEWAY_URL_LIMIT + 1);
|
||||
|
||||
const urls = rows
|
||||
.map((row) =>
|
||||
row.fullDomain
|
||||
? `${row.ssl ? "https" : "http"}://${row.fullDomain}`
|
||||
: null
|
||||
)
|
||||
.filter((url): url is string => Boolean(url));
|
||||
|
||||
return {
|
||||
urls: urls.slice(0, EMAIL_GATEWAY_URL_LIMIT),
|
||||
hasMore: rows.length > EMAIL_GATEWAY_URL_LIMIT
|
||||
};
|
||||
}
|
||||
|
||||
export async function listOrgInferenceGatewayUrls(orgId: string): Promise<{
|
||||
urls: string[];
|
||||
hasMore: boolean;
|
||||
}> {
|
||||
return listVirtualApiKeyGatewayUrls({
|
||||
orgId,
|
||||
allResources: true,
|
||||
virtualApiKeyId: ""
|
||||
});
|
||||
}
|
||||
|
||||
export async function resolveVirtualApiKeyEmailRecipients(params: {
|
||||
sendEmail: boolean;
|
||||
sendToAttributedUser: boolean;
|
||||
userId: string | null | undefined;
|
||||
emails: string[];
|
||||
}): Promise<
|
||||
{ ok: true; recipients: string[] } | { ok: false; message: string }
|
||||
> {
|
||||
if (!params.sendEmail) {
|
||||
return { ok: true, recipients: [] };
|
||||
}
|
||||
|
||||
if (!config.getRawConfig().email) {
|
||||
return {
|
||||
ok: false,
|
||||
message: "Email is not configured on this server"
|
||||
};
|
||||
}
|
||||
|
||||
const recipients = new Set(
|
||||
params.emails.map((email) => email.trim().toLowerCase()).filter(Boolean)
|
||||
);
|
||||
|
||||
if (params.sendToAttributedUser) {
|
||||
if (!params.userId) {
|
||||
return {
|
||||
ok: false,
|
||||
message: "Associate a user to email the key to that user"
|
||||
};
|
||||
}
|
||||
|
||||
const [user] = await db
|
||||
.select({ email: users.email })
|
||||
.from(users)
|
||||
.where(eq(users.userId, params.userId))
|
||||
.limit(1);
|
||||
|
||||
if (!user?.email) {
|
||||
return {
|
||||
ok: false,
|
||||
message: "The associated user does not have an email address"
|
||||
};
|
||||
}
|
||||
|
||||
recipients.add(user.email.toLowerCase());
|
||||
}
|
||||
|
||||
if (recipients.size === 0) {
|
||||
return {
|
||||
ok: false,
|
||||
message: "Select at least one email recipient"
|
||||
};
|
||||
}
|
||||
|
||||
return { ok: true, recipients: [...recipients] };
|
||||
}
|
||||
|
||||
export async function sendVirtualApiKeyEmails(params: {
|
||||
recipients: string[];
|
||||
orgName: string;
|
||||
orgId: string;
|
||||
keyName: string | null;
|
||||
virtualApiKeyId: string;
|
||||
secret: string;
|
||||
allResources: boolean;
|
||||
isIdentityKey?: boolean;
|
||||
accountLabel?: string | null;
|
||||
gatewayUrls?: { urls: string[]; hasMore: boolean };
|
||||
}): Promise<void> {
|
||||
if (params.recipients.length === 0) {
|
||||
return;
|
||||
}
|
||||
|
||||
const credential = formatVirtualApiKeyCredential(
|
||||
params.virtualApiKeyId,
|
||||
params.secret
|
||||
);
|
||||
const { urls, hasMore } =
|
||||
params.gatewayUrls ??
|
||||
(await listVirtualApiKeyGatewayUrls({
|
||||
orgId: params.orgId,
|
||||
allResources: params.allResources,
|
||||
virtualApiKeyId: params.virtualApiKeyId
|
||||
}));
|
||||
const from = config.getNoReplyEmail();
|
||||
const subject = params.isIdentityKey
|
||||
? `Your identity key for ${params.orgName}`
|
||||
: `Virtual API key for ${params.orgName}`;
|
||||
|
||||
await mapInBatches(params.recipients, async (to) => {
|
||||
await sendEmail(
|
||||
params.isIdentityKey
|
||||
? IdentityApiKeyGenerated({
|
||||
orgName: params.orgName,
|
||||
accountLabel: params.accountLabel,
|
||||
credential,
|
||||
resourceUrls: urls,
|
||||
hasMoreResources: hasMore
|
||||
})
|
||||
: VirtualApiKeyGenerated({
|
||||
orgName: params.orgName,
|
||||
keyName: params.keyName,
|
||||
credential,
|
||||
resourceUrls: urls,
|
||||
hasMoreResources: hasMore
|
||||
}),
|
||||
{
|
||||
to,
|
||||
from,
|
||||
subject
|
||||
}
|
||||
);
|
||||
});
|
||||
}
|
||||
+166
-20
@@ -1,6 +1,6 @@
|
||||
import { z } from "zod";
|
||||
import { db, logsDb, statusHistory } from "@server/db";
|
||||
import { and, eq, gte, lt, asc, desc } from "drizzle-orm";
|
||||
import { and, eq, gte, lt, asc, desc, inArray, max, sql } from "drizzle-orm";
|
||||
import { regionalCache as cache } from "#dynamic/lib/cache";
|
||||
|
||||
const STATUS_HISTORY_CACHE_TTL = 60; // seconds
|
||||
@@ -8,26 +8,42 @@ const STATUS_HISTORY_CACHE_TTL = 60; // seconds
|
||||
function statusHistoryCacheKey(
|
||||
entityType: string,
|
||||
entityId: number,
|
||||
days: number
|
||||
days: number,
|
||||
tzOffsetMinutes: number
|
||||
): string {
|
||||
return `statusHistory:${entityType}:${entityId}:${days}`;
|
||||
return `statusHistory:${entityType}:${entityId}:${days}:${tzOffsetMinutes}`;
|
||||
}
|
||||
|
||||
// Returns the epoch seconds of the most recent local-calendar-day midnight,
|
||||
// where "local" is defined by tzOffsetMinutes (minutes to ADD to UTC to get
|
||||
// local time, e.g. Australia/Sydney standard time is 600). Defaults to 0
|
||||
// (UTC) so callers that don't pass a timezone keep the original behavior.
|
||||
function localMidnightSec(tzOffsetMinutes: number): number {
|
||||
const localNow = new Date(Date.now() + tzOffsetMinutes * 60_000);
|
||||
localNow.setUTCHours(0, 0, 0, 0);
|
||||
return Math.floor(localNow.getTime() / 1000) - tzOffsetMinutes * 60;
|
||||
}
|
||||
|
||||
export async function getCachedStatusHistory(
|
||||
entityType: string,
|
||||
entityId: number,
|
||||
days: number
|
||||
days: number,
|
||||
tzOffsetMinutes: number = 0
|
||||
): Promise<StatusHistoryResponse> {
|
||||
const cacheKey = statusHistoryCacheKey(entityType, entityId, days);
|
||||
const cacheKey = statusHistoryCacheKey(
|
||||
entityType,
|
||||
entityId,
|
||||
days,
|
||||
tzOffsetMinutes
|
||||
);
|
||||
const cached = await cache.get<StatusHistoryResponse>(cacheKey);
|
||||
if (cached !== undefined) {
|
||||
return cached;
|
||||
}
|
||||
|
||||
// Anchor to UTC midnight so the query window aligns with stable calendar days
|
||||
const utcToday = new Date();
|
||||
utcToday.setUTCHours(0, 0, 0, 0);
|
||||
const todayMidnightSec = Math.floor(utcToday.getTime() / 1000);
|
||||
// Anchor to local midnight (UTC when tzOffsetMinutes is 0) so the query
|
||||
// window aligns with stable calendar days for the requesting client
|
||||
const todayMidnightSec = localMidnightSec(tzOffsetMinutes);
|
||||
const startSec = todayMidnightSec - days * 86400;
|
||||
|
||||
const events = await logsDb
|
||||
@@ -63,7 +79,8 @@ export async function getCachedStatusHistory(
|
||||
const { buckets, totalDowntime } = computeBuckets(
|
||||
events,
|
||||
days,
|
||||
priorStatus
|
||||
priorStatus,
|
||||
tzOffsetMinutes
|
||||
);
|
||||
const totalWindow = days * 86400;
|
||||
const overallUptime =
|
||||
@@ -99,11 +116,19 @@ export const statusHistoryQuerySchema = z
|
||||
days: z
|
||||
.string()
|
||||
.optional()
|
||||
.transform((v) => (v ? parseInt(v, 10) : 90))
|
||||
.transform((v) => (v ? parseInt(v, 10) : 90)),
|
||||
// Minutes to add to UTC to get the requesting client's local time
|
||||
// (e.g. Australia/Sydney standard time is 600). Optional and
|
||||
// defaults to 0 (UTC) so older clients keep the prior behavior.
|
||||
tzOffsetMinutes: z
|
||||
.string()
|
||||
.optional()
|
||||
.transform((v) => (v ? parseInt(v, 10) : 0))
|
||||
})
|
||||
.pipe(
|
||||
z.object({
|
||||
days: z.number().int().min(1).max(365)
|
||||
days: z.number().int().min(1).max(365),
|
||||
tzOffsetMinutes: z.number().int().min(-720).max(840)
|
||||
})
|
||||
);
|
||||
|
||||
@@ -133,15 +158,15 @@ export function computeBuckets(
|
||||
id: number;
|
||||
}[],
|
||||
days: number,
|
||||
priorStatus: string | null = null
|
||||
priorStatus: string | null = null,
|
||||
tzOffsetMinutes: number = 0
|
||||
): { buckets: StatusHistoryDayBucket[]; totalDowntime: number } {
|
||||
const nowSec = Math.floor(Date.now() / 1000);
|
||||
|
||||
// Anchor bucket boundaries to UTC midnight so dates are stable calendar days
|
||||
// and don't drift as the cache expires and is recomputed
|
||||
const utcToday = new Date();
|
||||
utcToday.setUTCHours(0, 0, 0, 0);
|
||||
const todayMidnightSec = Math.floor(utcToday.getTime() / 1000);
|
||||
// Anchor bucket boundaries to local midnight (UTC when tzOffsetMinutes is
|
||||
// 0) so dates are stable calendar days for the requesting client and
|
||||
// don't drift as the cache expires and is recomputed
|
||||
const todayMidnightSec = localMidnightSec(tzOffsetMinutes);
|
||||
|
||||
const buckets: StatusHistoryDayBucket[] = [];
|
||||
let totalDowntime = 0;
|
||||
@@ -237,7 +262,11 @@ export function computeBuckets(
|
||||
)
|
||||
: 100;
|
||||
|
||||
const dateStr = new Date(dayStartSec * 1000).toISOString().slice(0, 10);
|
||||
// Shift by the client's offset before formatting so the label reflects
|
||||
// their local calendar date rather than the UTC date of dayStartSec
|
||||
const dateStr = new Date((dayStartSec + tzOffsetMinutes * 60) * 1000)
|
||||
.toISOString()
|
||||
.slice(0, 10);
|
||||
|
||||
const hasAnyData = currentStatus !== null || dayEvents.length > 0;
|
||||
|
||||
@@ -270,6 +299,123 @@ export function computeBuckets(
|
||||
status
|
||||
});
|
||||
}
|
||||
|
||||
return { buckets, totalDowntime };
|
||||
}
|
||||
|
||||
export type BatchedStatusHistoryResponse = Record<
|
||||
string,
|
||||
StatusHistoryResponse
|
||||
>;
|
||||
|
||||
export async function getBatchedStatusHistory(
|
||||
entityType: string,
|
||||
entityIds: number[],
|
||||
days: number,
|
||||
tzOffsetMinutes: number = 0
|
||||
): Promise<BatchedStatusHistoryResponse> {
|
||||
// Anchor to local midnight (UTC when tzOffsetMinutes is 0) so the query
|
||||
// window aligns with stable calendar days for the requesting client
|
||||
const todayMidnightSec = localMidnightSec(tzOffsetMinutes);
|
||||
const startSec = todayMidnightSec - days * 86400;
|
||||
|
||||
const events = await logsDb
|
||||
.select()
|
||||
.from(statusHistory)
|
||||
.where(
|
||||
and(
|
||||
eq(statusHistory.entityType, entityType),
|
||||
inArray(statusHistory.entityId, entityIds),
|
||||
gte(statusHistory.timestamp, startSec)
|
||||
)
|
||||
)
|
||||
.orderBy(asc(statusHistory.timestamp));
|
||||
|
||||
// Fetch the last known state before the window so that entities that
|
||||
// haven't changed status recently still show the correct status rather
|
||||
// than appearing as "no_data".
|
||||
|
||||
/**
|
||||
* If we used only postgres, we would have used `SELECT DISTINCT ON` to get the
|
||||
* latest event for each `entityId`,
|
||||
* but it doesn't work on SQLite, so instead we use a subquery,
|
||||
* the `ROW_NUMBER() OVER PARTITION` allows to assign a number
|
||||
* to each row ordered by the timestamp, the number 1 is the first one appearing in
|
||||
* the specified order, then the next and more, we only want the highest timestamp,
|
||||
* so we get for `row_number=1`
|
||||
*/
|
||||
const lastKnowEventsSub = logsDb
|
||||
.select({
|
||||
entityId: statusHistory.entityId,
|
||||
status: statusHistory.status,
|
||||
timestamp: statusHistory.timestamp,
|
||||
row_number:
|
||||
sql<number>`ROW_NUMBER() OVER (PARTITION BY ${statusHistory.entityId} ORDER BY ${statusHistory.timestamp} DESC)`.as(
|
||||
"row_number"
|
||||
)
|
||||
})
|
||||
.from(statusHistory)
|
||||
.where(
|
||||
and(
|
||||
eq(statusHistory.entityType, entityType),
|
||||
inArray(statusHistory.entityId, entityIds),
|
||||
lt(statusHistory.timestamp, startSec)
|
||||
)
|
||||
)
|
||||
.as("sub");
|
||||
|
||||
const lastKnownEvents = await logsDb
|
||||
.select({
|
||||
entityId: lastKnowEventsSub.entityId,
|
||||
status: lastKnowEventsSub.status,
|
||||
timestamp: lastKnowEventsSub.timestamp
|
||||
})
|
||||
.from(lastKnowEventsSub)
|
||||
.where(eq(lastKnowEventsSub.row_number, 1));
|
||||
|
||||
const eventStatusMap: Record<
|
||||
number,
|
||||
{
|
||||
events: typeof events;
|
||||
lastKnownEvent: (typeof lastKnownEvents)[number] | null;
|
||||
}
|
||||
> = {};
|
||||
|
||||
for (const entityId of entityIds) {
|
||||
eventStatusMap[entityId] = {
|
||||
events: events.filter((ev) => ev.entityId === entityId),
|
||||
lastKnownEvent:
|
||||
lastKnownEvents.find((ev) => ev.entityId === entityId) ?? null
|
||||
};
|
||||
}
|
||||
|
||||
const result: BatchedStatusHistoryResponse = {};
|
||||
|
||||
for (const entityId in eventStatusMap) {
|
||||
const event = eventStatusMap[Number(entityId)];
|
||||
const priorStatus = event.lastKnownEvent?.status ?? null;
|
||||
|
||||
const { buckets, totalDowntime } = computeBuckets(
|
||||
event.events,
|
||||
days,
|
||||
priorStatus,
|
||||
tzOffsetMinutes
|
||||
);
|
||||
const totalWindow = days * 86400;
|
||||
const overallUptime =
|
||||
totalWindow > 0
|
||||
? Math.max(
|
||||
0,
|
||||
((totalWindow - totalDowntime) / totalWindow) * 100
|
||||
)
|
||||
: 100;
|
||||
|
||||
result[entityId] = {
|
||||
entityType,
|
||||
entityId: Number(entityId),
|
||||
days: buckets,
|
||||
overallUptimePercent: Math.round(overallUptime * 100) / 100,
|
||||
totalDowntimeSeconds: totalDowntime
|
||||
};
|
||||
}
|
||||
return result;
|
||||
}
|
||||
|
||||
+50
-1
@@ -3,6 +3,8 @@ import config from "./config";
|
||||
import { getHostMeta } from "./hostMeta";
|
||||
import logger from "@server/logger";
|
||||
import {
|
||||
aiProviders,
|
||||
aiUsageRecords,
|
||||
alertRules,
|
||||
apiKeys,
|
||||
blueprints,
|
||||
@@ -11,7 +13,16 @@ import {
|
||||
siteResources
|
||||
} from "@server/db";
|
||||
import { sites, users, orgs, resources, clients, idp } from "@server/db";
|
||||
import { eq, count, notInArray, and, isNotNull, isNull } from "drizzle-orm";
|
||||
import {
|
||||
eq,
|
||||
count,
|
||||
countDistinct,
|
||||
notInArray,
|
||||
and,
|
||||
isNotNull,
|
||||
isNull,
|
||||
gte
|
||||
} from "drizzle-orm";
|
||||
import { APP_VERSION } from "./consts";
|
||||
import crypto from "crypto";
|
||||
import { UserType } from "@server/types/UserTypes";
|
||||
@@ -172,6 +183,25 @@ class TelemetryClient {
|
||||
.select({ count: count() })
|
||||
.from(blueprints);
|
||||
|
||||
const [aiProvidersCount] = await db
|
||||
.select({ count: count() })
|
||||
.from(aiProviders);
|
||||
const [orgsWithAiProviders] = await db
|
||||
.select({ count: countDistinct(aiProviders.orgId) })
|
||||
.from(aiProviders);
|
||||
|
||||
const usageWindowStart =
|
||||
Math.floor(Date.now() / 1000) -
|
||||
this.collectionIntervalDays * 24 * 60 * 60;
|
||||
const [aiUsageRecordsRecent] = await db
|
||||
.select({ count: count() })
|
||||
.from(aiUsageRecords)
|
||||
.where(gte(aiUsageRecords.createdAt, usageWindowStart));
|
||||
const [orgsWithRecentAiUsage] = await db
|
||||
.select({ count: countDistinct(aiUsageRecords.orgId) })
|
||||
.from(aiUsageRecords)
|
||||
.where(gte(aiUsageRecords.createdAt, usageWindowStart));
|
||||
|
||||
const supporterKey = config.getSupporterData();
|
||||
|
||||
const allPrivateResources = await db.select().from(siteResources);
|
||||
@@ -182,6 +212,7 @@ class TelemetryClient {
|
||||
let numPrivResourceCidr = 0;
|
||||
let numPrivResourceHttp = 0;
|
||||
let numPrivResourceSsh = 0;
|
||||
let numPrivResourceInference = 0;
|
||||
for (const res of allPrivateResources) {
|
||||
if (res.mode === "host") {
|
||||
numPrivResourceHosts += 1;
|
||||
@@ -191,6 +222,8 @@ class TelemetryClient {
|
||||
numPrivResourceHttp += 1;
|
||||
} else if (res.mode === "ssh") {
|
||||
numPrivResourceSsh += 1;
|
||||
} else if (res.mode === "inference") {
|
||||
numPrivResourceInference += 1;
|
||||
}
|
||||
|
||||
if (res.alias) {
|
||||
@@ -211,6 +244,11 @@ class TelemetryClient {
|
||||
numPrivateResourceCidr: numPrivResourceCidr,
|
||||
numPrivateResourceHttp: numPrivResourceHttp,
|
||||
numPrivateResourceSsh: numPrivResourceSsh,
|
||||
numPrivateResourceInference: numPrivResourceInference,
|
||||
numAiProviders: aiProvidersCount.count,
|
||||
numOrgsWithAiProviders: orgsWithAiProviders.count,
|
||||
numAiUsageRecordsRecent: aiUsageRecordsRecent.count,
|
||||
numOrgsWithRecentAiUsage: orgsWithRecentAiUsage.count,
|
||||
numAlertRules: numAlertRules.count,
|
||||
numUserDevices: userDevicesCount.count,
|
||||
numMachineClients: machineClients.count,
|
||||
@@ -323,6 +361,17 @@ class TelemetryClient {
|
||||
num_resources_non_http: stats.resources.filter(
|
||||
(r) => r.mode !== "http"
|
||||
).length,
|
||||
num_resources_ai_gateway: stats.resources.filter(
|
||||
(r) => r.mode === "inference"
|
||||
).length,
|
||||
num_private_resources_ai_gateway:
|
||||
stats.numPrivateResourceInference,
|
||||
num_ai_providers: stats.numAiProviders,
|
||||
num_orgs_with_ai_providers: stats.numOrgsWithAiProviders,
|
||||
num_ai_usage_records_recent:
|
||||
stats.numAiUsageRecordsRecent,
|
||||
num_orgs_with_recent_ai_usage:
|
||||
stats.numOrgsWithRecentAiUsage,
|
||||
num_newt_sites: stats.sites.filter((s) => s.type === "newt")
|
||||
.length,
|
||||
num_local_sites: stats.sites.filter(
|
||||
|
||||
@@ -0,0 +1,15 @@
|
||||
import { gzipSync, gunzipSync } from "zlib";
|
||||
|
||||
/**
|
||||
* Gzip a string and return it as base64 so it can be stored in a TEXT column.
|
||||
*/
|
||||
export function compressText(value: string): string {
|
||||
return gzipSync(Buffer.from(value, "utf8")).toString("base64");
|
||||
}
|
||||
|
||||
/**
|
||||
* Reverse of compressText - base64-decode and gunzip back to the original string.
|
||||
*/
|
||||
export function decompressText(value: string): string {
|
||||
return gunzipSync(Buffer.from(value, "base64")).toString("utf8");
|
||||
}
|
||||
@@ -8,7 +8,7 @@ import { db, exitNodes } from "@server/db";
|
||||
import { eq } from "drizzle-orm";
|
||||
import { getCurrentExitNodeId } from "@server/lib/exitNodes";
|
||||
import { getTraefikConfig } from "#dynamic/lib/traefik";
|
||||
import { getValidCertificatesForDomains } from "#dynamic/lib/certificates";
|
||||
import { getValidCertificatesForDomains } from "@server/lib/certificates";
|
||||
import { sendToExitNode } from "#dynamic/lib/exitNodes";
|
||||
import { build } from "@server/build";
|
||||
|
||||
@@ -516,6 +516,11 @@ export class TraefikConfigManager {
|
||||
const maintenanceHost =
|
||||
config.getRawConfig().server.internal_hostname;
|
||||
const pangolinUIUrl = `http://${maintenanceHost}:${maintenancePort}`;
|
||||
const aiGatewayUrl =
|
||||
config.getRawConfig().server.ai_gateway_override ||
|
||||
`http://${maintenanceHost}:${
|
||||
config.getRawConfig().server.ai_gateway_port
|
||||
}`;
|
||||
|
||||
// logger.debug(`Fetching traefik config for exit node: ${currentExitNode}`);
|
||||
traefikConfig = await getTraefikConfig(
|
||||
@@ -528,7 +533,8 @@ export class TraefikConfigManager {
|
||||
? false
|
||||
: config.getRawConfig().traefik.allow_raw_resources, // dont allow raw resources on saas otherwise use config
|
||||
pangolinUIUrl, // generate maintenance pages on cloud and hybrid
|
||||
pangolinUIUrl // generate browser gateway targets on cloud and hybrid
|
||||
pangolinUIUrl, // generate browser gateway targets on cloud and hybrid
|
||||
aiGatewayUrl
|
||||
);
|
||||
|
||||
const domains = new Set<string>();
|
||||
@@ -599,7 +605,30 @@ export class TraefikConfigManager {
|
||||
|
||||
resourceSessionRequestParam:
|
||||
config.getRawConfig().server
|
||||
.resource_session_request_param
|
||||
.resource_session_request_param,
|
||||
|
||||
remoteUserIdHeader:
|
||||
config.getRawConfig().server.remote_headers
|
||||
.user_id,
|
||||
|
||||
remoteVirtualApiKeyIdHeader:
|
||||
config.getRawConfig().server.remote_headers
|
||||
.virtual_api_key_id,
|
||||
|
||||
remoteUserHeader:
|
||||
config.getRawConfig().server.remote_headers
|
||||
.user,
|
||||
|
||||
remoteEmailHeader:
|
||||
config.getRawConfig().server.remote_headers
|
||||
.email,
|
||||
|
||||
remoteNameHeader:
|
||||
config.getRawConfig().server.remote_headers
|
||||
.name,
|
||||
|
||||
remoteRoleHeader:
|
||||
config.getRawConfig().server.remote_headers.role
|
||||
}
|
||||
}
|
||||
};
|
||||
|
||||
@@ -0,0 +1,165 @@
|
||||
import config from "@server/lib/config";
|
||||
import {
|
||||
AI_GATEWAY_TRUST_HEADER,
|
||||
AI_GATEWAY_RESOURCE_TYPE_HEADER,
|
||||
AI_GATEWAY_CLIENT_IP_HEADER,
|
||||
getAiGatewayTrustToken
|
||||
} from "@server/lib/aiGatewayTrust";
|
||||
|
||||
// The trust token is the same for every inference route on an exit node, so
|
||||
// these middlewares are built once and attached to each inference router.
|
||||
// Two variants exist (public resource vs. siteResource) so the resource
|
||||
// type header lets the gateway know which kind of router the request came
|
||||
// through without re-deriving it from resourceId.
|
||||
export const AI_GATEWAY_TRUST_MIDDLEWARE_RESOURCE =
|
||||
"ai-gateway-trust-headers-resource";
|
||||
export const AI_GATEWAY_TRUST_MIDDLEWARE_SITE_RESOURCE =
|
||||
"ai-gateway-trust-headers-site-resource";
|
||||
|
||||
// Opt-in: a Badger instance with forward auth disabled, used only to stamp
|
||||
// the resolved client IP into a dedicated header before the request reaches
|
||||
// whatever sits between Traefik and the AI gateway. Only the site-resource
|
||||
// router needs this - it's the only path that resolves request identity
|
||||
// from the client IP (see resolveRequestUser in aiGateway/pipeline.ts) -
|
||||
// and it's the only inference router that doesn't already run Badger.
|
||||
export const AI_GATEWAY_CLIENT_IP_MIDDLEWARE_NAME = "ai-gateway-client-ip";
|
||||
|
||||
/**
|
||||
* The AI gateway may live on a different host than the inference resource
|
||||
* itself (e.g. a remote exit node forwarding to the central dashboard over
|
||||
* a tunnel), so callers use this to decide whether to pin the Host header
|
||||
* to the gateway's own host.
|
||||
*/
|
||||
export function getAiGatewayHost(aiGatewayUrl: string): string | undefined {
|
||||
try {
|
||||
return new URL(aiGatewayUrl).host;
|
||||
} catch {
|
||||
return undefined;
|
||||
}
|
||||
}
|
||||
|
||||
/**
|
||||
* Header middleware that pins the Host header to the AI gateway's own host
|
||||
* (when it differs from the resource's) and smuggles the original resource
|
||||
* host through in "p-host" instead, so passHostHeader can't leak the wrong
|
||||
* Host to a gateway that lives on a different host than the resource.
|
||||
*/
|
||||
export function buildAiGatewayHostHeaderMiddleware(
|
||||
aiGatewayHost: string | undefined,
|
||||
fullDomain: string
|
||||
): { headers: { customRequestHeaders: Record<string, string> } } {
|
||||
return {
|
||||
headers: {
|
||||
customRequestHeaders: {
|
||||
...(aiGatewayHost ? { Host: aiGatewayHost } : {}),
|
||||
"p-host": fullDomain
|
||||
}
|
||||
}
|
||||
};
|
||||
}
|
||||
|
||||
export function buildAiGatewayTrustMiddlewares(): Record<string, any> {
|
||||
const token = getAiGatewayTrustToken();
|
||||
return {
|
||||
[AI_GATEWAY_TRUST_MIDDLEWARE_RESOURCE]: {
|
||||
headers: {
|
||||
customRequestHeaders: {
|
||||
[AI_GATEWAY_TRUST_HEADER]: token,
|
||||
[AI_GATEWAY_RESOURCE_TYPE_HEADER]: "resource"
|
||||
}
|
||||
}
|
||||
},
|
||||
[AI_GATEWAY_TRUST_MIDDLEWARE_SITE_RESOURCE]: {
|
||||
headers: {
|
||||
customRequestHeaders: {
|
||||
[AI_GATEWAY_TRUST_HEADER]: token,
|
||||
[AI_GATEWAY_RESOURCE_TYPE_HEADER]: "site-resource"
|
||||
}
|
||||
}
|
||||
}
|
||||
};
|
||||
}
|
||||
|
||||
export function buildAiGatewayClientIpMiddleware(): Record<string, any> | null {
|
||||
const enabled =
|
||||
config.getRawConfig().server.enable_ai_gateway_client_ip_header;
|
||||
if (!enabled) {
|
||||
return null;
|
||||
}
|
||||
return {
|
||||
[AI_GATEWAY_CLIENT_IP_MIDDLEWARE_NAME]: {
|
||||
plugin: {
|
||||
badger: {
|
||||
disableForwardAuth: true,
|
||||
realIpHeader: AI_GATEWAY_CLIENT_IP_HEADER
|
||||
}
|
||||
}
|
||||
}
|
||||
};
|
||||
}
|
||||
|
||||
/**
|
||||
* Build the redirect (if ssl), main router, and single-server service for
|
||||
* an AI-gateway-backed inference router. Identical between the public
|
||||
* inference-resource and siteResource-inference cases, and between the OSS
|
||||
* and private config generators - only the rule/tls/middleware chain
|
||||
* differs, which callers resolve themselves beforehand.
|
||||
*/
|
||||
export function buildAiGatewayRouterAndService(params: {
|
||||
routerName: string;
|
||||
serviceName: string;
|
||||
rule: string;
|
||||
ssl: boolean | null;
|
||||
tls: any;
|
||||
priority: number;
|
||||
routerMiddlewares: string[];
|
||||
aiGatewayUrl: string;
|
||||
redirectHttpsMiddlewareName: string;
|
||||
}): { routers: Record<string, any>; services: Record<string, any> } {
|
||||
const {
|
||||
routerName,
|
||||
serviceName,
|
||||
rule,
|
||||
ssl,
|
||||
tls,
|
||||
priority,
|
||||
routerMiddlewares,
|
||||
aiGatewayUrl,
|
||||
redirectHttpsMiddlewareName
|
||||
} = params;
|
||||
|
||||
const routers: Record<string, any> = {};
|
||||
|
||||
if (ssl) {
|
||||
routers[`${routerName}-redirect`] = {
|
||||
entryPoints: [config.getRawConfig().traefik.http_entrypoint],
|
||||
middlewares: [redirectHttpsMiddlewareName],
|
||||
service: serviceName,
|
||||
rule,
|
||||
priority
|
||||
};
|
||||
}
|
||||
|
||||
routers[routerName] = {
|
||||
entryPoints: [
|
||||
ssl
|
||||
? config.getRawConfig().traefik.https_entrypoint
|
||||
: config.getRawConfig().traefik.http_entrypoint
|
||||
],
|
||||
middlewares: routerMiddlewares,
|
||||
service: serviceName,
|
||||
rule,
|
||||
priority,
|
||||
...(ssl ? { tls } : {})
|
||||
};
|
||||
|
||||
const services = {
|
||||
[serviceName]: {
|
||||
loadBalancer: {
|
||||
servers: [{ url: aiGatewayUrl }]
|
||||
}
|
||||
}
|
||||
};
|
||||
|
||||
return { routers, services };
|
||||
}
|
||||
@@ -0,0 +1,399 @@
|
||||
import config from "@server/lib/config";
|
||||
import { sanitize } from "./utils";
|
||||
|
||||
export type BrowserGatewayResourceRow = {
|
||||
resourceId: number;
|
||||
resourceName: string | null;
|
||||
mode: string;
|
||||
fullDomain: string | null;
|
||||
ssl: boolean | null;
|
||||
subdomain: string | null;
|
||||
domainId: string | null;
|
||||
enabled: boolean | null;
|
||||
wildcard: boolean | null;
|
||||
domainCertResolver: string | null;
|
||||
preferWildcardCert: boolean | null;
|
||||
maintenanceModeEnabled: boolean | null;
|
||||
maintenanceModeType: string | null;
|
||||
maintenanceTitle: string | null;
|
||||
maintenanceMessage: string | null;
|
||||
maintenanceEstimatedTime: string | null;
|
||||
targetId: number;
|
||||
siteId: number;
|
||||
siteType: string;
|
||||
siteOnline: boolean | null;
|
||||
subnet: string | null;
|
||||
// Cloud-only namespace field - absent on OSS rows, so the namespace
|
||||
// filter below naturally no-ops there.
|
||||
domainNamespaceId?: unknown;
|
||||
};
|
||||
|
||||
export type BrowserGatewayResourceEntry = {
|
||||
resourceId: number;
|
||||
name: string;
|
||||
fullDomain: string | null;
|
||||
ssl: boolean | null;
|
||||
subdomain: string | null;
|
||||
domainId: string | null;
|
||||
enabled: boolean | null;
|
||||
wildcard: boolean | null;
|
||||
domainCertResolver: string | null;
|
||||
preferWildcardCert: boolean | null;
|
||||
maintenanceModeEnabled: boolean | null;
|
||||
maintenanceModeType: string | null;
|
||||
maintenanceTitle: string | null;
|
||||
maintenanceMessage: string | null;
|
||||
maintenanceEstimatedTime: string | null;
|
||||
targets: {
|
||||
targetId: number;
|
||||
bgType: string;
|
||||
siteId: number;
|
||||
siteType: string;
|
||||
siteOnline: boolean | null;
|
||||
subnet: string | null;
|
||||
}[];
|
||||
};
|
||||
|
||||
/**
|
||||
* Group the raw resource/target/site rows into per-resource browser-gateway
|
||||
* entries (SSH/VNC/RDP-mode resources served through the browser gateway
|
||||
* web UI instead of a real backend target).
|
||||
*/
|
||||
export function buildBrowserGatewayResourcesMap(
|
||||
rows: BrowserGatewayResourceRow[],
|
||||
filterOutNamespaceDomains: boolean
|
||||
): Map<number, BrowserGatewayResourceEntry> {
|
||||
const map = new Map<number, BrowserGatewayResourceEntry>();
|
||||
|
||||
for (const row of rows) {
|
||||
if (!["ssh", "vnc", "rdp"].includes(row.mode)) {
|
||||
continue;
|
||||
}
|
||||
if (filterOutNamespaceDomains && row.domainNamespaceId) {
|
||||
continue;
|
||||
}
|
||||
if (!map.has(row.resourceId)) {
|
||||
map.set(row.resourceId, {
|
||||
resourceId: row.resourceId,
|
||||
name: sanitize(row.resourceName ?? undefined) || "",
|
||||
fullDomain: row.fullDomain,
|
||||
ssl: row.ssl,
|
||||
subdomain: row.subdomain,
|
||||
domainId: row.domainId,
|
||||
enabled: row.enabled,
|
||||
wildcard: row.wildcard,
|
||||
domainCertResolver: row.domainCertResolver,
|
||||
preferWildcardCert: row.preferWildcardCert,
|
||||
maintenanceModeEnabled: row.maintenanceModeEnabled,
|
||||
maintenanceModeType: row.maintenanceModeType,
|
||||
maintenanceTitle: row.maintenanceTitle,
|
||||
maintenanceMessage: row.maintenanceMessage,
|
||||
maintenanceEstimatedTime: row.maintenanceEstimatedTime,
|
||||
targets: []
|
||||
});
|
||||
}
|
||||
map.get(row.resourceId)!.targets.push({
|
||||
targetId: row.targetId,
|
||||
bgType: row.mode,
|
||||
siteId: row.siteId,
|
||||
siteType: row.siteType,
|
||||
siteOnline: row.siteOnline,
|
||||
subnet: row.subnet
|
||||
});
|
||||
}
|
||||
|
||||
return map;
|
||||
}
|
||||
|
||||
/**
|
||||
* Build the Traefik routers/services for browser-gateway resources
|
||||
* (SSH/VNC/RDP served via a browser-based client instead of a raw target),
|
||||
* mutating config_output. TLS/cert-resolver handling differs between the
|
||||
* OSS (always resolve directly) and private (pangolin-dns aware) config
|
||||
* generators, so callers resolve that themselves via resolveTls - returning
|
||||
* null skips the resource (no valid cert available yet).
|
||||
*/
|
||||
export function buildBrowserGatewayConfig(params: {
|
||||
config_output: any;
|
||||
browserGatewayResourcesMap: Map<number, BrowserGatewayResourceEntry>;
|
||||
browserGatewayUiUrl: string;
|
||||
maintenancePageUiUrl: string | null;
|
||||
badgerMiddlewareName: string;
|
||||
redirectHttpsMiddlewareName: string;
|
||||
resolveTls: (args: {
|
||||
fullDomain: string;
|
||||
hasSubdomain: boolean;
|
||||
domainCertResolver: string | null;
|
||||
preferWildcardCert: boolean | null;
|
||||
}) => any | null;
|
||||
}): void {
|
||||
const {
|
||||
config_output,
|
||||
browserGatewayResourcesMap,
|
||||
browserGatewayUiUrl,
|
||||
maintenancePageUiUrl,
|
||||
badgerMiddlewareName,
|
||||
redirectHttpsMiddlewareName,
|
||||
resolveTls
|
||||
} = params;
|
||||
|
||||
const bgRateLimitMiddlewareName = "bg-ratelimit";
|
||||
if (!config_output.http.middlewares) {
|
||||
config_output.http.middlewares = {};
|
||||
}
|
||||
if (!config_output.http.middlewares[bgRateLimitMiddlewareName]) {
|
||||
const traefikRateLimit = config.getRawConfig().traefik.rate_limit;
|
||||
config_output.http.middlewares[bgRateLimitMiddlewareName] = {
|
||||
rateLimit: {
|
||||
average: traefikRateLimit.average,
|
||||
burst: traefikRateLimit.burst
|
||||
}
|
||||
};
|
||||
}
|
||||
|
||||
const browserGatewayPort = 39999;
|
||||
|
||||
for (const [, bgResource] of browserGatewayResourcesMap.entries()) {
|
||||
if (!bgResource.enabled) continue;
|
||||
if (!bgResource.domainId) continue;
|
||||
if (!bgResource.fullDomain) continue;
|
||||
|
||||
if (!config_output.http.routers) config_output.http.routers = {};
|
||||
if (!config_output.http.services) config_output.http.services = {};
|
||||
|
||||
const fullDomain = bgResource.fullDomain;
|
||||
const additionalMiddlewares =
|
||||
config.getRawConfig().traefik.additional_middlewares || [];
|
||||
const routerMiddlewares = [
|
||||
badgerMiddlewareName,
|
||||
bgRateLimitMiddlewareName,
|
||||
...additionalMiddlewares
|
||||
];
|
||||
|
||||
const hostRule = `Host(\`${fullDomain}\`)`;
|
||||
|
||||
// Build TLS config
|
||||
const tls = resolveTls({
|
||||
fullDomain,
|
||||
hasSubdomain: !!bgResource.subdomain,
|
||||
domainCertResolver: bgResource.domainCertResolver,
|
||||
preferWildcardCert: bgResource.preferWildcardCert
|
||||
});
|
||||
if (tls === null) {
|
||||
continue;
|
||||
}
|
||||
|
||||
const bgUiServiceName = `bg-r${bgResource.resourceId}-ui-service`;
|
||||
|
||||
if (bgResource.ssl) {
|
||||
const redirectRouterName = `bg-r${bgResource.resourceId}-redirect`;
|
||||
config_output.http.routers![redirectRouterName] = {
|
||||
entryPoints: [config.getRawConfig().traefik.http_entrypoint],
|
||||
middlewares: [redirectHttpsMiddlewareName],
|
||||
service: bgUiServiceName,
|
||||
rule: hostRule,
|
||||
priority: 100
|
||||
};
|
||||
}
|
||||
|
||||
// Collect online sites for this resource (for any type)
|
||||
const anySiteOnline = bgResource.targets.some((t) => t.siteOnline);
|
||||
|
||||
// Maintenance page logic for browser gateway resources
|
||||
let showBgMaintenancePage = false;
|
||||
if (bgResource.maintenanceModeEnabled) {
|
||||
if (bgResource.maintenanceModeType === "forced") {
|
||||
showBgMaintenancePage = true;
|
||||
} else if (bgResource.maintenanceModeType === "automatic") {
|
||||
showBgMaintenancePage = !anySiteOnline;
|
||||
}
|
||||
}
|
||||
|
||||
if (showBgMaintenancePage && maintenancePageUiUrl) {
|
||||
const bgMaintenanceServiceName = `bg-r${bgResource.resourceId}-maintenance-service`;
|
||||
const bgMaintenanceRouterName = `bg-r${bgResource.resourceId}-maintenance-router`;
|
||||
const bgRewriteMiddlewareName = `bg-r${bgResource.resourceId}-maintenance-rewrite`;
|
||||
const bgMaintenanceHeadersMiddlewareName = `bg-r${bgResource.resourceId}-maintenance-headers`;
|
||||
|
||||
const entrypointHttp =
|
||||
config.getRawConfig().traefik.http_entrypoint;
|
||||
const entrypointHttps =
|
||||
config.getRawConfig().traefik.https_entrypoint;
|
||||
|
||||
if (!config_output.http.services) config_output.http.services = {};
|
||||
if (!config_output.http.middlewares)
|
||||
config_output.http.middlewares = {};
|
||||
if (!config_output.http.routers) config_output.http.routers = {};
|
||||
|
||||
config_output.http.services![bgMaintenanceServiceName] = {
|
||||
loadBalancer: {
|
||||
servers: [
|
||||
{
|
||||
url: maintenancePageUiUrl
|
||||
}
|
||||
],
|
||||
passHostHeader: true
|
||||
}
|
||||
};
|
||||
|
||||
config_output.http.middlewares![bgRewriteMiddlewareName] = {
|
||||
replacePathRegex: {
|
||||
regex: "^/(.*)",
|
||||
replacement: "/maintenance-screen"
|
||||
}
|
||||
};
|
||||
|
||||
config_output.http.middlewares![
|
||||
bgMaintenanceHeadersMiddlewareName
|
||||
] = {
|
||||
headers: {
|
||||
customRequestHeaders: {
|
||||
Host: "app.pangolin.net", // if we are sending to the cloud the host needs to be this but we will pull the p-host to find the resource
|
||||
"p-host": fullDomain
|
||||
}
|
||||
}
|
||||
};
|
||||
|
||||
config_output.http.routers![bgMaintenanceRouterName] = {
|
||||
entryPoints: [
|
||||
bgResource.ssl ? entrypointHttps : entrypointHttp
|
||||
],
|
||||
service: bgMaintenanceServiceName,
|
||||
middlewares: [
|
||||
bgRewriteMiddlewareName,
|
||||
bgMaintenanceHeadersMiddlewareName
|
||||
],
|
||||
rule: hostRule,
|
||||
priority: 2000,
|
||||
...(bgResource.ssl ? { tls } : {})
|
||||
};
|
||||
|
||||
// Router to allow Next.js assets to load without rewrite
|
||||
config_output.http.routers![`${bgMaintenanceRouterName}-assets`] = {
|
||||
entryPoints: [
|
||||
bgResource.ssl ? entrypointHttps : entrypointHttp
|
||||
],
|
||||
service: bgMaintenanceServiceName,
|
||||
middlewares: [bgMaintenanceHeadersMiddlewareName],
|
||||
rule: `${hostRule} && (PathPrefix(\`/_next\`) || PathRegexp(\`^/__nextjs*\`) || Path(\`/favicon.ico\`))`,
|
||||
priority: 2001,
|
||||
...(bgResource.ssl ? { tls } : {})
|
||||
};
|
||||
|
||||
continue;
|
||||
}
|
||||
|
||||
// Group targets by type and generate per-type websocket routers and services
|
||||
const typeMap = new Map<string, typeof bgResource.targets>();
|
||||
for (const t of bgResource.targets) {
|
||||
if (!typeMap.has(t.bgType)) typeMap.set(t.bgType, []);
|
||||
typeMap.get(t.bgType)!.push(t);
|
||||
}
|
||||
|
||||
for (const [bgType, typedTargets] of typeMap.entries()) {
|
||||
const bgKey = `bg-r${bgResource.resourceId}-${bgType}`;
|
||||
const bgRouterName = `${bgKey}-router`;
|
||||
const bgServiceName = `${bgKey}-service`;
|
||||
const bgRule = `${hostRule} && PathPrefix(\`/gateway/${bgType}\`)`;
|
||||
|
||||
const servers = typedTargets
|
||||
.filter((t) => {
|
||||
if (!t.siteOnline && anySiteOnline) return false;
|
||||
if (t.siteType === "newt") return !!t.subnet;
|
||||
return false; // browser gateway only supported on newt sites
|
||||
})
|
||||
.map((t) => ({
|
||||
url: `http://${t.subnet!.split("/")[0]}:${browserGatewayPort}`
|
||||
}))
|
||||
.filter((v, i, a) => a.findIndex((u) => u.url === v.url) === i);
|
||||
|
||||
config_output.http.routers![bgRouterName] = {
|
||||
entryPoints: [
|
||||
bgResource.ssl
|
||||
? config.getRawConfig().traefik.https_entrypoint
|
||||
: config.getRawConfig().traefik.http_entrypoint
|
||||
],
|
||||
middlewares: routerMiddlewares,
|
||||
service: bgServiceName,
|
||||
rule: bgRule,
|
||||
priority: 110, // highest - websocket path takes precedence
|
||||
...(bgResource.ssl ? { tls } : {})
|
||||
};
|
||||
|
||||
config_output.http.services![bgServiceName] = {
|
||||
loadBalancer: {
|
||||
servers
|
||||
}
|
||||
};
|
||||
}
|
||||
|
||||
// UI: serve the browser gateway page from the internal pangolin instance.
|
||||
// The primary type is used for the path rewrite (e.g. /rdp), mirroring
|
||||
// how the maintenance page rewrites everything to /maintenance-screen.
|
||||
const primaryType = typeMap.keys().next().value as string;
|
||||
const uiRewriteMiddlewareName = `bg-r${bgResource.resourceId}-ui-rewrite`;
|
||||
const uiHeadersMiddlewareName = `bg-r${bgResource.resourceId}-ui-headers`;
|
||||
const entrypoint = bgResource.ssl
|
||||
? config.getRawConfig().traefik.https_entrypoint
|
||||
: config.getRawConfig().traefik.http_entrypoint;
|
||||
|
||||
if (!config_output.http.middlewares) {
|
||||
config_output.http.middlewares = {};
|
||||
}
|
||||
|
||||
config_output.http.middlewares![uiRewriteMiddlewareName] = {
|
||||
replacePathRegex: {
|
||||
regex: "^/(.*)",
|
||||
replacement: `/${primaryType}`
|
||||
}
|
||||
};
|
||||
|
||||
config_output.http.middlewares![uiHeadersMiddlewareName] = {
|
||||
headers: {
|
||||
customRequestHeaders: {
|
||||
Host: "app.pangolin.net", // if we are sending to the cloud the host needs to be this but we will pull the p-host to find the resource
|
||||
"p-host": fullDomain
|
||||
}
|
||||
}
|
||||
};
|
||||
|
||||
config_output.http.services![bgUiServiceName] = {
|
||||
loadBalancer: {
|
||||
servers: [
|
||||
{
|
||||
url: browserGatewayUiUrl
|
||||
}
|
||||
]
|
||||
}
|
||||
};
|
||||
|
||||
// Assets router at higher priority so /_next files load without rewrite.
|
||||
// Do NOT apply the path-rewrite middleware here — static assets must
|
||||
// keep their original path; only the host headers are needed.
|
||||
config_output.http.routers![
|
||||
`bg-r${bgResource.resourceId}-assets-router`
|
||||
] = {
|
||||
entryPoints: [entrypoint],
|
||||
middlewares: [...routerMiddlewares, uiHeadersMiddlewareName],
|
||||
service: bgUiServiceName,
|
||||
rule: `${hostRule} && (PathPrefix(\`/_next\`) || PathRegexp(\`^/__nextjs*\`) || Path(\`/favicon.ico\`))`,
|
||||
priority: 101,
|
||||
...(bgResource.ssl ? { tls } : {})
|
||||
};
|
||||
|
||||
// Catch-all router rewrites everything on the domain to /{primaryType}
|
||||
config_output.http.routers![`bg-r${bgResource.resourceId}-ui-router`] =
|
||||
{
|
||||
entryPoints: [entrypoint],
|
||||
middlewares: [
|
||||
...routerMiddlewares,
|
||||
uiRewriteMiddlewareName,
|
||||
uiHeadersMiddlewareName
|
||||
],
|
||||
service: bgUiServiceName,
|
||||
rule: hostRule,
|
||||
priority: 100,
|
||||
...(bgResource.ssl ? { tls } : {})
|
||||
};
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,44 @@
|
||||
import config from "@server/lib/config";
|
||||
|
||||
/**
|
||||
* Build the Traefik `tls` block for a domain using the cert-resolver /
|
||||
* wildcard-cert logic shared by both the OSS and private Traefik config
|
||||
* generators (used whenever certs are obtained directly via ACME rather
|
||||
* than through pangolin-dns).
|
||||
*/
|
||||
export function buildWildcardTls(params: {
|
||||
fullDomain: string;
|
||||
hasSubdomain: boolean;
|
||||
domainCertResolver?: string | null;
|
||||
preferWildcardCert?: boolean | null;
|
||||
}): { certResolver: string | undefined; domains?: { main: string }[] } {
|
||||
const { fullDomain, hasSubdomain, domainCertResolver, preferWildcardCert } =
|
||||
params;
|
||||
|
||||
const domainParts = fullDomain.split(".");
|
||||
let wildCard =
|
||||
domainParts.length <= 2
|
||||
? `*.${domainParts.join(".")}`
|
||||
: `*.${domainParts.slice(1).join(".")}`;
|
||||
if (!hasSubdomain) {
|
||||
wildCard = fullDomain;
|
||||
}
|
||||
|
||||
const globalDefaultResolver = config.getRawConfig().traefik.cert_resolver;
|
||||
const globalDefaultPreferWildcard =
|
||||
config.getRawConfig().traefik.prefer_wildcard_cert;
|
||||
|
||||
const resolverName = domainCertResolver
|
||||
? domainCertResolver.trim()
|
||||
: globalDefaultResolver;
|
||||
|
||||
const preferWildcard =
|
||||
preferWildcardCert !== undefined && preferWildcardCert !== null
|
||||
? preferWildcardCert
|
||||
: globalDefaultPreferWildcard;
|
||||
|
||||
return {
|
||||
certResolver: resolverName,
|
||||
...(preferWildcard ? { domains: [{ main: wildCard }] } : {})
|
||||
};
|
||||
}
|
||||
@@ -1,4 +1,13 @@
|
||||
import { db, targetHealthCheck, domains } from "@server/db";
|
||||
import {
|
||||
db,
|
||||
targetHealthCheck,
|
||||
domains,
|
||||
aiProviders,
|
||||
resourceAiProviders,
|
||||
siteResources,
|
||||
siteNetworks,
|
||||
exitNodes
|
||||
} from "@server/db";
|
||||
import {
|
||||
and,
|
||||
eq,
|
||||
@@ -12,41 +21,64 @@ import {
|
||||
} from "drizzle-orm";
|
||||
import logger from "@server/logger";
|
||||
import config from "@server/lib/config";
|
||||
import { resources, sites, Target, targets } from "@server/db";
|
||||
import createPathRewriteMiddleware from "./middleware";
|
||||
import { resources, sites, targets } from "@server/db";
|
||||
import { applyPathRewriteMiddleware } from "./middleware";
|
||||
import { sanitize, encodePath, validatePathRewriteConfig } from "./utils";
|
||||
import regionalCache from "@server/lib/cache";
|
||||
import { TargetWithSite } from "./types";
|
||||
import { buildWildcardTls } from "./certResolver";
|
||||
import { buildHostRule, appendPathMatch, computeRoutePriority } from "./rule";
|
||||
import {
|
||||
buildHttpLoadBalancerServers,
|
||||
buildStickySessionCookie,
|
||||
buildTcpUdpLoadBalancerServers,
|
||||
buildStickySessionIp
|
||||
} from "./loadBalancer";
|
||||
import { buildCustomHeadersMiddleware } from "./headersMiddleware";
|
||||
import {
|
||||
AI_GATEWAY_TRUST_MIDDLEWARE_RESOURCE,
|
||||
AI_GATEWAY_TRUST_MIDDLEWARE_SITE_RESOURCE,
|
||||
AI_GATEWAY_CLIENT_IP_MIDDLEWARE_NAME,
|
||||
getAiGatewayHost,
|
||||
buildAiGatewayTrustMiddlewares,
|
||||
buildAiGatewayClientIpMiddleware,
|
||||
buildAiGatewayHostHeaderMiddleware,
|
||||
buildAiGatewayRouterAndService
|
||||
} from "./aiGatewayMiddlewares";
|
||||
import {
|
||||
buildBrowserGatewayResourcesMap,
|
||||
buildBrowserGatewayConfig
|
||||
} from "./browserGateway";
|
||||
import { buildSiteResourceAliasCertPlaceholders } from "./siteResourceAlias";
|
||||
|
||||
const redirectHttpsMiddlewareName = "redirect-to-https";
|
||||
const badgerMiddlewareName = "badger";
|
||||
|
||||
// Define extended target type with site information
|
||||
type TargetWithSite = Target & {
|
||||
resourceId: number;
|
||||
targetId: number;
|
||||
ip: string | null;
|
||||
method: string | null;
|
||||
port: number | null;
|
||||
internalPort: number | null;
|
||||
enabled: boolean;
|
||||
health: string | null;
|
||||
site: {
|
||||
siteId: number;
|
||||
type: string;
|
||||
subnet: string | null;
|
||||
exitNodeId: number | null;
|
||||
online: boolean;
|
||||
};
|
||||
};
|
||||
|
||||
export async function getTraefikConfig(
|
||||
exitNodeId: number,
|
||||
siteTypes: string[],
|
||||
filterOutNamespaceDomains = false, // UNUSED BUT USED IN PRIVATE
|
||||
generateLoginPageRouters = false, // UNUSED BUT USED IN PRIVATE
|
||||
allowRawResources = true,
|
||||
maintenancePageUiUrl: string | null = null, // UNUSED BUT USED IN PRIVATE
|
||||
browserGatewayUiUrl: string | null = null // UNUSED BUT USED IN PRIVATE
|
||||
maintenancePageUiUrl: string | null = null,
|
||||
browserGatewayUiUrl: string | null = null,
|
||||
aiGatewayUrl: string | null = null
|
||||
): Promise<any> {
|
||||
// Get the exit node but cache it for 5 minutes to avoid hitting the DB too often
|
||||
const exitNodeCacheKey = `exitNode:${exitNodeId}`;
|
||||
let exitNode =
|
||||
await regionalCache.get<typeof exitNodes.$inferSelect>(
|
||||
exitNodeCacheKey
|
||||
);
|
||||
if (!exitNode) {
|
||||
[exitNode] = await db
|
||||
.select()
|
||||
.from(exitNodes)
|
||||
.where(eq(exitNodes.exitNodeId, exitNodeId))
|
||||
.limit(1);
|
||||
await regionalCache.set(exitNodeCacheKey, exitNode, 300);
|
||||
}
|
||||
|
||||
// Get resources with their targets and sites in a single optimized query
|
||||
// Start from sites on this exit node, then join to targets and resources
|
||||
const resourcesWithTargetsAndSites = await db
|
||||
@@ -68,8 +100,15 @@ export async function getTraefikConfig(
|
||||
responseHeaders: resources.responseHeaders,
|
||||
proxyProtocol: resources.proxyProtocol,
|
||||
proxyProtocolVersion: resources.proxyProtocolVersion,
|
||||
wildcard: resources.wildcard,
|
||||
mode: resources.mode,
|
||||
|
||||
maintenanceModeEnabled: resources.maintenanceModeEnabled,
|
||||
maintenanceModeType: resources.maintenanceModeType,
|
||||
maintenanceTitle: resources.maintenanceTitle,
|
||||
maintenanceMessage: resources.maintenanceMessage,
|
||||
maintenanceEstimatedTime: resources.maintenanceEstimatedTime,
|
||||
|
||||
// Target fields
|
||||
targetId: targets.targetId,
|
||||
targetEnabled: targets.enabled,
|
||||
@@ -88,7 +127,7 @@ export async function getTraefikConfig(
|
||||
siteId: sites.siteId,
|
||||
siteType: sites.type,
|
||||
siteOnline: sites.online,
|
||||
subnet: sites.subnet,
|
||||
subnet: sites.exitNodeSubnet,
|
||||
exitNodeId: sites.exitNodeId,
|
||||
// Domain cert resolver fields
|
||||
domainCertResolver: domains.certResolver,
|
||||
@@ -116,8 +155,15 @@ export async function getTraefikConfig(
|
||||
),
|
||||
inArray(sites.type, siteTypes),
|
||||
allowRawResources
|
||||
? inArray(resources.mode, ["http", "udp", "tcp"]) // allow all three
|
||||
: eq(resources.mode, "http")
|
||||
? inArray(resources.mode, [
|
||||
"http",
|
||||
"udp",
|
||||
"tcp",
|
||||
"vnc",
|
||||
"ssh",
|
||||
"rdp"
|
||||
]) // allow all three, plus browser-gateway modes
|
||||
: inArray(resources.mode, ["http", "vnc", "ssh", "rdp"])
|
||||
)
|
||||
)
|
||||
.orderBy(desc(targets.priority), targets.targetId); // stable ordering
|
||||
@@ -126,6 +172,9 @@ export async function getTraefikConfig(
|
||||
const resourcesMap = new Map();
|
||||
|
||||
resourcesWithTargetsAndSites.forEach((row) => {
|
||||
if (!["http", "tcp", "udp"].includes(row.mode)) {
|
||||
return;
|
||||
}
|
||||
const resourceId = row.resourceId;
|
||||
const resourceName = sanitize(row.resourceName) || "";
|
||||
const targetPath = encodePath(row.path); // Use encodePath to avoid collisions (e.g. "/a/b" vs "/a-b")
|
||||
@@ -211,8 +260,81 @@ export async function getTraefikConfig(
|
||||
});
|
||||
});
|
||||
|
||||
// Group browser gateway targets by resource (SSH/VNC/RDP-mode resources
|
||||
// served through the browser gateway web UI instead of a real target).
|
||||
const browserGatewayResourcesMap = browserGatewayUiUrl
|
||||
? buildBrowserGatewayResourcesMap(
|
||||
resourcesWithTargetsAndSites,
|
||||
filterOutNamespaceDomains
|
||||
)
|
||||
: new Map();
|
||||
|
||||
// Query siteResources in HTTP mode with SSL enabled and aliases, so
|
||||
// Traefik generates TLS certificates for those domains even before a
|
||||
// matching resource exists.
|
||||
const siteResourcesWithFullDomain = await db
|
||||
.select({
|
||||
siteResourceId: siteResources.siteResourceId,
|
||||
fullDomain: siteResources.fullDomain
|
||||
})
|
||||
.from(siteResources)
|
||||
.innerJoin(
|
||||
siteNetworks,
|
||||
eq(siteResources.networkId, siteNetworks.networkId)
|
||||
)
|
||||
.innerJoin(sites, eq(siteNetworks.siteId, sites.siteId))
|
||||
.where(
|
||||
and(
|
||||
eq(siteResources.enabled, true),
|
||||
isNotNull(siteResources.fullDomain),
|
||||
eq(siteResources.mode, "http"), // important so we dont double get the inference siteResources below
|
||||
eq(siteResources.ssl, true),
|
||||
eq(sites.exitNodeId, exitNodeId),
|
||||
inArray(sites.type, siteTypes)
|
||||
)
|
||||
);
|
||||
|
||||
// Inference-mode resources have no targets/sites (their "backend" is the
|
||||
// central AI gateway), so they can't be reached via the targets->sites
|
||||
// join above - query them separately and include them on every exit node.
|
||||
const inferenceResources = await db
|
||||
.selectDistinct({
|
||||
resourceId: resources.resourceId,
|
||||
resourceName: resources.name,
|
||||
fullDomain: resources.fullDomain,
|
||||
ssl: resources.ssl,
|
||||
subdomain: resources.subdomain,
|
||||
domainId: resources.domainId,
|
||||
enabled: resources.enabled,
|
||||
wildcard: resources.wildcard,
|
||||
domainCertResolver: domains.certResolver,
|
||||
preferWildcardCert: domains.preferWildcardCert
|
||||
})
|
||||
.from(resources)
|
||||
// .innerJoin(
|
||||
// resourceAiProviders,
|
||||
// eq(resources.resourceId, resourceAiProviders.resourceId)
|
||||
// )
|
||||
// .innerJoin(
|
||||
// aiProviders,
|
||||
// eq(resourceAiProviders.providerId, aiProviders.providerId)
|
||||
// )
|
||||
.leftJoin(domains, eq(domains.domainId, resources.domainId))
|
||||
.where(
|
||||
and(
|
||||
eq(resources.mode, "inference"),
|
||||
eq(resources.enabled, true)
|
||||
// eq(aiProviders.enabled, true)
|
||||
)
|
||||
);
|
||||
|
||||
// make sure we have at least one resource
|
||||
if (resourcesMap.size === 0) {
|
||||
if (
|
||||
resourcesMap.size === 0 &&
|
||||
inferenceResources.length === 0 &&
|
||||
browserGatewayResourcesMap.size === 0 &&
|
||||
siteResourcesWithFullDomain.length === 0
|
||||
) {
|
||||
return {};
|
||||
}
|
||||
|
||||
@@ -256,56 +378,12 @@ export async function getTraefikConfig(
|
||||
config_output.http.services = {};
|
||||
}
|
||||
|
||||
const domainParts = fullDomain.split(".");
|
||||
let wildCard;
|
||||
if (domainParts.length <= 2) {
|
||||
wildCard = `*.${domainParts.join(".")}`;
|
||||
} else {
|
||||
wildCard = `*.${domainParts.slice(1).join(".")}`;
|
||||
}
|
||||
|
||||
if (!resource.subdomain) {
|
||||
wildCard = resource.fullDomain;
|
||||
}
|
||||
|
||||
const globalDefaultResolver =
|
||||
config.getRawConfig().traefik.cert_resolver;
|
||||
const globalDefaultPreferWildcard =
|
||||
config.getRawConfig().traefik.prefer_wildcard_cert;
|
||||
|
||||
const domainCertResolver = resource.domainCertResolver;
|
||||
const preferWildcardCert = resource.preferWildcardCert;
|
||||
|
||||
let resolverName: string | undefined;
|
||||
let preferWildcard: boolean | undefined;
|
||||
// Handle both letsencrypt & custom cases
|
||||
if (domainCertResolver) {
|
||||
resolverName = domainCertResolver.trim();
|
||||
} else {
|
||||
resolverName = globalDefaultResolver;
|
||||
}
|
||||
|
||||
if (
|
||||
preferWildcardCert !== undefined &&
|
||||
preferWildcardCert !== null
|
||||
) {
|
||||
preferWildcard = preferWildcardCert;
|
||||
} else {
|
||||
preferWildcard = globalDefaultPreferWildcard;
|
||||
}
|
||||
|
||||
const tls = {
|
||||
certResolver: resolverName,
|
||||
...(preferWildcard
|
||||
? {
|
||||
domains: [
|
||||
{
|
||||
main: wildCard
|
||||
}
|
||||
]
|
||||
}
|
||||
: {})
|
||||
};
|
||||
const tls = buildWildcardTls({
|
||||
fullDomain,
|
||||
hasSubdomain: !!resource.subdomain,
|
||||
domainCertResolver: resource.domainCertResolver,
|
||||
preferWildcardCert: resource.preferWildcardCert
|
||||
});
|
||||
|
||||
const additionalMiddlewares =
|
||||
config.getRawConfig().traefik.additional_middlewares || [];
|
||||
@@ -316,155 +394,41 @@ export async function getTraefikConfig(
|
||||
];
|
||||
|
||||
// Handle path rewriting middleware
|
||||
if (
|
||||
resource.rewritePath !== null &&
|
||||
resource.path !== null &&
|
||||
resource.pathMatchType &&
|
||||
resource.rewritePathType
|
||||
) {
|
||||
// Create a unique middleware name
|
||||
const rewriteMiddlewareName = `rewrite-r${resource.resourceId}-${key}`;
|
||||
|
||||
try {
|
||||
const rewriteResult = createPathRewriteMiddleware(
|
||||
rewriteMiddlewareName,
|
||||
resource.path,
|
||||
resource.pathMatchType,
|
||||
resource.rewritePath,
|
||||
resource.rewritePathType
|
||||
);
|
||||
|
||||
// Initialize middlewares object if it doesn't exist
|
||||
if (!config_output.http.middlewares) {
|
||||
config_output.http.middlewares = {};
|
||||
}
|
||||
|
||||
// the middleware to the config
|
||||
Object.assign(
|
||||
config_output.http.middlewares,
|
||||
rewriteResult.middlewares
|
||||
);
|
||||
|
||||
// middlewares to the router middleware chain
|
||||
if (rewriteResult.chain) {
|
||||
// For chained middlewares (like stripPrefix + addPrefix)
|
||||
routerMiddlewares.push(...rewriteResult.chain);
|
||||
} else {
|
||||
// Single middleware
|
||||
routerMiddlewares.push(rewriteMiddlewareName);
|
||||
}
|
||||
|
||||
// logger.debug(
|
||||
// `Created path rewrite middleware ${rewriteMiddlewareName}: ${resource.pathMatchType}(${resource.path}) -> ${resource.rewritePathType}(${resource.rewritePath})`
|
||||
// );
|
||||
} catch (error) {
|
||||
logger.error(
|
||||
`Failed to create path rewrite middleware for resource ${resource.resourceId}: ${error}`
|
||||
);
|
||||
}
|
||||
}
|
||||
applyPathRewriteMiddleware(
|
||||
config_output,
|
||||
resource.resourceId,
|
||||
key,
|
||||
resource.path,
|
||||
resource.pathMatchType,
|
||||
resource.rewritePath,
|
||||
resource.rewritePathType,
|
||||
routerMiddlewares
|
||||
);
|
||||
|
||||
// Handle custom headers middleware
|
||||
if (resource.requestHeaders || resource.responseHeaders || resource.setHostHeader) {
|
||||
const requestHeadersObj: { [key: string]: string } = {};
|
||||
const responseHeadersObj: { [key: string]: string } = {};
|
||||
|
||||
if (resource.requestHeaders) {
|
||||
let requestHeadersArr: { name: string; value: string }[] = [];
|
||||
try {
|
||||
requestHeadersArr = JSON.parse(resource.requestHeaders) as {
|
||||
name: string;
|
||||
value: string;
|
||||
}[];
|
||||
} catch (e) {
|
||||
logger.warn(
|
||||
`Failed to parse requestHeaders for resource ${resource.resourceId}: ${e}`
|
||||
);
|
||||
}
|
||||
requestHeadersArr.forEach((header) => {
|
||||
requestHeadersObj[header.name] = header.value;
|
||||
});
|
||||
}
|
||||
|
||||
if (resource.setHostHeader) {
|
||||
requestHeadersObj["Host"] = resource.setHostHeader;
|
||||
}
|
||||
|
||||
if (resource.responseHeaders) {
|
||||
let responseHeadersArr: { name: string; value: string }[] = [];
|
||||
try {
|
||||
responseHeadersArr = JSON.parse(resource.responseHeaders) as {
|
||||
name: string;
|
||||
value: string;
|
||||
}[];
|
||||
} catch (e) {
|
||||
logger.warn(
|
||||
`Failed to parse responseHeaders for resource ${resource.resourceId}: ${e}`
|
||||
);
|
||||
}
|
||||
responseHeadersArr.forEach((header) => {
|
||||
responseHeadersObj[header.name] = header.value;
|
||||
});
|
||||
}
|
||||
|
||||
const hasRequestHeaders = Object.keys(requestHeadersObj).length > 0;
|
||||
const hasResponseHeaders = Object.keys(responseHeadersObj).length > 0;
|
||||
|
||||
if (hasRequestHeaders || hasResponseHeaders) {
|
||||
if (!config_output.http.middlewares) {
|
||||
config_output.http.middlewares = {};
|
||||
}
|
||||
config_output.http.middlewares[headersMiddlewareName] = {
|
||||
headers: {
|
||||
...(hasRequestHeaders && { customRequestHeaders: requestHeadersObj }),
|
||||
...(hasResponseHeaders && { customResponseHeaders: responseHeadersObj })
|
||||
}
|
||||
};
|
||||
|
||||
routerMiddlewares.push(headersMiddlewareName);
|
||||
const customHeadersMiddleware = buildCustomHeadersMiddleware(
|
||||
resource.requestHeaders,
|
||||
resource.responseHeaders,
|
||||
resource.setHostHeader,
|
||||
resource.resourceId
|
||||
);
|
||||
if (customHeadersMiddleware) {
|
||||
if (!config_output.http.middlewares) {
|
||||
config_output.http.middlewares = {};
|
||||
}
|
||||
config_output.http.middlewares[headersMiddlewareName] =
|
||||
customHeadersMiddleware;
|
||||
routerMiddlewares.push(headersMiddlewareName);
|
||||
}
|
||||
|
||||
// Build routing rules
|
||||
let rule = `Host(\`${fullDomain}\`)`;
|
||||
|
||||
// priority logic
|
||||
let priority: number;
|
||||
if (resource.priority && resource.priority != 100) {
|
||||
priority = resource.priority;
|
||||
} else {
|
||||
priority = 100;
|
||||
if (resource.path && resource.pathMatchType) {
|
||||
priority += 10;
|
||||
if (resource.pathMatchType === "exact") {
|
||||
priority += 5;
|
||||
} else if (resource.pathMatchType === "prefix") {
|
||||
priority += 3;
|
||||
} else if (resource.pathMatchType === "regex") {
|
||||
priority += 2;
|
||||
}
|
||||
if (resource.path === "/") {
|
||||
priority = 1; // lowest for catch-all
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
if (resource.path && resource.pathMatchType) {
|
||||
// priority += 1;
|
||||
// add path to rule based on match type
|
||||
let path = resource.path;
|
||||
// if the path doesn't start with a /, add it
|
||||
if (!path.startsWith("/")) {
|
||||
path = `/${path}`;
|
||||
}
|
||||
if (resource.pathMatchType === "exact") {
|
||||
rule += ` && Path(\`${path}\`)`;
|
||||
} else if (resource.pathMatchType === "prefix") {
|
||||
rule += ` && PathPrefix(\`${path}\`)`;
|
||||
} else if (resource.pathMatchType === "regex") {
|
||||
rule += ` && PathRegexp(\`${resource.path}\`)`; // this is the raw path because it's a regex
|
||||
}
|
||||
}
|
||||
let rule = buildHostRule(fullDomain);
|
||||
const priority = computeRoutePriority(
|
||||
resource.priority,
|
||||
resource.path,
|
||||
resource.pathMatchType
|
||||
);
|
||||
rule = appendPathMatch(rule, resource.path, resource.pathMatchType);
|
||||
|
||||
config_output.http.routers![routerName] = {
|
||||
entryPoints: [
|
||||
@@ -493,90 +457,9 @@ export async function getTraefikConfig(
|
||||
|
||||
config_output.http.services![serviceName] = {
|
||||
loadBalancer: {
|
||||
servers: (() => {
|
||||
// Check if any sites are online
|
||||
// THIS IS SO THAT THERE IS SOME IMMEDIATE FEEDBACK
|
||||
// EVEN IF THE SITES HAVE NOT UPDATED YET FROM THE
|
||||
// RECEIVE BANDWIDTH ENDPOINT.
|
||||
|
||||
// TODO: HOW TO HANDLE ^^^^^^ BETTER
|
||||
const anySitesOnline = targets.some(
|
||||
(target) => target.site.online
|
||||
);
|
||||
|
||||
return (
|
||||
targets
|
||||
.filter((target) => {
|
||||
if (!target.enabled) {
|
||||
return false;
|
||||
}
|
||||
|
||||
if (target.health == "unhealthy") {
|
||||
return false;
|
||||
}
|
||||
|
||||
// If any sites are online, exclude offline sites
|
||||
if (anySitesOnline && !target.site.online) {
|
||||
return false;
|
||||
}
|
||||
|
||||
if (
|
||||
target.site.type === "local" ||
|
||||
target.site.type === "wireguard"
|
||||
) {
|
||||
if (
|
||||
!target.ip ||
|
||||
!target.port ||
|
||||
!target.method
|
||||
) {
|
||||
return false;
|
||||
}
|
||||
} else if (target.site.type === "newt") {
|
||||
if (
|
||||
!target.internalPort ||
|
||||
!target.method ||
|
||||
!target.site.subnet
|
||||
) {
|
||||
return false;
|
||||
}
|
||||
}
|
||||
return true;
|
||||
})
|
||||
.map((target) => {
|
||||
if (
|
||||
target.site.type === "local" ||
|
||||
target.site.type === "wireguard"
|
||||
) {
|
||||
return {
|
||||
url: `${target.method}://${target.ip}:${target.port}`
|
||||
};
|
||||
} else if (target.site.type === "newt") {
|
||||
const ip =
|
||||
target.site.subnet!.split("/")[0];
|
||||
return {
|
||||
url: `${target.method}://${ip}:${target.internalPort}`
|
||||
};
|
||||
}
|
||||
})
|
||||
// filter out duplicates
|
||||
.filter(
|
||||
(v, i, a) =>
|
||||
a.findIndex(
|
||||
(t) => t && v && t.url === v.url
|
||||
) === i
|
||||
)
|
||||
);
|
||||
})(),
|
||||
servers: buildHttpLoadBalancerServers(targets),
|
||||
...(resource.stickySession
|
||||
? {
|
||||
sticky: {
|
||||
cookie: {
|
||||
name: "p_sticky", // TODO: make this configurable via config.yml like other cookies
|
||||
secure: resource.ssl,
|
||||
httpOnly: true
|
||||
}
|
||||
}
|
||||
}
|
||||
? buildStickySessionCookie(resource.ssl)
|
||||
: {})
|
||||
}
|
||||
};
|
||||
@@ -626,75 +509,249 @@ export async function getTraefikConfig(
|
||||
|
||||
config_output[protocol].services[serviceName] = {
|
||||
loadBalancer: {
|
||||
servers: (() => {
|
||||
// Check if any sites are online
|
||||
const anySitesOnline = targets.some(
|
||||
(target) => target.site.online
|
||||
);
|
||||
|
||||
return targets
|
||||
.filter((target) => {
|
||||
if (!target.enabled) {
|
||||
return false;
|
||||
}
|
||||
|
||||
// If any sites are online, exclude offline sites
|
||||
if (anySitesOnline && !target.site.online) {
|
||||
return false;
|
||||
}
|
||||
|
||||
if (
|
||||
target.site.type === "local" ||
|
||||
target.site.type === "wireguard"
|
||||
) {
|
||||
if (!target.ip || !target.port) {
|
||||
return false;
|
||||
}
|
||||
} else if (target.site.type === "newt") {
|
||||
if (
|
||||
!target.internalPort ||
|
||||
!target.site.subnet
|
||||
) {
|
||||
return false;
|
||||
}
|
||||
}
|
||||
return true;
|
||||
})
|
||||
.map((target) => {
|
||||
if (
|
||||
target.site.type === "local" ||
|
||||
target.site.type === "wireguard"
|
||||
) {
|
||||
return {
|
||||
address: `${target.ip}:${target.port}`
|
||||
};
|
||||
} else if (target.site.type === "newt") {
|
||||
const ip =
|
||||
target.site.subnet!.split("/")[0];
|
||||
return {
|
||||
address: `${ip}:${target.internalPort}`
|
||||
};
|
||||
}
|
||||
});
|
||||
})(),
|
||||
servers: buildTcpUdpLoadBalancerServers(targets),
|
||||
...(resource.proxyProtocol && protocol == "tcp"
|
||||
? {
|
||||
serversTransport: `${ppPrefix}${resource.proxyProtocolVersion || 1}@file` // TODO: does @file here cause issues?
|
||||
}
|
||||
: {}),
|
||||
...(resource.stickySession
|
||||
? {
|
||||
sticky: {
|
||||
ipStrategy: {
|
||||
depth: 0,
|
||||
sourcePort: true
|
||||
}
|
||||
}
|
||||
}
|
||||
: {})
|
||||
...(resource.stickySession ? buildStickySessionIp() : {})
|
||||
}
|
||||
};
|
||||
}
|
||||
}
|
||||
|
||||
if (browserGatewayUiUrl) {
|
||||
buildBrowserGatewayConfig({
|
||||
config_output,
|
||||
browserGatewayResourcesMap,
|
||||
browserGatewayUiUrl,
|
||||
maintenancePageUiUrl,
|
||||
badgerMiddlewareName,
|
||||
redirectHttpsMiddlewareName,
|
||||
resolveTls: ({
|
||||
fullDomain,
|
||||
hasSubdomain,
|
||||
domainCertResolver,
|
||||
preferWildcardCert
|
||||
}) =>
|
||||
buildWildcardTls({
|
||||
fullDomain,
|
||||
hasSubdomain,
|
||||
domainCertResolver,
|
||||
preferWildcardCert
|
||||
})
|
||||
});
|
||||
}
|
||||
|
||||
// Add Traefik routes for siteResource aliases (HTTP mode + SSL) so that
|
||||
// Traefik generates TLS certificates for those domains even when no
|
||||
// matching resource exists yet.
|
||||
if (siteResourcesWithFullDomain.length > 0) {
|
||||
// Build a set of domains already covered by normal resources
|
||||
const existingFullDomains = new Set<string>();
|
||||
for (const resource of resourcesMap.values()) {
|
||||
if (resource.fullDomain) {
|
||||
existingFullDomains.add(resource.fullDomain);
|
||||
}
|
||||
}
|
||||
|
||||
buildSiteResourceAliasCertPlaceholders({
|
||||
config_output,
|
||||
siteResourcesWithFullDomain,
|
||||
existingFullDomains,
|
||||
maintenancePageUiUrl,
|
||||
redirectHttpsMiddlewareName,
|
||||
resolveTls: (fullDomain) =>
|
||||
buildWildcardTls({
|
||||
fullDomain,
|
||||
hasSubdomain: true
|
||||
})
|
||||
});
|
||||
}
|
||||
|
||||
if (aiGatewayUrl) {
|
||||
// The AI gateway may live on a different host than the inference
|
||||
// resource itself (e.g. a remote exit node forwarding to the
|
||||
// central dashboard over a tunnel). passHostHeader would forward
|
||||
// the resource's own Host, which that external host won't
|
||||
// recognize, so we pin the Host header to the gateway's own host
|
||||
// and smuggle the original resource host through in "p-host"
|
||||
// instead.
|
||||
const aiGatewayHost = getAiGatewayHost(aiGatewayUrl);
|
||||
|
||||
if (!config_output.http.middlewares) {
|
||||
config_output.http.middlewares = {};
|
||||
}
|
||||
Object.assign(
|
||||
config_output.http.middlewares,
|
||||
buildAiGatewayTrustMiddlewares()
|
||||
);
|
||||
|
||||
const aiGatewayClientIpMiddleware = buildAiGatewayClientIpMiddleware();
|
||||
const enableAiGatewayClientIpHeader = !!aiGatewayClientIpMiddleware;
|
||||
if (aiGatewayClientIpMiddleware) {
|
||||
Object.assign(
|
||||
config_output.http.middlewares,
|
||||
aiGatewayClientIpMiddleware
|
||||
);
|
||||
}
|
||||
|
||||
// Public inference resources: same TLS/cert-resolver handling as
|
||||
// plain http-mode resources, but the service points at the AI
|
||||
// gateway instead of any real backend targets.
|
||||
//
|
||||
// Inference-mode resources are allowed to share a fullDomain with
|
||||
// each other (see createResource.ts), and a siteResource inference
|
||||
// alias can share that domain too - all of them proxy to the same
|
||||
// aiGatewayUrl, so dedupe by fullDomain here (lowest resourceId
|
||||
// wins, for stable output across regenerations) and skip the
|
||||
// siteResource alias router for any domain already covered below.
|
||||
const eligibleInferenceResources = inferenceResources
|
||||
.filter((ir) => ir.enabled && ir.domainId && ir.fullDomain)
|
||||
.sort((a, b) => a.resourceId - b.resourceId);
|
||||
const dedupedInferenceResources = new Map<
|
||||
string,
|
||||
(typeof eligibleInferenceResources)[number]
|
||||
>();
|
||||
for (const ir of eligibleInferenceResources) {
|
||||
if (!dedupedInferenceResources.has(ir.fullDomain!)) {
|
||||
dedupedInferenceResources.set(ir.fullDomain!, ir);
|
||||
}
|
||||
}
|
||||
|
||||
const publicInferenceDomains = new Set<string>();
|
||||
for (const ir of dedupedInferenceResources.values()) {
|
||||
if (!config_output.http.routers) config_output.http.routers = {};
|
||||
if (!config_output.http.services) config_output.http.services = {};
|
||||
|
||||
const fullDomain = ir.fullDomain!;
|
||||
const irKey = `inference-r${ir.resourceId}`;
|
||||
const routerName = `${irKey}-router`;
|
||||
const serviceName = `${irKey}-service`;
|
||||
|
||||
const rule = buildHostRule(fullDomain, ir.wildcard);
|
||||
|
||||
const tls = buildWildcardTls({
|
||||
fullDomain,
|
||||
hasSubdomain: !!ir.subdomain,
|
||||
domainCertResolver: ir.domainCertResolver,
|
||||
preferWildcardCert: ir.preferWildcardCert
|
||||
});
|
||||
|
||||
const irHeadersMiddlewareName = `${irKey}-headers-middleware`;
|
||||
config_output.http.middlewares[irHeadersMiddlewareName] =
|
||||
buildAiGatewayHostHeaderMiddleware(aiGatewayHost, fullDomain);
|
||||
|
||||
const additionalMiddlewares =
|
||||
config.getRawConfig().traefik.additional_middlewares || [];
|
||||
const routerMiddlewares = [
|
||||
badgerMiddlewareName,
|
||||
AI_GATEWAY_TRUST_MIDDLEWARE_RESOURCE,
|
||||
irHeadersMiddlewareName,
|
||||
...additionalMiddlewares
|
||||
];
|
||||
|
||||
const { routers, services } = buildAiGatewayRouterAndService({
|
||||
routerName,
|
||||
serviceName,
|
||||
rule,
|
||||
ssl: ir.ssl,
|
||||
tls,
|
||||
priority: 100,
|
||||
routerMiddlewares,
|
||||
aiGatewayUrl,
|
||||
redirectHttpsMiddlewareName
|
||||
});
|
||||
Object.assign(config_output.http.routers, routers);
|
||||
Object.assign(config_output.http.services, services);
|
||||
publicInferenceDomains.add(fullDomain);
|
||||
}
|
||||
|
||||
// Private (siteResource) inference resources: routed by their alias
|
||||
// instead of a public fullDomain, and deliberately WITHOUT the
|
||||
// badger middleware - no per-user auth/policy stack exists for
|
||||
// siteResources today, so gating here is reachability-only for now.
|
||||
const siteResourcesInference = await db
|
||||
.selectDistinct({
|
||||
siteResourceId: siteResources.siteResourceId,
|
||||
fullDomain: siteResources.fullDomain,
|
||||
ssl: siteResources.ssl,
|
||||
enabled: siteResources.enabled
|
||||
})
|
||||
.from(siteResources)
|
||||
.where(
|
||||
and(
|
||||
eq(siteResources.mode, "inference"),
|
||||
eq(siteResources.enabled, true),
|
||||
isNotNull(siteResources.fullDomain)
|
||||
)
|
||||
);
|
||||
|
||||
if (exitNode) {
|
||||
for (const sr of siteResourcesInference) {
|
||||
if (!sr.enabled || !sr.fullDomain) continue;
|
||||
|
||||
// A public inference resource already owns a router for
|
||||
// this exact fullDomain - both point at the same AI gateway,
|
||||
// so avoid registering a duplicate router for it here.
|
||||
if (publicInferenceDomains.has(sr.fullDomain)) continue;
|
||||
|
||||
if (!config_output.http.routers)
|
||||
config_output.http.routers = {};
|
||||
if (!config_output.http.services)
|
||||
config_output.http.services = {};
|
||||
|
||||
const fullDomain = sr.fullDomain;
|
||||
const srKey = `inference-sr${sr.siteResourceId}`;
|
||||
const routerName = `${srKey}-router`;
|
||||
const serviceName = `${srKey}-service`;
|
||||
const rule = `Host(\`${fullDomain}\`) && ClientIP(\`${exitNode.address}\`)`; // restrict to coming from the exit node ip range that the client is connected to
|
||||
|
||||
// siteResource aliases don't have a per-domain cert resolver
|
||||
// stored, so always fall back to the global defaults.
|
||||
const tls = buildWildcardTls({
|
||||
fullDomain,
|
||||
hasSubdomain: true
|
||||
});
|
||||
|
||||
const srHeadersMiddlewareName = `${srKey}-headers-middleware`;
|
||||
if (!config_output.http.middlewares) {
|
||||
config_output.http.middlewares = {};
|
||||
}
|
||||
config_output.http.middlewares[srHeadersMiddlewareName] =
|
||||
buildAiGatewayHostHeaderMiddleware(
|
||||
aiGatewayHost,
|
||||
fullDomain
|
||||
);
|
||||
|
||||
const additionalMiddlewares =
|
||||
config.getRawConfig().traefik.additional_middlewares || [];
|
||||
const routerMiddlewares = [
|
||||
...(enableAiGatewayClientIpHeader
|
||||
? [AI_GATEWAY_CLIENT_IP_MIDDLEWARE_NAME]
|
||||
: []),
|
||||
AI_GATEWAY_TRUST_MIDDLEWARE_SITE_RESOURCE,
|
||||
srHeadersMiddlewareName,
|
||||
...additionalMiddlewares
|
||||
];
|
||||
|
||||
const { routers, services } = buildAiGatewayRouterAndService({
|
||||
routerName,
|
||||
serviceName,
|
||||
rule,
|
||||
ssl: sr.ssl,
|
||||
tls,
|
||||
priority: 200, // we want to match on the site resource first because the clientIP rule is more specific than the public inference resource rule, which is just the exit node IP range. so we give it a higher priority to ensure it matches first.
|
||||
routerMiddlewares,
|
||||
aiGatewayUrl,
|
||||
redirectHttpsMiddlewareName
|
||||
});
|
||||
Object.assign(config_output.http.routers, routers);
|
||||
Object.assign(config_output.http.services, services);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
return config_output;
|
||||
}
|
||||
|
||||
@@ -0,0 +1,78 @@
|
||||
import logger from "@server/logger";
|
||||
|
||||
function parseHeaders(
|
||||
headers: string,
|
||||
label: string,
|
||||
resourceId: number
|
||||
): { name: string; value: string }[] {
|
||||
try {
|
||||
return JSON.parse(headers) as {
|
||||
name: string;
|
||||
value: string;
|
||||
}[];
|
||||
} catch (e) {
|
||||
logger.warn(
|
||||
`Failed to parse ${label} for resource ${resourceId}: ${e}`
|
||||
);
|
||||
return [];
|
||||
}
|
||||
}
|
||||
|
||||
/**
|
||||
* Build the custom headers middleware definition for a resource's
|
||||
* custom request/response headers + setHostHeader config. Returns null when
|
||||
* there are no headers to set, so the caller can skip attaching the
|
||||
* middleware.
|
||||
*/
|
||||
export function buildCustomHeadersMiddleware(
|
||||
requestHeaders: string | null | undefined,
|
||||
responseHeaders: string | null | undefined,
|
||||
setHostHeader: string | null | undefined,
|
||||
resourceId: number
|
||||
): {
|
||||
headers: {
|
||||
customRequestHeaders?: { [key: string]: string };
|
||||
customResponseHeaders?: { [key: string]: string };
|
||||
};
|
||||
} | null {
|
||||
const requestHeadersObj: { [key: string]: string } = {};
|
||||
const responseHeadersObj: { [key: string]: string } = {};
|
||||
|
||||
if (requestHeaders) {
|
||||
parseHeaders(requestHeaders, "requestHeaders", resourceId).forEach(
|
||||
(header) => {
|
||||
requestHeadersObj[header.name] = header.value;
|
||||
}
|
||||
);
|
||||
}
|
||||
|
||||
if (setHostHeader) {
|
||||
requestHeadersObj["Host"] = setHostHeader;
|
||||
}
|
||||
|
||||
if (responseHeaders) {
|
||||
parseHeaders(responseHeaders, "responseHeaders", resourceId).forEach(
|
||||
(header) => {
|
||||
responseHeadersObj[header.name] = header.value;
|
||||
}
|
||||
);
|
||||
}
|
||||
|
||||
const hasRequestHeaders = Object.keys(requestHeadersObj).length > 0;
|
||||
const hasResponseHeaders = Object.keys(responseHeadersObj).length > 0;
|
||||
|
||||
if (!hasRequestHeaders && !hasResponseHeaders) {
|
||||
return null;
|
||||
}
|
||||
|
||||
return {
|
||||
headers: {
|
||||
...(hasRequestHeaders && {
|
||||
customRequestHeaders: requestHeadersObj
|
||||
}),
|
||||
...(hasResponseHeaders && {
|
||||
customResponseHeaders: responseHeadersObj
|
||||
})
|
||||
}
|
||||
};
|
||||
}
|
||||
@@ -0,0 +1,134 @@
|
||||
import { TargetWithSite } from "./types";
|
||||
|
||||
/**
|
||||
* Build the loadBalancer.servers list for an HTTP-mode resource, preferring
|
||||
* currently-online sites but falling back to all enabled/healthy targets if
|
||||
* none are online yet (so there's still some feedback before sites report
|
||||
* back over the receive-bandwidth endpoint).
|
||||
*/
|
||||
export function buildHttpLoadBalancerServers(targets: TargetWithSite[]) {
|
||||
const anySitesOnline = targets.some((target) => target.site.online);
|
||||
|
||||
return targets
|
||||
.filter((target) => {
|
||||
if (!target.enabled) {
|
||||
return false;
|
||||
}
|
||||
|
||||
if (target.health == "unhealthy") {
|
||||
return false;
|
||||
}
|
||||
|
||||
// If any sites are online, exclude offline sites
|
||||
if (anySitesOnline && !target.site.online) {
|
||||
return false;
|
||||
}
|
||||
|
||||
if (
|
||||
target.site.type === "local" ||
|
||||
target.site.type === "wireguard"
|
||||
) {
|
||||
if (!target.ip || !target.port || !target.method) {
|
||||
return false;
|
||||
}
|
||||
} else if (target.site.type === "newt") {
|
||||
if (
|
||||
!target.internalPort ||
|
||||
!target.method ||
|
||||
!target.site.subnet
|
||||
) {
|
||||
return false;
|
||||
}
|
||||
}
|
||||
return true;
|
||||
})
|
||||
.map((target) => {
|
||||
if (
|
||||
target.site.type === "local" ||
|
||||
target.site.type === "wireguard"
|
||||
) {
|
||||
return {
|
||||
url: `${target.method}://${target.ip}:${target.port}`
|
||||
};
|
||||
} else if (target.site.type === "newt") {
|
||||
const ip = target.site.subnet!.split("/")[0];
|
||||
return {
|
||||
url: `${target.method}://${ip}:${target.internalPort}`
|
||||
};
|
||||
}
|
||||
})
|
||||
.filter(
|
||||
(v, i, a) => a.findIndex((t) => t && v && t.url === v.url) === i
|
||||
);
|
||||
}
|
||||
|
||||
export function buildStickySessionCookie(ssl: boolean | null) {
|
||||
return {
|
||||
sticky: {
|
||||
cookie: {
|
||||
name: "p_sticky", // TODO: make this configurable via config.yml like other cookies
|
||||
secure: ssl,
|
||||
httpOnly: true
|
||||
}
|
||||
}
|
||||
};
|
||||
}
|
||||
|
||||
/**
|
||||
* Build the loadBalancer.servers list for a TCP/UDP-mode resource.
|
||||
*/
|
||||
export function buildTcpUdpLoadBalancerServers(targets: TargetWithSite[]) {
|
||||
const anySitesOnline = targets.some((target) => target.site.online);
|
||||
|
||||
return targets
|
||||
.filter((target) => {
|
||||
if (!target.enabled) {
|
||||
return false;
|
||||
}
|
||||
|
||||
// If any sites are online, exclude offline sites
|
||||
if (anySitesOnline && !target.site.online) {
|
||||
return false;
|
||||
}
|
||||
|
||||
if (
|
||||
target.site.type === "local" ||
|
||||
target.site.type === "wireguard"
|
||||
) {
|
||||
if (!target.ip || !target.port) {
|
||||
return false;
|
||||
}
|
||||
} else if (target.site.type === "newt") {
|
||||
if (!target.internalPort || !target.site.subnet) {
|
||||
return false;
|
||||
}
|
||||
}
|
||||
return true;
|
||||
})
|
||||
.map((target) => {
|
||||
if (
|
||||
target.site.type === "local" ||
|
||||
target.site.type === "wireguard"
|
||||
) {
|
||||
return {
|
||||
address: `${target.ip}:${target.port}`
|
||||
};
|
||||
} else if (target.site.type === "newt") {
|
||||
const ip = target.site.subnet!.split("/")[0];
|
||||
return {
|
||||
address: `${ip}:${target.internalPort}`
|
||||
};
|
||||
}
|
||||
});
|
||||
}
|
||||
|
||||
export function buildStickySessionIp() {
|
||||
return {
|
||||
sticky: {
|
||||
ipStrategy: {
|
||||
depth: 0,
|
||||
sourcePort: true
|
||||
}
|
||||
}
|
||||
};
|
||||
}
|
||||
@@ -1,5 +1,64 @@
|
||||
import logger from "@server/logger";
|
||||
|
||||
/**
|
||||
* Create (if configured) and attach a path-rewrite middleware for a
|
||||
* resource, mutating both config_output.http.middlewares and the
|
||||
* router's middleware chain. Shared by the OSS and private Traefik config
|
||||
* generators, which apply it identically.
|
||||
*/
|
||||
export function applyPathRewriteMiddleware(
|
||||
config_output: any,
|
||||
resourceId: number,
|
||||
key: string,
|
||||
path: string | null,
|
||||
pathMatchType: string | null,
|
||||
rewritePath: string | null,
|
||||
rewritePathType: string | null,
|
||||
routerMiddlewares: string[]
|
||||
) {
|
||||
if (
|
||||
rewritePath === null ||
|
||||
path === null ||
|
||||
!pathMatchType ||
|
||||
!rewritePathType
|
||||
) {
|
||||
return;
|
||||
}
|
||||
|
||||
const rewriteMiddlewareName = `rewrite-r${resourceId}-${key}`;
|
||||
|
||||
try {
|
||||
const rewriteResult = createPathRewriteMiddleware(
|
||||
rewriteMiddlewareName,
|
||||
path,
|
||||
pathMatchType,
|
||||
rewritePath,
|
||||
rewritePathType
|
||||
);
|
||||
|
||||
if (!config_output.http.middlewares) {
|
||||
config_output.http.middlewares = {};
|
||||
}
|
||||
|
||||
Object.assign(
|
||||
config_output.http.middlewares,
|
||||
rewriteResult.middlewares
|
||||
);
|
||||
|
||||
if (rewriteResult.chain) {
|
||||
// For chained middlewares (like stripPrefix + addPrefix)
|
||||
routerMiddlewares.push(...rewriteResult.chain);
|
||||
} else {
|
||||
// Single middleware
|
||||
routerMiddlewares.push(rewriteMiddlewareName);
|
||||
}
|
||||
} catch (error) {
|
||||
logger.error(
|
||||
`Failed to create path rewrite middleware for resource ${resourceId}: ${error}`
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
export default function createPathRewriteMiddleware(
|
||||
middlewareName: string,
|
||||
path: string,
|
||||
|
||||
@@ -0,0 +1,71 @@
|
||||
/**
|
||||
* Build the Host()/HostRegexp() Traefik rule for a resource's domain.
|
||||
* Wildcard resources match any single subdomain via HostRegexp.
|
||||
*/
|
||||
export function buildHostRule(
|
||||
fullDomain: string,
|
||||
wildcard?: boolean | null
|
||||
): string {
|
||||
if (wildcard && fullDomain.startsWith("*.")) {
|
||||
// Convert *.foo.bar.com -> HostRegexp(`^[^.]+\.foo\.bar\.com$`)
|
||||
const escaped = fullDomain.slice(2).replace(/\./g, "\\.");
|
||||
return `HostRegexp(\`^[^.]+\\.${escaped}$\`)`;
|
||||
}
|
||||
return `Host(\`${fullDomain}\`)`;
|
||||
}
|
||||
|
||||
/**
|
||||
* Append a path-matching clause to a Traefik rule based on the resource's
|
||||
* configured path and pathMatchType.
|
||||
*/
|
||||
export function appendPathMatch(
|
||||
rule: string,
|
||||
path: string | null | undefined,
|
||||
pathMatchType: string | null | undefined
|
||||
): string {
|
||||
if (!path || !pathMatchType) return rule;
|
||||
|
||||
let p = path;
|
||||
if (!p.startsWith("/")) {
|
||||
p = `/${p}`;
|
||||
}
|
||||
|
||||
if (pathMatchType === "exact") {
|
||||
return `${rule} && Path(\`${p}\`)`;
|
||||
} else if (pathMatchType === "prefix") {
|
||||
return `${rule} && PathPrefix(\`${p}\`)`;
|
||||
} else if (pathMatchType === "regex") {
|
||||
return `${rule} && PathRegexp(\`${path}\`)`; // this is the raw path because it's a regex
|
||||
}
|
||||
return rule;
|
||||
}
|
||||
|
||||
/**
|
||||
* Compute the router priority for a resource, favoring an explicit override
|
||||
* and otherwise deriving it from the path match specificity.
|
||||
*/
|
||||
export function computeRoutePriority(
|
||||
priority: number | null | undefined,
|
||||
path: string | null | undefined,
|
||||
pathMatchType: string | null | undefined
|
||||
): number {
|
||||
if (priority && priority != 100) {
|
||||
return priority;
|
||||
}
|
||||
|
||||
let p = 100;
|
||||
if (path && pathMatchType) {
|
||||
p += 10;
|
||||
if (pathMatchType === "exact") {
|
||||
p += 5;
|
||||
} else if (pathMatchType === "prefix") {
|
||||
p += 3;
|
||||
} else if (pathMatchType === "regex") {
|
||||
p += 2;
|
||||
}
|
||||
if (path === "/") {
|
||||
p = 1; // lowest for catch-all
|
||||
}
|
||||
}
|
||||
return p;
|
||||
}
|
||||
@@ -0,0 +1,114 @@
|
||||
import config from "@server/lib/config";
|
||||
|
||||
export type SiteResourceAliasRow = {
|
||||
siteResourceId: number;
|
||||
fullDomain: string | null;
|
||||
};
|
||||
|
||||
/**
|
||||
* Add placeholder Traefik routes for siteResource HTTP aliases so Traefik
|
||||
* generates TLS certificates for those domains even before a matching
|
||||
* resource exists. Requests that land on these routes before a real
|
||||
* resource is created are served the placeholder page. TLS/cert-resolver
|
||||
* handling differs between the OSS and private (pangolin-dns aware) config
|
||||
* generators, so callers resolve that themselves via resolveTls - returning
|
||||
* null skips the alias (no valid cert available yet).
|
||||
*/
|
||||
export function buildSiteResourceAliasCertPlaceholders(params: {
|
||||
config_output: any;
|
||||
siteResourcesWithFullDomain: SiteResourceAliasRow[];
|
||||
existingFullDomains: Set<string>;
|
||||
maintenancePageUiUrl: string | null;
|
||||
redirectHttpsMiddlewareName: string;
|
||||
resolveTls: (fullDomain: string) => any | null;
|
||||
}): void {
|
||||
const {
|
||||
config_output,
|
||||
siteResourcesWithFullDomain,
|
||||
existingFullDomains,
|
||||
maintenancePageUiUrl,
|
||||
redirectHttpsMiddlewareName,
|
||||
resolveTls
|
||||
} = params;
|
||||
|
||||
if (siteResourcesWithFullDomain.length === 0 || !maintenancePageUiUrl) {
|
||||
return;
|
||||
}
|
||||
|
||||
for (const sr of siteResourcesWithFullDomain) {
|
||||
if (!sr.fullDomain) continue;
|
||||
|
||||
// Skip if this alias is already handled by a resource router
|
||||
if (existingFullDomains.has(sr.fullDomain)) continue;
|
||||
|
||||
const fullDomain = sr.fullDomain;
|
||||
const srKey = `site-resource-cert-${sr.siteResourceId}`;
|
||||
const siteResourceServiceName = `${srKey}-service`;
|
||||
const siteResourceRouterName = `${srKey}-router`;
|
||||
const siteResourceRewriteMiddlewareName = `${srKey}-rewrite`;
|
||||
|
||||
if (!config_output.http.routers) {
|
||||
config_output.http.routers = {};
|
||||
}
|
||||
if (!config_output.http.services) {
|
||||
config_output.http.services = {};
|
||||
}
|
||||
if (!config_output.http.middlewares) {
|
||||
config_output.http.middlewares = {};
|
||||
}
|
||||
|
||||
// Service pointing at the internal maintenance/Next.js page
|
||||
config_output.http.services[siteResourceServiceName] = {
|
||||
loadBalancer: {
|
||||
servers: [
|
||||
{
|
||||
url: maintenancePageUiUrl
|
||||
}
|
||||
],
|
||||
passHostHeader: true
|
||||
}
|
||||
};
|
||||
|
||||
// Middleware that rewrites any path to /private-maintenance-screen
|
||||
config_output.http.middlewares[siteResourceRewriteMiddlewareName] = {
|
||||
replacePathRegex: {
|
||||
regex: "^/(.*)",
|
||||
replacement: "/private-maintenance-screen"
|
||||
}
|
||||
};
|
||||
|
||||
// HTTP -> HTTPS redirect so the ACME challenge can be served
|
||||
config_output.http.routers[`${siteResourceRouterName}-redirect`] = {
|
||||
entryPoints: [config.getRawConfig().traefik.http_entrypoint],
|
||||
middlewares: [redirectHttpsMiddlewareName],
|
||||
service: siteResourceServiceName,
|
||||
rule: `Host(\`${fullDomain}\`)`,
|
||||
priority: 100
|
||||
};
|
||||
|
||||
// Determine TLS / cert-resolver configuration
|
||||
const tls = resolveTls(fullDomain);
|
||||
if (tls === null) {
|
||||
continue;
|
||||
}
|
||||
|
||||
// HTTPS router - presence of this entry triggers cert generation
|
||||
config_output.http.routers[siteResourceRouterName] = {
|
||||
entryPoints: [config.getRawConfig().traefik.https_entrypoint],
|
||||
service: siteResourceServiceName,
|
||||
middlewares: [siteResourceRewriteMiddlewareName],
|
||||
rule: `Host(\`${fullDomain}\`)`,
|
||||
priority: 100,
|
||||
tls
|
||||
};
|
||||
|
||||
// Assets bypass router - lets Next.js static files load without rewrite
|
||||
config_output.http.routers[`${siteResourceRouterName}-assets`] = {
|
||||
entryPoints: [config.getRawConfig().traefik.https_entrypoint],
|
||||
service: siteResourceServiceName,
|
||||
rule: `Host(\`${fullDomain}\`) && (PathPrefix(\`/_next\`) || PathRegexp(\`^/__nextjs*\`) || Path(\`/favicon.ico\`))`,
|
||||
priority: 101,
|
||||
tls
|
||||
};
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,21 @@
|
||||
import { Target } from "@server/db";
|
||||
|
||||
// Extended target type with site information, shared between the OSS and
|
||||
// private getTraefikConfig implementations.
|
||||
export type TargetWithSite = Target & {
|
||||
resourceId: number;
|
||||
targetId: number;
|
||||
ip: string | null;
|
||||
method: string | null;
|
||||
port: number | null;
|
||||
internalPort: number | null;
|
||||
enabled: boolean;
|
||||
health: string | null;
|
||||
site: {
|
||||
siteId: number;
|
||||
type: string;
|
||||
subnet: string | null;
|
||||
exitNodeId: number | null;
|
||||
online: boolean;
|
||||
};
|
||||
};
|
||||
@@ -2,6 +2,7 @@ import {
|
||||
db,
|
||||
Org,
|
||||
orgs,
|
||||
resourceAccessToken,
|
||||
resources,
|
||||
siteResources,
|
||||
sites,
|
||||
@@ -83,6 +84,15 @@ export async function removeUserFromOrg(
|
||||
.delete(userOrgs)
|
||||
.where(and(eq(userOrgs.userId, userId), eq(userOrgs.orgId, org.orgId)));
|
||||
|
||||
await trx
|
||||
.delete(resourceAccessToken)
|
||||
.where(
|
||||
and(
|
||||
eq(resourceAccessToken.userId, userId),
|
||||
eq(resourceAccessToken.orgId, org.orgId)
|
||||
)
|
||||
);
|
||||
|
||||
await trx.delete(userResources).where(
|
||||
and(
|
||||
eq(userResources.userId, userId),
|
||||
|
||||
@@ -1,10 +1,29 @@
|
||||
import {
|
||||
getResourceRuleValueValidationError,
|
||||
isValidDomain,
|
||||
isValidUrlGlobPattern
|
||||
} from "./validators";
|
||||
import { assertEquals } from "@test/assert";
|
||||
|
||||
function runTests() {
|
||||
console.log("Running domain validation tests...");
|
||||
|
||||
assertEquals(
|
||||
isValidDomain("example.com"),
|
||||
true,
|
||||
"Standard ASCII domain should be valid"
|
||||
);
|
||||
assertEquals(
|
||||
isValidDomain("xn--e1afmkfd.xn--p1ai"),
|
||||
true,
|
||||
"Punycode IDN domain should be valid"
|
||||
);
|
||||
assertEquals(
|
||||
isValidDomain("example.invalid-tld"),
|
||||
false,
|
||||
"Domain with unknown TLD should be invalid"
|
||||
);
|
||||
|
||||
console.log("Running URL pattern validation tests...");
|
||||
|
||||
// Test valid patterns
|
||||
|
||||
@@ -74,6 +74,7 @@ export const RESOURCE_RULE_MATCH_TYPES = [
|
||||
"IP",
|
||||
"PATH",
|
||||
"COUNTRY",
|
||||
"COUNTRY_IS_NOT",
|
||||
"ASN",
|
||||
"REGION"
|
||||
] as const;
|
||||
@@ -96,6 +97,7 @@ export function getResourceRuleValueValidationError(
|
||||
case "REGION":
|
||||
return isValidRegionId(value) ? null : "Invalid region ID provided";
|
||||
case "COUNTRY":
|
||||
case "COUNTRY_IS_NOT":
|
||||
return COUNTRIES.some((country) => country.code === value)
|
||||
? null
|
||||
: "Invalid country code provided";
|
||||
@@ -169,9 +171,10 @@ export function isValidDomain(domain: string): boolean {
|
||||
if (!/^[a-zA-Z0-9-]+$/.test(label)) return false;
|
||||
}
|
||||
|
||||
// TLD should be at least 2 characters and contain only letters
|
||||
// TLD should be at least 2 characters. Punycode TLDs can contain digits
|
||||
// and hyphens, so validity is ultimately enforced by the TLD allowlist.
|
||||
const tld = labels[labels.length - 1];
|
||||
if (tld.length < 2 || !/^[a-zA-Z]+$/.test(tld)) return false;
|
||||
if (tld.length < 2) return false;
|
||||
|
||||
// Check if TLD is in the list of valid TLDs
|
||||
if (!validTlds.includes(tld.toUpperCase())) return false;
|
||||
|
||||
@@ -0,0 +1,228 @@
|
||||
import {
|
||||
generateId,
|
||||
generateIdFromEntropySize
|
||||
} from "@server/auth/sessions/app";
|
||||
import {
|
||||
db,
|
||||
resources,
|
||||
User,
|
||||
virtualApiKeyResources,
|
||||
virtualApiKeys,
|
||||
type Transaction,
|
||||
type VirtualApiKey
|
||||
} from "@server/db";
|
||||
import config from "@server/lib/config";
|
||||
import { decrypt, encrypt } from "@server/lib/crypto";
|
||||
import { and, eq, inArray } from "drizzle-orm";
|
||||
|
||||
export {
|
||||
VIRTUAL_API_KEY_PREFIX,
|
||||
formatVirtualApiKeyCredential,
|
||||
formatVirtualApiKeyPreview,
|
||||
looksLikeVirtualApiKeyCredential,
|
||||
stripVirtualApiKeyAuthHeaders
|
||||
} from "@app/lib/virtualApiKeyFormat";
|
||||
|
||||
export type MintedVirtualApiKeySecret = {
|
||||
virtualApiKeyId: string;
|
||||
secret: string;
|
||||
lastChars: string;
|
||||
};
|
||||
|
||||
export type PublicVirtualApiKey = Omit<VirtualApiKey, "token"> & {
|
||||
secret?: string;
|
||||
};
|
||||
|
||||
export function mintVirtualApiKeySecret(): MintedVirtualApiKeySecret {
|
||||
const secret = generateIdFromEntropySize(16);
|
||||
return {
|
||||
virtualApiKeyId: generateId(8),
|
||||
secret,
|
||||
lastChars: secret.slice(-4)
|
||||
};
|
||||
}
|
||||
|
||||
export function encryptVirtualApiKeyToken(secret: string): string {
|
||||
return encrypt(secret, config.getRawConfig().server.secret!);
|
||||
}
|
||||
|
||||
export function decryptVirtualApiKeyToken(ciphertext: string): string {
|
||||
return decrypt(ciphertext, config.getRawConfig().server.secret!);
|
||||
}
|
||||
|
||||
export function toPublicVirtualApiKey(
|
||||
row: VirtualApiKey,
|
||||
options?: { includeSecret?: boolean }
|
||||
): PublicVirtualApiKey {
|
||||
const { token, ...rest } = row;
|
||||
if (!options?.includeSecret) {
|
||||
return rest;
|
||||
}
|
||||
return {
|
||||
...rest,
|
||||
secret: decryptVirtualApiKeyToken(token)
|
||||
};
|
||||
}
|
||||
|
||||
export async function assertManualKeyResourcesInOrg(params: {
|
||||
allResources: boolean;
|
||||
resourceIds: number[];
|
||||
orgId: string;
|
||||
}): Promise<{ ok: true } | { ok: false; message: string }> {
|
||||
const { allResources, resourceIds, orgId } = params;
|
||||
|
||||
if (allResources) {
|
||||
return { ok: true };
|
||||
}
|
||||
|
||||
if (resourceIds.length === 0) {
|
||||
return {
|
||||
ok: false,
|
||||
message:
|
||||
"Select at least one public inference resource, or enable all public inference resources"
|
||||
};
|
||||
}
|
||||
|
||||
const uniqueIds = [...new Set(resourceIds)];
|
||||
const rows = await db
|
||||
.select({ resourceId: resources.resourceId })
|
||||
.from(resources)
|
||||
.where(
|
||||
and(
|
||||
eq(resources.orgId, orgId),
|
||||
eq(resources.mode, "inference"),
|
||||
inArray(resources.resourceId, uniqueIds)
|
||||
)
|
||||
);
|
||||
|
||||
if (rows.length !== uniqueIds.length) {
|
||||
return {
|
||||
ok: false,
|
||||
message:
|
||||
"One or more resources are invalid public inference resources for this organization"
|
||||
};
|
||||
}
|
||||
|
||||
return { ok: true };
|
||||
}
|
||||
|
||||
export async function replaceVirtualApiKeyResources(
|
||||
trx: Transaction | typeof db,
|
||||
virtualApiKeyId: string,
|
||||
resourceIds: number[]
|
||||
): Promise<void> {
|
||||
await trx
|
||||
.delete(virtualApiKeyResources)
|
||||
.where(eq(virtualApiKeyResources.virtualApiKeyId, virtualApiKeyId));
|
||||
|
||||
const uniqueIds = [...new Set(resourceIds)];
|
||||
if (uniqueIds.length === 0) {
|
||||
return;
|
||||
}
|
||||
|
||||
await trx.insert(virtualApiKeyResources).values(
|
||||
uniqueIds.map((resourceId) => ({
|
||||
virtualApiKeyId,
|
||||
resourceId
|
||||
}))
|
||||
);
|
||||
}
|
||||
|
||||
async function selectUserVirtualApiKey(
|
||||
orgId: string,
|
||||
userId: string
|
||||
): Promise<VirtualApiKey | null> {
|
||||
const [existing] = await db
|
||||
.select()
|
||||
.from(virtualApiKeys)
|
||||
.where(
|
||||
and(
|
||||
eq(virtualApiKeys.orgId, orgId),
|
||||
eq(virtualApiKeys.userId, userId),
|
||||
eq(virtualApiKeys.kind, "user")
|
||||
)
|
||||
)
|
||||
.limit(1);
|
||||
|
||||
return existing ?? null;
|
||||
}
|
||||
|
||||
export async function getOrCreateUserVirtualApiKey(params: {
|
||||
orgId: string;
|
||||
user: User;
|
||||
createdByUserId?: string | null;
|
||||
}): Promise<{ key: VirtualApiKey; secret: string }> {
|
||||
const { orgId, user, createdByUserId } = params;
|
||||
|
||||
const existing = await selectUserVirtualApiKey(orgId, user.userId);
|
||||
if (existing) {
|
||||
return {
|
||||
key: existing,
|
||||
secret: decryptVirtualApiKeyToken(existing.token)
|
||||
};
|
||||
}
|
||||
|
||||
const minted = mintVirtualApiKeySecret();
|
||||
const now = Date.now();
|
||||
|
||||
try {
|
||||
const [created] = await db
|
||||
.insert(virtualApiKeys)
|
||||
.values({
|
||||
virtualApiKeyId: minted.virtualApiKeyId,
|
||||
orgId,
|
||||
kind: "user",
|
||||
userId: user.userId,
|
||||
name: `${user.name ?? user.username}'s API Key`,
|
||||
description: null,
|
||||
token: encryptVirtualApiKeyToken(minted.secret),
|
||||
lastChars: minted.lastChars,
|
||||
allResources: false,
|
||||
expiresAt: null,
|
||||
lastUsedAt: null,
|
||||
createdAt: now,
|
||||
createdByUserId: createdByUserId ?? null
|
||||
})
|
||||
.returning();
|
||||
|
||||
return { key: created, secret: minted.secret };
|
||||
} catch {
|
||||
const raced = await selectUserVirtualApiKey(orgId, user.userId);
|
||||
if (raced) {
|
||||
return {
|
||||
key: raced,
|
||||
secret: decryptVirtualApiKeyToken(raced.token)
|
||||
};
|
||||
}
|
||||
throw new Error("Failed to create user virtual API key");
|
||||
}
|
||||
}
|
||||
|
||||
export async function rotateUserVirtualApiKey(params: {
|
||||
orgId: string;
|
||||
user: User;
|
||||
createdByUserId?: string | null;
|
||||
}): Promise<{ key: VirtualApiKey; secret: string }> {
|
||||
const { orgId, user, createdByUserId } = params;
|
||||
const existing = await selectUserVirtualApiKey(orgId, user.userId);
|
||||
|
||||
if (!existing) {
|
||||
return getOrCreateUserVirtualApiKey(params);
|
||||
}
|
||||
|
||||
const minted = mintVirtualApiKeySecret();
|
||||
const [updated] = await db
|
||||
.update(virtualApiKeys)
|
||||
.set({
|
||||
token: encryptVirtualApiKeyToken(minted.secret),
|
||||
lastChars: minted.lastChars,
|
||||
createdByUserId:
|
||||
createdByUserId !== undefined
|
||||
? createdByUserId
|
||||
: existing.createdByUserId
|
||||
})
|
||||
.where(eq(virtualApiKeys.virtualApiKeyId, existing.virtualApiKeyId))
|
||||
.returning();
|
||||
|
||||
return { key: updated, secret: minted.secret };
|
||||
}
|
||||
Reference in New Issue
Block a user