/** * Type-safe SQL Expression Builders * * Provides builder functions for creating type-safe SQL expressions. * All comparisons are validated at compile time. */ import type { SQLExpression, ColumnRef, ComparisonOperator, AggregateExpr, TableDefinition, ColumnRefs, } from './types.js' // ============================================================================ // Expression Implementation // ============================================================================ /** * Internal implementation of SQL expression */ class SQLExpressionImpl implements SQLExpression { readonly _type!: TType readonly _brand = 'sql_expression' as const constructor( private sqlGen: () => { sql: string; params: unknown[] } ) {} toSQL(): { sql: string; params: unknown[] } { return this.sqlGen() } } /** * Column reference implementation */ class ColumnRefImpl< TTable extends string, TColumn extends string, TType, > extends SQLExpressionImpl implements ColumnRef { readonly table: TTable readonly column: TColumn constructor(table: TTable, column: TColumn) { super(() => ({ sql: `"${this.table}"."${this.column}"`, params: [], })) this.table = table this.column = column } } /** * Literal value implementation */ class LiteralImpl extends SQLExpressionImpl { readonly value: TType constructor(value: TType, private paramIndex: () => number) { super(() => ({ sql: `$${this.paramIndex()}`, params: [this.value], })) this.value = value } } // ============================================================================ // Comparison Operators // ============================================================================ /** * Create an equality comparison (=) */ export function eq( left: SQLExpression, right: TType | SQLExpression ): SQLExpression { return createComparison(left, '=', right) } /** * Create an inequality comparison (!=) */ export function ne( left: SQLExpression, right: TType | SQLExpression ): SQLExpression { return createComparison(left, '!=', right) } /** * Create a less than comparison (<) */ export function lt( left: SQLExpression, right: TType | SQLExpression ): SQLExpression { return createComparison(left, '<', right) } /** * Create a less than or equal comparison (<=) */ export function lte( left: SQLExpression, right: TType | SQLExpression ): SQLExpression { return createComparison(left, '<=', right) } /** * Create a greater than comparison (>) */ export function gt( left: SQLExpression, right: TType | SQLExpression ): SQLExpression { return createComparison(left, '>', right) } /** * Create a greater than or equal comparison (>=) */ export function gte( left: SQLExpression, right: TType | SQLExpression ): SQLExpression { return createComparison(left, '>=', right) } /** * Create a LIKE comparison */ export function like( column: SQLExpression, pattern: string ): SQLExpression { return createComparison(column, 'LIKE', pattern) } /** * Create a case-insensitive LIKE comparison (ILIKE) */ export function ilike( column: SQLExpression, pattern: string ): SQLExpression { return createComparison(column, 'ILIKE', pattern) } /** * Create an IN comparison */ export function inArray( column: SQLExpression, values: TType[] ): SQLExpression { return new SQLExpressionImpl(() => { const colSql = column.toSQL() const placeholders = values.map((_, i) => `$${colSql.params.length + i + 1}`).join(', ') return { sql: `${colSql.sql} IN (${placeholders})`, params: [...colSql.params, ...values], } }) } /** * Create a NOT IN comparison */ export function notInArray( column: SQLExpression, values: TType[] ): SQLExpression { return new SQLExpressionImpl(() => { const colSql = column.toSQL() const placeholders = values.map((_, i) => `$${colSql.params.length + i + 1}`).join(', ') return { sql: `${colSql.sql} NOT IN (${placeholders})`, params: [...colSql.params, ...values], } }) } /** * Create an IS NULL comparison */ export function isNull( column: SQLExpression ): SQLExpression { return new SQLExpressionImpl(() => { const colSql = column.toSQL() return { sql: `${colSql.sql} IS NULL`, params: colSql.params, } }) } /** * Create an IS NOT NULL comparison */ export function isNotNull( column: SQLExpression ): SQLExpression { return new SQLExpressionImpl(() => { const colSql = column.toSQL() return { sql: `${colSql.sql} IS NOT NULL`, params: colSql.params, } }) } /** * Create a BETWEEN comparison */ export function between( column: SQLExpression, min: TType, max: TType ): SQLExpression { return new SQLExpressionImpl(() => { const colSql = column.toSQL() return { sql: `${colSql.sql} BETWEEN $${colSql.params.length + 1} AND $${colSql.params.length + 2}`, params: [...colSql.params, min, max], } }) } /** * Create a NOT BETWEEN comparison */ export function notBetween( column: SQLExpression, min: TType, max: TType ): SQLExpression { return new SQLExpressionImpl(() => { const colSql = column.toSQL() return { sql: `${colSql.sql} NOT BETWEEN $${colSql.params.length + 1} AND $${colSql.params.length + 2}`, params: [...colSql.params, min, max], } }) } // ============================================================================ // Logical Operators // ============================================================================ /** * Create an AND expression */ export function and( ...expressions: SQLExpression[] ): SQLExpression { if (expressions.length === 0) { return new SQLExpressionImpl(() => ({ sql: 'TRUE', params: [] })) } if (expressions.length === 1) { return expressions[0]! } return new SQLExpressionImpl(() => { const parts: string[] = [] const params: unknown[] = [] for (const expr of expressions) { const { sql, params: exprParams } = expr.toSQL() // Adjust parameter placeholders const adjustedSql = sql.replace(/\$(\d+)/g, (_, n) => `$${parseInt(n) + params.length}`) parts.push(`(${adjustedSql})`) params.push(...exprParams) } return { sql: parts.join(' AND '), params, } }) } /** * Create an OR expression */ export function or( ...expressions: SQLExpression[] ): SQLExpression { if (expressions.length === 0) { return new SQLExpressionImpl(() => ({ sql: 'FALSE', params: [] })) } if (expressions.length === 1) { return expressions[0]! } return new SQLExpressionImpl(() => { const parts: string[] = [] const params: unknown[] = [] for (const expr of expressions) { const { sql, params: exprParams } = expr.toSQL() // Adjust parameter placeholders const adjustedSql = sql.replace(/\$(\d+)/g, (_, n) => `$${parseInt(n) + params.length}`) parts.push(`(${adjustedSql})`) params.push(...exprParams) } return { sql: parts.join(' OR '), params, } }) } /** * Create a NOT expression */ export function not( expression: SQLExpression ): SQLExpression { return new SQLExpressionImpl(() => { const { sql, params } = expression.toSQL() return { sql: `NOT (${sql})`, params, } }) } // ============================================================================ // Aggregate Functions // ============================================================================ /** * Create a COUNT aggregate */ export function count(): AggregateExpr export function count(column: SQLExpression): AggregateExpr export function count(column?: SQLExpression): AggregateExpr { return createAggregate('COUNT', column ?? '*') } /** * Create a COUNT DISTINCT aggregate */ export function countDistinct(column: SQLExpression): AggregateExpr { return createAggregate('COUNT', column, true) } /** * Create a SUM aggregate */ export function sum(column: SQLExpression): AggregateExpr { return createAggregate('SUM', column) } /** * Create an AVG aggregate */ export function avg(column: SQLExpression): AggregateExpr { return createAggregate('AVG', column) } /** * Create a MIN aggregate */ export function min(column: SQLExpression): AggregateExpr { return createAggregate('MIN', column) } /** * Create a MAX aggregate */ export function max(column: SQLExpression): AggregateExpr { return createAggregate('MAX', column) } // ============================================================================ // String Functions // ============================================================================ /** * Create a CONCAT expression */ export function concat( ...values: (SQLExpression | string)[] ): SQLExpression { return new SQLExpressionImpl(() => { const parts: string[] = [] const params: unknown[] = [] for (const value of values) { if (typeof value === 'string') { parts.push(`$${params.length + 1}`) params.push(value) } else { const { sql, params: exprParams } = value.toSQL() const adjustedSql = sql.replace(/\$(\d+)/g, (_, n) => `$${parseInt(n) + params.length}`) parts.push(adjustedSql) params.push(...exprParams) } } return { sql: `CONCAT(${parts.join(', ')})`, params, } }) } /** * Create a LOWER expression */ export function lower(value: SQLExpression): SQLExpression { return new SQLExpressionImpl(() => { const { sql, params } = value.toSQL() return { sql: `LOWER(${sql})`, params, } }) } /** * Create an UPPER expression */ export function upper(value: SQLExpression): SQLExpression { return new SQLExpressionImpl(() => { const { sql, params } = value.toSQL() return { sql: `UPPER(${sql})`, params, } }) } /** * Create a TRIM expression */ export function trim(value: SQLExpression): SQLExpression { return new SQLExpressionImpl(() => { const { sql, params } = value.toSQL() return { sql: `TRIM(${sql})`, params, } }) } /** * Create a COALESCE expression */ export function coalesce( ...values: (SQLExpression | TType)[] ): SQLExpression { return new SQLExpressionImpl(() => { const parts: string[] = [] const params: unknown[] = [] for (const value of values) { if (typeof value === 'object' && value !== null && '_brand' in value) { const expr = value as SQLExpression const { sql, params: exprParams } = expr.toSQL() const adjustedSql = sql.replace(/\$(\d+)/g, (_, n) => `$${parseInt(n) + params.length}`) parts.push(adjustedSql) params.push(...exprParams) } else { parts.push(`$${params.length + 1}`) params.push(value) } } return { sql: `COALESCE(${parts.join(', ')})`, params, } }) } // ============================================================================ // Helper Functions // ============================================================================ /** * Create a column reference */ export function col< TTable extends string, TColumn extends string, TType = unknown, >(table: TTable, column: TColumn): ColumnRef { return new ColumnRefImpl(table, column) } /** * Create column references for a table definition */ export function columnsOf(table: T): ColumnRefs { const refs = {} as ColumnRefs for (const key of Object.keys(table.columns)) { const columnDef = table.columns[key]! refs[key as keyof typeof refs] = new ColumnRefImpl( table.tableName, columnDef.name ) as ColumnRefs[keyof typeof refs] } return refs } /** * Create a raw SQL expression (use with caution!) */ export function raw( sql: string, params: unknown[] = [] ): SQLExpression { return new SQLExpressionImpl(() => ({ sql, params })) } /** * Create a literal value expression */ export function literal(value: TType): SQLExpression { let paramIdx = 1 return new LiteralImpl(value, () => paramIdx++) } /** * Create a SQL expression from a template literal */ export function sql( strings: TemplateStringsArray, ...values: unknown[] ): SQLExpression { return new SQLExpressionImpl(() => { const parts: string[] = [] const params: unknown[] = [] for (let i = 0; i < strings.length; i++) { parts.push(strings[i]!) if (i < values.length) { const value = values[i] if (typeof value === 'object' && value !== null && '_brand' in value) { const expr = value as SQLExpression const { sql, params: exprParams } = expr.toSQL() const adjustedSql = sql.replace(/\$(\d+)/g, (_, n) => `$${parseInt(n) + params.length}`) parts.push(adjustedSql) params.push(...exprParams) } else { parts.push(`$${params.length + 1}`) params.push(value) } } } return { sql: parts.join(''), params, } }) } // ============================================================================ // Internal Helpers // ============================================================================ function createComparison( left: SQLExpression, operator: ComparisonOperator, right: TType | SQLExpression ): SQLExpression { return new SQLExpressionImpl(() => { const leftSql = left.toSQL() let rightSql: { sql: string; params: unknown[] } if (typeof right === 'object' && right !== null && '_brand' in right) { const rightExpr = right as SQLExpression const { sql, params } = rightExpr.toSQL() // Adjust parameter indices rightSql = { sql: sql.replace(/\$(\d+)/g, (_, n) => `$${parseInt(n) + leftSql.params.length}`), params, } } else { rightSql = { sql: `$${leftSql.params.length + 1}`, params: [right], } } return { sql: `${leftSql.sql} ${operator} ${rightSql.sql}`, params: [...leftSql.params, ...rightSql.params], } }) } function createAggregate( fn: 'COUNT' | 'SUM' | 'AVG' | 'MIN' | 'MAX', column: SQLExpression | '*', distinct = false ): AggregateExpr { const baseExpr = new SQLExpressionImpl(() => { if (column === '*') { return { sql: `${fn}(*)`, params: [], } } const { sql, params } = column.toSQL() const distinctClause = distinct ? 'DISTINCT ' : '' return { sql: `${fn}(${distinctClause}${sql})`, params, } }) // Create aggregate expression with required properties const aggregateExpr = baseExpr as unknown as AggregateExpr // Add aggregate properties Object.defineProperties(aggregateExpr, { fn: { value: fn, writable: false, enumerable: true }, column: { value: column, writable: false, enumerable: true }, distinct: { value: distinct, writable: false, enumerable: true }, }) return aggregateExpr }