fix: harden workspace diagnostic protocols

This commit is contained in:
2026-08-03 23:59:46 +02:00
parent 9e2eafb66c
commit ca97bbb9c2
5 changed files with 371 additions and 17 deletions
+136 -14
View File
@@ -1,7 +1,8 @@
import { randomUUID } from "node:crypto";
import { readFile } from "node:fs/promises";
import { readFile, realpath } from "node:fs/promises";
import { createConnection } from "node:net";
import { once } from "node:events";
import { Client } from "pg";
import { MAX_WORKSPACE_DIAGNOSTIC_TIMEOUT_MS } from "../config.js";
import { buildInstallationContract } from "./contracts.js";
import type { RuntimeBindings } from "./runtime-renderer.js";
@@ -90,6 +91,10 @@ export interface VectorDiagnosticRequest {
collection: string;
dimensions?: number;
distance?: "cosine" | "l2" | "inner_product";
host?: string;
port?: number;
user?: string;
resource?: DiagnosticResource;
timeoutMs: number;
signal: AbortSignal;
}
@@ -130,6 +135,17 @@ export interface DirectProtocolFactory {
probe(request: ConnectorDiagnosticRequest): Promise<ConnectorDiagnosticResult>;
}
export interface DatabaseDiagnosticClient {
query(sql: string, values: readonly unknown[]): Promise<{ rows: Array<Record<string, unknown>> }>;
end(): Promise<void>;
}
export interface DatabaseDiagnosticClientFactory {
connect(request: {
host: string; port: number; database: string; user: string; credentialFile: string; tlsCaFile: string; signal: AbortSignal;
}): Promise<DatabaseDiagnosticClient>;
}
export interface SshProcessFactory {
start(request: SshTunnelRequest, args: readonly string[]): Promise<{
tunnel: LoopbackTunnel;
@@ -140,6 +156,7 @@ export interface SshProcessFactory {
export interface ConcreteDiagnosticAdapterDependencies {
directProtocol?: DirectProtocolFactory;
sshProcess?: SshProcessFactory;
databaseClient?: DatabaseDiagnosticClientFactory;
}
/**
@@ -176,6 +193,24 @@ async function secretPresent(file: string): Promise<boolean> {
return (await readFile(file, "utf8")).trim().length > 0;
}
async function sameSecretFile(first: string, second: string): Promise<boolean> {
try {
return await realpath(first) === await realpath(second);
} catch {
return first === second;
}
}
async function restHeaders(
diagnostic: RestDiagnosticRequest,
credentialFile: string | undefined,
): Promise<Record<string, string>> {
if (diagnostic.auth === "none") return {};
if (!credentialFile || !(await secretPresent(credentialFile))) throw new Error("REST probe failed");
const secret = (await readFile(credentialFile, "utf8")).trim();
return diagnostic.auth === "bearer" ? { authorization: `Bearer ${secret}` } : { "x-api-key": secret };
}
/**
* Concrete production adapters deliberately retain only probe metadata. Protocol failures and
* response bodies are discarded at this boundary; callers receive fixed diagnostics instead.
@@ -183,23 +218,57 @@ async function secretPresent(file: string): Promise<boolean> {
export function createConcreteDiagnosticAdapters(
dependencies: ConcreteDiagnosticAdapterDependencies = {},
): DiagnosticAdapters {
const databaseClient = dependencies.databaseClient ?? {
async connect(request: { host: string; port: number; database: string; user: string; credentialFile: string; tlsCaFile: string; signal: AbortSignal }) {
const client = new Client({
host: request.host, port: request.port, database: request.database, user: request.user,
password: (await readFile(request.credentialFile, "utf8")).trim(),
ssl: { ca: await readFile(request.tlsCaFile, "utf8"), rejectUnauthorized: true },
connectionTimeoutMillis: 5_000,
});
const abort = () => { void client.end(); };
request.signal.addEventListener("abort", abort, { once: true });
try {
await client.connect();
return { query: async (sql: string, values: readonly unknown[]) => await client.query(sql, [...values]), end: async () => { request.signal.removeEventListener("abort", abort); await client.end(); } };
} catch (error) {
request.signal.removeEventListener("abort", abort);
await client.end().catch(() => undefined);
throw error;
}
},
};
const directProtocol = dependencies.directProtocol ?? {
async probe(request: ConnectorDiagnosticRequest): Promise<ConnectorDiagnosticResult> {
if (!request.host || !request.port || !request.credentialFile || !(await secretPresent(request.credentialFile))) {
if (!request.host || !request.port || !request.user || !request.credentialFile || !request.tlsCaFile
|| !(await secretPresent(request.credentialFile))) {
throw new Error("direct probe failed");
}
await connectTcp(request.host, request.port, request.signal);
return { resolved: true, tlsVerified: request.tlsCaFile !== undefined, authenticated: true, resource: request.resource };
const database = request.resource.database;
const schema = request.resource.schema;
if (!database || !schema) throw new Error("direct probe failed");
const client = await databaseClient.connect({
host: request.host, port: request.port, database, user: request.user,
credentialFile: request.credentialFile, tlsCaFile: request.tlsCaFile, signal: request.signal,
});
try {
const result = await client.query("SELECT current_database() AS database, current_schema() AS schema", []);
const row = result.rows[0];
if (row?.database !== database || row.schema !== schema) throw new Error("direct probe failed");
return { resolved: true, tlsVerified: true, authenticated: true, resource: request.resource };
} finally {
await client.end().catch(() => undefined);
}
},
};
return {
async probeConnector(request) {
if (request.transport === "rest_api") {
if (!request.baseUrl || !request.diagnostic || !request.credentialFile || !(await secretPresent(request.credentialFile))) throw new Error("REST probe failed");
if (!request.baseUrl || !request.diagnostic || request.tlsCaFile) throw new Error("REST probe failed");
const endpoint = resolveDiagnosticUrl(request.baseUrl, request.diagnostic.path);
const response = await fetch(endpoint.toString(), {
method: request.diagnostic.method,
headers: { authorization: `Bearer ${(await readFile(request.credentialFile, "utf8")).trim()}` },
headers: await restHeaders(request.diagnostic, request.credentialFile),
signal: request.signal,
redirect: "error",
});
@@ -241,6 +310,28 @@ export function createConcreteDiagnosticAdapters(
}
},
async inspectVector(request) {
if (request.transport === "pgvector_direct" || request.transport === "ssh_tunnel") {
const resource = request.resource;
if (!request.host || !request.port || !request.user || !request.credentialFile || !request.tlsCaFile
|| !resource?.database || !resource.schema || !(await secretPresent(request.credentialFile))) {
throw new Error("vector metadata adapter is unavailable");
}
const client = await databaseClient.connect({
host: request.host, port: request.port, database: resource.database, user: request.user,
credentialFile: request.credentialFile, tlsCaFile: request.tlsCaFile, signal: request.signal,
});
try {
const metadata = await client.query(
"SELECT a.atttypmod - 4 AS dimensions, CASE WHEN pg_get_indexdef(i.indexrelid) LIKE '%vector_cosine_ops%' THEN 'cosine' WHEN pg_get_indexdef(i.indexrelid) LIKE '%vector_l2_ops%' THEN 'l2' WHEN pg_get_indexdef(i.indexrelid) LIKE '%vector_ip_ops%' THEN 'inner_product' END AS distance FROM pg_attribute a JOIN pg_class c ON c.oid = a.attrelid JOIN pg_namespace n ON n.oid = c.relnamespace LEFT JOIN pg_index i ON i.indrelid = c.oid WHERE n.nspname = $1 AND c.relname = $2 AND a.attnum > 0 AND NOT a.attisdropped AND a.atttypid = (SELECT oid FROM pg_type WHERE typname = 'vector') LIMIT 1",
[resource.schema, request.collection],
);
const row = metadata.rows[0];
if (!row || !Number.isInteger(row.dimensions) || (row.distance !== "cosine" && row.distance !== "l2" && row.distance !== "inner_product")) throw new Error("vector metadata adapter is unavailable");
return { collection: request.collection, dimensions: row.dimensions as number, distance: row.distance as VectorDiagnosticResult["distance"] };
} finally {
await client.end().catch(() => undefined);
}
}
if (request.transport !== "rest_api" || !request.baseUrl || !request.credentialFile || !request.diagnostic
|| !(await secretPresent(request.credentialFile))) throw new Error("vector metadata adapter is unavailable");
const response = await fetch(resolveDiagnosticUrl(request.baseUrl, request.diagnostic.path).toString(), {
@@ -514,6 +605,7 @@ export function createWorkspaceDiagnoser(
const dwhTimeout = boundedTimeout(canonical.dwh.timeout_ms, fallbackTimeout);
const vectorTimeout = boundedTimeout(canonical.semantic_index.vector_store.timeout_ms, fallbackTimeout);
const embeddingTimeout = boundedTimeout(canonical.semantic_index.embedding.timeout_ms, fallbackTimeout);
let tunneledVectorMetadata: VectorDiagnosticResult | undefined;
for (const role of ["dwh", "vector"] as const) {
const timeoutMs = role === "dwh" ? dwhTimeout : vectorTimeout;
@@ -526,9 +618,22 @@ export function createWorkspaceDiagnoser(
const result = "sshHost" in request
? await withTimeout(timeoutMs, (signal) => adapters.withSshTunnel(
{ ...request, signal },
(tunnel) => adapters.probeConnector(tunnelProbeRequest(
canonical, role, bindings, timeoutMs, tunnel, signal,
)),
async (tunnel) => {
const tunneledRequest = tunnelProbeRequest(canonical, role, bindings, timeoutMs, tunnel, signal);
const connector = await adapters.probeConnector(tunneledRequest);
if (role === "vector") {
tunneledVectorMetadata = await adapters.inspectVector({
transport: "ssh_tunnel", host: tunneledRequest.host, port: tunneledRequest.port,
user: tunneledRequest.user, credentialFile: tunneledRequest.credentialFile,
tlsCaFile: tunneledRequest.tlsCaFile, resource: tunneledRequest.resource,
collection: canonical.semantic_index.vector_store.collection,
dimensions: canonical.semantic_index.vector_store.dimensions,
distance: canonical.semantic_index.vector_store.distance,
timeoutMs: vectorTimeout, signal,
});
}
return connector;
},
))
: await withTimeout(timeoutMs, (signal) => adapters.probeConnector({ ...request, signal }));
const resource = role === "dwh"
@@ -549,7 +654,9 @@ export function createWorkspaceDiagnoser(
try {
const vectorBinding = bindings.vector;
const vectorRest = canonical.diagnostics?.vector_rest?.metadata;
const vector = await withTimeout(vectorTimeout, (signal) => adapters.inspectVector({
const vectorDirect = connectorRequest(canonical, "vector", bindings, vectorTimeout);
const directVectorRequest = vectorDirect && !("sshHost" in vectorDirect) ? vectorDirect : undefined;
const vector = tunneledVectorMetadata ?? await withTimeout(vectorTimeout, (signal) => adapters.inspectVector({
transport: vectorBinding.transport === "pgvector_direct" || vectorBinding.transport === "rest_api"
|| vectorBinding.transport === "ssh_tunnel" ? vectorBinding.transport : undefined,
baseUrl: vectorBinding.values[bindingName(canonical, "VECTOR", "BASE_URL")],
@@ -559,6 +666,14 @@ export function createWorkspaceDiagnoser(
collection: canonical.semantic_index.vector_store.collection,
dimensions: canonical.semantic_index.vector_store.dimensions,
distance: canonical.semantic_index.vector_store.distance,
...(directVectorRequest ? {
host: directVectorRequest.host,
port: directVectorRequest.port,
user: directVectorRequest.user,
credentialFile: directVectorRequest.credentialFile,
tlsCaFile: directVectorRequest.tlsCaFile,
resource: directVectorRequest.resource,
} : {}),
timeoutMs: vectorTimeout,
signal,
}));
@@ -601,6 +716,11 @@ export function createWorkspaceDiagnoser(
) {
const credentialFile = bindings.vector.values[bindingName(canonical, "VECTOR_WRITER", "API_KEY_FILE")];
if (!credentialFile) return { activatable: true, diagnostics };
const readerCredentialFile = bindings.vector.values[bindingName(canonical, "VECTOR", "API_KEY_FILE")];
if (readerCredentialFile && await sameSecretFile(credentialFile, readerCredentialFile)) {
diagnostics.push(diagnosticError("binding_missing", bindingName(canonical, "VECTOR_WRITER", "API_KEY_FILE")));
return { activatable: false, diagnostics };
}
const request: WriteDiagnosticRecordRequest = {
collection: canonical.semantic_index.vector_store.collection,
id: `diagnostic:${randomUUID()}`,
@@ -611,20 +731,22 @@ export function createWorkspaceDiagnoser(
baseUrl: bindings.vector.values[bindingName(canonical, "VECTOR", "BASE_URL")],
diagnostic: canonical.diagnostics.vector_rest.reversible_probe,
};
let writeSucceeded = false;
let writeStarted = false;
let cleanupAttempted = false;
let cleanupFailed = false;
try {
writeStarted = true;
await withTimeout(vectorTimeout, (signal) => adapters.writeDiagnosticRecord({ ...request, signal }));
writeSucceeded = true;
cleanupAttempted = true;
await withTimeout(vectorTimeout, (signal) => adapters.removeDiagnosticRecord({ ...request, signal }));
} catch {
cleanupFailed = true;
} finally {
if (writeSucceeded && cleanupFailed) {
if (writeStarted && (!cleanupAttempted || cleanupFailed)) {
try {
await withTimeout(vectorTimeout, (signal) => adapters.removeDiagnosticRecord({ ...request, signal }));
} catch {
// The cleanup attempt is deliberately best-effort and remains redacted.
cleanupFailed = true;
}
}
}
+1 -1
View File
@@ -10,7 +10,7 @@ export type VectorTransport = (typeof VECTOR_TRANSPORTS)[number];
export const REST_DIAGNOSTIC_METHODS = ["GET", "POST"] as const;
export type RestDiagnosticMethod = (typeof REST_DIAGNOSTIC_METHODS)[number];
export const DIAGNOSTIC_AUTH_MODES = ["none", "bearer"] as const;
export const DIAGNOSTIC_AUTH_MODES = ["none", "bearer", "x-api-key"] as const;
export type DiagnosticAuthMode = (typeof DIAGNOSTIC_AUTH_MODES)[number];
export interface RestDiagnosticRequest {