diff --git a/frontend/src/components/DataViewer.primary-key.test.tsx b/frontend/src/components/DataViewer.primary-key.test.tsx index ae8d0315..c9e1db07 100644 --- a/frontend/src/components/DataViewer.primary-key.test.tsx +++ b/frontend/src/components/DataViewer.primary-key.test.tsx @@ -165,6 +165,52 @@ describe('DataViewer safe editing locator', () => { renderer!.unmount(); }); + it('does not block the initial Kingbase table query on edit-locator metadata', async () => { + storeState.connections[0].config.type = 'kingbase'; + storeState.connections[0].config.database = 'ldf_server_dbs_dev'; + + let resolveColumns!: (value: any) => void; + let resolveIndexes!: (value: any) => void; + backendApp.DBGetColumns.mockReturnValue(new Promise((resolve) => { + resolveColumns = resolve; + })); + backendApp.DBGetIndexes.mockReturnValue(new Promise((resolve) => { + resolveIndexes = resolve; + })); + + let renderer: ReactTestRenderer; + await act(async () => { + renderer = create(); + }); + await flushPromises(); + + expect(backendApp.DBQuery).toHaveBeenCalled(); + expect(dataGridState.latestProps?.data).toEqual([ + expect.objectContaining({ ID: 7, NAME: 'old-name' }), + ]); + + await act(async () => { + resolveColumns({ + success: true, + data: [{ name: 'ID', key: 'PRI' }, { name: 'NAME', key: '' }], + }); + resolveIndexes({ success: true, data: [] }); + }); + await flushPromises(); + + expect(dataGridState.latestProps?.editLocator).toMatchObject({ + strategy: 'primary-key', + columns: ['ID'], + readOnly: false, + }); + renderer!.unmount(); + }); + it('enables table preview editing after primary keys are loaded', async () => { backendApp.DBGetColumns.mockResolvedValue({ success: true, diff --git a/frontend/src/components/DataViewer.tsx b/frontend/src/components/DataViewer.tsx index bb7de2c0..ba3b992a 100644 --- a/frontend/src/components/DataViewer.tsx +++ b/frontend/src/components/DataViewer.tsx @@ -644,24 +644,28 @@ const DataViewer: React.FC<{ tab: TabData; isActive?: boolean }> = React.memo(({ if (pkKeyRef.current !== locatorKey || !editLocatorForQuery) { pkKeyRef.current = locatorKey; const locatorSeq = ++pkSeqRef.current; - try { - const [resCols, resIndexes] = await Promise.all([ - DBGetColumns(buildRpcConnectionConfig(config) as any, dbName, tableName), - DBGetIndexes(buildRpcConnectionConfig(config) as any, dbName, tableName) - .catch((error: any) => ({ success: false, message: String(error?.message || error || 'Failed to load indexes'), data: [] })), - ]); - if (fetchSeqRef.current !== seq) return; - if (pkSeqRef.current !== locatorSeq) return; - if (pkKeyRef.current !== locatorKey) return; + const loadEditLocator = async (): Promise<{ + primaryKeys: string[]; + locator: EditRowLocator; + } | null> => { + try { + const [resCols, resIndexes] = await Promise.all([ + DBGetColumns(buildRpcConnectionConfig(config) as any, dbName, tableName), + DBGetIndexes(buildRpcConnectionConfig(config) as any, dbName, tableName) + .catch((error: any) => ({ success: false, message: String(error?.message || error || 'Failed to load indexes'), data: [] })), + ]); + if (fetchSeqRef.current !== seq) return null; + if (pkSeqRef.current !== locatorSeq) return null; + if (pkKeyRef.current !== locatorKey) return null; + + if (!resCols?.success || !Array.isArray(resCols.data)) { + const nextLocator = buildAllColumnsLocator([], { translate: tr }); + setPkColumns([]); + setEditLocator(nextLocator); + if (nextLocator.reason) message.info(nextLocator.reason); + return { primaryKeys: [], locator: nextLocator }; + } - if (!resCols?.success || !Array.isArray(resCols.data)) { - const nextLocator = buildAllColumnsLocator([], { translate: tr }); - pkColumnsForQuery = []; - editLocatorForQuery = nextLocator; - setPkColumns([]); - setEditLocator(nextLocator); - if (nextLocator.reason) message.info(nextLocator.reason); - } else { const columnDefs = resCols.data as ColumnDefinition[]; const primaryKeys = columnDefs .filter((column: any) => getColumnDefinitionKey(column) === 'PRI') @@ -686,8 +690,6 @@ const DataViewer: React.FC<{ tab: TabData; isActive?: boolean }> = React.memo(({ translate: tr, }), tr); - pkColumnsForQuery = primaryKeys; - editLocatorForQuery = nextLocator; setPkColumns(primaryKeys); setEditLocator(nextLocator); if (nextLocator.readOnly) { @@ -695,17 +697,28 @@ const DataViewer: React.FC<{ tab: TabData; isActive?: boolean }> = React.memo(({ } else if (nextLocator.strategy === 'all-columns' && nextLocator.reason) { message.info(nextLocator.reason); } + return { primaryKeys, locator: nextLocator }; + } catch { + if (fetchSeqRef.current !== seq) return null; + if (pkSeqRef.current !== locatorSeq) return null; + if (pkKeyRef.current !== locatorKey) return null; + const nextLocator = buildAllColumnsLocator([], { translate: tr }); + setPkColumns([]); + setEditLocator(nextLocator); + if (nextLocator.reason) message.info(nextLocator.reason); + return { primaryKeys: [], locator: nextLocator }; } - } catch { - if (fetchSeqRef.current !== seq) return; - if (pkSeqRef.current !== locatorSeq) return; - if (pkKeyRef.current !== locatorKey) return; - const nextLocator = buildAllColumnsLocator([], { translate: tr }); - pkColumnsForQuery = []; - editLocatorForQuery = nextLocator; - setPkColumns([]); - setEditLocator(nextLocator); - if (nextLocator.reason) message.info(nextLocator.reason); + }; + + if (dbTypeLower === 'kingbase') { + // Kingbase catalog metadata can be noticeably slower than the page query. + // Keep the grid read-only briefly and enable editing when the locator arrives. + void loadEditLocator(); + } else { + const locatorMetadata = await loadEditLocator(); + if (!locatorMetadata) return; + pkColumnsForQuery = locatorMetadata.primaryKeys; + editLocatorForQuery = locatorMetadata.locator; } } } diff --git a/frontend/src/components/useDataGridErDiagram.test.ts b/frontend/src/components/useDataGridErDiagram.test.ts index 3506a248..f4b953fe 100644 --- a/frontend/src/components/useDataGridErDiagram.test.ts +++ b/frontend/src/components/useDataGridErDiagram.test.ts @@ -1,11 +1,13 @@ import React from 'react'; import { act, create, type ReactTestRenderer } from 'react-test-renderer'; import { beforeEach, describe, expect, it, vi } from 'vitest'; +import type { ForeignKeyDefinition } from '../types'; import type { ErDiagramTableSnapshot } from './dataGridErDiagramModel'; import { collectErDiagramNeighborhood, useDataGridErDiagram } from './useDataGridErDiagram'; const backendApp = vi.hoisted(() => ({ DBGetColumns: vi.fn(), + DBGetDatabaseForeignKeys: vi.fn(), DBGetForeignKeys: vi.fn(), DBGetIndexes: vi.fn(), DBGetTables: vi.fn(), @@ -128,12 +130,62 @@ describe('collectErDiagramNeighborhood', () => { ); expect(twoHop.canExpandRelations).toBe(false); }); + + it('uses a prefetched foreign-key snapshot instead of scanning every table', async () => { + const unrelatedTables = Array.from({ length: 500 }, (_, index) => `unrelated_${index}`); + const schemaTableNames = ['orders', 'customers', 'order_items', ...unrelatedTables]; + const loadSnapshot = vi.fn(async (tableName: string) => ({ + tableName, + columns: [], + foreignKeys: [], + uniqueKeyGroups: [], + })); + const loadForeignKeys = vi.fn(async () => []); + const prefetchedForeignKeysByTable = new Map( + schemaTableNames.map((tableName) => [tableName, []]), + ); + prefetchedForeignKeysByTable.set('orders', [{ + name: 'fk_orders_customer', + columnName: 'customer_id', + refTableName: 'customers', + refColumnName: 'id', + constraintName: 'fk_orders_customer', + }]); + prefetchedForeignKeysByTable.set('order_items', [{ + name: 'fk_items_order', + columnName: 'order_id', + refTableName: 'orders', + refColumnName: 'id', + constraintName: 'fk_items_order', + }]); + + const result = await collectErDiagramNeighborhood({ + currentSnapshot: { + tableName: 'orders', + columns: [], + foreignKeys: [], + uniqueKeyGroups: [], + }, + schemaTableNames, + relationDepth: 1, + loadSnapshot, + loadForeignKeys, + resolveTableName: (tableName) => tableName, + prefetchedForeignKeysByTable, + }); + + expect(result.relations.map((relation) => `${relation.sourceTableName}->${relation.targetTableName}`)).toEqual( + expect.arrayContaining(['orders->customers', 'order_items->orders']), + ); + expect(loadForeignKeys).not.toHaveBeenCalled(); + }); }); describe('useDataGridErDiagram cache invalidation', () => { beforeEach(() => { Object.values(backendApp).forEach((mock) => mock.mockReset()); backendApp.DBGetColumns.mockResolvedValue({ success: true, data: [] }); + backendApp.DBGetDatabaseForeignKeys.mockResolvedValue({ success: true, data: {} }); backendApp.DBGetForeignKeys.mockResolvedValue({ success: true, data: [] }); backendApp.DBGetIndexes.mockResolvedValue({ success: true, data: [] }); backendApp.DBGetTables.mockResolvedValue({ success: true, data: [{ table: 'orders' }] }); @@ -201,4 +253,50 @@ describe('useDataGridErDiagram cache invalidation', () => { renderer?.unmount(); }); }); + + it('loads one Kingbase foreign-key snapshot instead of querying every table', async () => { + const unrelatedTables = Array.from({ length: 500 }, (_, index) => ({ + table: `ldf_server.unrelated_${index}`, + })); + backendApp.DBGetTables.mockResolvedValue({ + success: true, + data: [{ table: 'ldf_server.orders' }, ...unrelatedTables], + }); + + let controller: ReturnType | null = null; + let renderer: ReactTestRenderer | null = null; + const params = { + connections: [{ + id: 'kingbase-er-snapshot-test', + config: { + type: 'kingbase', + host: '127.0.0.1', + port: 54321, + database: 'ldf_server_dbs_dev', + }, + }], + connectionId: 'kingbase-er-snapshot-test', + dbName: 'ldf_server_dbs_dev', + tableName: 'ldf_server.orders', + }; + const Harness = () => { + controller = useDataGridErDiagram(params); + return null; + }; + + await act(async () => { + renderer = create(React.createElement(Harness)); + await vi.waitFor(() => { + expect(controller?.loading).toBe(false); + expect(controller?.graph).not.toBeNull(); + }); + }); + + expect(backendApp.DBGetDatabaseForeignKeys).toHaveBeenCalledTimes(1); + expect(backendApp.DBGetForeignKeys).not.toHaveBeenCalled(); + + act(() => { + renderer?.unmount(); + }); + }); }); diff --git a/frontend/src/components/useDataGridErDiagram.ts b/frontend/src/components/useDataGridErDiagram.ts index b1f244cc..39a539e1 100644 --- a/frontend/src/components/useDataGridErDiagram.ts +++ b/frontend/src/components/useDataGridErDiagram.ts @@ -1,9 +1,16 @@ import { useCallback, useEffect, useMemo, useRef, useState } from 'react'; -import { DBGetColumns, DBGetForeignKeys, DBGetIndexes, DBGetTables } from '../../wailsjs/go/app/App'; +import { + DBGetColumns, + DBGetDatabaseForeignKeys, + DBGetForeignKeys, + DBGetIndexes, + DBGetTables, +} from '../../wailsjs/go/app/App'; import type { ColumnDefinition, ForeignKeyDefinition } from '../types'; import { createBoundedAsyncCache } from '../utils/boundedAsyncCache'; import { buildRpcConnectionConfig } from '../utils/connectionRpcConfig'; import { normalizeColumnDefinitions } from '../utils/columnDefinition'; +import { resolveDataSourceType } from '../utils/dataSourceCapabilities'; import { resolveUniqueKeyGroupsFromIndexes } from './dataGridCopyInsert'; import { buildErDiagramGraph, @@ -38,6 +45,7 @@ const ER_SCHEMA_CACHE_MAX_ENTRIES = 32; const ER_TABLE_METADATA_CACHE_MAX_ENTRIES = 256; const schemaTableNamesCache = createBoundedAsyncCache(ER_SCHEMA_CACHE_MAX_ENTRIES); +const databaseForeignKeysCache = createBoundedAsyncCache>(ER_SCHEMA_CACHE_MAX_ENTRIES); const tableColumnsCache = createBoundedAsyncCache(ER_TABLE_METADATA_CACHE_MAX_ENTRIES); const tableForeignKeysCache = createBoundedAsyncCache(ER_TABLE_METADATA_CACHE_MAX_ENTRIES); const tableUniqueKeyGroupsCache = createBoundedAsyncCache(ER_TABLE_METADATA_CACHE_MAX_ENTRIES); @@ -62,7 +70,13 @@ const normalizeConnectionConfig = (connection: any) => ({ }); const invalidateCacheByPrefix = (prefix: string) => { - [schemaTableNamesCache, tableColumnsCache, tableForeignKeysCache, tableUniqueKeyGroupsCache].forEach((cache) => { + [ + schemaTableNamesCache, + databaseForeignKeysCache, + tableColumnsCache, + tableForeignKeysCache, + tableUniqueKeyGroupsCache, + ].forEach((cache) => { cache.invalidatePrefix(prefix); }); }; @@ -139,6 +153,27 @@ const loadTableForeignKeys = async ( return normalizeForeignKeyDefinitions(response.data); }); +const loadDatabaseForeignKeys = async ( + config: any, + dbName: string, + cacheKey: string, +): Promise> => databaseForeignKeysCache.getOrLoad(cacheKey, async () => { + const response = await DBGetDatabaseForeignKeys(buildRpcConnectionConfig(config) as any, dbName); + if (!response?.success || !response.data || typeof response.data !== 'object' || Array.isArray(response.data)) { + throw new Error(response?.message || 'Failed to load database foreign keys'); + } + + const result = new Map(); + Object.entries(response.data as Record).forEach(([sourceTableName, rawForeignKeys]) => { + const key = normalizeErQualifiedName(sourceTableName); + if (!key) { + return; + } + result.set(key, normalizeForeignKeyDefinitions(rawForeignKeys)); + }); + return result; +}); + const loadTableUniqueKeyGroups = async ( config: any, dbName: string, @@ -157,10 +192,12 @@ const loadTableSnapshot = async ( dbName: string, tableName: string, tableCacheKey: string, + prefetchedForeignKeys?: Promise, ): Promise => { const [columnsResult, foreignKeysResult, uniqueKeyGroupsResult] = await Promise.allSettled([ loadTableColumns(config, dbName, tableName, `${tableCacheKey}|columns`), - loadTableForeignKeys(config, dbName, tableName, `${tableCacheKey}|foreignKeys`), + prefetchedForeignKeys + || loadTableForeignKeys(config, dbName, tableName, `${tableCacheKey}|foreignKeys`), loadTableUniqueKeyGroups(config, dbName, tableName, `${tableCacheKey}|uniqueKeys`), ]); @@ -225,6 +262,7 @@ type CollectErDiagramNeighborhoodParams = { loadSnapshot: (tableName: string) => Promise; loadForeignKeys: (tableName: string) => Promise; resolveTableName: (tableName: string) => string; + prefetchedForeignKeysByTable?: ReadonlyMap; }; type CollectErDiagramNeighborhoodResult = { @@ -263,14 +301,25 @@ export const collectErDiagramNeighborhood = async ( registerTableName(currentTableName); params.schemaTableNames.forEach(registerTableName); + params.prefetchedForeignKeysByTable?.forEach((foreignKeys, tableName) => { + const actualTableName = registerTableName(tableName); + const tableKey = normalizeErQualifiedName(actualTableName); + if (tableKey) { + foreignKeysByKey.set(tableKey, foreignKeys); + } + }); params.currentSnapshot.foreignKeys.forEach((foreignKey) => { registerRelationTarget(foreignKey.refTableName); }); + const currentForeignKeys = foreignKeysByKey.has(currentKey) + ? foreignKeysByKey.get(currentKey) || [] + : params.currentSnapshot.foreignKeys || []; snapshotByKey.set(currentKey, { ...params.currentSnapshot, tableName: registerTableName(params.currentSnapshot.tableName), + foreignKeys: currentForeignKeys, }); - foreignKeysByKey.set(currentKey, params.currentSnapshot.foreignKeys || []); + foreignKeysByKey.set(currentKey, currentForeignKeys); visitedKeys.add(currentKey); const loadSnapshotByKey = async (tableKey: string): Promise => { @@ -284,14 +333,18 @@ export const collectErDiagramNeighborhood = async ( const snapshot = await params.loadSnapshot(tableName); const actualTableName = registerTableName(snapshot.tableName || tableName); const normalizedActualTableName = normalizeErQualifiedName(actualTableName); + const snapshotForeignKeys = foreignKeysByKey.has(tableKey) + ? foreignKeysByKey.get(tableKey) || [] + : snapshot.foreignKeys || []; const nextSnapshot = { ...snapshot, tableName: actualTableName, + foreignKeys: snapshotForeignKeys, }; snapshotByKey.set(tableKey, nextSnapshot); - foreignKeysByKey.set(tableKey, nextSnapshot.foreignKeys || []); + foreignKeysByKey.set(tableKey, snapshotForeignKeys); snapshotByKey.set(normalizedActualTableName, nextSnapshot); - foreignKeysByKey.set(normalizedActualTableName, nextSnapshot.foreignKeys || []); + foreignKeysByKey.set(normalizedActualTableName, snapshotForeignKeys); return nextSnapshot; } catch { warningCount += 1; @@ -506,7 +559,9 @@ export const useDataGridErDiagram = (params: DataGridErDiagramParams) => { const seq = ++requestSeqRef.current; const config = normalizeConnectionConfig(connection); const schemaCacheKey = `${cachePrefix}schemaTables`; + const databaseForeignKeysCacheKey = `${cachePrefix}databaseForeignKeys`; const currentTableCacheKey = `${cachePrefix}${normalizedTableName}`; + const isKingbase = resolveDataSourceType(config) === 'kingbase'; setState((prev) => ({ ...prev, @@ -520,7 +575,31 @@ export const useDataGridErDiagram = (params: DataGridErDiagramParams) => { const loadGraph = async () => { let warningCount = 0; - const currentSnapshot = await loadTableSnapshot(config, normalizedDbName, normalizedTableName, currentTableCacheKey); + const databaseForeignKeysResultPromise = isKingbase + ? loadDatabaseForeignKeys(config, normalizedDbName, databaseForeignKeysCacheKey) + .then((value) => ({ value, failed: false as const })) + .catch(() => ({ value: null, failed: true as const })) + : null; + const prefetchedCurrentForeignKeys = databaseForeignKeysResultPromise + ? databaseForeignKeysResultPromise.then((result) => { + if (result.value) { + return result.value.get(normalizeErQualifiedName(normalizedTableName)) || []; + } + return loadTableForeignKeys( + config, + normalizedDbName, + normalizedTableName, + `${currentTableCacheKey}|foreignKeys`, + ); + }) + : undefined; + const currentSnapshot = await loadTableSnapshot( + config, + normalizedDbName, + normalizedTableName, + currentTableCacheKey, + prefetchedCurrentForeignKeys, + ); let schemaTableNames = [currentSnapshot.tableName]; try { @@ -536,6 +615,23 @@ export const useDataGridErDiagram = (params: DataGridErDiagramParams) => { ]); const resolveTableName = (name: string) => resolveErActualTableName(name, resolvedSchemaTableNames); + let prefetchedForeignKeysByTable: Map | undefined; + if (databaseForeignKeysResultPromise) { + const databaseForeignKeysResult = await databaseForeignKeysResultPromise; + prefetchedForeignKeysByTable = new Map( + resolvedSchemaTableNames.map((name) => [normalizeErQualifiedName(name), []]), + ); + databaseForeignKeysResult.value?.forEach((foreignKeys, tableKey) => { + prefetchedForeignKeysByTable?.set(tableKey, foreignKeys); + }); + prefetchedForeignKeysByTable.set( + normalizeErQualifiedName(currentSnapshot.tableName), + currentSnapshot.foreignKeys, + ); + if (databaseForeignKeysResult.failed) { + warningCount += 1; + } + } const neighborhood = await collectErDiagramNeighborhood({ currentSnapshot, @@ -543,13 +639,25 @@ export const useDataGridErDiagram = (params: DataGridErDiagramParams) => { relationDepth, loadSnapshot: async (relatedTableName) => { const tableCacheKey = `${cachePrefix}${relatedTableName}`; - return loadTableSnapshot(config, normalizedDbName, relatedTableName, tableCacheKey); + const prefetchedForeignKeys = prefetchedForeignKeysByTable + ? Promise.resolve( + prefetchedForeignKeysByTable.get(normalizeErQualifiedName(relatedTableName)) || [], + ) + : undefined; + return loadTableSnapshot( + config, + normalizedDbName, + relatedTableName, + tableCacheKey, + prefetchedForeignKeys, + ); }, loadForeignKeys: async (relatedTableName) => { const tableCacheKey = `${cachePrefix}${relatedTableName}`; return loadTableForeignKeys(config, normalizedDbName, relatedTableName, `${tableCacheKey}|foreignKeys`); }, resolveTableName, + prefetchedForeignKeysByTable, }); warningCount += neighborhood.warningCount; diff --git a/frontend/src/main.tsx b/frontend/src/main.tsx index 74020eaf..131e0dcf 100644 --- a/frontend/src/main.tsx +++ b/frontend/src/main.tsx @@ -420,6 +420,7 @@ if ( DBGetDatabases: async () => ({ success: true, data: ['missav_bot'] }), DBGetTables: async () => ({ success: true, data: cloneBrowserMockValue(mockQueryTables) }), DBGetAllColumns: async () => ({ success: true, data: cloneBrowserMockValue(mockQueryColumns) }), + DBGetDatabaseForeignKeys: async () => ({ success: true, data: {} }), DBGetColumns: async (_config: any, _dbName: string, tableName: string) => ({ success: true, data: cloneBrowserMockValue( diff --git a/frontend/wailsjs/go/app/App.d.ts b/frontend/wailsjs/go/app/App.d.ts index 87da284e..c25dce4a 100755 --- a/frontend/wailsjs/go/app/App.d.ts +++ b/frontend/wailsjs/go/app/App.d.ts @@ -62,6 +62,8 @@ export function DBGetColumns(arg1:connection.ConnectionConfig,arg2:string,arg3:s export function DBGetDatabases(arg1:connection.ConnectionConfig):Promise; +export function DBGetDatabaseForeignKeys(arg1:connection.ConnectionConfig,arg2:string):Promise; + export function DBGetForeignKeys(arg1:connection.ConnectionConfig,arg2:string,arg3:string):Promise; export function DBGetIndexes(arg1:connection.ConnectionConfig,arg2:string,arg3:string):Promise; diff --git a/frontend/wailsjs/go/app/App.js b/frontend/wailsjs/go/app/App.js index 5944bf2f..05ba248b 100755 --- a/frontend/wailsjs/go/app/App.js +++ b/frontend/wailsjs/go/app/App.js @@ -110,6 +110,10 @@ export function DBGetDatabases(arg1) { return window['go']['app']['App']['DBGetDatabases'](arg1); } +export function DBGetDatabaseForeignKeys(arg1, arg2) { + return window['go']['app']['App']['DBGetDatabaseForeignKeys'](arg1, arg2); +} + export function DBGetForeignKeys(arg1, arg2, arg3) { return window['go']['app']['App']['DBGetForeignKeys'](arg1, arg2, arg3); } diff --git a/internal/app/methods_db.go b/internal/app/methods_db.go index c027f322..a7962e5a 100644 --- a/internal/app/methods_db.go +++ b/internal/app/methods_db.go @@ -2876,6 +2876,29 @@ func (a *App) DBGetForeignKeys(config connection.ConnectionConfig, dbName string return connection.QueryResult{Success: true, Data: ensureNonNilSlice(fks)} } +func (a *App) DBGetDatabaseForeignKeys(config connection.ConnectionConfig, dbName string) connection.QueryResult { + runConfig := normalizeRunConfig(config, dbName) + + dbInst, err := a.getDatabase(runConfig) + if err != nil { + return connection.QueryResult{Success: false, Message: err.Error()} + } + + provider, ok := dbInst.(db.DatabaseForeignKeyProvider) + if !ok { + return connection.QueryResult{Success: false, Message: "database-wide foreign-key metadata is not supported"} + } + + foreignKeysByTable, err := provider.GetDatabaseForeignKeys(dbName) + if err != nil { + return connection.QueryResult{Success: false, Message: err.Error()} + } + if foreignKeysByTable == nil { + foreignKeysByTable = make(map[string][]connection.ForeignKeyDefinition) + } + return connection.QueryResult{Success: true, Data: foreignKeysByTable} +} + func (a *App) DBGetTriggers(config connection.ConnectionConfig, dbName string, tableName string) connection.QueryResult { runConfig := normalizeRunConfig(config, dbName) diff --git a/internal/db/database.go b/internal/db/database.go index c6f51d97..dda380ee 100644 --- a/internal/db/database.go +++ b/internal/db/database.go @@ -44,6 +44,13 @@ type Database interface { GetTriggers(dbName, tableName string) ([]connection.TriggerDefinition, error) } +// DatabaseForeignKeyProvider is an optional metadata interface for drivers that +// can load a database-wide foreign-key snapshot more efficiently than one table +// at a time. +type DatabaseForeignKeyProvider interface { + GetDatabaseForeignKeys(dbName string) (map[string][]connection.ForeignKeyDefinition, error) +} + // TableRowCounter is an optional metadata interface for drivers that can // provide exact table row counts alongside a table list. type TableRowCounter interface { diff --git a/internal/db/kingbase_impl.go b/internal/db/kingbase_impl.go index d8882e63..14b3fa0a 100644 --- a/internal/db/kingbase_impl.go +++ b/internal/db/kingbase_impl.go @@ -34,6 +34,7 @@ type kingbaseSessionExecer struct { var _ QueryMessageExecer = (*KingbaseDB)(nil) var _ StatementQueryMessageExecer = (*kingbaseSessionExecer)(nil) +var _ DatabaseForeignKeyProvider = (*KingbaseDB)(nil) func quoteConnValue(v string) string { if v == "" { @@ -612,6 +613,86 @@ func (k *KingbaseDB) GetForeignKeys(dbName, tableName string) ([]connection.Fore return fks, nil } +func buildKingbaseDatabaseForeignKeysQuery() string { + return ` + SELECT + source_ns.nspname AS table_schema, + source_table.relname AS table_name, + con.conname AS constraint_name, + source_column.attname AS column_name, + target_ns.nspname AS foreign_table_schema, + target_table.relname AS foreign_table_name, + target_column.attname AS foreign_column_name + FROM pg_catalog.pg_constraint AS con + JOIN pg_catalog.pg_class AS source_table + ON source_table.oid = con.conrelid + JOIN pg_catalog.pg_namespace AS source_ns + ON source_ns.oid = source_table.relnamespace + JOIN pg_catalog.pg_class AS target_table + ON target_table.oid = con.confrelid + JOIN pg_catalog.pg_namespace AS target_ns + ON target_ns.oid = target_table.relnamespace + JOIN LATERAL pg_catalog.generate_subscripts(con.conkey, 1) AS key_position(position) + ON TRUE + JOIN pg_catalog.pg_attribute AS source_column + ON source_column.attrelid = source_table.oid + AND source_column.attnum = con.conkey[key_position.position] + JOIN pg_catalog.pg_attribute AS target_column + ON target_column.attrelid = target_table.oid + AND target_column.attnum = con.confkey[key_position.position] + WHERE con.contype = 'f' + AND source_ns.nspname NOT IN ('pg_catalog', 'information_schema') + AND source_ns.nspname NOT LIKE 'pg|_%' ESCAPE '|' + ORDER BY source_ns.nspname, source_table.relname, con.conname, key_position.position` +} + +func buildKingbaseDatabaseForeignKeys( + data []map[string]interface{}, +) map[string][]connection.ForeignKeyDefinition { + foreignKeysByTable := make(map[string][]connection.ForeignKeyDefinition) + for _, row := range data { + sourceSchema := strings.TrimSpace(fmt.Sprint(row["table_schema"])) + sourceTable := strings.TrimSpace(fmt.Sprint(row["table_name"])) + targetSchema := strings.TrimSpace(fmt.Sprint(row["foreign_table_schema"])) + targetTable := strings.TrimSpace(fmt.Sprint(row["foreign_table_name"])) + if sourceTable == "" || targetTable == "" { + continue + } + + sourceTableName := sourceTable + if sourceSchema != "" { + sourceTableName = sourceSchema + "." + sourceTable + } + targetTableName := targetTable + if targetSchema != "" { + targetTableName = targetSchema + "." + targetTable + } + constraintName := strings.TrimSpace(fmt.Sprint(row["constraint_name"])) + foreignKeysByTable[sourceTableName] = append( + foreignKeysByTable[sourceTableName], + connection.ForeignKeyDefinition{ + Name: constraintName, + ColumnName: strings.TrimSpace(fmt.Sprint(row["column_name"])), + RefTableName: targetTableName, + RefColumnName: strings.TrimSpace(fmt.Sprint(row["foreign_column_name"])), + ConstraintName: constraintName, + }, + ) + } + return foreignKeysByTable +} + +// GetDatabaseForeignKeys loads all user-schema foreign keys in one catalog +// query. The ER diagram uses this snapshot to avoid issuing one metadata query +// per table on large Kingbase schemas. +func (k *KingbaseDB) GetDatabaseForeignKeys(_ string) (map[string][]connection.ForeignKeyDefinition, error) { + data, _, err := k.Query(buildKingbaseDatabaseForeignKeysQuery()) + if err != nil { + return nil, err + } + return buildKingbaseDatabaseForeignKeys(data), nil +} + func (k *KingbaseDB) GetTriggers(dbName, tableName string) ([]connection.TriggerDefinition, error) { // 解析 schema.table 格式 schema := strings.TrimSpace(dbName) diff --git a/internal/db/kingbase_impl_test.go b/internal/db/kingbase_impl_test.go index c754d435..8a078e33 100644 --- a/internal/db/kingbase_impl_test.go +++ b/internal/db/kingbase_impl_test.go @@ -148,6 +148,110 @@ func TestSplitKingbaseQualifiedTable(t *testing.T) { } } +func TestBuildKingbaseDatabaseForeignKeysQueryUsesOneCatalogSnapshot(t *testing.T) { + query := buildKingbaseDatabaseForeignKeysQuery() + + if strings.Contains(query, "information_schema") && !strings.Contains(query, "NOT IN ('pg_catalog', 'information_schema')") { + t.Fatalf("expected pg_catalog query, got %s", query) + } + for _, fragment := range []string{ + "pg_catalog.pg_constraint", + "pg_catalog.generate_subscripts(con.conkey, 1) AS key_position(position)", + "con.contype = 'f'", + "source_ns.nspname", + "target_ns.nspname", + } { + if !strings.Contains(query, fragment) { + t.Fatalf("expected query to contain %q, got %s", fragment, query) + } + } +} + +func TestBuildKingbaseDatabaseForeignKeysGroupsQualifiedTables(t *testing.T) { + foreignKeysByTable := buildKingbaseDatabaseForeignKeys([]map[string]interface{}{ + { + "table_schema": "ldf_server", + "table_name": "orders", + "constraint_name": "fk_orders_customer", + "column_name": "customer_id", + "foreign_table_schema": "crm", + "foreign_table_name": "customers", + "foreign_column_name": "id", + }, + { + "table_schema": "ldf_server", + "table_name": "order_items", + "constraint_name": "fk_items_order", + "column_name": "order_id", + "foreign_table_schema": "ldf_server", + "foreign_table_name": "orders", + "foreign_column_name": "id", + }, + }) + + orders := foreignKeysByTable["ldf_server.orders"] + if len(orders) != 1 { + t.Fatalf("expected one orders foreign key, got %+v", orders) + } + if orders[0].RefTableName != "crm.customers" || orders[0].ColumnName != "customer_id" { + t.Fatalf("unexpected orders foreign key: %+v", orders[0]) + } + if len(foreignKeysByTable["ldf_server.order_items"]) != 1 { + t.Fatalf("expected qualified order_items foreign key, got %+v", foreignKeysByTable) + } +} + +func TestKingbaseGetDatabaseForeignKeysUsesSingleQuery(t *testing.T) { + registerFakeKingbaseDriverOnce.Do(func() { + sql.Register(fakeKingbaseDriverName, fakeKingbaseDriver{}) + }) + + sqlDB, err := sql.Open(fakeKingbaseDriverName, "") + if err != nil { + t.Fatalf("open fake kingbase db failed: %v", err) + } + defer sqlDB.Close() + + query := buildKingbaseDatabaseForeignKeysQuery() + fakeKingbaseStateMu.Lock() + fakeKingbaseState.queryErr = nil + fakeKingbaseState.queryResults = map[string]fakeKingbaseQueryResult{ + query: { + columns: []string{ + "table_schema", + "table_name", + "constraint_name", + "column_name", + "foreign_table_schema", + "foreign_table_name", + "foreign_column_name", + }, + rows: [][]driver.Value{ + {"ldf_server", "orders", "fk_orders_customer", "customer_id", "crm", "customers", "id"}, + }, + }, + } + fakeKingbaseState.lastQuery = "" + fakeKingbaseState.queries = nil + fakeKingbaseStateMu.Unlock() + + client := &KingbaseDB{conn: sqlDB} + foreignKeysByTable, err := client.GetDatabaseForeignKeys("ldf_server_dbs_dev") + if err != nil { + t.Fatalf("GetDatabaseForeignKeys returned error: %v", err) + } + if len(foreignKeysByTable["ldf_server.orders"]) != 1 { + t.Fatalf("unexpected foreign keys: %+v", foreignKeysByTable) + } + + fakeKingbaseStateMu.Lock() + queries := append([]string(nil), fakeKingbaseState.queries...) + fakeKingbaseStateMu.Unlock() + if len(queries) != 1 || queries[0] != query { + t.Fatalf("expected one database-wide foreign-key query, got %v", queries) + } +} + func TestKingbaseGetDatabasesFallsBackToCurrentDatabase(t *testing.T) { registerFakeKingbaseDriverOnce.Do(func() { sql.Register(fakeKingbaseDriverName, fakeKingbaseDriver{})