diff --git a/src/runtime/index.ts b/src/runtime/index.ts index 67fd14d..ab39e27 100644 --- a/src/runtime/index.ts +++ b/src/runtime/index.ts @@ -16,6 +16,8 @@ export type * from "./stdlib/fileAttachment.js"; export {FileAttachment, requireFileRegistration, registerFile} from "./stdlib/fileAttachment.js"; export type * from "./stdlib/interpreter.js"; export {Interpreter} from "./stdlib/interpreter.js"; +export type * from "./stdlib/sql.js"; +export {sql} from "./stdlib/sql.js"; export class NotebookRuntime { readonly runtime: Runtime & {fileAttachments: typeof fileAttachments}; diff --git a/src/runtime/stdlib/databaseClient.ts b/src/runtime/stdlib/databaseClient.ts index 14da584..68eb5af 100644 --- a/src/runtime/stdlib/databaseClient.ts +++ b/src/runtime/stdlib/databaseClient.ts @@ -6,7 +6,7 @@ import {hash, nameHash} from "../../lib/hash.js"; export type QueryParam = any; /** @see https://observablehq.com/@observablehq/database-client-specification#%C2%A71 */ -export type QueryResult = Record[] & {schema: ColumnSchema[]; date: Date}; +export type QueryResult> = T[] & {schema: ColumnSchema[]; date: Date}; /** @see https://observablehq.com/@observablehq/database-client-specification#%C2%A72.2 */ export interface ColumnSchema { @@ -39,10 +39,24 @@ export interface QueryOptions extends QueryOptionsSpec { since?: Date; } +export type SqlDialect = + | "bigquery" + | "databricks" + | "duckdb" + | "mongosql" + | "mssql" + | "mysql" + | "oracle" + | "postgres" + | "snowflake" + | "sql" + | "sqlite"; + export interface DatabaseClient { readonly name: string; readonly options: QueryOptions; - sql(strings: readonly string[], ...params: QueryParam[]): Promise; + readonly dialect?: SqlDialect; + sql>(strings: readonly string[], ...params: QueryParam[]): Promise>; // prettier-ignore } export const DatabaseClient = (name: string, options?: QueryOptionsSpec): DatabaseClient => { @@ -65,11 +79,11 @@ class DatabaseClientImpl implements DatabaseClient { options: {value: options, enumerable: true} }); } - async sql(strings: readonly string[], ...params: QueryParam[]): Promise { + async sql>(strings: readonly string[], ...params: QueryParam[]): Promise> { const path = await this.cachePath(strings, ...params); const response = await fetch(path); if (!response.ok) throw new Error(`failed to fetch: ${path}`); - return await response.json().then(revive); + return (await response.json().then(revive)) as QueryResult; } async cachePath(strings: readonly string[], ...params: QueryParam[]): Promise { return `.observable/cache/${await nameHash(this.name)}-${await hash(strings, ...params)}.json`; diff --git a/src/runtime/stdlib/index.ts b/src/runtime/stdlib/index.ts index 84a6a38..b95429b 100644 --- a/src/runtime/stdlib/index.ts +++ b/src/runtime/stdlib/index.ts @@ -8,12 +8,14 @@ import {Mutable} from "./mutable.js"; import * as Promises from "./promises/index.js"; import * as recommendedLibraries from "./recommendedLibraries.js"; import * as sampleDatasets from "./sampleDatasets.js"; +import {sql} from "./sql.js"; export const root = document.querySelector("main") ?? document.body; export const library = { dark: () => Generators.dark(), now: () => Generators.now(), + sql: () => sql, width: () => Generators.width(root), DatabaseClient: () => DatabaseClient, FileAttachment: () => FileAttachment, diff --git a/src/runtime/stdlib/recommendedLibraries.ts b/src/runtime/stdlib/recommendedLibraries.ts index d8d85b0..1c6bd27 100644 --- a/src/runtime/stdlib/recommendedLibraries.ts +++ b/src/runtime/stdlib/recommendedLibraries.ts @@ -17,7 +17,6 @@ export const mermaid = () => import("./mermaid.js").then((_) => _.mermaid); export const Plot = () => import("https://cdn.jsdelivr.net/npm/@observablehq/plot/+esm"); export const React = () => import("https://cdn.jsdelivr.net/npm/react/+esm"); export const ReactDOM = () => import("https://cdn.jsdelivr.net/npm/react-dom/+esm"); -// export const sql = () => import("observablehq:stdlib/duckdb").then((_) => _.sql); // export const SQLite = () => import("observablehq:stdlib/sqlite").then((_) => _.default); // export const SQLiteDatabaseClient = () => import("observablehq:stdlib/sqlite").then((_) => _.SQLiteDatabaseClient); export const tex = () => import("./tex.js").then((_) => _.tex); diff --git a/src/runtime/stdlib/sql.test.ts b/src/runtime/stdlib/sql.test.ts new file mode 100644 index 0000000..ee896e6 --- /dev/null +++ b/src/runtime/stdlib/sql.test.ts @@ -0,0 +1,452 @@ +import {assert, describe, test} from "vitest"; +import {sql} from "./sql.js"; +import type {QueryResult, SqlDialect} from "./databaseClient.js"; + +describe("sql`…`", () => { + test("defines a SQL fragment", () => { + assert.deepStrictEqual(sql`PURCHASES`, new sql.Fragment(["PURCHASES"], [])); + assert.deepStrictEqual(sql`FOO = ${42}`, new sql.Fragment(["FOO = ", ""], [42])); + assert.deepStrictEqual(sql`${1} = ${2}`, new sql.Fragment(["", " = ", ""], [1, 2])); + assert.deepStrictEqual(sql`FOO = ${sql`42`}`, new sql.Fragment(["FOO = ", ""], [sql`42`])); + }); + test("maintains reference equality with SQL params", () => { + const val = sql`42`; + assert.strictEqual(sql`FOO = ${val}`.params[0], val); + }); +}); + +describe("sql.view`…`", () => { + test("defines a SQL view", () => { + assert.deepStrictEqual(sql.view`SELECT * FROM FOO`, new sql.View(["SELECT * FROM FOO"], [])); + assert.deepStrictEqual(sql.view`SELECT * FROM FOO WHERE BAR = ${42}`, new sql.View(["SELECT * FROM FOO WHERE BAR = ", ""], [42])); // prettier-ignore + assert.deepStrictEqual(sql.view`SELECT * FROM FOO WHERE BAR = ${sql`42`}`, new sql.View(["SELECT * FROM FOO WHERE BAR = ", ""], [sql`42`])); // prettier-ignore + }); + test("maintains reference equality with SQL params", () => { + const bar = sql`42`; + const foo = sql.view`SELECT * FROM FOO WHERE BAR = ${bar}`; + assert.strictEqual(foo.params[0], bar); + assert.strictEqual(sql.view`SELECT * FROM ${foo}`.params[0], foo); + }); + test("sql.view`…`.flat() returns a sql.View", () => { + assert.instanceOf(sql.view`SELECT ${sql`1`}`.flat(), sql.View); + }); + test("sql.view`…`.toDialect() returns a sql.View", () => { + assert.instanceOf(sql.view`SELECT ${sql.variant({default: sql`1`})}`.toDialect(), sql.View); + }); +}); + +describe("sql`…`.flat()", () => { + test("returns a flattened SQL fragment", () => { + assert.deepStrictEqual(sql`FOO = 42`.flat(), new sql.Fragment(["FOO = 42"], [])); + assert.deepStrictEqual(sql`FOO = ${42}`.flat(), new sql.Fragment(["FOO = ", ""], [42])); + assert.deepStrictEqual(sql`FOO = ${sql`42`}`.flat(), new sql.Fragment(["FOO = 42"], [])); + assert.deepStrictEqual(sql`FOO = ${sql`${42}`}`.flat(), new sql.Fragment(["FOO = ", ""], [42])); + }); + test("handles arbitrarily nested SQL", () => { + assert.deepStrictEqual(sql`ORDER BY 3, ${sql`2 ${sql`DESC`}`}`.flat(), new sql.Fragment(["ORDER BY 3, 2 DESC"], [])); + }); + test("adds a WITH clause for selected views", () => { + const view = sql.view`SELECT * FROM PURCHASES`; + assert.deepStrictEqual( + sql`SELECT * FROM ${view}`.flat(), + new sql.Fragment( + [ + `WITH +"_1" AS (SELECT * FROM PURCHASES) +SELECT * FROM "_1"` + ], + [] + ) + ); + }); + test("handles multiple views", () => { + const view1 = sql.view`SELECT * FROM PURCHASES`; + const view2 = sql.view`SELECT * FROM SURVEYS`; + assert.deepStrictEqual( + sql`SELECT * FROM ${view1} UNION ALL SELECT * FROM ${view2}`.flat(), + new sql.Fragment( + [ + `WITH +"_1" AS (SELECT * FROM PURCHASES), +"_2" AS (SELECT * FROM SURVEYS) +SELECT * FROM "_1" UNION ALL SELECT * FROM "_2"` + ], + [] + ) + ); + }); + test("respects the specified dialect", () => { + const view1 = sql.view`SELECT * FROM PURCHASES`; + const view2 = sql.view`SELECT * FROM SURVEYS`; + assert.deepStrictEqual( + sql`SELECT * FROM ${view1} UNION ALL SELECT * FROM ${view2}`.flat("databricks"), + new sql.Fragment( + [ + `WITH +\`_1\` AS (SELECT * FROM PURCHASES), +\`_2\` AS (SELECT * FROM SURVEYS) +SELECT * FROM \`_1\` UNION ALL SELECT * FROM \`_2\`` + ], + [] + ) + ); + }); + test("combines with an existing WITH clause", () => { + const view = sql.view`SELECT * FROM PURCHASES`; + assert.deepStrictEqual( + sql`WITH FOO AS (SELECT * FROM BAR) +SELECT * FROM ${view} +UNION ALL SELECT * FROM FOO`.flat(), + new sql.Fragment( + [ + `WITH +"_1" AS (SELECT * FROM PURCHASES), + FOO AS (SELECT * FROM BAR) +SELECT * FROM "_1" +UNION ALL SELECT * FROM FOO` + ], + [] + ) + ); + }); + test("chains a view into a recursive query (flight reachability)", () => { + const routes = sql.view`SELECT origin, dest FROM routes WHERE airline = ${"AA"}`; + const reachable = sql` +WITH RECURSIVE trip AS ( + SELECT ${sql.text("SFO")} AS airport, 0 AS hops + UNION + SELECT r.dest, t.hops + 1 + FROM ${routes} r + JOIN trip t ON r.origin = t.airport +) +SELECT airport, min(hops) AS hops FROM trip GROUP BY airport ORDER BY hops`; + assert.deepStrictEqual( + reachable.flat(), + new sql.Fragment( + [ + ` +WITH RECURSIVE +"_1" AS (SELECT origin, dest FROM routes WHERE airline = `, + `), + trip AS ( + SELECT 'SFO' AS airport, 0 AS hops + UNION + SELECT r.dest, t.hops + 1 + FROM "_1" r + JOIN trip t ON r.origin = t.airport +) +SELECT airport, min(hops) AS hops FROM trip GROUP BY airport ORDER BY hops` + ], + ["AA"] + ) + ); + }); + test("consolidates multiple references to the same view", () => { + const view = sql.view`SELECT * FROM PURCHASES`; + assert.deepStrictEqual( + sql`SELECT * FROM ${view} UNION ALL SELECT * FROM ${view}`.flat(), + new sql.Fragment( + [ + `WITH +"_1" AS (SELECT * FROM PURCHASES) +SELECT * FROM "_1" UNION ALL SELECT * FROM "_1"` + ], + [] + ) + ); + }); + test("hoists chained views in topological order", () => { + const purchases = sql.view`SELECT * FROM PURCHASES`; + const filtered = sql.view`SELECT * FROM ${purchases} WHERE FOO = 42`; + assert.deepStrictEqual( + sql`SELECT * FROM ${filtered}`.flat(), + new sql.Fragment( + [ + `WITH +"_2" AS (SELECT * FROM PURCHASES), +"_1" AS (SELECT * FROM "_2" WHERE FOO = 42) +SELECT * FROM "_1"` + ], + [] + ) + ); + }); + test("orders params by CTE definition when chaining views", () => { + const inner = sql.view`SELECT * FROM PURCHASES WHERE X = ${1}`; + const outer = sql.view`SELECT * FROM ${inner} WHERE Y = ${2}`; + assert.deepStrictEqual( + sql`SELECT * FROM ${outer}`.flat(), + new sql.Fragment( + [ + `WITH +"_2" AS (SELECT * FROM PURCHASES WHERE X = `, + `), +"_1" AS (SELECT * FROM "_2" WHERE Y = `, + `) +SELECT * FROM "_1"` + ], + [1, 2] + ) + ); + }); + test("avoids conflicts with existing unquoted table names", () => { + const view = sql.view`SELECT * FROM PURCHASES`; + assert.deepStrictEqual( + sql`SELECT * FROM ${view} UNION ALL SELECT * FROM _1`.flat(), + new sql.Fragment( + [ + `WITH +"_2" AS (SELECT * FROM PURCHASES) +SELECT * FROM "_2" UNION ALL SELECT * FROM _1` + ], + [] + ) + ); + }); + test("avoids conflicts with existing quoted table names", () => { + const view = sql.view`SELECT * FROM PURCHASES`; + assert.deepStrictEqual( + sql`SELECT * FROM ${view} UNION ALL SELECT * FROM "_1"`.flat(), + new sql.Fragment( + [ + `WITH +"_2" AS (SELECT * FROM PURCHASES) +SELECT * FROM "_2" UNION ALL SELECT * FROM "_1"` + ], + [] + ) + ); + }); + test("returns itself if the SQL fragment is already flat", () => { + const input = sql`42`; + const output = input.flat(); + assert.strictEqual(input, output); + }); +}); + +describe("sql`…`.toDialect(dialect)", () => { + test("converts the SQL fragment to the specified dialect", () => { + assert.deepStrictEqual(sql`SELECT ${sql.variant({duckdb: "DUCKDB"})}`.toDialect("duckdb"), sql`SELECT ${"DUCKDB"}`); + }); + test("handles the default dialect", () => { + const fragment = sql`SELECT ${sql.variant({duckdb: "DUCKDB", default: "DEFAULT"})}`; + assert.deepStrictEqual(fragment.toDialect("duckdb"), sql`SELECT ${"DUCKDB"}`); + assert.deepStrictEqual(fragment.toDialect("postgres"), sql`SELECT ${"DEFAULT"}`); + assert.deepStrictEqual(fragment.toDialect("unknown" as SqlDialect), sql`SELECT ${"DEFAULT"}`); + assert.deepStrictEqual(fragment.toDialect(), sql`SELECT ${"DEFAULT"}`); + }); + test("throws an error for unknown dialects", () => { + assert.throws(() => sql`SELECT ${sql.variant({duckdb: "DUCKDB"})}`.toDialect("postgres"), /missing variant/); + }); + test("converts nested SQL fragments", () => { + assert.deepStrictEqual( + sql`SELECT COUNT(*) FROM ${sql.view`SELECT ${sql.variant({duckdb: "DUCKDB"})}`}`.toDialect("duckdb"), // prettier-ignore + sql`SELECT COUNT(*) FROM ${sql.view`SELECT ${"DUCKDB"}`}` + ); + }); + test("converts nested SQL variants", () => { + const v1 = sql.variant({duckdb: "DUCKDB"}); + const v2 = sql.variant({duckdb: sql`${v1}`}); + assert.deepStrictEqual(v2.toDialect("duckdb"), sql`${"DUCKDB"}`); + }); + test("converts nested SQL variants (2)", () => { + const v1 = sql.variant({duckdb: "DUCKDB"}); + const v2 = sql.variant({duckdb: v1}); + assert.deepStrictEqual(v2.toDialect("duckdb"), "DUCKDB"); + }); + test("handles dialect- and non-dialect-specific fragments", () => { + assert.deepStrictEqual( + sql`SELECT ${sql`COUNT(*)`} FROM ${sql.view`SELECT ${sql.variant({duckdb: "DUCKDB"})}`}`.toDialect("duckdb"), // prettier-ignore + sql`SELECT ${sql`COUNT(*)`} FROM ${sql.view`SELECT ${"DUCKDB"}`}` + ); + assert.deepStrictEqual( + sql`SELECT COUNT(*) FROM ${sql.view`SELECT ${sql.variant({duckdb: "DUCKDB"})}`} WHERE ${sql`FOO`}`.toDialect("duckdb"), // prettier-ignore + sql`SELECT COUNT(*) FROM ${sql.view`SELECT ${"DUCKDB"}`} WHERE ${sql`FOO`}` + ); + }); + test("returns itself if the fragment is not dialect-specific", () => { + const fragment = sql`SELECT ${sql`1 + ${2}`}`; + assert.strictEqual(fragment.toDialect("duckdb"), fragment); + assert.strictEqual(fragment.toDialect(), fragment); + }); +}); + +describe("sql`…`.query(database)", () => { + test("queries the specified database via database.sql", async () => { + const args: unknown[] = []; + const result: QueryResult = Object.assign([], {schema: [], date: new Date()}); + const output = await sql`SELECT * FROM ${sql`PURCHASES`}`.query({ + name: "", + options: {}, + async sql(strings: Readonly, ...params: unknown[]) { + args.push(strings, ...params); + return result as QueryResult; + } + }); + assert.strictEqual(output, result); + assert.deepStrictEqual(args, [["SELECT * FROM PURCHASES"]]); + }); + test("flattens any views", async () => { + const view = sql.view`SELECT * FROM PURCHASES WHERE FOO = ${42}`; + const args: unknown[] = []; + const result: QueryResult = Object.assign([], {schema: [], date: new Date()}); + const output = await sql`SELECT * FROM ${view}`.query({ + name: "", + options: {}, + async sql(strings: Readonly, ...params: unknown[]) { + args.push(strings, ...params); + return result as QueryResult; + } + }); + assert.strictEqual(output, result); + assert.deepStrictEqual(args, [ + [ + `WITH +"_1" AS (SELECT * FROM PURCHASES WHERE FOO = `, + `) +SELECT * FROM "_1"` + ], + 42 + ]); + }); + test("flattens chained views", async () => { + const inner = sql.view`SELECT * FROM PURCHASES WHERE X = ${1}`; + const outer = sql.view`SELECT * FROM ${inner}`; + const args: unknown[] = []; + const result: QueryResult = Object.assign([], {schema: [], date: new Date()}); + const output = await sql`SELECT * FROM ${outer}`.query({ + name: "", + options: {}, + async sql(strings: Readonly, ...params: unknown[]) { + args.push(strings, ...params); + return result as QueryResult; + } + }); + assert.strictEqual(output, result); + assert.deepStrictEqual(args, [ + [ + `WITH +"_2" AS (SELECT * FROM PURCHASES WHERE X = `, + `), +"_1" AS (SELECT * FROM "_2") +SELECT * FROM "_1"` + ], + 1 + ]); + }); + test("handles matching dialect-specific variants", async () => { + const args: unknown[] = []; + const result: QueryResult = Object.assign([], {schema: [], date: new Date()}); + const query = sql`SELECT * FROM ${sql.variant({duckdb: sql`PURCHASES_DUCKDB`, default: sql`PURCHASES`})}`; + const output = await query.query({ + name: "", + options: {}, + dialect: "duckdb", + async sql(strings: Readonly, ...params: unknown[]) { + args.push(strings, ...params); + return result as QueryResult; + } + }); + assert.strictEqual(output, result); + assert.deepStrictEqual(args, [["SELECT * FROM PURCHASES_DUCKDB"]]); + }); + test("handles default dialect-specific variants", async () => { + const args: unknown[] = []; + const result: QueryResult = Object.assign([], {schema: [], date: new Date()}); + const query = sql`SELECT * FROM ${sql.variant({duckdb: sql`PURCHASES_DUCKDB`, default: sql`PURCHASES`})}`; + const output = await query.query({ + name: "", + options: {}, + dialect: "postgres", + async sql(strings: Readonly, ...params: unknown[]) { + args.push(strings, ...params); + return result as QueryResult; + } + }); + assert.strictEqual(output, result); + assert.deepStrictEqual(args, [["SELECT * FROM PURCHASES"]]); + }); + test("does not duplicate dialect-specific views", async () => { + const view = sql.view`SELECT v FROM ${sql.ident("table")}`; + const query = sql`SELECT COUNT(*) FROM ${view} +UNION ALL SELECT SUM(v) FROM ${view}`.flat("duckdb"); + assert.deepStrictEqual( + query, + sql`WITH +\"_1\" AS (SELECT v FROM \"table\") +SELECT COUNT(*) FROM \"_1\" +UNION ALL SELECT SUM(v) FROM \"_1\"` + ); + }); +}); + +describe("sql.ident(name)", () => { + test("quotes a name", () => { + assert.deepStrictEqual(sql.ident("foo").toDialect(), sql`"foo"`); + assert.deepStrictEqual(sql.ident("foo").toDialect("databricks"), sql`\`foo\``); + }); + test("quotes a name with quotes", () => { + assert.deepStrictEqual(sql.ident('fo"c"sle').toDialect(), sql`"fo""c""sle"`); + assert.deepStrictEqual(sql.ident('fo"c"sle').toDialect("databricks"), sql`\`fo"c"sle\``); + assert.deepStrictEqual(sql.ident("fo`c`sle").toDialect(), sql`"fo\`c\`sle"`); + assert.deepStrictEqual(sql.ident("fo`c`sle").toDialect("databricks"), sql`\`fo\`\`c\`\`sle\``); + }); +}); + +describe("sql.text(value)", () => { + test("quotes a value", () => { + assert.deepStrictEqual(sql.text("foo").toDialect(), sql`'foo'`); + assert.deepStrictEqual(sql.text("foo").toDialect("databricks"), sql`'foo'`); + }); + test("quotes a value with quotes", () => { + assert.deepStrictEqual(sql.text("fo'c'sle").toDialect(), sql`'fo''c''sle'`); + assert.deepStrictEqual(sql.text("fo'c'sle").toDialect("databricks"), sql`'fo''c''sle'`); + }); +}); + +describe("sql.variant(variants)", () => { + test("defines dialect-specific SQL variants", () => { + const duckdb = sql`APPROX_QUANTILE`; + const snowflake = sql`APPROX_PERCENTILE`; + const variant = sql.variant({duckdb, snowflake}); + assert.deepStrictEqual(variant.toDialect("duckdb"), duckdb); + assert.deepStrictEqual(variant.toDialect("snowflake"), snowflake); + }); + test("allows a default dialect", () => { + const duckdb = sql`APPROX_QUANTILE`; + const snowflake = sql`APPROX_PERCENTILE`; + const variant = sql.variant({default: duckdb, snowflake}); + assert.deepStrictEqual(variant.toDialect("duckdb"), duckdb); + assert.deepStrictEqual(variant.toDialect("snowflake"), snowflake); + assert.deepStrictEqual(variant.toDialect("postgres"), duckdb); + assert.deepStrictEqual(variant.toDialect("unknown" as SqlDialect), duckdb); + assert.deepStrictEqual(variant.toDialect(), duckdb); + }); + test("supports variants with parameters", () => { + const value = sql`PRICE_PER_UNIT`; + const p = 0.5; + const duckdb = sql`APPROX_QUANTILE(${value}, ${p})`; + const snowflake = sql`APPROX_PERCENTILE(${value}, ${p})`; + const postgres = sql`percentile_disc(${p}) WITHIN GROUP (ORDER BY ${value})`; + const variant = sql.variant({duckdb, snowflake, postgres}); + assert.deepStrictEqual(variant.toDialect("duckdb"), duckdb); + assert.deepStrictEqual(variant.toDialect("postgres"), postgres); + assert.deepStrictEqual(variant.toDialect("snowflake"), snowflake); + }); + test("supports literal variants", () => { + const variant = sql.variant({duckdb: 0, snowflake: 1, postgres: 2}); + assert.strictEqual(variant.toDialect("duckdb"), 0); + assert.strictEqual(variant.toDialect("snowflake"), 1); + assert.strictEqual(variant.toDialect("postgres"), 2); + }); + test("throws an error given an unsupported dialect", () => { + const duckdb = sql`APPROX_QUANTILE`; + const snowflake = sql`APPROX_PERCENTILE`; + const variant = sql.variant({duckdb, snowflake}); + assert.throws(() => variant.toDialect("postgres"), /missing variant: postgres/); + assert.throws(() => variant.toDialect(), /missing dialect/); + }); + test("implements toString", () => { + assert.strictEqual(String(sql.variant({default: sql`"Shipping Address State"`})), '"Shipping Address State"'); + }); +}); diff --git a/src/runtime/stdlib/sql.ts b/src/runtime/stdlib/sql.ts new file mode 100644 index 0000000..93bf5a8 --- /dev/null +++ b/src/runtime/stdlib/sql.ts @@ -0,0 +1,297 @@ +import type {DatabaseClient, QueryResult, SqlDialect} from "./databaseClient.js"; + +export function sql(strings: Readonly, ...params: unknown[]): SqlFragment { + return new SqlFragment(strings, params); +} + +sql.view = function view( + strings: Readonly, + ...params: unknown[] +): SqlView { + return new SqlView(strings, params); +}; + +sql.variant = function variant( + variants: Partial> +): SqlVariant { + return new SqlVariant(variants); +}; + +sql.ident = function ident(name: string): SqlVariant { + return new SqlVariant({ + get databricks() { + return sql([tquote(name)]); + }, + get bigquery() { + return sql([tquote(name)]); + }, + get default() { + return sql([dquote(name)]); + } + }); +}; + +sql.text = function text(value: string): SqlFragment { + return sql([squote(value)]); +}; + +class SqlFragment { + readonly strings: Readonly; + readonly params: unknown[]; + constructor(strings: Readonly, params: unknown[]) { + this.strings = strings; + this.params = params; + } + flat(dialect?: SqlDialect): typeof this { + const iquote = getIquote(dialect); + const that = this.toDialect(dialect); + const {strings, params} = that; + const names = findUndernames(strings, params); // in-use table names, including ctes + const map = new Map(); // from view to cte name + const views: Record = {}; + let cteIndex = 0; + + // Bakes any SQL view or fragment params into the query, recursively. + // Has the side-effect of populating the views map (of CTEs). + function flatParams( + istrings: Readonly, + iparams: unknown[] + ): [Readonly, unknown[]] { + let ostrings: string[] | undefined; + let oparams!: unknown[]; + for (let i = 0; i < iparams.length; ++i) { + const param = iparams[i]; + const string = istrings[i + 1]; + if (param instanceof SqlView) { + if (ostrings === undefined) { + ostrings = istrings.slice(0, i + 1); + oparams = iparams.slice(0, i); + } + let key = map.get(param); + if (key == null) { + do key = `_${++cteIndex}`; + while (names.has(key)); + names.add(key); + map.set(param, key); + const [pstrings, pparams] = flatParams(param.strings, param.params); + views[key] = new SqlFragment(pstrings, pparams); + } + ostrings[ostrings.length - 1] += iquote(key) + string; + } else if (param instanceof SqlFragment) { + if (ostrings === undefined) { + ostrings = istrings.slice(0, i + 1); + oparams = iparams.slice(0, i); + } + const [pstrings, pparams] = flatParams(param.strings, param.params); + ostrings[ostrings.length - 1] += pstrings[0]; + ostrings.push(...pstrings.slice(1)); + ostrings[ostrings.length - 1] += string; + oparams.push(...pparams); + } else if (ostrings !== undefined) { + oparams.push(param); + ostrings.push(string); + } + } + return ostrings === undefined + ? [istrings, iparams] + : [ostrings, oparams]; + } + + const [fstrings, fparams] = flatParams(strings, params); + const [vstrings, vparams] = withViews(fstrings, fparams, views, iquote); + return vstrings === strings && vparams === params + ? that + : reconstruct(that, vstrings, vparams); + } + query(database: DatabaseClient): Promise> { + const {strings, params} = this.flat(database.dialect); + return database.sql(strings, ...params); + } + toDialect(dialect?: SqlDialect): typeof this { + return fragmentToDialect(this, dialect, new Map()); + } + toString() { + return this.flat().strings.join("?"); + } +} + +class SqlView extends SqlFragment { + constructor(strings: Readonly, params: unknown[]) { + super(strings, params); + } +} + +class SqlVariant { + readonly variants: Partial>; + constructor(variants: Partial>) { + this.variants = variants; + } + toDialect(dialect?: SqlDialect): unknown { + return variantToDialect(this, dialect, new Map()); + } + toString() { + return String(this.toDialect()); + } +} + +sql.Fragment = SqlFragment; +sql.View = SqlView; +sql.Variant = SqlVariant; + +function fragmentToDialect( + fragment: T, + dialect: SqlDialect | undefined, + cache: Map +): T { + const {strings, params} = fragment; + let ostrings: string[] | undefined; + let oparams!: unknown[]; + for (let i = 0; i < params.length; ++i) { + const param = params[i]; + const string = strings[i + 1]; + if (param instanceof SqlFragment) { + let dparam: SqlFragment; + if (cache.has(param)) dparam = cache.get(param) as SqlFragment; + else cache.set(param, (dparam = fragmentToDialect(param, dialect, cache))); + if (dparam !== param) { + if (ostrings === undefined) { + ostrings = strings.slice(0, i + 1); + oparams = params.slice(0, i); + } + ostrings.push(string); + oparams.push(dparam); + } else if (ostrings !== undefined) { + ostrings.push(string); + oparams.push(param); + } + } else if (param instanceof SqlVariant) { + let dparam: unknown; + if (cache.has(param)) dparam = cache.get(param); + else cache.set(param, (dparam = variantToDialect(param, dialect, cache))); + if (ostrings === undefined) { + ostrings = strings.slice(0, i + 1); + oparams = params.slice(0, i); + } + ostrings.push(string); + oparams.push(dparam); + } else if (ostrings !== undefined) { + ostrings.push(string); + oparams.push(param); + } + } + return ostrings === undefined + ? fragment + : reconstruct(fragment, ostrings, oparams); +} + +function reconstruct( + fragment: T, + strings: Readonly, + params: unknown[] +): T { + return new (fragment.constructor as typeof SqlFragment)(strings, params) as T; +} + +function variantToDialect( + variant: SqlVariant, + dialect: SqlDialect | undefined, + cache: Map +): unknown { + let v: unknown; + if (dialect !== undefined && dialect in variant.variants) v = variant.variants[dialect]; + else if ("default" in variant.variants) v = variant.variants.default; + else throw new Error(dialect ? `missing variant: ${dialect}` : `missing dialect`); + return v instanceof SqlFragment + ? fragmentToDialect(v, dialect, cache) + : v instanceof SqlVariant + ? variantToDialect(v, dialect, cache) + : v; +} + +// Assumptions: +// - the views’ names are unique and non-conflicting +// - the views are defined in topological order +// - the views do not reference any other views +// - the views do not reference any other tables +function withViews( + istrings: Readonly, + iparams: unknown[], + views: Record, + iquote: (name: string) => string +): [Readonly, unknown[]] { + const entries = Object.entries(views); + if (!entries.length) return [istrings, iparams]; + const input = istrings[0]; + const withIndex = findWith(input); + const ostrings: string[] = []; + const oparams: unknown[] = []; + let first = true; + ostrings[0] = `${withIndex >= 0 ? input.slice(0, withIndex) : "WITH"}\n`; + for (const [name, view] of entries) { + if (view.params.some((p) => p instanceof SqlFragment)) + throw new Error("nested fragment"); + if (first) first = false; + else ostrings[ostrings.length - 1] += ",\n"; + ostrings[ostrings.length - 1] += `${iquote(name)} AS (${view.strings[0]}`; + ostrings.push(...view.strings.slice(1)); + ostrings[ostrings.length - 1] += ")"; + oparams.push(...view.params); + } + ostrings[ostrings.length - 1] += + withIndex >= 0 ? `,\n${input.slice(withIndex)}` : `\n${input}`; + ostrings.push(...istrings.slice(1)); + oparams.push(...iparams); + return [ostrings, oparams]; +} + +function findWith(input: string): number { + const match = /^\s*(--.*\n|\/\*[\s\S]*?\*\/|\s)*with\b(\s+recursive\b)?/i.exec(input); + return match ? match[0].length : -1; +} + +function findUndernames( + strings: Readonly, + params: unknown[], + names = new Set() +): Set { + const string = strings.join(" "); + const pattern = /\b_\d+\b/g; + for ( + let match: RegExpExecArray | null; + (match = pattern.exec(string)) !== null; + ) { + names.add(match[0]); + } + for (const param of params) { + if (param instanceof SqlFragment) { + findUndernames(param.strings, param.params, names); + } + } + return names; +} + +/** Quotes the specified SQL identifier. */ +function getIquote(dialect?: SqlDialect): (name: string) => string { + switch (dialect) { + case "databricks": + case "bigquery": + return tquote; + default: + return dquote; + } +} + +/** Quotes the specified name with double quotes. */ +function dquote(name: string): string { + return `"${name.replace(/"/g, '""')}"`; +} + +/** Quotes the specified name with backticks. */ +function tquote(name: string): string { + return `\`${name.replace(/`/g, "``")}\``; +} + +/** Quotes the specified name with single quotes. */ +function squote(name: string): string { + return `'${name.replace(/'/g, "''")}'`; +}