Files
ThothII/backend/src/catalog/postgres-access.ts
T

263 lines
9.8 KiB
TypeScript

import { spawn, type ChildProcessWithoutNullStreams } from "node:child_process";
import { readFile } from "node:fs/promises";
import { fileURLToPath } from "node:url";
import { Duplex } from "node:stream";
import { setTimeout as delay } from "node:timers/promises";
import { Client, type ClientConfig } from "pg";
import type { WorkspaceSecretStore } from "../workspaces/secret-store.js";
import { CATALOG_SECRET_IDS } from "./secrets.js";
import { CatalogConnectorError, type WorkspaceDatabase } from "./types.js";
export interface CatalogDatabaseClient {
query(sql: string, values: readonly unknown[]): Promise<{ rows: Array<Record<string, unknown>> }>;
end(): Promise<void>;
}
export interface CatalogPostgresAccess {
connect(database: WorkspaceDatabase, signal: AbortSignal): Promise<CatalogDatabaseClient>;
}
type SpawnSsh = (
command: string,
args: readonly string[],
options: { env: NodeJS.ProcessEnv },
) => ChildProcessWithoutNullStreams;
interface AccessDependencies {
createClient?: (config: ClientConfig) => Client;
spawnSsh?: SpawnSsh;
sshBinary?: string;
askpassPath?: string;
connectTimeoutMs?: number;
}
function required(value: string | number | undefined): string | number {
if (value === undefined || value === "") throw new CatalogConnectorError("Database binding is incomplete");
return value;
}
function connectionSsl(ca: string | undefined, servername: string | undefined): ClientConfig["ssl"] {
if (!ca && !servername) return false;
return {
...(ca ? { ca } : {}),
...(servername ? { servername } : {}),
rejectUnauthorized: true,
};
}
export function buildSshArguments(input: {
sshHost: string;
sshPort: number;
sshUsername: string;
targetHost: string;
targetPort: number;
privateKeyFile: string;
knownHostsFile: string;
passphraseFile?: string;
connectTimeoutMs: number;
}): string[] {
const batchMode = input.passphraseFile ? "no" : "yes";
return [
"-F", "/dev/null",
"-T",
"-o", `BatchMode=${batchMode}`,
"-o", "StrictHostKeyChecking=yes",
"-o", `UserKnownHostsFile=${input.knownHostsFile}`,
"-o", "GlobalKnownHostsFile=/dev/null",
"-o", "IdentitiesOnly=yes",
"-o", "IdentityAgent=none",
"-o", `IdentityFile=${input.privateKeyFile}`,
"-o", "PreferredAuthentications=publickey",
"-o", "PasswordAuthentication=no",
"-o", "KbdInteractiveAuthentication=no",
"-o", "ConnectionAttempts=1",
"-o", `ConnectTimeout=${Math.max(1, Math.ceil(input.connectTimeoutMs / 1_000))}`,
"-o", "ServerAliveInterval=5",
"-o", "ServerAliveCountMax=1",
"-o", "NumberOfPasswordPrompts=1",
"-o", "RequestTTY=no",
"-o", "LogLevel=ERROR",
"-p", String(input.sshPort),
"-W", `${input.targetHost}:${input.targetPort}`,
"--", `${input.sshUsername}@${input.sshHost}`,
];
}
async function stopChild(child: ChildProcessWithoutNullStreams): Promise<void> {
if (child.exitCode !== null || child.signalCode !== null) return;
child.kill("SIGTERM");
await Promise.race([
new Promise<void>((resolve) => child.once("exit", () => resolve())),
delay(500).then(() => undefined),
]);
if (child.exitCode === null && child.signalCode === null) child.kill("SIGKILL");
}
function sshDuplex(child: ChildProcessWithoutNullStreams): Duplex {
let ended = false;
let stream: Duplex;
const forward = () => {
let chunk: Buffer | string | null;
while ((chunk = child.stdout.read() as Buffer | string | null) !== null) {
if (!stream.push(chunk)) break;
}
};
const finish = () => {
if (ended) return;
ended = true;
stream.push(null);
};
const fail = (error: Error) => stream.destroy(error);
stream = new Duplex({
read: forward,
write: (chunk, encoding, callback) => child.stdin.write(chunk, encoding, callback),
final: (callback) => child.stdin.end(callback),
destroy: (error, callback) => {
child.stdout.off("readable", forward);
child.stdout.off("end", finish);
child.stdout.off("error", fail);
child.stdin.off("error", fail);
callback(error);
},
});
child.stdout.on("readable", forward);
child.stdout.once("end", finish);
child.stdout.once("error", fail);
child.stdin.once("error", fail);
return stream;
}
/**
* Deep connection module for direct and SSH-forwarded PostgreSQL access. It owns secret leases,
* TLS, OpenSSH lifecycle, abort propagation, and pg cleanup behind one connect interface.
*/
export class ConcreteCatalogPostgresAccess implements CatalogPostgresAccess {
private readonly createClient: (config: ClientConfig) => Client;
private readonly spawnSsh: SpawnSsh;
private readonly sshBinary: string;
private readonly askpassPath: string;
private readonly connectTimeoutMs: number;
constructor(
private readonly secretStore: WorkspaceSecretStore,
dependencies: AccessDependencies = {},
) {
this.createClient = dependencies.createClient ?? ((config) => new Client(config));
this.spawnSsh = dependencies.spawnSsh ?? ((command, args, options) => (
spawn(command, [...args], { ...options, stdio: ["pipe", "pipe", "pipe"] })
));
this.sshBinary = dependencies.sshBinary ?? process.env.THT_SSH_BIN ?? "ssh";
this.askpassPath = dependencies.askpassPath
?? process.env.THT_SSH_ASKPASS_BIN
?? fileURLToPath(new URL("../../scripts/ssh-askpass.mjs", import.meta.url));
this.connectTimeoutMs = dependencies.connectTimeoutMs ?? 5_000;
}
async connect(database: WorkspaceDatabase, signal: AbortSignal): Promise<CatalogDatabaseClient> {
if (database.binding.transport === "rest_api") {
throw new CatalogConnectorError("REST is not a PostgreSQL wire binding");
}
const ids: string[] = [CATALOG_SECRET_IDS.password, CATALOG_SECRET_IDS.tlsCa];
if (database.binding.transport === "ssh_tunnel") {
ids.push(
CATALOG_SECRET_IDS.sshPrivateKey,
CATALOG_SECRET_IDS.sshPrivateKeyPassphrase,
CATALOG_SECRET_IDS.sshKnownHosts,
);
}
const materialized = this.secretStore.materialize(database.workspaceId, ids);
let child: ChildProcessWithoutNullStreams | undefined;
let stream: Duplex | undefined;
let client: Client | undefined;
let ended = false;
const close = async () => {
if (ended) return;
ended = true;
signal.removeEventListener("abort", abort);
if (client) await client.end().catch(() => undefined);
stream?.destroy();
if (child) await stopChild(child);
materialized.release();
};
const abort = () => { void close(); };
signal.addEventListener("abort", abort, { once: true });
try {
if (signal.aborted) throw new CatalogConnectorError("PostgreSQL connector aborted");
const passwordFile = materialized.files.get(CATALOG_SECRET_IDS.password);
if (!passwordFile) throw new CatalogConnectorError("Database password is not configured");
const password = await readFile(passwordFile, "utf8");
if (signal.aborted) throw new CatalogConnectorError("PostgreSQL connector aborted");
const tlsCaFile = materialized.files.get(CATALOG_SECRET_IDS.tlsCa);
const tlsCa = tlsCaFile ? await readFile(tlsCaFile, "utf8") : undefined;
if (signal.aborted) throw new CatalogConnectorError("PostgreSQL connector aborted");
let host: string;
let port: number;
if (database.binding.transport === "ssh_tunnel") {
const privateKeyFile = materialized.files.get(CATALOG_SECRET_IDS.sshPrivateKey);
const knownHostsFile = materialized.files.get(CATALOG_SECRET_IDS.sshKnownHosts);
if (!privateKeyFile || !knownHostsFile) {
throw new CatalogConnectorError("SSH private key and known hosts are required");
}
const passphraseFile = materialized.files.get(CATALOG_SECRET_IDS.sshPrivateKeyPassphrase);
host = String(required(database.binding.sshTargetHost));
port = Number(required(database.binding.sshTargetPort));
const args = buildSshArguments({
sshHost: String(required(database.binding.sshHost)),
sshPort: Number(required(database.binding.sshPort)),
sshUsername: String(required(database.binding.sshUsername)),
targetHost: host,
targetPort: port,
privateKeyFile,
knownHostsFile,
passphraseFile,
connectTimeoutMs: this.connectTimeoutMs,
});
child = this.spawnSsh(this.sshBinary, args, {
env: {
...process.env,
LC_ALL: "C",
...(passphraseFile ? {
DISPLAY: "thothii",
SSH_ASKPASS: this.askpassPath,
SSH_ASKPASS_REQUIRE: "force",
THT_SSH_PASSPHRASE_FILE: passphraseFile,
} : {}),
},
});
stream = sshDuplex(child);
child.once("error", () => stream?.destroy(new CatalogConnectorError("SSH process failed")));
child.once("exit", (code) => {
if (!ended && code !== 0) stream?.destroy(new CatalogConnectorError("SSH tunnel failed"));
});
child.stderr.on("data", () => undefined);
} else {
host = String(required(database.binding.host));
port = Number(required(database.binding.port));
}
client = this.createClient({
host,
port,
database: database.databaseName,
user: String(required(database.binding.username)),
password,
ssl: connectionSsl(tlsCa, database.binding.tlsServername),
connectionTimeoutMillis: this.connectTimeoutMs,
...(stream ? { stream: () => stream } : {}),
});
await client.connect();
if (signal.aborted) throw new CatalogConnectorError("PostgreSQL connector aborted");
return {
query: async (sql, values) => await client!.query(sql, [...values]),
end: close,
};
} catch (error) {
await close();
if (error instanceof CatalogConnectorError) throw error;
throw new CatalogConnectorError("PostgreSQL connector failed");
}
}
}