diff --git a/backend/src/workspaces/diagnostics.ts b/backend/src/workspaces/diagnostics.ts index d25dfa02..6c652696 100644 --- a/backend/src/workspaces/diagnostics.ts +++ b/backend/src/workspaces/diagnostics.ts @@ -5,7 +5,13 @@ import { once } from "node:events"; import { MAX_WORKSPACE_DIAGNOSTIC_TIMEOUT_MS } from "../config.js"; import { buildInstallationContract } from "./contracts.js"; import type { RuntimeBindings } from "./runtime-renderer.js"; -import { validateCanonicalWorkspace, type CanonicalWorkspace, type WorkspaceDescriptor } from "./schema.js"; +import { + resolveDiagnosticUrl, + validateCanonicalWorkspace, + type CanonicalWorkspace, + type RestDiagnosticRequest, + type WorkspaceDescriptor, +} from "./schema.js"; import type { WorkspaceErrorCode } from "./types.js"; export interface Diagnostic { @@ -28,6 +34,10 @@ interface DiagnosticResource { collection?: string; } +type RestConnectorDiagnostic = RestDiagnosticRequest & { + response?: Record; +}; + export interface ConnectorDiagnosticRequest { role: ConnectorRole; transport: "postgres_direct" | "pgvector_direct" | "rest_api" | "ssh_tunnel"; @@ -35,11 +45,12 @@ export interface ConnectorDiagnosticRequest { port?: number; baseUrl?: string; user?: string; - credentialFile: string; + credentialFile?: string; tlsCaFile?: string; resource: DiagnosticResource; timeoutMs: number; signal: AbortSignal; + diagnostic?: RestConnectorDiagnostic; } export interface ConnectorDiagnosticResult { @@ -69,7 +80,16 @@ export interface LoopbackTunnel { } export interface VectorDiagnosticRequest { + transport?: "pgvector_direct" | "rest_api" | "ssh_tunnel"; + baseUrl?: string; + credentialFile?: string; + tlsCaFile?: string; + diagnostic?: RestDiagnosticRequest & { + response: { collection: string; dimensions: string; distance: string }; + }; collection: string; + dimensions?: number; + distance?: "cosine" | "l2" | "inner_product"; timeoutMs: number; signal: AbortSignal; } @@ -87,6 +107,7 @@ export interface EmbeddingDiagnosticRequest { model: string; timeoutMs: number; signal: AbortSignal; + diagnostic?: RestDiagnosticRequest & { response: { model: string; dimensions: string } }; } export interface EmbeddingDiagnosticResult { @@ -100,6 +121,25 @@ export interface WriteDiagnosticRecordRequest { dimensions: number; timeoutMs: number; signal: AbortSignal; + credentialFile?: string; + baseUrl?: string; + diagnostic?: RestDiagnosticRequest; +} + +export interface DirectProtocolFactory { + probe(request: ConnectorDiagnosticRequest): Promise; +} + +export interface SshProcessFactory { + start(request: SshTunnelRequest, args: readonly string[]): Promise<{ + tunnel: LoopbackTunnel; + close(): Promise; + }>; +} + +export interface ConcreteDiagnosticAdapterDependencies { + directProtocol?: DirectProtocolFactory; + sshProcess?: SshProcessFactory; } /** @@ -140,18 +180,36 @@ async function secretPresent(file: string): Promise { * Concrete production adapters deliberately retain only probe metadata. Protocol failures and * response bodies are discarded at this boundary; callers receive fixed diagnostics instead. */ -export function createConcreteDiagnosticAdapters(): DiagnosticAdapters { +export function createConcreteDiagnosticAdapters( + dependencies: ConcreteDiagnosticAdapterDependencies = {}, +): DiagnosticAdapters { + const directProtocol = dependencies.directProtocol ?? { + async probe(request: ConnectorDiagnosticRequest): Promise { + if (!request.host || !request.port || !request.credentialFile || !(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 }; + }, + }; return { async probeConnector(request) { if (request.transport === "rest_api") { - if (!request.baseUrl || !(await secretPresent(request.credentialFile))) throw new Error("REST probe failed"); - const response = await fetch(request.baseUrl, { - method: "HEAD", - headers: { authorization: `Bearer ${await readFile(request.credentialFile, "utf8")}` }, + if (!request.baseUrl || !request.diagnostic || !request.credentialFile || !(await secretPresent(request.credentialFile))) 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()}` }, signal: request.signal, redirect: "error", }); if (!response.ok) throw new Error("REST probe failed"); + if (request.diagnostic && "response" in request.diagnostic) { + const payload = await response.json().catch(() => undefined) as Record | undefined; + const declared = request.diagnostic.response as { database?: string; schema?: string }; + if (!payload || (declared.database && payload[declared.database] !== request.resource.database) + || (declared.schema && payload[declared.schema] !== request.resource.schema)) throw new Error("REST probe failed"); + } return { resolved: true, tlsVerified: new URL(request.baseUrl).protocol === "https:", @@ -159,18 +217,7 @@ export function createConcreteDiagnosticAdapters(): DiagnosticAdapters { resource: request.resource, }; } - if (!request.host || !request.port || !(await secretPresent(request.credentialFile))) { - throw new Error("direct probe failed"); - } - await connectTcp(request.host, request.port, request.signal); - return { - resolved: true, - // Direct TLS verification requires an explicit CA file. A plain TCP success alone is - // intentionally insufficient for activation. - tlsVerified: request.tlsCaFile !== undefined, - authenticated: true, - resource: request.resource, - }; + return await directProtocol.probe(request); }, async withSshTunnel(request, probe) { // The image supplies OpenSSH for the registry's SSH implementation. This adapter refuses @@ -179,27 +226,83 @@ export function createConcreteDiagnosticAdapters(): DiagnosticAdapters { if (!request.knownHostsFile || !(await secretPresent(request.privateKeyFile))) { throw new Error("SSH probe failed"); } - throw new Error("SSH tunnel process is unavailable"); + if (!dependencies.sshProcess) throw new Error("SSH tunnel process is unavailable"); + const args = [ + "-N", "-o", "BatchMode=yes", "-o", "StrictHostKeyChecking=yes", + "-o", `UserKnownHostsFile=${request.knownHostsFile}`, "-i", request.privateKeyFile, + "-p", String(request.sshPort), "-L", `127.0.0.1:0:${request.targetHost}:${request.targetPort}`, + `${request.sshUser}@${request.sshHost}`, + ]; + const tunnel = await dependencies.sshProcess.start(request, args); + try { + return await probe(tunnel.tunnel); + } finally { + await tunnel.close().catch(() => undefined); + } }, - async inspectVector() { - throw new Error("vector metadata adapter is unavailable"); + async inspectVector(request) { + 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(), { + method: request.diagnostic.method, + headers: { authorization: `Bearer ${(await readFile(request.credentialFile, "utf8")).trim()}` }, + signal: request.signal, + redirect: "error", + }); + const payload = await response.json().catch(() => undefined) as Record | undefined; + const fields = request.diagnostic.response; + if (!response.ok || !payload || typeof payload[fields.collection] !== "string" + || !Number.isInteger(payload[fields.dimensions]) || typeof payload[fields.distance] !== "string") { + throw new Error("vector metadata adapter is unavailable"); + } + return { + collection: payload[fields.collection] as string, + dimensions: payload[fields.dimensions] as number, + distance: payload[fields.distance] as VectorDiagnosticResult["distance"], + }; }, async probeEmbedding(request) { - if (!(await secretPresent(request.credentialFile ?? ""))) throw new Error("embedding probe failed"); - const response = await fetch(request.baseUrl, { - method: "HEAD", + if (!request.diagnostic || !(await secretPresent(request.credentialFile ?? ""))) throw new Error("embedding probe failed"); + const response = await fetch(resolveDiagnosticUrl(request.baseUrl, request.diagnostic.path).toString(), { + method: request.diagnostic.method, headers: { authorization: `Bearer ${await readFile(request.credentialFile!, "utf8")}` }, signal: request.signal, redirect: "error", }); - if (!response.ok) throw new Error("embedding probe failed"); - return { available: true, dimensions: undefined }; + const payload = await response.json().catch(() => undefined) as Record | undefined; + if (!response.ok || !payload || payload[request.diagnostic.response.model] !== request.model + || !Number.isInteger(payload[request.diagnostic.response.dimensions])) throw new Error("embedding probe failed"); + return { available: true, dimensions: payload[request.diagnostic.response.dimensions] as number }; }, - async writeDiagnosticRecord() { - throw new Error("vector write adapter is unavailable"); + async writeDiagnosticRecord(request) { + if (!request.baseUrl || !request.credentialFile || !request.diagnostic + || !(await secretPresent(request.credentialFile))) throw new Error("vector write adapter is unavailable"); + const response = await fetch(resolveDiagnosticUrl(request.baseUrl, request.diagnostic.path).toString(), { + method: request.diagnostic.method, + headers: { + authorization: `Bearer ${(await readFile(request.credentialFile, "utf8")).trim()}`, + "content-type": "application/json", + }, + body: JSON.stringify({ operation: "create", id: request.id, collection: request.collection, dimensions: request.dimensions }), + signal: request.signal, + redirect: "error", + }); + if (!response.ok) throw new Error("vector write adapter is unavailable"); }, - async removeDiagnosticRecord() { - throw new Error("vector write adapter is unavailable"); + async removeDiagnosticRecord(request) { + if (!request.baseUrl || !request.credentialFile || !request.diagnostic + || !(await secretPresent(request.credentialFile))) throw new Error("vector write adapter is unavailable"); + const response = await fetch(resolveDiagnosticUrl(request.baseUrl, request.diagnostic.path).toString(), { + method: request.diagnostic.method, + headers: { + authorization: `Bearer ${(await readFile(request.credentialFile, "utf8")).trim()}`, + "content-type": "application/json", + }, + body: JSON.stringify({ operation: "remove", id: request.id, collection: request.collection }), + signal: request.signal, + redirect: "error", + }); + if (!response.ok) throw new Error("vector write adapter is unavailable"); }, }; } @@ -256,7 +359,7 @@ function diagnosticError(code: WorkspaceErrorCode, field?: string): Diagnostic { function bindingName( workspace: CanonicalWorkspace, - role: "DWH" | "VECTOR" | "EMBEDDING", + role: "DWH" | "VECTOR" | "VECTOR_WRITER" | "EMBEDDING", suffix: string, ): string { const entry = buildInstallationContract(workspace).variables.find((variable) => ( @@ -299,14 +402,21 @@ function connectorRequest( const values = binding.values; const resource: DiagnosticResource = role === "dwh" ? { database: workspace.dwh.database, schema: workspace.dwh.schema } - : { collection: workspace.semantic_index.vector_store.collection }; + : { + database: workspace.semantic_index.vector_store.database, + schema: workspace.semantic_index.vector_store.schema, + 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")]; - if (baseUrl === undefined) return undefined; + const diagnostic = role === "dwh" + ? workspace.diagnostics?.dwh_rest + : workspace.diagnostics?.vector_rest?.metadata; + if (baseUrl === undefined || diagnostic === undefined) return undefined; return { role, transport: "rest_api", @@ -316,6 +426,7 @@ function connectorRequest( resource, timeoutMs, signal: new AbortController().signal, + diagnostic, }; } @@ -375,7 +486,11 @@ function tunnelProbeRequest( tlsCaFile: binding.values[bindingName(workspace, contractRole, "TLS_CA_FILE")], resource: role === "dwh" ? { database: workspace.dwh.database, schema: workspace.dwh.schema } - : { collection: workspace.semantic_index.vector_store.collection }, + : { + database: workspace.semantic_index.vector_store.database, + schema: workspace.semantic_index.vector_store.schema, + collection: workspace.semantic_index.vector_store.collection, + }, timeoutMs, signal, }; @@ -418,7 +533,11 @@ export function createWorkspaceDiagnoser( : await withTimeout(timeoutMs, (signal) => adapters.probeConnector({ ...request, signal })); const resource = role === "dwh" ? { database: canonical.dwh.database, schema: canonical.dwh.schema } - : { collection: canonical.semantic_index.vector_store.collection }; + : { + database: canonical.semantic_index.vector_store.database, + schema: canonical.semantic_index.vector_store.schema, + collection: canonical.semantic_index.vector_store.collection, + }; if (!hasRequiredConnectorChecks(result, resource)) diagnostics.push(diagnosticError("connector_unavailable")); else diagnostics.push({ level: "info", code: "binding_ok", message: `${role === "dwh" ? "DWH" : "Vector"} binding diagnostic passed.` }); } catch { @@ -428,8 +547,18 @@ export function createWorkspaceDiagnoser( if (!diagnostics.some((diagnostic) => diagnostic.level === "error")) { try { + const vectorBinding = bindings.vector; + const vectorRest = canonical.diagnostics?.vector_rest?.metadata; const vector = 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")], + credentialFile: vectorBinding.values[bindingName(canonical, "VECTOR", "API_KEY_FILE")], + tlsCaFile: vectorBinding.values[bindingName(canonical, "VECTOR", "TLS_CA_FILE")], + diagnostic: vectorRest, collection: canonical.semantic_index.vector_store.collection, + dimensions: canonical.semantic_index.vector_store.dimensions, + distance: canonical.semantic_index.vector_store.distance, timeoutMs: vectorTimeout, signal, })); @@ -453,6 +582,7 @@ export function createWorkspaceDiagnoser( model: canonical.semantic_index.embedding.model, timeoutMs: embeddingTimeout, signal, + diagnostic: canonical.diagnostics?.embedding, })); if (!embedding.available || embedding.dimensions !== canonical.semantic_index.embedding.dimensions) { diagnostics.push(diagnosticError("semantic_index_incompatible")); @@ -462,13 +592,24 @@ export function createWorkspaceDiagnoser( } } - if (options.writeProbe && !diagnostics.some((diagnostic) => diagnostic.level === "error")) { + if ( + options.writeProbe + && canonical.semantic_index.vector_writer + && canonical.diagnostics?.vector_rest?.reversible_probe + && bindings.vector.transport === "rest_api" + && !diagnostics.some((diagnostic) => diagnostic.level === "error") + ) { + const credentialFile = bindings.vector.values[bindingName(canonical, "VECTOR_WRITER", "API_KEY_FILE")]; + if (!credentialFile) return { activatable: true, diagnostics }; const request: WriteDiagnosticRecordRequest = { collection: canonical.semantic_index.vector_store.collection, id: `diagnostic:${randomUUID()}`, dimensions: canonical.semantic_index.vector_store.dimensions, timeoutMs: vectorTimeout, signal: new AbortController().signal, + credentialFile, + baseUrl: bindings.vector.values[bindingName(canonical, "VECTOR", "BASE_URL")], + diagnostic: canonical.diagnostics.vector_rest.reversible_probe, }; let writeSucceeded = false; let cleanupFailed = false; diff --git a/backend/test/workspaces-diagnostics.test.ts b/backend/test/workspaces-diagnostics.test.ts index db2e3058..7ca1c5dc 100644 --- a/backend/test/workspaces-diagnostics.test.ts +++ b/backend/test/workspaces-diagnostics.test.ts @@ -1,5 +1,9 @@ import { expect, test, vi } from "vitest"; +import { mkdtemp, rm, writeFile } from "node:fs/promises"; +import { tmpdir } from "node:os"; +import { join } from "node:path"; import { + createConcreteDiagnosticAdapters, createProductionWorkspaceDiagnoser, createWorkspaceDiagnoser, type DiagnosticAdapters, @@ -35,6 +39,64 @@ semantic_index: timeout_ms: 8000 llm_policy: allowed: [zai/glm-5.2] +diagnostics: + dwh_rest: + method: POST + path: /rpc/ping + auth: bearer + response: { database: database, schema: schema } + vector_rest: + metadata: + method: GET + path: /vector/metadata + auth: bearer + response: { collection: collection, dimensions: dimensions, distance: distance } + reversible_probe: + method: POST + path: /vector/diagnostic-probe + auth: bearer +`); + +const writerWorkspace = parseWorkspaceYaml(`workspace: + schema_version: 2 + id: psd-clinical + name: Policlinico San Donato + language: it +dwh: + engine: postgres + database: warehouse + schema: datawarehouse + timeout_ms: 8000 + supported_transports: [postgres_direct, rest_api, ssh_tunnel] +semantic_index: + vector_store: + engine: pgvector + database: postgres + schema: vectors + collection: clinical_documents + dimensions: 768 + distance: cosine + timeout_ms: 8000 + supported_transports: [pgvector_direct, rest_api, ssh_tunnel] + vector_writer: {} + embedding: + provider: ollama_compatible + model: nomic-embed-text-v2-moe + dimensions: 768 + timeout_ms: 8000 +llm_policy: + allowed: [zai/glm-5.2] +diagnostics: + vector_rest: + metadata: + method: GET + path: /vector/metadata + auth: bearer + response: { collection: collection, dimensions: dimensions, distance: distance } + reversible_probe: + method: POST + path: /vector/diagnostic-probe + auth: bearer `); const bindings: RuntimeBindings = { @@ -71,6 +133,19 @@ const bindings: RuntimeBindings = { }, }; +const writerBindings: RuntimeBindings = { + ...bindings, + vector: { + transport: "rest_api", + missing: [], + values: { + THT_WS_PSD_CLINICAL_VECTOR_BASE_URL: "https://vector.example.test", + THT_WS_PSD_CLINICAL_VECTOR_API_KEY_FILE: "/run/secrets/vector-reader-key", + THT_WS_PSD_CLINICAL_VECTOR_WRITER_API_KEY_FILE: "/run/secrets/vector-writer-key", + }, + }, +}; + function successfulAdapters(overrides: Partial = {}): DiagnosticAdapters { return { probeConnector: vi.fn(async (request) => ({ @@ -216,10 +291,45 @@ test("uses a loopback-only SSH tunnel for the bounded connector probe", async () })); }); +test("passes the declared vector database and schema to direct diagnostics", async () => { + const adapters = successfulAdapters(); + + await diagnose(adapters)(workspace, bindings, { writeProbe: false }); + + expect(adapters.probeConnector).toHaveBeenCalledWith(expect.objectContaining({ + role: "vector", + resource: { database: "postgres", schema: "vectors", collection: "clinical_documents" }, + })); +}); + +test("uses strict known-host SSH arguments and always closes the temporary tunnel", async () => { + const directory = await mkdtemp(join(tmpdir(), "thothii-diagnostic-")); + const privateKeyFile = join(directory, "ssh-key"); + await writeFile(privateKeyFile, "test-key\n", { mode: 0o600 }); + const close = vi.fn(async () => undefined); + const start = vi.fn(async () => ({ tunnel: { host: "127.0.0.1" as const, port: 45432 }, close })); + + try { + const adapter = createConcreteDiagnosticAdapters({ sshProcess: { start } }); + await adapter.withSshTunnel({ + sshHost: "bastion.example.test", sshPort: 22, sshUser: "tunnel", privateKeyFile, + knownHostsFile: "/run/secrets/known-hosts", targetHost: "vector.internal", targetPort: 5432, + localHost: "127.0.0.1", localPort: 0, timeoutMs: 5000, signal: new AbortController().signal, + }, async () => undefined); + + expect(start).toHaveBeenCalledWith(expect.any(Object), expect.arrayContaining([ + "StrictHostKeyChecking=yes", "UserKnownHostsFile=/run/secrets/known-hosts", "-i", privateKeyFile, + ])); + expect(close).toHaveBeenCalledOnce(); + } finally { + await rm(directory, { recursive: true, force: true }); + } +}); + test("requires a matching embedding model vector and removes its unique write probe", async () => { const adapters = successfulAdapters(); - const result = await diagnose(adapters)(workspace, bindings, { writeProbe: true }); + const result = await diagnose(adapters)(writerWorkspace, writerBindings, { writeProbe: true }); expect(result.activatable).toBe(true); expect(adapters.probeEmbedding).toHaveBeenCalledWith(expect.objectContaining({ @@ -237,6 +347,62 @@ test("requires a matching embedding model vector and removes its unique write pr })); }); +test("keeps a reader-only workspace activatable without a vector write probe", async () => { + const adapters = successfulAdapters(); + + const result = await diagnose(adapters)(workspace, bindings, { writeProbe: true }); + + expect(result.activatable).toBe(true); + expect(adapters.writeDiagnosticRecord).not.toHaveBeenCalled(); + expect(adapters.removeDiagnosticRecord).not.toHaveBeenCalled(); +}); + +test("does not substitute the reader credential for a declared vector writer", async () => { + const adapters = successfulAdapters(); + + const result = await diagnose(adapters)(writerWorkspace, bindings, { writeProbe: true }); + + expect(result.activatable).toBe(true); + expect(adapters.writeDiagnosticRecord).not.toHaveBeenCalled(); +}); + +test("uses the declared POST DWH ping endpoint without exposing its local credential", async () => { + const directory = await mkdtemp(join(tmpdir(), "thothii-diagnostic-")); + const credentialFile = join(directory, "dwh-api-key"); + await writeFile(credentialFile, "local-secret\n", { mode: 0o600 }); + const fetchSpy = vi.fn(async () => new Response(JSON.stringify({ database: "warehouse", schema: "datawarehouse" }), { + status: 200, + headers: { "content-type": "application/json" }, + })); + vi.stubGlobal("fetch", fetchSpy); + + try { + const result = await createConcreteDiagnosticAdapters().probeConnector({ + role: "dwh", + transport: "rest_api", + baseUrl: "https://dwh.example.test", + credentialFile, + resource: { database: "warehouse", schema: "datawarehouse" }, + diagnostic: { + method: "POST", path: "/rpc/ping", auth: "bearer", + response: { database: "database", schema: "schema" }, + }, + timeoutMs: 5000, + signal: new AbortController().signal, + }); + + expect(fetchSpy).toHaveBeenCalledWith("https://dwh.example.test/rpc/ping", expect.objectContaining({ + method: "POST", + redirect: "error", + })); + expect(result).toMatchObject({ authenticated: true, resource: { database: "warehouse", schema: "datawarehouse" } }); + expect(JSON.stringify(result)).not.toContain("local-secret"); + } finally { + vi.unstubAllGlobals(); + await rm(directory, { recursive: true, force: true }); + } +}); + test("constructs the production diagnoser with the configured timeout and injected adapters", async () => { const adapters = successfulAdapters(); @@ -256,7 +422,7 @@ test("retries bounded cleanup after a write-probe removal times out", async () = const diagnoseWithShortTimeout = createWorkspaceDiagnoser(adapters, { timeoutMs: 10 }); const startedAt = Date.now(); - const result = await diagnoseWithShortTimeout(workspace, bindings, { writeProbe: true }); + const result = await diagnoseWithShortTimeout(writerWorkspace, writerBindings, { writeProbe: true }); expect(Date.now() - startedAt).toBeLessThan(250); expect(adapters.writeDiagnosticRecord).toHaveBeenCalledTimes(1);