fix: complete diagnostic extension remediation

This commit is contained in:
2026-08-04 00:46:44 +02:00
parent 565e93a456
commit 49fa7030a5
6 changed files with 192 additions and 19 deletions
+21 -11
View File
@@ -50,6 +50,7 @@ export interface ConnectorDiagnosticRequest {
user?: string;
credentialFile?: string;
tlsCaFile?: string;
tlsServername?: string;
resource: DiagnosticResource;
timeoutMs: number;
signal: AbortSignal;
@@ -87,6 +88,7 @@ export interface VectorDiagnosticRequest {
baseUrl?: string;
credentialFile?: string;
tlsCaFile?: string;
tlsServername?: string;
diagnostic?: RestDiagnosticRequest & {
response: { collection: string; dimensions: string; distance: string };
};
@@ -147,7 +149,7 @@ export interface DatabaseDiagnosticClient {
export interface DatabaseDiagnosticClientFactory {
connect(request: {
host: string; port: number; database: string; user: string; credentialFile: string; tlsCaFile?: string; signal: AbortSignal;
host: string; port: number; database: string; user: string; credentialFile: string; tlsCaFile?: string; tlsServername?: string; signal: AbortSignal;
}): Promise<DatabaseDiagnosticClient>;
}
@@ -298,13 +300,15 @@ export function createConcreteDiagnosticAdapters(
},
};
const databaseClient = dependencies.databaseClient ?? {
async connect(request: { host: string; port: number; database: string; user: string; credentialFile: string; tlsCaFile?: string; signal: AbortSignal }) {
async connect(request: { host: string; port: number; database: string; user: string; credentialFile: string; tlsCaFile?: string; tlsServername?: 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: request.tlsCaFile
? { ca: await readFile(request.tlsCaFile, "utf8"), rejectUnauthorized: true }
: { rejectUnauthorized: true },
ssl: {
...(request.tlsCaFile ? { ca: await readFile(request.tlsCaFile, "utf8") } : {}),
...(request.tlsServername ? { servername: request.tlsServername } : {}),
rejectUnauthorized: true,
},
connectionTimeoutMillis: 5_000,
});
const abort = () => { void client.end(); };
@@ -330,7 +334,8 @@ export function createConcreteDiagnosticAdapters(
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,
credentialFile: request.credentialFile, tlsCaFile: request.tlsCaFile,
tlsServername: request.tlsServername, signal: request.signal,
});
try {
const result = await client.query("SELECT current_database() AS database, current_schema() AS schema", []);
@@ -398,7 +403,8 @@ export function createConcreteDiagnosticAdapters(
}
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,
credentialFile: request.credentialFile, tlsCaFile: request.tlsCaFile,
tlsServername: request.tlsServername, signal: request.signal,
});
try {
const metadata = await client.query(
@@ -584,15 +590,14 @@ function connectorRequest(
collection: workspace.semantic_index.vector_store.collection,
};
const field = (suffix: string) => bindingName(workspace, contractRole, suffix);
const credentialFile = values[field(binding.transport === "rest_api" ? "API_KEY_FILE" : "PASSWORD_FILE")];
if (credentialFile === undefined) return undefined;
if (binding.transport === "rest_api") {
const baseUrl = values[field("BASE_URL")];
const diagnostic = role === "dwh"
? workspace.diagnostics?.dwh_rest
: workspace.diagnostics?.vector_rest?.metadata;
if (baseUrl === undefined || diagnostic === undefined) return undefined;
const credentialFile = diagnostic.auth === "none" ? undefined : values[field("API_KEY_FILE")];
if (diagnostic.auth !== "none" && credentialFile === undefined) return undefined;
return {
role,
transport: "rest_api",
@@ -621,6 +626,9 @@ function connectorRequest(
};
}
const credentialFile = values[field("PASSWORD_FILE")];
if (credentialFile === undefined) return undefined;
const host = values[field("HOST")];
const port = numericBinding(values, field("PORT"));
const user = values[field("USER")];
@@ -660,6 +668,7 @@ function tunnelProbeRequest(
user,
credentialFile: password,
tlsCaFile: binding.values[bindingName(workspace, contractRole, "TLS_CA_FILE")],
tlsServername: binding.values[bindingName(workspace, contractRole, "SSH_TARGET_HOST")],
resource: role === "dwh"
? { database: workspace.dwh.database, schema: workspace.dwh.schema }
: {
@@ -710,7 +719,8 @@ export function createWorkspaceDiagnoser(
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,
tlsCaFile: tunneledRequest.tlsCaFile, tlsServername: tunneledRequest.tlsServername,
resource: tunneledRequest.resource,
collection: canonical.semantic_index.vector_store.collection,
dimensions: canonical.semantic_index.vector_store.dimensions,
distance: canonical.semantic_index.vector_store.distance,
+4 -3
View File
@@ -45,11 +45,12 @@ function legacyDirectConnection(
function legacyRestEndpoint(
binding: ResolvedBinding,
names: { baseUrl: string; apiKeyFile: string; tlsCaFile: string },
requiresCredential: boolean,
): Record<string, unknown> {
const endpoint: Record<string, unknown> = {
base_url: requireBinding(binding, names.baseUrl),
api_key_file: requireBinding(binding, names.apiKeyFile),
};
if (requiresCredential) endpoint.api_key_file = requireBinding(binding, names.apiKeyFile);
const tlsCaFile = bindingValue(binding, names.tlsCaFile);
if (tlsCaFile !== undefined) endpoint.ssl_ca_file = tlsCaFile;
return endpoint;
@@ -124,7 +125,7 @@ export function renderRuntimeConfig(
const rest = legacyRestEndpoint(bindings.dwh, {
baseUrl: name("DWH", "BASE_URL"), apiKeyFile: name("DWH", "API_KEY_FILE"),
tlsCaFile: name("DWH", "TLS_CA_FILE"),
});
}, canonical.diagnostics?.dwh_rest?.auth !== "none");
rendered.rest = rest;
rendered.dwh = { type: "thoth_rest", database: dwhIdentity, endpoint: rest };
} else {
@@ -136,7 +137,7 @@ export function renderRuntimeConfig(
const vectorRest = legacyRestEndpoint(bindings.vector, {
baseUrl: name("VECTOR", "BASE_URL"), apiKeyFile: name("VECTOR", "API_KEY_FILE"),
tlsCaFile: name("VECTOR", "TLS_CA_FILE"),
});
}, canonical.diagnostics?.vector_rest?.metadata.auth !== "none");
rendered.vector_rest = vectorRest;
rendered.vectors = { type: "thoth_vector_http", reader: vectorRest };
} else {
+123 -1
View File
@@ -10,7 +10,7 @@ import {
type DiagnosticAdapters,
} from "../src/workspaces/diagnostics.js";
import { resolveRuntimeBindings } from "../src/workspaces/bindings.js";
import type { RuntimeBindings } from "../src/workspaces/runtime-renderer.js";
import { renderRuntimeConfig, type RuntimeBindings } from "../src/workspaces/runtime-renderer.js";
import { parseWorkspaceYaml } from "../src/workspaces/schema.js";
const workspace = parseWorkspaceYaml(`workspace:
@@ -268,6 +268,74 @@ test("checks direct and REST resolution, TLS, authentication, and resource metad
expect(JSON.stringify(result)).not.toContain("/run/secrets/dwh-api-key");
});
test("carries auth-none REST bindings from resolver through runtime rendering to diagnostics without a key", async () => {
const unauthenticatedWorkspace = parseWorkspaceYaml(`workspace:
schema_version: 2
id: psd-clinical
name: Policlinico San Donato
language: it
dwh:
engine: postgres
database: warehouse
schema: datawarehouse
supported_transports: [rest_api]
semantic_index:
vector_store:
engine: pgvector
database: postgres
schema: vectors
collection: clinical_documents
dimensions: 768
distance: cosine
supported_transports: [rest_api]
embedding:
provider: ollama_compatible
model: nomic-embed-text-v2-moe
dimensions: 768
diagnostics:
dwh_rest:
method: POST
path: /rpc/ping
auth: none
response: { database: database, schema: schema }
vector_rest:
metadata:
method: GET
path: /vector/metadata
auth: none
response: { collection: collection, dimensions: dimensions, distance: distance }
embedding:
method: GET
path: /models
auth: none
response: { model: model, dimensions: dimensions }
llm_policy:
allowed: [zai/glm-5.2]
`);
const resolved = resolveRuntimeBindings(unauthenticatedWorkspace, {
THT_WS_PSD_CLINICAL_DWH_TRANSPORT: "rest_api",
THT_WS_PSD_CLINICAL_DWH_BASE_URL: "https://dwh.example.test",
THT_WS_PSD_CLINICAL_VECTOR_TRANSPORT: "rest_api",
THT_WS_PSD_CLINICAL_VECTOR_BASE_URL: "https://vector.example.test",
THT_WS_PSD_CLINICAL_EMBEDDING_BASE_URL: "https://embedding.example.test",
}, ["/run/secrets"]);
const adapters = successfulAdapters();
const runtime = renderRuntimeConfig(unauthenticatedWorkspace, resolved, {
sessions: "/data/sessions", artifacts: "/data/artifacts", indexes: "/data/indexes",
});
const result = await diagnose(adapters)(unauthenticatedWorkspace, resolved, { writeProbe: false });
expect(runtime).not.toContain("api_key_file");
expect(result.activatable).toBe(true);
expect(adapters.probeConnector).toHaveBeenCalledWith(expect.objectContaining({
role: "dwh", credentialFile: undefined,
}));
expect(adapters.probeConnector).toHaveBeenCalledWith(expect.objectContaining({
role: "vector", credentialFile: undefined,
}));
});
test("uses a loopback-only SSH tunnel for the bounded connector probe", async () => {
const adapters = successfulAdapters();
const sshBindings: RuntimeBindings = {
@@ -305,6 +373,35 @@ test("uses a loopback-only SSH tunnel for the bounded connector probe", async ()
}));
});
test("retains the SSH target hostname for forwarded PostgreSQL TLS validation", async () => {
const adapters = successfulAdapters();
const sshBindings: RuntimeBindings = {
...bindings,
dwh: {
transport: "ssh_tunnel",
missing: [],
values: {
THT_WS_PSD_CLINICAL_DWH_USER: "reader",
THT_WS_PSD_CLINICAL_DWH_PASSWORD_FILE: "/run/secrets/dwh-password",
THT_WS_PSD_CLINICAL_DWH_SSH_HOST: "bastion.example.test",
THT_WS_PSD_CLINICAL_DWH_SSH_PORT: "22",
THT_WS_PSD_CLINICAL_DWH_SSH_USER: "tunnel",
THT_WS_PSD_CLINICAL_DWH_SSH_PRIVATE_KEY_FILE: "/run/secrets/ssh-key",
THT_WS_PSD_CLINICAL_DWH_SSH_KNOWN_HOSTS_FILE: "/run/secrets/known-hosts",
THT_WS_PSD_CLINICAL_DWH_SSH_TARGET_HOST: "dwh.internal",
THT_WS_PSD_CLINICAL_DWH_SSH_TARGET_PORT: "5432",
},
},
};
await diagnose(adapters)(workspace, sshBindings, { writeProbe: false });
expect(adapters.probeConnector).toHaveBeenCalledWith(expect.objectContaining({
host: "127.0.0.1",
tlsServername: "dwh.internal",
}));
});
test("passes the declared vector database and schema to direct diagnostics", async () => {
const adapters = successfulAdapters();
@@ -622,6 +719,31 @@ test("uses system trust for direct and SSH PostgreSQL diagnostics when no CA bin
}
});
test("passes the original target hostname to the PostgreSQL TLS client", async () => {
const directory = await mkdtemp(join(tmpdir(), "thothii-diagnostic-"));
const passwordFile = join(directory, "password");
await writeFile(passwordFile, "password\n", { mode: 0o600 });
const connect = vi.fn(async () => ({
query: vi.fn(async () => ({ rows: [{ database: "warehouse", schema: "datawarehouse" }] })),
end: vi.fn(async () => undefined),
}));
try {
await createConcreteDiagnosticAdapters({ databaseClient: { connect } } as any).probeConnector({
role: "dwh", transport: "ssh_tunnel", host: "127.0.0.1", port: 5432, user: "reader",
credentialFile: passwordFile, tlsServername: "dwh.internal",
resource: { database: "warehouse", schema: "datawarehouse" }, timeoutMs: 5000,
signal: new AbortController().signal,
});
expect(connect).toHaveBeenCalledWith(expect.objectContaining({
host: "127.0.0.1",
tlsServername: "dwh.internal",
}));
} finally {
await rm(directory, { recursive: true, force: true });
}
});
test("selects the vector index containing the declared vector column for direct metadata", async () => {
const directory = await mkdtemp(join(tmpdir(), "thothii-diagnostic-"));
const passwordFile = join(directory, "password");