Files
ThothII/backend/src/catalog/description-source-sampler.ts
T

253 lines
9.3 KiB
TypeScript

import { readFile } from "node:fs/promises";
import type { WorkspaceSecretStore } from "../workspaces/secret-store.js";
import type { CatalogPostgresAccess } from "./postgres-access.js";
import { CATALOG_SECRET_IDS } from "./secrets.js";
import { CatalogConnectorError, type WorkspaceDatabase } from "./types.js";
const MAX_SOURCE_ROWS = 5;
const MAX_REPRESENTATIVE_VALUES = 5;
const MAX_SOURCE_COLUMNS_PER_TARGET = 8;
const MAX_SOURCE_VALUE_BYTES = 256;
export type DescriptionSourceSampleValue = string | number | boolean | null;
export interface DescriptionSourceSampleField {
name: string;
value: DescriptionSourceSampleValue;
}
export interface DescriptionSourceSampleRow {
fields: readonly DescriptionSourceSampleField[];
}
export interface DescriptionSourceRepresentativeValues {
column: string;
values: readonly Exclude<DescriptionSourceSampleValue, null>[];
}
export interface DescriptionTargetSourceSample {
targetId: string;
tableName: string;
rows: readonly DescriptionSourceSampleRow[];
representativeValues: readonly DescriptionSourceRepresentativeValues[];
}
export interface DescriptionSourceSamplingTarget {
targetId: string;
tableName: string;
columnNames: readonly string[];
}
/** Optional, transient source context for one model-completion batch. */
export interface DescriptionSourceSampler {
sample(
database: WorkspaceDatabase,
targets: readonly DescriptionSourceSamplingTarget[],
signal: AbortSignal,
): Promise<readonly DescriptionTargetSourceSample[]>;
}
function quoteIdentifier(identifier: string): string {
return `"${identifier.replaceAll('"', '""')}"`;
}
function boundedUtf8(value: string, maxBytes: number): string {
const normalized = value
.normalize("NFC")
.replace(/\r\n?/g, "\n")
.replace(/[\u0000-\u0008\u000b\u000c\u000e-\u001f\u007f]/g, " ");
if (Buffer.byteLength(normalized, "utf8") <= maxBytes) return normalized;
let result = "";
let bytes = 0;
for (const character of normalized) {
const characterBytes = Buffer.byteLength(character, "utf8");
if (bytes + characterBytes > maxBytes) break;
result += character;
bytes += characterBytes;
}
return result;
}
function normalizeValue(value: unknown): DescriptionSourceSampleValue | undefined {
if (value === null) return null;
if (typeof value === "string") return boundedUtf8(value, MAX_SOURCE_VALUE_BYTES);
if (typeof value === "boolean") return value;
if (typeof value === "number") return Number.isFinite(value) ? value : undefined;
if (typeof value === "bigint") return boundedUtf8(String(value), MAX_SOURCE_VALUE_BYTES);
if (value instanceof Date && !Number.isNaN(value.valueOf())) return value.toISOString();
return undefined;
}
function distinctKey(value: Exclude<DescriptionSourceSampleValue, null>): string {
return `${typeof value}:${String(value)}`;
}
function columnsFor(target: DescriptionSourceSamplingTarget): string[] {
return [...new Set(target.columnNames)].slice(0, MAX_SOURCE_COLUMNS_PER_TARGET);
}
function normalizedSample(
target: DescriptionSourceSamplingTarget,
columnNames: readonly string[],
sourceRows: readonly Record<string, unknown>[],
): DescriptionTargetSourceSample {
const rows = sourceRows.slice(0, MAX_SOURCE_ROWS).map((row) => ({
fields: columnNames.flatMap((name) => {
const value = normalizeValue(row[name]);
return value === undefined ? [] : [{ name, value }];
}),
}));
const valuesByColumn = new Map<string, Exclude<DescriptionSourceSampleValue, null>[]>();
const seenByColumn = new Map<string, Set<string>>();
let representativeValueCount = 0;
for (const row of rows) {
for (const field of row.fields) {
if (representativeValueCount === MAX_REPRESENTATIVE_VALUES) break;
if (field.value === null) continue;
const seen = seenByColumn.get(field.name) ?? new Set<string>();
const key = distinctKey(field.value);
if (seen.has(key)) continue;
seen.add(key);
seenByColumn.set(field.name, seen);
const values = valuesByColumn.get(field.name) ?? [];
values.push(field.value);
valuesByColumn.set(field.name, values);
representativeValueCount += 1;
}
if (representativeValueCount === MAX_REPRESENTATIVE_VALUES) break;
}
return {
targetId: target.targetId,
tableName: target.tableName,
rows,
representativeValues: columnNames.flatMap((column) => {
const values = valuesByColumn.get(column);
return values && values.length > 0 ? [{ column, values }] : [];
}),
};
}
function samplingSql(
database: WorkspaceDatabase,
target: DescriptionSourceSamplingTarget,
columns: readonly string[],
): string {
const projections = columns.map((columnName) => {
const identifier = quoteIdentifier(columnName);
return `LEFT((${identifier})::text, ${MAX_SOURCE_VALUE_BYTES}) AS ${identifier}`;
});
return [
`SELECT ${projections.join(", ")}`,
`FROM ${quoteIdentifier(database.schema)}.${quoteIdentifier(target.tableName)}`,
`LIMIT ${MAX_SOURCE_ROWS}`,
].join(" ");
}
/** Bounded source sampler that follows the database's PostgreSQL-wire or REST binding. */
export class ConcreteDescriptionSourceSampler implements DescriptionSourceSampler {
constructor(
private readonly access: CatalogPostgresAccess,
private readonly secretStore?: Pick<WorkspaceSecretStore, "materialize">,
) {}
async sample(
database: WorkspaceDatabase,
targets: readonly DescriptionSourceSamplingTarget[],
signal: AbortSignal,
): Promise<readonly DescriptionTargetSourceSample[]> {
if (database.binding.transport === "rest_api") {
return await this.sampleRest(database, targets, signal);
}
const client = await this.access.connect(database, signal);
let transactionOpen = false;
try {
await client.query("BEGIN TRANSACTION READ ONLY", []);
transactionOpen = true;
const samples: DescriptionTargetSourceSample[] = [];
for (const target of targets) {
const columnNames = columnsFor(target);
if (columnNames.length === 0) {
samples.push({
targetId: target.targetId,
tableName: target.tableName,
rows: [],
representativeValues: [],
});
continue;
}
const projections = columnNames.map((columnName) => {
const identifier = quoteIdentifier(columnName);
return `LEFT((${identifier})::text, $1) AS ${identifier}`;
});
const sql = [
`SELECT ${projections.join(", ")}`,
`FROM ${quoteIdentifier(database.schema)}.${quoteIdentifier(target.tableName)}`,
"LIMIT $2",
].join(" ");
const result = await client.query(sql, [MAX_SOURCE_VALUE_BYTES, MAX_SOURCE_ROWS]);
samples.push(normalizedSample(target, columnNames, result.rows));
}
return samples;
} finally {
if (transactionOpen) await client.query("ROLLBACK", []).catch(() => undefined);
await client.end().catch(() => undefined);
}
}
private async sampleRest(
database: WorkspaceDatabase,
targets: readonly DescriptionSourceSamplingTarget[],
signal: AbortSignal,
): Promise<readonly DescriptionTargetSourceSample[]> {
if (!this.secretStore) throw new CatalogConnectorError("REST source sampling is not configured");
const auth = database.binding.restAuth ?? "bearer";
const materialized = this.secretStore.materialize(
database.workspaceId,
auth === "none" ? [] : [CATALOG_SECRET_IDS.apiKey],
);
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 = database.binding.baseUrl?.replace(/\/+$/, "");
if (!baseUrl) throw new CatalogConnectorError("Database binding is incomplete");
const samples: DescriptionTargetSourceSample[] = [];
for (const target of targets) {
const columnNames = columnsFor(target);
if (columnNames.length === 0) {
samples.push(normalizedSample(target, columnNames, []));
continue;
}
const response = await fetch(`${baseUrl}/rpc/run_query`, {
method: "POST",
headers,
body: JSON.stringify({ query_text: samplingSql(database, target, columnNames) }),
signal,
});
if (!response.ok) throw new CatalogConnectorError("REST source sampling 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 source sampling response is invalid");
}
samples.push(normalizedSample(
target,
columnNames,
body as Array<Record<string, unknown>>,
));
}
return samples;
} catch (error) {
if (error instanceof CatalogConnectorError) throw error;
throw new CatalogConnectorError("REST source sampling failed");
} finally {
materialized.release();
}
}
}