diff --git a/precise/src/AsyncTrinoClient.tsx b/precise/src/AsyncTrinoClient.tsx index 8043819..e343e24 100644 --- a/precise/src/AsyncTrinoClient.tsx +++ b/precise/src/AsyncTrinoClient.tsx @@ -26,6 +26,9 @@ class TrinoQueryRunner { private trinoSchema: string | null = null private setHeadersCallback: (catalog: string | null, schema: string | null) => void = () => {} + // Authentication: custom headers to include in every Trino request (e.g. Authorization, X-Trino-User) + private requestHeaders: Record = {} + SetAllResultsCallback(setAllResults: (n: any[], error: boolean) => any): TrinoQueryRunner { this.setAllResults = setAllResults return this @@ -47,6 +50,19 @@ class TrinoQueryRunner { return this } + // Set custom headers to include in every Trino request (e.g. Authorization, X-Trino-User) + SetRequestHeaders(headers: Record): TrinoQueryRunner { + this.requestHeaders = headers + return this + } + + // Resolve headers; falls back to X-Trino-User: system if none provided + private resolveHeaders(): Record { + return Object.keys(this.requestHeaders).length > 0 + ? { ...this.requestHeaders } + : { 'X-Trino-User': 'system' } + } + // Add getters for catalog and schema GetCatalog(): string | null { return this.trinoCatalog @@ -71,9 +87,7 @@ class TrinoQueryRunner { // cancel query fetch(nextUri, { method: 'DELETE', - headers: { - 'X-Trino-User': 'system', - }, + headers: this.resolveHeaders(), }) .then((response) => response) .then((data) => { @@ -157,16 +171,11 @@ class TrinoQueryRunner { const controller = new AbortController() const timeoutId = setTimeout(() => controller.abort('Timeout: Trino is not responding'), 15000) - // Prepare headers for the request - const headers: Record = { - 'X-Trino-User': 'system', - } - - // Add catalog and schema headers if they exist + // Merge authentication headers with catalog/schema headers + const headers: Record = { ...this.resolveHeaders() } if (this.trinoCatalog) { headers['X-Trino-Catalog'] = this.trinoCatalog } - if (this.trinoSchema) { headers['X-Trino-Schema'] = this.trinoSchema } @@ -256,12 +265,10 @@ class TrinoQueryRunner { async NextPage(previous: any) { try { // fix cors for testing - const nextUri = await previous.nextUri.replace(/^https?:\/\/[^/]+/, '') + const nextUri = previous.nextUri.replace(/^https?:\/\/[^/]+/, '') const response = await fetch(nextUri, { method: 'GET', - headers: { - 'X-Trino-User': 'system', - }, + headers: this.resolveHeaders(), }) if (!response.ok) { diff --git a/precise/src/QueryCell.tsx b/precise/src/QueryCell.tsx index e1fb7e8..fbfea4a 100644 --- a/precise/src/QueryCell.tsx +++ b/precise/src/QueryCell.tsx @@ -34,6 +34,7 @@ interface QueryCellProps { height: number onDrawerToggle: () => void theme?: string + requestHeaders?: Record } class QueryCell extends React.Component { @@ -70,6 +71,7 @@ class QueryCell extends React.Component { return ( this.props.drawerOpen !== nextProps.drawerOpen || this.props.height !== nextProps.height || + this.props.requestHeaders !== nextProps.requestHeaders || this.state.results !== nextState.results || this.state.columns !== nextState.columns || this.state.response !== nextState.response || @@ -84,6 +86,12 @@ class QueryCell extends React.Component { ) } + componentDidUpdate(prevProps: QueryCellProps) { + if (prevProps.requestHeaders !== this.props.requestHeaders && this.props.requestHeaders) { + this.queryRunner.SetRequestHeaders(this.props.requestHeaders) + } + } + handleQueriesChange = () => { this.setState({ currentQuery: this.props.queries.getCurrentQuery() }) } @@ -119,6 +127,10 @@ class QueryCell extends React.Component { schema: schema ?? undefined, }) }) + + if (this.props.requestHeaders) { + this.queryRunner.SetRequestHeaders(this.props.requestHeaders) + } } setRunningQueryId = (queryId: string | null) => { diff --git a/precise/src/QueryEditor.tsx b/precise/src/QueryEditor.tsx index 2bad7a7..e8cfdcf 100644 --- a/precise/src/QueryEditor.tsx +++ b/precise/src/QueryEditor.tsx @@ -1,4 +1,4 @@ -import React, { useRef, useState } from 'react' +import React, { useEffect, useRef, useState } from 'react' import { styled } from '@mui/material/styles' import { Box, Drawer, useMediaQuery } from '@mui/material' import CssBaseline from '@mui/material/CssBaseline' @@ -9,11 +9,13 @@ import { darkTheme, lightTheme } from './theme' import Queries from './schema/Queries' import QueryInfo from './schema/QueryInfo' import CatalogViewer from './controls/catalog_viewer/CatalogViewer' +import SchemaProvider from './sql/SchemaProvider' interface IQueryEditor { height: number theme?: 'dark' | 'light' enableCatalogSearchColumns?: boolean + requestHeaders?: Record } const DRAWER_WIDTH = 260 @@ -71,7 +73,7 @@ const AppBar = styled(MuiAppBar, { ], })) -export const QueryEditor = ({ height, theme, enableCatalogSearchColumns }: IQueryEditor) => { +export const QueryEditor = ({ height, theme, enableCatalogSearchColumns, requestHeaders }: IQueryEditor) => { const [queries, setQueries] = useState(() => new Queries()) const [drawerOpen, setDrawerOpen] = useState(true) const [queryRunning, setQueryRunning] = useState(false) @@ -79,6 +81,11 @@ export const QueryEditor = ({ height, theme, enableCatalogSearchColumns }: IQuer const prefersDarkMode = useMediaQuery('(prefers-color-scheme: dark)') const containerRef = useRef(null) + // Propagate request headers to SchemaProvider so catalog browsing is authenticated + useEffect(() => { + SchemaProvider.setRequestHeaders(requestHeaders ?? {}) + }, [requestHeaders]) + const muiThemeToUse = () => { if (theme === 'dark') { return darkTheme @@ -188,6 +195,7 @@ export const QueryEditor = ({ height, theme, enableCatalogSearchColumns }: IQuer onAppendQuery={appendQueryContent} onDrawerToggle={() => setDrawerOpen(false)} enableSearchColumns={enableCatalogSearchColumns} + requestHeaders={requestHeaders} /> @@ -198,6 +206,7 @@ export const QueryEditor = ({ height, theme, enableCatalogSearchColumns }: IQuer height={height} onDrawerToggle={() => setDrawerOpen(true)} theme={theme} + requestHeaders={requestHeaders} /> diff --git a/precise/src/controls/catalog_viewer/CatalogViewer.tsx b/precise/src/controls/catalog_viewer/CatalogViewer.tsx index 32aac0b..5094ccf 100644 --- a/precise/src/controls/catalog_viewer/CatalogViewer.tsx +++ b/precise/src/controls/catalog_viewer/CatalogViewer.tsx @@ -31,6 +31,7 @@ interface CatalogViewerProps { onAppendQuery?: (query: string, catalog?: string, schema?: string) => void onDrawerToggle?: () => void enableSearchColumns?: boolean + requestHeaders?: Record } const CatalogViewer: React.FC = ({ @@ -39,6 +40,7 @@ const CatalogViewer: React.FC = ({ onAppendQuery, onDrawerToggle, enableSearchColumns, + requestHeaders, }) => { // Basic state const [catalogs, setCatalogs] = useState>(new Map()) @@ -89,6 +91,7 @@ const CatalogViewer: React.FC = ({ }, [debouncedFilterText, searchColumns, catalogs]) const loadCatalogs = useCallback(async () => { + SchemaProvider.setRequestHeaders(requestHeaders ?? {}) setIsLoading(true) setErrorMessage(undefined) @@ -107,7 +110,7 @@ const CatalogViewer: React.FC = ({ setErrorMessage(error instanceof Error ? error.message : 'An unknown error occurred') setIsLoading(false) } - }, []) + }, [requestHeaders]) const handleToggle = async (path: string) => { if (!viewerState.current) return diff --git a/precise/src/sql/SchemaProvider.ts b/precise/src/sql/SchemaProvider.ts index 627a4ed..7d90658 100644 --- a/precise/src/sql/SchemaProvider.ts +++ b/precise/src/sql/SchemaProvider.ts @@ -13,6 +13,17 @@ class SchemaProvider { // map of fully qualified table name to tables static tables: Map = new Map() + // Configurable request headers for authentication + private static requestHeaders: Record = {} + + static setRequestHeaders(headers: Record) { + this.requestHeaders = { ...headers } + } + + private static createRunner(): TrinoQueryRunner { + return new TrinoQueryRunner().SetRequestHeaders(this.requestHeaders) + } + static getTableNameList(catalogFilter: string | undefined, schemaFilter: string | undefined): string[] { // get list from catalogs, because tables may not be resolved const tableNames: string[] = [] @@ -68,7 +79,7 @@ class SchemaProvider { errorCallback: ((error: string) => void) | null = null ) { // refresh catalogs - new TrinoQueryRunner() + this.createRunner() .SetAllResultsCallback((results: any[], isError: boolean) => { for (let i = 0; i < results.length; i++) { const catalog: Catalog = new Catalog(results[i][0], results[i][1]) @@ -78,7 +89,7 @@ class SchemaProvider { this.lastSchemaFetchError = undefined // refresh tables and schemas for this catalog - new TrinoQueryRunner() + this.createRunner() .SetAllResultsCallback((results: any[], isError: boolean) => { for (let i = 0; i < results.length; i++) { const schemaName = results[i][0] @@ -121,7 +132,7 @@ class SchemaProvider { /* callback returns a table type */ static async getTableRefreshCache(tableRef: TableReference, callback: (table: Table) => void) { // First try to load all tables in the schema at once - const query = new TrinoQueryRunner() + const query = this.createRunner() query .SetAllResultsCallback((results: any[]) => { // Create a temporary map to hold all tables in this schema @@ -186,7 +197,7 @@ class SchemaProvider { } private static fallbackToDescribe(tableRef: TableReference, callback: (table: Table) => void) { - const fallbackQuery = new TrinoQueryRunner() + const fallbackQuery = this.createRunner() fallbackQuery .SetAllResultsCallback((results: any[]) => { const table = new Table(tableRef.tableName)