Files
ThothII/backend/src/catalog/sensitivity-value-source.ts
T

292 lines
12 KiB
TypeScript

import { readFile } from "node:fs/promises";
import type { WorkspaceSecretStore } from "../workspaces/secret-store.js";
import { CATALOG_SECRET_IDS } from "./secrets.js";
import type { CatalogPostgresAccess } from "./postgres-access.js";
import type {
SensitivityScanCoverage,
SensitivityScanRequest,
SensitivityValueObservation,
SensitivityValueSource,
} from "./sensitivity-classifier.js";
import { CatalogConnectorError, type CatalogColumn } from "./types.js";
const MAX_VALUE_CHARACTERS = 501;
const MAX_COLUMNS_PER_QUERY = 25;
const SAMPLE_OVERSCAN_FACTOR = 10;
function quoteIdentifier(identifier: string): string {
return `"${identifier.replaceAll('"', '""')}"`;
}
function chunks<T>(items: readonly T[], size: number): T[][] {
const result: T[][] = [];
for (let offset = 0; offset < items.length; offset += size) {
result.push(items.slice(offset, offset + size));
}
return result;
}
function tableReference(request: SensitivityScanRequest): string {
return `${quoteIdentifier(request.database.schema)}.${quoteIdentifier(request.table.name)}`;
}
function samplePercentage(valuesPerColumn: number): number {
if (valuesPerColumn <= 300) return 30;
if (valuesPerColumn <= 700) return 70;
return 100;
}
function flatValueQuery(
request: SensitivityScanRequest,
columns: readonly CatalogColumn[],
options: { complete: boolean; randomized: boolean },
): string {
const projections = columns.map((column) => quoteIdentifier(column.name)).join(", ");
const perColumnLimit = options.complete
? request.fullScanThreshold ?? request.valuesPerColumn
: request.valuesPerColumn;
const rowLimit = Math.max(perColumnLimit, perColumnLimit * SAMPLE_OVERSCAN_FACTOR);
const sample = options.complete
? `SELECT ${projections} FROM ${tableReference(request)}`
: [
`SELECT ${projections} FROM ${tableReference(request)}`,
...(options.randomized
? [`TABLESAMPLE SYSTEM (${samplePercentage(request.valuesPerColumn)}) REPEATABLE (${request.sampleSeed})`]
: []),
`LIMIT ${rowLimit} OFFSET ${request.sampleOffset}`,
].join(" ");
const values = columns.map((column, index) => {
const identifier = quoteIdentifier(column.name);
return [
`(${index}, LEFT((sampled.${identifier})::text, ${MAX_VALUE_CHARACTERS}),`,
`CASE WHEN sampled.${identifier} IS NULL THEN NULL`,
`ELSE char_length((sampled.${identifier})::text) END)`,
].join(" ");
}).join(", ");
return [
`WITH sampled AS MATERIALIZED (${sample}),`,
"ranked AS (",
"SELECT value.__column_index, value.__value, value.__length,",
"row_number() OVER (PARTITION BY value.__column_index) AS __rank",
"FROM sampled",
`CROSS JOIN LATERAL (VALUES ${values}) AS value(__column_index, __value, __length)`,
"WHERE value.__value IS NOT NULL",
")",
"SELECT __column_index, __value, __length FROM ranked",
`WHERE __rank <= ${perColumnLimit}`,
].join(" ");
}
function observations(
columns: readonly CatalogColumn[],
rows: readonly Record<string, unknown>[],
): SensitivityValueObservation[] {
return rows.flatMap((row) => {
const index = Number(row.__column_index);
const column = Number.isSafeInteger(index) && index >= 0 ? columns[index] : undefined;
if (!column || row.__value === null || row.__value === undefined) return [];
const value = String(row.__value);
const parsedLength = row.__length === null || row.__length === undefined
? null
: Number(row.__length);
return [{
columnId: column.id,
value,
characterLength: parsedLength !== null && Number.isSafeInteger(parsedLength) && parsedLength >= 0
? parsedLength
: value.length,
}];
});
}
function cancelled(error: unknown): boolean {
return Boolean(error && typeof error === "object" && "code" in error && error.code === "57014");
}
/**
* Database-specific sampling adapter. Policy stays in SensitivityClassifier; this module only
* produces bounded, normalized non-null observations without persisting or logging values.
*/
export class ConcreteSensitivityValueSource implements SensitivityValueSource {
constructor(
private readonly access: CatalogPostgresAccess,
private readonly secretStore?: Pick<WorkspaceSecretStore, "materialize">,
) {}
async scanTable(
request: SensitivityScanRequest,
consume: (batch: readonly SensitivityValueObservation[]) => void | Promise<void>,
signal: AbortSignal,
): Promise<SensitivityScanCoverage> {
if (request.columns.length === 0) return { kind: "complete", observedValues: 0 };
if (request.database.binding.transport === "rest_api") {
return await this.scanRest(request, consume, signal);
}
return await this.scanPostgres(request, consume, signal);
}
private async scanPostgres(
request: SensitivityScanRequest,
consume: (batch: readonly SensitivityValueObservation[]) => void | Promise<void>,
signal: AbortSignal,
): Promise<SensitivityScanCoverage> {
const client = await this.access.connect(request.database, signal);
let transactionOpen = false;
let savepointSequence = 0;
let observedValues = 0;
try {
signal.throwIfAborted();
await client.query("BEGIN TRANSACTION READ ONLY", []);
transactionOpen = true;
await client.query("SELECT set_config('statement_timeout', $1, true)", [
`${Math.max(1, Math.floor(request.queryTimeoutMs))}ms`,
]);
const boundedQuery = async (sql: string): Promise<Array<Record<string, unknown>> | undefined> => {
signal.throwIfAborted();
savepointSequence += 1;
const savepoint = `sensitivity_scan_${savepointSequence}`;
await client.query(`SAVEPOINT ${savepoint}`, []);
try {
return (await client.query(sql, [])).rows;
} catch (error) {
if (!cancelled(error)) throw error;
await client.query(`ROLLBACK TO SAVEPOINT ${savepoint}`, []);
return undefined;
} finally {
await client.query(`RELEASE SAVEPOINT ${savepoint}`, []).catch(() => undefined);
}
};
let complete = false;
if (request.fullScanThreshold !== undefined) {
const probe = await boundedQuery(
`SELECT 1 AS __present FROM ${tableReference(request)} LIMIT ${request.fullScanThreshold + 1}`,
);
complete = probe !== undefined && probe.length <= request.fullScanThreshold;
}
for (const columnChunk of chunks(request.columns, MAX_COLUMNS_PER_QUERY)) {
signal.throwIfAborted();
let rows = await boundedQuery(flatValueQuery(request, columnChunk, {
complete,
randomized: !complete,
}));
if (rows === undefined && complete) {
complete = false;
rows = await boundedQuery(flatValueQuery(request, columnChunk, {
complete: false,
randomized: true,
}));
}
if (!complete && (rows === undefined || rows.length === 0)) {
rows = await boundedQuery(flatValueQuery(request, columnChunk, {
complete: false,
randomized: false,
}));
}
if (rows === undefined) throw new CatalogConnectorError("Sensitivity sample query timed out");
const batch = observations(columnChunk, rows);
observedValues += batch.length;
if (batch.length > 0) await consume(batch);
}
return { kind: complete ? "complete" : "sampled", observedValues };
} catch (error) {
if (error instanceof CatalogConnectorError) throw error;
throw new CatalogConnectorError("Sensitivity source scan failed");
} finally {
if (transactionOpen) await client.query("ROLLBACK", []).catch(() => undefined);
await client.end().catch(() => undefined);
}
}
private async scanRest(
request: SensitivityScanRequest,
consume: (batch: readonly SensitivityValueObservation[]) => void | Promise<void>,
signal: AbortSignal,
): Promise<SensitivityScanCoverage> {
if (!this.secretStore) throw new CatalogConnectorError("REST sensitivity scanning is not configured");
const auth = request.database.binding.restAuth ?? "bearer";
const materialized = this.secretStore.materialize(
request.database.workspaceId,
auth === "none" ? [] : [CATALOG_SECRET_IDS.apiKey],
);
let observedValues = 0;
try {
const headers: Record<string, string> = { "content-type": "application/json" };
if (auth !== "none") {
const credentialFile = materialized.files.get(CATALOG_SECRET_IDS.apiKey);
if (!credentialFile) throw new CatalogConnectorError("REST API key is not configured");
const credential = (await readFile(credentialFile, "utf8")).trim();
if (auth === "bearer") headers.authorization = `Bearer ${credential}`;
else headers["x-api-key"] = credential;
}
const baseUrl = request.database.binding.baseUrl?.replace(/\/+$/u, "");
if (!baseUrl) throw new CatalogConnectorError("Database binding is incomplete");
const runQuery = async (sql: string): Promise<Array<Record<string, unknown>> | undefined> => {
const timeout = AbortSignal.timeout(Math.max(1, Math.floor(request.queryTimeoutMs)));
try {
const response = await fetch(`${baseUrl}/rpc/run_query`, {
method: "POST",
headers,
body: JSON.stringify({ query_text: sql }),
signal: AbortSignal.any([signal, timeout]),
});
if (!response.ok) throw new CatalogConnectorError("REST sensitivity source scan failed");
const body: unknown = await response.json();
if (!Array.isArray(body)
|| body.some((row) => !row || typeof row !== "object" || Array.isArray(row))) {
throw new CatalogConnectorError("REST sensitivity source response is invalid");
}
return body as Array<Record<string, unknown>>;
} catch (error) {
if (signal.aborted) throw error;
if (timeout.aborted) return undefined;
throw error;
}
};
let complete = false;
if (request.fullScanThreshold !== undefined) {
const probe = await runQuery(
`SELECT 1 AS __present FROM ${tableReference(request)} LIMIT ${request.fullScanThreshold + 1}`,
);
complete = probe !== undefined && probe.length <= request.fullScanThreshold;
}
let requestCount = request.fullScanThreshold === undefined ? 0 : 1;
for (const columnChunk of chunks(request.columns, MAX_COLUMNS_PER_QUERY)) {
signal.throwIfAborted();
let rows = await runQuery(flatValueQuery(request, columnChunk, {
complete,
randomized: !complete,
}));
requestCount += 1;
if (rows === undefined && complete) {
complete = false;
rows = await runQuery(flatValueQuery(request, columnChunk, {
complete: false,
randomized: true,
}));
requestCount += 1;
}
if (!complete && (rows === undefined || rows.length === 0)) {
rows = await runQuery(flatValueQuery(request, columnChunk, {
complete: false,
randomized: false,
}));
requestCount += 1;
}
if (rows === undefined) throw new CatalogConnectorError("REST sensitivity sample query timed out");
const batch = observations(columnChunk, rows);
observedValues += batch.length;
if (batch.length > 0) await consume(batch);
}
// Multiple HTTP requests cannot share a source snapshot, so only one-request reads are complete.
return { kind: complete && requestCount === 1 ? "complete" : "sampled", observedValues };
} catch (error) {
if (error instanceof CatalogConnectorError) throw error;
throw new CatalogConnectorError("REST sensitivity source scan failed");
} finally {
materialized.release();
}
}
}