feat: refine metadata catalog workflows
This commit is contained in:
+2
-2
@@ -59,7 +59,7 @@ import { DescriptionGenerationWorker } from "./catalog/description-generation-wo
|
||||
import { SensitiveDataSuggester } from "./catalog/sensitive-data-suggester.js";
|
||||
import { SensitiveDataSuggestionRunner } from "./catalog/sensitive-data-suggestion-runner.js";
|
||||
import {
|
||||
PostgresDescriptionSourceSampler,
|
||||
ConcreteDescriptionSourceSampler,
|
||||
type DescriptionSourceSampler,
|
||||
} from "./catalog/description-source-sampler.js";
|
||||
import { catalogDescriptionGenerationRoutes } from "./routes/catalog-description-generation.js";
|
||||
@@ -177,7 +177,7 @@ export function buildApp(config: AppConfig, deps?: BuildAppDeps): FastifyInstanc
|
||||
{ connectTimeoutMs: config.workspaceDiagnosticTimeoutMs },
|
||||
);
|
||||
const descriptionSourceSampler = deps?.descriptionSourceSampler
|
||||
?? new PostgresDescriptionSourceSampler(catalogPostgresAccess);
|
||||
?? new ConcreteDescriptionSourceSampler(catalogPostgresAccess, workspaceSecretStore);
|
||||
const descriptionGenerationWorker = new DescriptionGenerationWorker(
|
||||
catalogRepository,
|
||||
workspaceRegistry,
|
||||
|
||||
@@ -1,7 +1,7 @@
|
||||
import { z } from "zod";
|
||||
import type { WorkspaceRegistry } from "../workspaces/registry.js";
|
||||
import type { MetadataGenerationModels, ResolvedMetadataGenerationModel } from "./metadata-generation-models.js";
|
||||
import type { ModelCompleter, ModelCompletionMessage } from "./model-completer.js";
|
||||
import type { ModelCompleter, ModelCompletionMessage, ModelCompletionResult } from "./model-completer.js";
|
||||
import {
|
||||
ModelCompletionCancelledError,
|
||||
ModelCompletionProviderError,
|
||||
@@ -38,6 +38,7 @@ const MAX_SAMPLE_FIELDS_PER_ROW = 4;
|
||||
const MAX_SAMPLE_COLUMNS = 4;
|
||||
const MAX_REPRESENTATIVE_VALUES_PER_REQUEST = 5;
|
||||
const MAX_TARGET_SAMPLE_JSON_BYTES = 8 * 1024;
|
||||
const MAX_COMPLETION_ATTEMPTS_PER_BATCH = 2;
|
||||
const generatedOutcomeSchema = z.object({
|
||||
targetId: z.uuid(),
|
||||
outcome: z.literal("generated"),
|
||||
@@ -55,7 +56,9 @@ const completionResponseSchema = z.object({
|
||||
results: z.array(outcomeSchema).min(1).max(MAX_TARGETS_PER_BATCH),
|
||||
}).strict();
|
||||
|
||||
class InvalidModelOutcomeError extends Error {}
|
||||
class InvalidModelJsonError extends Error {}
|
||||
class InvalidModelSchemaError extends Error {}
|
||||
class MissingModelTargetsError extends Error {}
|
||||
class DescriptionGenerationBatchError extends Error {
|
||||
constructor(
|
||||
readonly failure: unknown,
|
||||
@@ -143,6 +146,9 @@ interface DescriptionGenerationCounters {
|
||||
generated: number;
|
||||
nonGeneratable: number;
|
||||
failed: number;
|
||||
inputTokens: number;
|
||||
cacheReadTokens: number;
|
||||
outputTokens: number;
|
||||
consecutiveTechnicalFailures: number;
|
||||
}
|
||||
|
||||
@@ -152,6 +158,9 @@ function persistedCounters(counters: DescriptionGenerationCounters) {
|
||||
generated: counters.generated,
|
||||
nonGeneratable: counters.nonGeneratable,
|
||||
failed: counters.failed,
|
||||
inputTokens: counters.inputTokens,
|
||||
cacheReadTokens: counters.cacheReadTokens,
|
||||
outputTokens: counters.outputTokens,
|
||||
};
|
||||
}
|
||||
|
||||
@@ -625,36 +634,47 @@ function messagesFor(
|
||||
}
|
||||
|
||||
function parseOutcomes(content: string, expectedTargetIds: readonly string[]): Map<string, ParsedOutcome> {
|
||||
const trimmed = content.trim();
|
||||
const fenced = /^```(?:json)?[ \t]*\r?\n([\s\S]*?)\r?\n```$/iu.exec(trimmed);
|
||||
let parsed: unknown;
|
||||
try {
|
||||
const trimmed = content.trim();
|
||||
const fenced = /^```(?:json)?[ \t]*\r?\n([\s\S]*?)\r?\n```$/iu.exec(trimmed);
|
||||
const outcomes = completionResponseSchema.parse(JSON.parse(fenced?.[1] ?? trimmed)).results;
|
||||
const expected = new Set(expectedTargetIds);
|
||||
if (outcomes.length !== expectedTargetIds.length || expected.size !== expectedTargetIds.length) {
|
||||
throw new InvalidModelOutcomeError();
|
||||
}
|
||||
const mapped = new Map<string, ParsedOutcome>();
|
||||
for (const outcome of outcomes) {
|
||||
if (!expected.has(outcome.targetId) || mapped.has(outcome.targetId)) {
|
||||
throw new InvalidModelOutcomeError();
|
||||
}
|
||||
mapped.set(outcome.targetId, outcome.outcome === "generated"
|
||||
? { ...outcome, description: outcome.description.trim() }
|
||||
: outcome);
|
||||
}
|
||||
if (mapped.size !== expected.size) throw new InvalidModelOutcomeError();
|
||||
return mapped;
|
||||
} catch (error) {
|
||||
if (error instanceof InvalidModelOutcomeError) throw error;
|
||||
throw new InvalidModelOutcomeError();
|
||||
parsed = JSON.parse(fenced?.[1] ?? trimmed);
|
||||
} catch {
|
||||
throw new InvalidModelJsonError();
|
||||
}
|
||||
|
||||
const response = completionResponseSchema.safeParse(parsed);
|
||||
if (!response.success) throw new InvalidModelSchemaError();
|
||||
|
||||
const outcomes = response.data.results;
|
||||
const expected = new Set(expectedTargetIds);
|
||||
if (outcomes.length !== expectedTargetIds.length || expected.size !== expectedTargetIds.length) {
|
||||
throw new MissingModelTargetsError();
|
||||
}
|
||||
const mapped = new Map<string, ParsedOutcome>();
|
||||
for (const outcome of outcomes) {
|
||||
if (!expected.has(outcome.targetId) || mapped.has(outcome.targetId)) {
|
||||
throw new MissingModelTargetsError();
|
||||
}
|
||||
mapped.set(outcome.targetId, outcome.outcome === "generated"
|
||||
? { ...outcome, description: outcome.description.trim() }
|
||||
: outcome);
|
||||
}
|
||||
if (mapped.size !== expected.size) throw new MissingModelTargetsError();
|
||||
return mapped;
|
||||
}
|
||||
|
||||
function safeFailure(error: unknown): string {
|
||||
if (error instanceof DescriptionGenerationFailureStreakError) return error.message;
|
||||
const failure = error instanceof DescriptionGenerationBatchError ? error.failure : error;
|
||||
if (failure instanceof ModelCompletionProviderError) return "The model provider request failed.";
|
||||
if (failure instanceof InvalidModelOutcomeError) return "The model response was invalid.";
|
||||
if (failure instanceof InvalidModelJsonError) return "The model response was not valid JSON.";
|
||||
if (failure instanceof InvalidModelSchemaError) {
|
||||
return "The model response did not match the required schema.";
|
||||
}
|
||||
if (failure instanceof MissingModelTargetsError) {
|
||||
return "The model response was missing one or more requested targets.";
|
||||
}
|
||||
return "Description generation failed.";
|
||||
}
|
||||
|
||||
@@ -907,6 +927,9 @@ export class DescriptionGenerationWorker {
|
||||
generated: 0,
|
||||
nonGeneratable: 0,
|
||||
failed: 0,
|
||||
inputTokens: 0,
|
||||
cacheReadTokens: 0,
|
||||
outputTokens: 0,
|
||||
consecutiveTechnicalFailures: 0,
|
||||
};
|
||||
await this.processTargets(run, database, plan.columnTargets, model, counters, signal);
|
||||
@@ -973,20 +996,41 @@ export class DescriptionGenerationWorker {
|
||||
);
|
||||
}
|
||||
throwIfCancelled(signal);
|
||||
let outcomes: Map<string, ParsedOutcome>;
|
||||
try {
|
||||
const content = await this.completer.complete({
|
||||
model,
|
||||
messages: messagesFor(database, batch, run.language, sourceSamples),
|
||||
signal,
|
||||
});
|
||||
throwIfCancelled(signal);
|
||||
outcomes = parseOutcomes(content, batch.map((target) => (
|
||||
target.kind === "column" ? target.column.id : target.table.id
|
||||
)));
|
||||
} catch (error) {
|
||||
if (error instanceof ModelCompletionCancelledError) throw error;
|
||||
const batchError = new DescriptionGenerationBatchError(error, batch.map(failureTarget));
|
||||
const expectedTargetIds = batch.map((target) => (
|
||||
target.kind === "column" ? target.column.id : target.table.id
|
||||
));
|
||||
const messages = messagesFor(database, batch, run.language, sourceSamples);
|
||||
let outcomes: Map<string, ParsedOutcome> | undefined;
|
||||
let terminalFailure: unknown;
|
||||
for (let attempt = 1; attempt <= MAX_COMPLETION_ATTEMPTS_PER_BATCH; attempt += 1) {
|
||||
try {
|
||||
const completion = await this.completer.complete({ model, messages, signal });
|
||||
const result: ModelCompletionResult = typeof completion === "string"
|
||||
? { content: completion, usage: { input: 0, cacheRead: 0, output: 0 } }
|
||||
: completion;
|
||||
counters.inputTokens += result.usage.input;
|
||||
counters.cacheReadTokens += result.usage.cacheRead;
|
||||
counters.outputTokens += result.usage.output;
|
||||
throwIfCancelled(signal);
|
||||
outcomes = parseOutcomes(result.content, expectedTargetIds);
|
||||
break;
|
||||
} catch (error) {
|
||||
if (error instanceof ModelCompletionCancelledError) throw error;
|
||||
terminalFailure = error;
|
||||
if (attempt < MAX_COMPLETION_ATTEMPTS_PER_BATCH) {
|
||||
await this.appendEvent(
|
||||
run.id,
|
||||
"warning",
|
||||
`${safeFailure(error)} Retrying batch (attempt ${attempt + 1} of ${MAX_COMPLETION_ATTEMPTS_PER_BATCH}).`,
|
||||
);
|
||||
}
|
||||
}
|
||||
}
|
||||
if (!outcomes) {
|
||||
const batchError = new DescriptionGenerationBatchError(
|
||||
terminalFailure,
|
||||
batch.map(failureTarget),
|
||||
);
|
||||
counters.processed += batch.length;
|
||||
counters.failed += batch.length;
|
||||
counters.consecutiveTechnicalFailures += 1;
|
||||
@@ -1009,7 +1053,10 @@ export class DescriptionGenerationWorker {
|
||||
const targetId = target.kind === "column" ? target.column.id : target.table.id;
|
||||
const outcome = outcomes.get(targetId);
|
||||
if (!outcome) {
|
||||
throw new DescriptionGenerationBatchError(new InvalidModelOutcomeError(), [failureTarget(target)]);
|
||||
throw new DescriptionGenerationBatchError(
|
||||
new MissingModelTargetsError(),
|
||||
[failureTarget(target)],
|
||||
);
|
||||
}
|
||||
const generatedDescription = outcome.outcome === "generated"
|
||||
? outcome.description
|
||||
|
||||
@@ -1,5 +1,8 @@
|
||||
import { readFile } from "node:fs/promises";
|
||||
import type { WorkspaceSecretStore } from "../workspaces/secret-store.js";
|
||||
import type { CatalogPostgresAccess } from "./postgres-access.js";
|
||||
import type { WorkspaceDatabase } from "./types.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;
|
||||
@@ -79,15 +82,82 @@ function distinctKey(value: Exclude<DescriptionSourceSampleValue, null>): string
|
||||
return `${typeof value}:${String(value)}`;
|
||||
}
|
||||
|
||||
/** PostgreSQL-wire sampler. REST bindings remain unsupported by CatalogPostgresAccess. */
|
||||
export class PostgresDescriptionSourceSampler implements DescriptionSourceSampler {
|
||||
constructor(private readonly access: CatalogPostgresAccess) {}
|
||||
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 {
|
||||
@@ -95,7 +165,7 @@ export class PostgresDescriptionSourceSampler implements DescriptionSourceSample
|
||||
transactionOpen = true;
|
||||
const samples: DescriptionTargetSourceSample[] = [];
|
||||
for (const target of targets) {
|
||||
const columnNames = [...new Set(target.columnNames)].slice(0, MAX_SOURCE_COLUMNS_PER_TARGET);
|
||||
const columnNames = columnsFor(target);
|
||||
if (columnNames.length === 0) {
|
||||
samples.push({
|
||||
targetId: target.targetId,
|
||||
@@ -115,44 +185,7 @@ export class PostgresDescriptionSourceSampler implements DescriptionSourceSample
|
||||
"LIMIT $2",
|
||||
].join(" ");
|
||||
const result = await client.query(sql, [MAX_SOURCE_VALUE_BYTES, MAX_SOURCE_ROWS]);
|
||||
const rows = result.rows.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;
|
||||
}
|
||||
const representativeValues = columnNames.flatMap((column) => {
|
||||
const values = valuesByColumn.get(column);
|
||||
return values && values.length > 0 ? [{ column, values }] : [];
|
||||
});
|
||||
samples.push({
|
||||
targetId: target.targetId,
|
||||
tableName: target.tableName,
|
||||
rows,
|
||||
representativeValues,
|
||||
});
|
||||
samples.push(normalizedSample(target, columnNames, result.rows));
|
||||
}
|
||||
return samples;
|
||||
} finally {
|
||||
@@ -160,4 +193,60 @@ export class PostgresDescriptionSourceSampler implements DescriptionSourceSample
|
||||
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();
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
@@ -335,7 +335,10 @@ export class MemoryCatalogRepository implements CatalogRepository {
|
||||
processed: 0,
|
||||
generated: 0,
|
||||
nonGeneratable: 0,
|
||||
failed: 0,
|
||||
failed: 0,
|
||||
inputTokens: 0,
|
||||
cacheReadTokens: 0,
|
||||
outputTokens: 0,
|
||||
createdAt: now,
|
||||
startedAt: null,
|
||||
updatedAt: now,
|
||||
@@ -444,7 +447,10 @@ export class MemoryCatalogRepository implements CatalogRepository {
|
||||
status: "running",
|
||||
total: 0,
|
||||
suggestedSensitive: 0,
|
||||
suggestedNonSensitive: 0,
|
||||
suggestedNonSensitive: 0,
|
||||
inputTokens: 0,
|
||||
cacheReadTokens: 0,
|
||||
outputTokens: 0,
|
||||
createdAt: now,
|
||||
startedAt: now,
|
||||
updatedAt: now,
|
||||
|
||||
@@ -11,6 +11,7 @@ import * as descriptionGenerationRunsMigration from "./migrations/005_descriptio
|
||||
import * as sensitiveDataFlagMigration from "./migrations/006_sensitive_data_flag.js";
|
||||
import * as sensitiveDataSuggestionRunsMigration from "./migrations/007_sensitive_data_suggestion_runs.js";
|
||||
import * as catalogLogicalRelationshipsMigration from "./migrations/008_catalog_logical_relationships.js";
|
||||
import * as aiTokenUsageMigration from "./migrations/009_ai_token_usage.js";
|
||||
|
||||
const connectionString = process.env.THT_CATALOG_MIGRATOR_DATABASE_URL;
|
||||
const host = process.env.THT_CATALOG_DB_HOST;
|
||||
@@ -44,6 +45,7 @@ const provider: MigrationProvider = {
|
||||
"006_sensitive_data_flag": sensitiveDataFlagMigration,
|
||||
"007_sensitive_data_suggestion_runs": sensitiveDataSuggestionRunsMigration,
|
||||
"008_catalog_logical_relationships": catalogLogicalRelationshipsMigration,
|
||||
"009_ai_token_usage": aiTokenUsageMigration,
|
||||
};
|
||||
},
|
||||
};
|
||||
|
||||
@@ -0,0 +1,20 @@
|
||||
import type { Kysely } from "kysely";
|
||||
import type { CatalogDatabase } from "../repository.js";
|
||||
|
||||
export async function up(db: Kysely<CatalogDatabase>): Promise<void> {
|
||||
for (const table of ["description_generation_runs", "sensitive_data_suggestion_runs"] as const) {
|
||||
await db.schema.alterTable(table)
|
||||
.addColumn("input_tokens", "integer", (col) => col.notNull().defaultTo(0))
|
||||
.addColumn("cache_read_tokens", "integer", (col) => col.notNull().defaultTo(0))
|
||||
.addColumn("output_tokens", "integer", (col) => col.notNull().defaultTo(0))
|
||||
.execute();
|
||||
}
|
||||
}
|
||||
|
||||
export async function down(db: Kysely<CatalogDatabase>): Promise<void> {
|
||||
for (const table of ["sensitive_data_suggestion_runs", "description_generation_runs"] as const) {
|
||||
await db.schema.alterTable(table)
|
||||
.dropColumn("input_tokens").dropColumn("cache_read_tokens").dropColumn("output_tokens")
|
||||
.execute();
|
||||
}
|
||||
}
|
||||
@@ -4,7 +4,7 @@ import type { ResolvedMetadataGenerationModel } from "./metadata-generation-mode
|
||||
|
||||
const MAX_HELPER_OUTPUT_BYTES = 64 * 1024;
|
||||
const helperOutputSchema = z.discriminatedUnion("ok", [
|
||||
z.object({ ok: z.literal(true), content: z.string() }).strict(),
|
||||
z.object({ ok: z.literal(true), content: z.string(), usage: z.object({ input: z.number().int().nonnegative(), cacheRead: z.number().int().nonnegative(), output: z.number().int().nonnegative() }).strict().optional() }).strict(),
|
||||
z.object({ ok: z.literal(false), error: z.literal("provider_failure") }).strict(),
|
||||
]);
|
||||
|
||||
@@ -18,10 +18,12 @@ export interface ModelCompletionRequest {
|
||||
messages: readonly ModelCompletionMessage[];
|
||||
signal: AbortSignal;
|
||||
}
|
||||
export interface ModelCompletionUsage { input: number; cacheRead: number; output: number; }
|
||||
export interface ModelCompletionResult { content: string; usage: ModelCompletionUsage; }
|
||||
|
||||
/** The provider boundary used by Description Generation. */
|
||||
export interface ModelCompleter {
|
||||
complete(request: ModelCompletionRequest): Promise<string>;
|
||||
complete(request: ModelCompletionRequest): Promise<string | ModelCompletionResult>;
|
||||
}
|
||||
|
||||
export class ModelCompletionProviderError extends Error {
|
||||
@@ -47,7 +49,7 @@ export class PythonModelCompleter implements ModelCompleter {
|
||||
terminationGraceMs?: number;
|
||||
}) {}
|
||||
|
||||
async complete(request: ModelCompletionRequest): Promise<string> {
|
||||
async complete(request: ModelCompletionRequest): Promise<ModelCompletionResult> {
|
||||
if (request.signal.aborted) throw new ModelCompletionCancelledError();
|
||||
const payload = {
|
||||
model: `${request.model.provider}/${request.model.model}`,
|
||||
@@ -62,7 +64,7 @@ export class PythonModelCompleter implements ModelCompleter {
|
||||
...(request.model.disableThinking === true ? { disable_thinking: true } : {}),
|
||||
};
|
||||
|
||||
return await new Promise<string>((resolve, reject) => {
|
||||
return await new Promise<ModelCompletionResult>((resolve, reject) => {
|
||||
const child = spawn(
|
||||
this.options.pythonExecutable,
|
||||
["-m", this.options.helperModule ?? "tht.internal.litellm_completion"],
|
||||
@@ -135,7 +137,7 @@ export class PythonModelCompleter implements ModelCompleter {
|
||||
if (!output.ok) return fail();
|
||||
settled = true;
|
||||
cleanup();
|
||||
resolve(output.content);
|
||||
resolve({ content: output.content, usage: output.usage ?? { input: 0, cacheRead: 0, output: 0 } });
|
||||
} catch {
|
||||
fail();
|
||||
}
|
||||
|
||||
@@ -171,6 +171,9 @@ interface DescriptionGenerationRunTable {
|
||||
generated: number;
|
||||
nonGeneratable: number;
|
||||
failed: number;
|
||||
inputTokens: number;
|
||||
cacheReadTokens: number;
|
||||
outputTokens: number;
|
||||
createdAt: Timestamp;
|
||||
startedAt: Timestamp | null;
|
||||
updatedAt: Timestamp;
|
||||
@@ -195,6 +198,9 @@ interface SensitiveDataSuggestionRunTable {
|
||||
total: number;
|
||||
suggestedSensitive: number;
|
||||
suggestedNonSensitive: number;
|
||||
inputTokens: number;
|
||||
cacheReadTokens: number;
|
||||
outputTokens: number;
|
||||
createdAt: Timestamp;
|
||||
startedAt: Timestamp;
|
||||
updatedAt: Timestamp;
|
||||
@@ -816,6 +822,9 @@ export class KyselyCatalogRepository implements CatalogRepository {
|
||||
generated: 0,
|
||||
nonGeneratable: 0,
|
||||
failed: 0,
|
||||
inputTokens: 0,
|
||||
cacheReadTokens: 0,
|
||||
outputTokens: 0,
|
||||
startedAt: null,
|
||||
finishedAt: null,
|
||||
errorSummary: null,
|
||||
@@ -943,6 +952,9 @@ export class KyselyCatalogRepository implements CatalogRepository {
|
||||
total: 0,
|
||||
suggestedSensitive: 0,
|
||||
suggestedNonSensitive: 0,
|
||||
inputTokens: 0,
|
||||
cacheReadTokens: 0,
|
||||
outputTokens: 0,
|
||||
finishedAt: null,
|
||||
errorSummary: null,
|
||||
}).returningAll().executeTakeFirstOrThrow();
|
||||
|
||||
@@ -1,6 +1,6 @@
|
||||
import { z } from "zod";
|
||||
import type { MetadataGenerationModels } from "./metadata-generation-models.js";
|
||||
import type { ModelCompleter, ModelCompletionMessage } from "./model-completer.js";
|
||||
import type { ModelCompleter, ModelCompletionMessage, ModelCompletionResult, ModelCompletionUsage } from "./model-completer.js";
|
||||
import type {
|
||||
CatalogColumn,
|
||||
CatalogRepository,
|
||||
@@ -201,6 +201,8 @@ export class SensitiveDataSuggester {
|
||||
targetIds: readonly string[],
|
||||
signal: AbortSignal,
|
||||
onPrepared?: (total: number) => void | Promise<void>,
|
||||
onProgress?: (processed: number, suggestions: readonly SensitiveDataSuggestion[]) => void | Promise<void>,
|
||||
onUsage?: (usage: ModelCompletionUsage) => void | Promise<void>,
|
||||
): Promise<readonly SensitiveDataSuggestion[]> {
|
||||
const database = await this.repository.get(databaseId);
|
||||
if (!database) throw new SensitiveDataSuggestionTargetNotFoundError("database");
|
||||
@@ -212,11 +214,16 @@ export class SensitiveDataSuggester {
|
||||
for (const batch of batchesFor(database, columns)) {
|
||||
let received: Map<string, { columnId: string; sensitive: boolean }> | undefined;
|
||||
for (let attempt = 0; attempt < 2 && !received; attempt += 1) {
|
||||
const content = await this.completer.complete({
|
||||
const completion = await this.completer.complete({
|
||||
model,
|
||||
signal,
|
||||
messages: [systemMessage, { role: "user", content: userContent(database, batch) }],
|
||||
});
|
||||
const result: ModelCompletionResult = typeof completion === "string"
|
||||
? { content: completion, usage: { input: 0, cacheRead: 0, output: 0 } }
|
||||
: completion;
|
||||
await onUsage?.(result.usage);
|
||||
const content = result.content;
|
||||
try {
|
||||
const parsed = responseSchema.parse(JSON.parse(content));
|
||||
const expected = new Set(batch.map((column) => column.columnId));
|
||||
@@ -240,6 +247,7 @@ export class SensitiveDataSuggester {
|
||||
currentSensitive: column.currentSensitive,
|
||||
sensitive: received!.get(column.columnId)!.sensitive,
|
||||
})));
|
||||
await onProgress?.(suggestions.length, suggestions.slice(-batch.length));
|
||||
}
|
||||
return suggestions;
|
||||
}
|
||||
|
||||
@@ -10,6 +10,7 @@ import type {
|
||||
SensitiveDataSuggestionRun,
|
||||
SensitiveDataSuggestionScope,
|
||||
} from "./types.js";
|
||||
import type { ModelCompletionUsage } from "./model-completer.js";
|
||||
|
||||
const interruptedMessage = "Sensitive-field suggestion generation was interrupted by backend restart.";
|
||||
const failedMessage = "Sensitive-field suggestion generation failed.";
|
||||
@@ -73,6 +74,31 @@ export class SensitiveDataSuggestionRunner {
|
||||
});
|
||||
if (!prepared) throw new Error("Sensitive Data Suggestion Run disappeared");
|
||||
},
|
||||
async (processed, batch) => {
|
||||
const suggestedSensitive = batch.filter((suggestion) => suggestion.sensitive).length;
|
||||
const suggestedNonSensitive = batch.length - suggestedSensitive;
|
||||
const current = await this.repository.getSensitiveDataSuggestionRun(started.id);
|
||||
if (!current) throw new Error("Sensitive Data Suggestion Run disappeared");
|
||||
const progress = await this.repository.updateSensitiveDataSuggestionRun(started.id, {
|
||||
suggestedSensitive: current.suggestedSensitive + suggestedSensitive,
|
||||
suggestedNonSensitive: current.suggestedNonSensitive + suggestedNonSensitive,
|
||||
});
|
||||
if (!progress) throw new Error("Sensitive Data Suggestion Run disappeared");
|
||||
await this.repository.appendSensitiveDataSuggestionEvent(
|
||||
started.id,
|
||||
"info",
|
||||
`Classified ${processed} of ${progress.total} columns.`,
|
||||
);
|
||||
},
|
||||
async (usage: ModelCompletionUsage) => {
|
||||
const current = await this.repository.getSensitiveDataSuggestionRun(started.id);
|
||||
if (!current) throw new Error("Sensitive Data Suggestion Run disappeared");
|
||||
await this.repository.updateSensitiveDataSuggestionRun(started.id, {
|
||||
inputTokens: current.inputTokens + usage.input,
|
||||
cacheReadTokens: current.cacheReadTokens + usage.cacheRead,
|
||||
outputTokens: current.outputTokens + usage.output,
|
||||
});
|
||||
},
|
||||
);
|
||||
const suggestedSensitive = suggestions.filter((suggestion) => suggestion.sensitive).length;
|
||||
const suggestedNonSensitive = suggestions.length - suggestedSensitive;
|
||||
|
||||
@@ -234,6 +234,9 @@ export interface DescriptionGenerationRun {
|
||||
generated: number;
|
||||
nonGeneratable: number;
|
||||
failed: number;
|
||||
inputTokens: number;
|
||||
cacheReadTokens: number;
|
||||
outputTokens: number;
|
||||
createdAt: string;
|
||||
startedAt: string | null;
|
||||
updatedAt: string;
|
||||
@@ -250,6 +253,9 @@ export interface DescriptionGenerationRunUpdate {
|
||||
startedAt?: string | null;
|
||||
finishedAt?: string | null;
|
||||
errorSummary?: string | null;
|
||||
inputTokens?: number;
|
||||
cacheReadTokens?: number;
|
||||
outputTokens?: number;
|
||||
}
|
||||
|
||||
export interface DescriptionGenerationEvent {
|
||||
@@ -272,6 +278,9 @@ export interface SensitiveDataSuggestionRun {
|
||||
total: number;
|
||||
suggestedSensitive: number;
|
||||
suggestedNonSensitive: number;
|
||||
inputTokens: number;
|
||||
cacheReadTokens: number;
|
||||
outputTokens: number;
|
||||
createdAt: string;
|
||||
startedAt: string;
|
||||
updatedAt: string;
|
||||
@@ -286,6 +295,9 @@ export interface SensitiveDataSuggestionRunUpdate {
|
||||
suggestedNonSensitive?: number;
|
||||
finishedAt?: string | null;
|
||||
errorSummary?: string | null;
|
||||
inputTokens?: number;
|
||||
cacheReadTokens?: number;
|
||||
outputTokens?: number;
|
||||
}
|
||||
|
||||
export interface SensitiveDataSuggestionEvent {
|
||||
|
||||
@@ -101,6 +101,9 @@ function publicRun(run: DescriptionGenerationRun) {
|
||||
generated: run.generated,
|
||||
nonGeneratable: run.nonGeneratable,
|
||||
failed: run.failed,
|
||||
inputTokens: run.inputTokens,
|
||||
cacheReadTokens: run.cacheReadTokens,
|
||||
outputTokens: run.outputTokens,
|
||||
createdAt: run.createdAt,
|
||||
startedAt: run.startedAt,
|
||||
updatedAt: run.updatedAt,
|
||||
@@ -129,6 +132,9 @@ function publicSensitiveDataSuggestionRun(run: SensitiveDataSuggestionRun) {
|
||||
total: run.total,
|
||||
suggestedSensitive: run.suggestedSensitive,
|
||||
suggestedNonSensitive: run.suggestedNonSensitive,
|
||||
inputTokens: run.inputTokens,
|
||||
cacheReadTokens: run.cacheReadTokens,
|
||||
outputTokens: run.outputTokens,
|
||||
createdAt: run.createdAt,
|
||||
startedAt: run.startedAt,
|
||||
updatedAt: run.updatedAt,
|
||||
|
||||
@@ -218,6 +218,12 @@ test("suggests sensitive flags from structural metadata without persisting them"
|
||||
runId: responseBody.run.id,
|
||||
sequence: 2,
|
||||
level: "info",
|
||||
message: "Classified 1 of 1 columns.",
|
||||
},
|
||||
{
|
||||
runId: responseBody.run.id,
|
||||
sequence: 3,
|
||||
level: "info",
|
||||
message: "Sensitive-field suggestion generation completed for 1 column.",
|
||||
},
|
||||
]);
|
||||
@@ -671,6 +677,7 @@ test("generates one selected Catalog Column from a single JSON code fence", asyn
|
||||
errorSummary: null,
|
||||
});
|
||||
expect(Object.keys(start.json()).sort()).toEqual([
|
||||
"cacheReadTokens",
|
||||
"createdAt",
|
||||
"databaseId",
|
||||
"errorSummary",
|
||||
@@ -678,9 +685,11 @@ test("generates one selected Catalog Column from a single JSON code fence", asyn
|
||||
"finishedAt",
|
||||
"generated",
|
||||
"id",
|
||||
"inputTokens",
|
||||
"language",
|
||||
"modelId",
|
||||
"nonGeneratable",
|
||||
"outputTokens",
|
||||
"processed",
|
||||
"scope",
|
||||
"startedAt",
|
||||
@@ -1272,7 +1281,7 @@ test("Stop aborts source sampling before any model request", async () => {
|
||||
test("an isolated exhausted technical batch failure allows completion with errors", async () => {
|
||||
const modelCompleter: ModelCompleter = {
|
||||
complete: vi.fn(async (request) => {
|
||||
if (vi.mocked(modelCompleter.complete).mock.calls.length === 1) {
|
||||
if (vi.mocked(modelCompleter.complete).mock.calls.length <= 2) {
|
||||
throw new ModelCompletionProviderError();
|
||||
}
|
||||
const context = JSON.parse(request.messages[1]!.content.split("\n").slice(1).join("\n"));
|
||||
@@ -1320,7 +1329,7 @@ test("an isolated exhausted technical batch failure allows completion with error
|
||||
failed: 10,
|
||||
errorSummary: "Description generation completed with errors.",
|
||||
});
|
||||
expect(modelCompleter.complete).toHaveBeenCalledTimes(2);
|
||||
expect(modelCompleter.complete).toHaveBeenCalledTimes(3);
|
||||
const updated = new Map(
|
||||
(await repository.listColumns(database.id, table.id)).map((column) => [column.id, column]),
|
||||
);
|
||||
@@ -1346,10 +1355,10 @@ test("success resets the technical-failure streak and the third later failure st
|
||||
const modelCompleter: ModelCompleter = {
|
||||
complete: vi.fn(async (request) => {
|
||||
const call = vi.mocked(modelCompleter.complete).mock.calls.length;
|
||||
if ([1, 2, 4, 5, 6].includes(call)) {
|
||||
if ([1, 2, 4, 5, 6, 7, 8, 9].includes(call)) {
|
||||
throw Object.assign(new ModelCompletionProviderError(), { message: sensitiveDiagnostic });
|
||||
}
|
||||
if (call > 6) throw new Error("a later batch must not start");
|
||||
if (call > 9) throw new Error("a later batch must not start");
|
||||
const context = JSON.parse(request.messages[1]!.content.split("\n").slice(1).join("\n"));
|
||||
return JSON.stringify({
|
||||
results: context.targets.map((target: { targetId: string }) => ({
|
||||
@@ -1389,25 +1398,25 @@ test("success resets the technical-failure streak and the third later failure st
|
||||
expect(run).toMatchObject({
|
||||
status: "failed",
|
||||
total: 61,
|
||||
processed: 60,
|
||||
processed: 50,
|
||||
generated: 10,
|
||||
nonGeneratable: 0,
|
||||
failed: 50,
|
||||
failed: 40,
|
||||
errorSummary: "Description generation stopped after three consecutive technical batch failures.",
|
||||
});
|
||||
expect(modelCompleter.complete).toHaveBeenCalledTimes(6);
|
||||
expect(modelCompleter.complete).toHaveBeenCalledTimes(9);
|
||||
const requests = vi.mocked(modelCompleter.complete).mock.calls.map(([request]) => request);
|
||||
expect(new Set(requests.map((request) => request.signal)).size).toBe(1);
|
||||
expect(new Set(requests.map((request) => request.model.id))).toEqual(new Set([configuredModel.id]));
|
||||
const updated = new Map(
|
||||
(await repository.listColumns(database.id, table.id)).map((column) => [column.id, column]),
|
||||
);
|
||||
targetIds.slice(20, 30).forEach((targetId) => {
|
||||
targetIds.slice(10, 20).forEach((targetId) => {
|
||||
expect(updated.get(targetId)?.generatedDescription).toBe("Successful reset batch.");
|
||||
});
|
||||
expect(updated.get(targetIds[60]!)?.generatedDescription).toBeNull();
|
||||
expect(updated.get(targetIds[50]!)?.generatedDescription).toBeNull();
|
||||
const events = await repository.listDescriptionGenerationEvents(run.id);
|
||||
expect(events.filter((event) => event.message.includes("model provider request failed"))).toHaveLength(5);
|
||||
expect(events.filter((event) => event.message.includes("model provider request failed"))).toHaveLength(8);
|
||||
expect(events.at(-1)).toEqual(expect.objectContaining({
|
||||
level: "error",
|
||||
message: "Description generation stopped after three consecutive technical batch failures.",
|
||||
@@ -1660,6 +1669,7 @@ test("Description Generation history is newest-first, bounded, and exposes only
|
||||
expect(response.statusCode).toBe(200);
|
||||
expect(response.json().map((run: { id: string }) => run.id)).toEqual(ids.slice(1).reverse());
|
||||
expect(Object.keys(response.json()[0]).sort()).toEqual([
|
||||
"cacheReadTokens",
|
||||
"createdAt",
|
||||
"databaseId",
|
||||
"errorSummary",
|
||||
@@ -1667,9 +1677,11 @@ test("Description Generation history is newest-first, bounded, and exposes only
|
||||
"finishedAt",
|
||||
"generated",
|
||||
"id",
|
||||
"inputTokens",
|
||||
"language",
|
||||
"modelId",
|
||||
"nonGeneratable",
|
||||
"outputTokens",
|
||||
"processed",
|
||||
"scope",
|
||||
"startedAt",
|
||||
@@ -1945,7 +1957,7 @@ test("retains completed batch writes when a later batch response is malformed",
|
||||
});
|
||||
const { run } = await waitForTerminalRun(app, start.json().id);
|
||||
|
||||
expect(modelCompleter.complete).toHaveBeenCalledTimes(2);
|
||||
expect(modelCompleter.complete).toHaveBeenCalledTimes(3);
|
||||
expect(run).toMatchObject({
|
||||
status: "completed_with_errors",
|
||||
total: 11,
|
||||
@@ -1972,9 +1984,14 @@ test("retains completed batch writes when a later batch response is malformed",
|
||||
expect(events.slice(2, 12).map((event) => event.message)).toEqual(
|
||||
orderedIds.slice(0, 10).map((targetId) => `Generated description for Catalog Column ${targetId}.`),
|
||||
);
|
||||
expect(events.find((event) => event.level === "warning" && event.message.includes("Retrying batch"))).toEqual(
|
||||
expect.objectContaining({
|
||||
message: "The model response did not match the required schema. Retrying batch (attempt 2 of 2).",
|
||||
}),
|
||||
);
|
||||
expect(events.find((event) => event.level === "error")).toEqual(expect.objectContaining({
|
||||
level: "error",
|
||||
message: `The model response was invalid. Affected Catalog Column target: ${orderedIds[10]}.`,
|
||||
message: `The model response did not match the required schema. Affected Catalog Column target: ${orderedIds[10]}.`,
|
||||
}));
|
||||
} finally {
|
||||
await app.close();
|
||||
@@ -2186,7 +2203,7 @@ test("Generate Missing skips prior partial results and includes null, empty, and
|
||||
const metadata = JSON.parse(request.messages[1]!.content.split("\n").slice(1).join("\n"));
|
||||
if (mode === "partial") {
|
||||
partialCall += 1;
|
||||
if (partialCall === 2) throw new ModelCompletionProviderError();
|
||||
if (partialCall === 2 || partialCall === 3) throw new ModelCompletionProviderError();
|
||||
return JSON.stringify({
|
||||
results: metadata.targets.map((target: { targetId: string }) => ({
|
||||
targetId: target.targetId,
|
||||
@@ -2512,6 +2529,58 @@ test("localizes valid non-generatable Catalog Column results in English", async
|
||||
}
|
||||
});
|
||||
|
||||
test("retries invalid JSON once and completes the batch when the second response is valid", async () => {
|
||||
const modelCompleter: ModelCompleter = {
|
||||
complete: vi.fn(async (request) => {
|
||||
if (vi.mocked(modelCompleter.complete).mock.calls.length === 1) {
|
||||
return { content: "not-json", usage: { input: 11, cacheRead: 3, output: 2 } };
|
||||
}
|
||||
const context = JSON.parse(request.messages[1]!.content.split("\n").slice(1).join("\n"));
|
||||
return {
|
||||
content: JSON.stringify({
|
||||
results: context.targets.map((target: { targetId: string }) => ({
|
||||
targetId: target.targetId,
|
||||
outcome: "generated",
|
||||
description: "Generated after the application retry.",
|
||||
})),
|
||||
}),
|
||||
usage: { input: 7, cacheRead: 1, output: 5 },
|
||||
};
|
||||
}),
|
||||
};
|
||||
const { app, repository, database, table, column } = await setup(modelCompleter);
|
||||
try {
|
||||
const start = await app.inject({
|
||||
method: "POST",
|
||||
url: `/catalog/databases/${database.id}/description-generation-runs`,
|
||||
payload: { modelId: configuredModel.id, scope: "selected_columns", targetIds: [column.id] },
|
||||
});
|
||||
const { run } = await waitForTerminalRun(app, start.json().id);
|
||||
|
||||
expect(run).toMatchObject({
|
||||
status: "completed",
|
||||
processed: 1,
|
||||
generated: 1,
|
||||
failed: 0,
|
||||
inputTokens: 18,
|
||||
cacheReadTokens: 4,
|
||||
outputTokens: 7,
|
||||
});
|
||||
expect(modelCompleter.complete).toHaveBeenCalledTimes(2);
|
||||
expect(await repository.getColumn(database.id, table.id, column.id)).toMatchObject({
|
||||
generatedDescription: "Generated after the application retry.",
|
||||
});
|
||||
expect(await repository.listDescriptionGenerationEvents(run.id)).toEqual(expect.arrayContaining([
|
||||
expect.objectContaining({
|
||||
level: "warning",
|
||||
message: "The model response was not valid JSON. Retrying batch (attempt 2 of 2).",
|
||||
}),
|
||||
]));
|
||||
} finally {
|
||||
await app.close();
|
||||
}
|
||||
});
|
||||
|
||||
test("fails safely when the provider fails and redacts provider diagnostics", async () => {
|
||||
const sensitiveDiagnostic = "test-provider-secret private prompt raw provider payload";
|
||||
const modelCompleter: ModelCompleter = {
|
||||
@@ -2528,6 +2597,7 @@ test("fails safely when the provider fails and redacts provider diagnostics", as
|
||||
});
|
||||
const { run } = await waitForTerminalRun(app, start.json().id);
|
||||
|
||||
expect(modelCompleter.complete).toHaveBeenCalledTimes(2);
|
||||
expect(run).toMatchObject({
|
||||
status: "completed_with_errors",
|
||||
processed: 1,
|
||||
@@ -2565,7 +2635,7 @@ test.each([
|
||||
["duplicate mappings", (targetIds: readonly string[]) => ({ results: [
|
||||
{ targetId: targetIds[0], outcome: "generated", description: "First valid value" },
|
||||
{ targetId: targetIds[0], outcome: "non_generatable" },
|
||||
] })],
|
||||
] }), "The model response was missing one or more requested targets."],
|
||||
["unknown mappings", (targetIds: readonly string[]) => ({ results: [
|
||||
{ targetId: targetIds[0], outcome: "non_generatable" },
|
||||
{
|
||||
@@ -2573,15 +2643,15 @@ test.each([
|
||||
outcome: "generated",
|
||||
description: "Unknown target value",
|
||||
},
|
||||
] })],
|
||||
] }), "The model response was missing one or more requested targets."],
|
||||
["missing mappings", (targetIds: readonly string[]) => ({ results: [
|
||||
{ targetId: targetIds[0], outcome: "generated", description: "Only one result" },
|
||||
] })],
|
||||
] }), "The model response was missing one or more requested targets."],
|
||||
["malformed mappings", (targetIds: readonly string[]) => ({ results: [
|
||||
{ targetId: targetIds[0], outcome: "generated", description: "First valid value" },
|
||||
{ targetId: targetIds[1], outcome: "generated", description: " " },
|
||||
] })],
|
||||
] as const)("rejects %s without applying any result from the batch", async (_name, responseFor) => {
|
||||
] }), "The model response did not match the required schema."],
|
||||
] as const)("rejects %s without applying any result from the batch", async (_name, responseFor, failureMessage) => {
|
||||
let selectedColumnIds: string[] = [];
|
||||
const modelCompleter: ModelCompleter = {
|
||||
complete: vi.fn(async () => JSON.stringify(responseFor(selectedColumnIds))),
|
||||
@@ -2625,6 +2695,7 @@ test.each([
|
||||
});
|
||||
const { run } = await waitForTerminalRun(app, start.json().id);
|
||||
|
||||
expect(modelCompleter.complete).toHaveBeenCalledTimes(2);
|
||||
expect(run).toMatchObject({
|
||||
status: "completed_with_errors",
|
||||
processed: 2,
|
||||
@@ -2644,7 +2715,7 @@ test.each([
|
||||
expect((await repository.listDescriptionGenerationEvents(run.id)).find((event) => event.level === "error")).toEqual(
|
||||
expect.objectContaining({
|
||||
level: "error",
|
||||
message: `The model response was invalid. Affected Catalog Column targets: ${selectedColumnIds.join(", ")}.`,
|
||||
message: `${failureMessage} Affected Catalog Column targets: ${selectedColumnIds.join(", ")}.`,
|
||||
}),
|
||||
);
|
||||
} finally {
|
||||
|
||||
@@ -13,6 +13,7 @@ import { up as upSchemaSync } from "../src/catalog/migrations/003_catalog_schema
|
||||
import { up as upDescriptionGeneration } from "../src/catalog/migrations/005_description_generation_runs.js";
|
||||
import { up as upSensitiveDataFlag } from "../src/catalog/migrations/006_sensitive_data_flag.js";
|
||||
import { up as upSensitiveSuggestionRuns } from "../src/catalog/migrations/007_sensitive_data_suggestion_runs.js";
|
||||
import { up as upAiTokenUsage } from "../src/catalog/migrations/009_ai_token_usage.js";
|
||||
import { KyselyCatalogRepository, type CatalogDatabase } from "../src/catalog/repository.js";
|
||||
import { loadConfig } from "../src/config.js";
|
||||
import type { WorkspaceRegistry } from "../src/workspaces/registry.js";
|
||||
@@ -48,6 +49,7 @@ test.skipIf(!dockerAvailable)("Fastify persists Description Generation success a
|
||||
await upSensitiveDataFlag(db);
|
||||
await upDescriptionGeneration(db);
|
||||
await upSensitiveSuggestionRuns(db);
|
||||
await upAiTokenUsage(db);
|
||||
const repository = new KyselyCatalogRepository(db);
|
||||
const database = await repository.create({
|
||||
workspaceId: "psd-clinical",
|
||||
@@ -124,7 +126,7 @@ test.skipIf(!dockerAvailable)("Fastify persists Description Generation success a
|
||||
description: "Elenco dei pazienti e dei loro dati clinici.",
|
||||
}] });
|
||||
}
|
||||
if (call === 4) {
|
||||
if (call === 5) {
|
||||
return JSON.stringify({ results: [
|
||||
{
|
||||
targetId: birthDate.id,
|
||||
@@ -138,14 +140,14 @@ test.skipIf(!dockerAvailable)("Fastify persists Description Generation success a
|
||||
},
|
||||
] });
|
||||
}
|
||||
if (call === 5) {
|
||||
if (call === 6) {
|
||||
return JSON.stringify({ results: [{
|
||||
targetId: table.id,
|
||||
outcome: "generated",
|
||||
description: "Descrizione rigenerata della tabella pazienti.",
|
||||
}] });
|
||||
}
|
||||
if (call === 6) {
|
||||
if (call === 7) {
|
||||
return JSON.stringify({ results: [{
|
||||
targetId: birthDate.id,
|
||||
outcome: "generated",
|
||||
@@ -340,7 +342,7 @@ test.skipIf(!dockerAvailable)("Fastify persists Description Generation success a
|
||||
expect(await repository.getColumn(database.id, table.id, birthDate.id)).toMatchObject({
|
||||
generatedDescription: "Descrizione recuperata della data di nascita.",
|
||||
});
|
||||
expect(modelCompleter.complete).toHaveBeenCalledTimes(6);
|
||||
expect(modelCompleter.complete).toHaveBeenCalledTimes(7);
|
||||
expect(JSON.stringify(vi.mocked(modelCompleter.complete).mock.calls)).toContain(persistedSampleSecret);
|
||||
|
||||
const runIds = [
|
||||
|
||||
@@ -1,6 +1,9 @@
|
||||
import { expect, test, vi } from "vitest";
|
||||
import { mkdtempSync, rmSync, writeFileSync } from "node:fs";
|
||||
import { tmpdir } from "node:os";
|
||||
import { join } from "node:path";
|
||||
import {
|
||||
PostgresDescriptionSourceSampler,
|
||||
ConcreteDescriptionSourceSampler,
|
||||
type DescriptionSourceSamplingTarget,
|
||||
} from "../src/catalog/description-source-sampler.js";
|
||||
import type {
|
||||
@@ -8,6 +11,8 @@ import type {
|
||||
CatalogPostgresAccess,
|
||||
} from "../src/catalog/postgres-access.js";
|
||||
import type { WorkspaceDatabase } from "../src/catalog/types.js";
|
||||
import type { WorkspaceSecretStore } from "../src/workspaces/secret-store.js";
|
||||
import { CATALOG_SECRET_IDS } from "../src/catalog/secrets.js";
|
||||
|
||||
const database: WorkspaceDatabase = {
|
||||
id: "11111111-1111-4111-8111-111111111111",
|
||||
@@ -51,7 +56,7 @@ test("samples at most five source rows and five distinct non-null examples in a
|
||||
const access: CatalogPostgresAccess = {
|
||||
connect: vi.fn(async () => ({ query, end }) as CatalogDatabaseClient),
|
||||
};
|
||||
const sampler = new PostgresDescriptionSourceSampler(access);
|
||||
const sampler = new ConcreteDescriptionSourceSampler(access);
|
||||
const controller = new AbortController();
|
||||
|
||||
const samples = await sampler.sample(database, [target], controller.signal);
|
||||
@@ -92,13 +97,71 @@ test("samples at most five source rows and five distinct non-null examples in a
|
||||
expect(end).toHaveBeenCalledOnce();
|
||||
});
|
||||
|
||||
test("samples source rows through the configured REST run_query binding", async () => {
|
||||
const root = mkdtempSync(join(tmpdir(), "tht-source-rest-"));
|
||||
const credentialFile = join(root, "api-key");
|
||||
writeFileSync(credentialFile, "test-api-key\n", { mode: 0o600 });
|
||||
const release = vi.fn();
|
||||
const secretStore = {
|
||||
materialize: vi.fn(() => ({
|
||||
files: new Map([[CATALOG_SECRET_IDS.apiKey, credentialFile]]),
|
||||
release,
|
||||
})),
|
||||
} as unknown as WorkspaceSecretStore;
|
||||
const fetchMock = vi.fn(async () => new Response(JSON.stringify([
|
||||
{ 'status"code': "active", ward: null },
|
||||
{ 'status"code': "pending", ward: "A" },
|
||||
]), { status: 200, headers: { "content-type": "application/json" } }));
|
||||
vi.stubGlobal("fetch", fetchMock);
|
||||
const access: CatalogPostgresAccess = {
|
||||
connect: vi.fn(async () => { throw new Error("PostgreSQL access must not be used"); }),
|
||||
};
|
||||
const sampler = new ConcreteDescriptionSourceSampler(access, secretStore);
|
||||
const restDatabase: WorkspaceDatabase = {
|
||||
...database,
|
||||
binding: {
|
||||
transport: "rest_api",
|
||||
baseUrl: "https://dwh.example.test/root/",
|
||||
restPath: "/health",
|
||||
restAuth: "x-api-key",
|
||||
},
|
||||
};
|
||||
|
||||
try {
|
||||
await expect(sampler.sample(restDatabase, [target], new AbortController().signal)).resolves.toEqual([{
|
||||
targetId: target.targetId,
|
||||
tableName: target.tableName,
|
||||
rows: [
|
||||
{ fields: [{ name: 'status"code', value: "active" }, { name: "ward", value: null }] },
|
||||
{ fields: [{ name: 'status"code', value: "pending" }, { name: "ward", value: "A" }] },
|
||||
],
|
||||
representativeValues: [
|
||||
{ column: 'status"code', values: ["active", "pending"] },
|
||||
{ column: "ward", values: ["A"] },
|
||||
],
|
||||
}]);
|
||||
expect(access.connect).not.toHaveBeenCalled();
|
||||
expect(fetchMock).toHaveBeenCalledWith("https://dwh.example.test/root/rpc/run_query", expect.objectContaining({
|
||||
method: "POST",
|
||||
headers: { "content-type": "application/json", "x-api-key": "test-api-key" },
|
||||
body: JSON.stringify({
|
||||
query_text: 'SELECT LEFT(("status""code")::text, 256) AS "status""code", LEFT(("ward")::text, 256) AS "ward" FROM "clinical""data"."patient""facts" LIMIT 5',
|
||||
}),
|
||||
}));
|
||||
expect(release).toHaveBeenCalledOnce();
|
||||
} finally {
|
||||
vi.unstubAllGlobals();
|
||||
rmSync(root, { recursive: true, force: true });
|
||||
}
|
||||
});
|
||||
|
||||
test("does not issue a SELECT when a protected target has no source columns", async () => {
|
||||
const query = vi.fn(async () => ({ rows: [] }));
|
||||
const end = vi.fn(async () => undefined);
|
||||
const access: CatalogPostgresAccess = {
|
||||
connect: vi.fn(async () => ({ query, end }) as CatalogDatabaseClient),
|
||||
};
|
||||
const sampler = new PostgresDescriptionSourceSampler(access);
|
||||
const sampler = new ConcreteDescriptionSourceSampler(access);
|
||||
|
||||
const samples = await sampler.sample(database, [{
|
||||
targetId: target.targetId,
|
||||
@@ -129,7 +192,7 @@ test("rolls back and closes the source connection when sampling fails", async ()
|
||||
const access: CatalogPostgresAccess = {
|
||||
connect: vi.fn(async () => ({ query, end }) as CatalogDatabaseClient),
|
||||
};
|
||||
const sampler = new PostgresDescriptionSourceSampler(access);
|
||||
const sampler = new ConcreteDescriptionSourceSampler(access);
|
||||
const controller = new AbortController();
|
||||
|
||||
await expect(sampler.sample(database, [target], controller.signal)).rejects.toThrow();
|
||||
|
||||
@@ -13,6 +13,7 @@ import { up as upDescriptionGeneration } from "../src/catalog/migrations/005_des
|
||||
import { up as upSensitiveDataFlag } from "../src/catalog/migrations/006_sensitive_data_flag.js";
|
||||
import { up as upSensitiveSuggestionRuns } from "../src/catalog/migrations/007_sensitive_data_suggestion_runs.js";
|
||||
import { up as upLogicalRelationships } from "../src/catalog/migrations/008_catalog_logical_relationships.js";
|
||||
import { up as upAiTokenUsage } from "../src/catalog/migrations/009_ai_token_usage.js";
|
||||
|
||||
const dockerAvailable = spawnSync("docker", ["info"], { stdio: "ignore" }).status === 0;
|
||||
|
||||
@@ -28,6 +29,9 @@ test.skipIf(!dockerAvailable)("PostgreSQL migration enforces one database per wo
|
||||
await upSchemaSync(db);
|
||||
await upSensitiveDataFlag(db);
|
||||
await upLogicalRelationships(db);
|
||||
await upDescriptionGeneration(db);
|
||||
await upSensitiveSuggestionRuns(db);
|
||||
await upAiTokenUsage(db);
|
||||
await sql`CREATE ROLE thothii_catalog_runtime`.execute(db);
|
||||
await upRuntimeSequencePrivileges(db);
|
||||
const sequencePrivilege = await sql<{ allowed: boolean }>`
|
||||
@@ -386,6 +390,7 @@ test.skipIf(!dockerAvailable)("PostgreSQL repository persists description and se
|
||||
await upLogicalRelationships(db);
|
||||
await upDescriptionGeneration(db);
|
||||
await upSensitiveSuggestionRuns(db);
|
||||
await upAiTokenUsage(db);
|
||||
const repository = new KyselyCatalogRepository(db);
|
||||
const firstDatabase = await repository.create({
|
||||
workspaceId: "generation-one",
|
||||
|
||||
@@ -127,7 +127,7 @@ test("resolves only a configured selection for the later generation boundary", (
|
||||
expect(() => models.resolve("unknown-model")).toThrow(MetadataGenerationModelUnavailableError);
|
||||
});
|
||||
|
||||
test("loads DeepSeek, GLM, and an explicit keyless Qwen endpoint from installation setup", () => {
|
||||
test("loads DeepSeek models, GLM, and an explicit keyless Qwen endpoint from installation setup", () => {
|
||||
const { installationFile, secretsFile } = metadataConfiguration(`metadataGeneration:
|
||||
default: glm-53
|
||||
models:
|
||||
@@ -135,6 +135,10 @@ test("loads DeepSeek, GLM, and an explicit keyless Qwen endpoint from installati
|
||||
label: DeepSeek V4 Pro
|
||||
litellm: {provider: deepseek, model: deepseek-v4-pro}
|
||||
apiKeyEnv: DEEPSEEK_API_KEY
|
||||
- id: deepseek-v4-flash
|
||||
label: DeepSeek V4 Flash
|
||||
litellm: {provider: deepseek, model: deepseek-v4-flash}
|
||||
apiKeyEnv: DEEPSEEK_API_KEY
|
||||
- id: glm-53
|
||||
label: GLM 5.3
|
||||
litellm:
|
||||
@@ -156,6 +160,7 @@ test("loads DeepSeek, GLM, and an explicit keyless Qwen endpoint from installati
|
||||
expect(models.catalog()).toEqual({
|
||||
models: [
|
||||
{ id: "deepseek-v4-pro", label: "DeepSeek V4 Pro" },
|
||||
{ id: "deepseek-v4-flash", label: "DeepSeek V4 Flash" },
|
||||
{ id: "glm-53", label: "GLM 5.3" },
|
||||
{ id: "qwen-36", label: "Qwen 3.6" },
|
||||
],
|
||||
|
||||
@@ -64,7 +64,7 @@ sys.stdout.write(json.dumps({"ok": True, "content": "Descrizione italiana"}))
|
||||
signal: new AbortController().signal,
|
||||
});
|
||||
|
||||
expect(content).toBe("Descrizione italiana");
|
||||
expect(content).toEqual({ content: "Descrizione italiana", usage: { input: 0, cacheRead: 0, output: 0 } });
|
||||
const captured = JSON.parse(readFileSync(join(roots[0]!, "request.json"), "utf8"));
|
||||
expect(captured.request).toEqual({
|
||||
model: "openai/gpt-4.1-mini",
|
||||
@@ -100,7 +100,7 @@ sys.stdout.write(json.dumps({"ok": True, "content": "Descrizione Qwen"}))
|
||||
},
|
||||
messages: [{ role: "user", content: "Describe invented metadata." }],
|
||||
signal: new AbortController().signal,
|
||||
})).resolves.toBe("Descrizione Qwen");
|
||||
})).resolves.toEqual({ content: "Descrizione Qwen", usage: { input: 0, cacheRead: 0, output: 0 } });
|
||||
|
||||
expect(JSON.parse(readFileSync(join(roots[0]!, "request.json"), "utf8"))).toEqual({
|
||||
model: "openai/qwen3.6-35b-a3b",
|
||||
|
||||
Reference in New Issue
Block a user