fix: harden workspace diagnostic protocols
This commit is contained in:
@@ -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;
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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 {
|
||||
|
||||
Reference in New Issue
Block a user