From 10a8622e7ceb40ff39aa66beec7127a1404bf210 Mon Sep 17 00:00:00 2001 From: Maxime Duchene-Savard Date: Tue, 23 Jun 2026 23:44:01 -0400 Subject: [PATCH] added sql builder framework --- src/main/kotlin/dev/mduchene/sql/Sql.kt | 1021 +++++++++++++++++ .../kotlin/dev/mduchene/sql/SqlBuilderTest.kt | 135 +++ 2 files changed, 1156 insertions(+) create mode 100644 src/main/kotlin/dev/mduchene/sql/Sql.kt create mode 100644 src/test/kotlin/dev/mduchene/sql/SqlBuilderTest.kt diff --git a/src/main/kotlin/dev/mduchene/sql/Sql.kt b/src/main/kotlin/dev/mduchene/sql/Sql.kt new file mode 100644 index 0000000..70a10a3 --- /dev/null +++ b/src/main/kotlin/dev/mduchene/sql/Sql.kt @@ -0,0 +1,1021 @@ +package dev.mduchene.sql + +data class SqlStatement( + val sql: String, + val parameters: List, +) + +interface SqlDialect { + val name: String + val supportsReturning: Boolean + + fun placeholder(parameterIndex: Int): String + fun quoteIdentifier(identifierPart: String): String +} + +object PostgresDialect : SqlDialect { + override val name = "postgres" + override val supportsReturning = true + + override fun placeholder(parameterIndex: Int): String = "\$$parameterIndex" + + override fun quoteIdentifier(identifierPart: String): String = + "\"${identifierPart.replace("\"", "\"\"")}\"" +} + +object AnsiDialect : SqlDialect { + override val name = "ansi" + override val supportsReturning = false + + override fun placeholder(parameterIndex: Int): String = "?" + + override fun quoteIdentifier(identifierPart: String): String = + "\"${identifierPart.replace("\"", "\"\"")}\"" +} + +class SqlRenderContext( + val dialect: SqlDialect, +) { + private val values = mutableListOf() + + val parameters: List + get() = values.toList() + + fun bind(value: Any?): String { + values += value + return dialect.placeholder(values.size) + } + + fun quoteIdentifier(identifier: String): String = + identifierParts(identifier).joinToString(".") { part -> + if (part == "*") "*" else dialect.quoteIdentifier(part) + } +} + +interface SqlFragment { + fun render(ctx: SqlRenderContext, out: StringBuilder) +} + +interface SqlExpression : SqlFragment + +interface SqlTable : SqlFragment + +interface SqlStatementBuilder : SqlFragment { + fun toSql(dialect: SqlDialect = PostgresDialect): SqlStatement { + val ctx = SqlRenderContext(dialect) + val out = StringBuilder() + render(ctx, out) + return SqlStatement(out.toString(), ctx.parameters) + } +} + +object Sql { + fun select(): SelectBuilder = SelectBuilder() + + fun select(vararg columns: String): SelectBuilder = + SelectBuilder().select(*columns) + + fun select(vararg expressions: SqlExpression): SelectBuilder = + SelectBuilder().select(*expressions) + + fun insertInto(table: String): InsertBuilder = InsertBuilder(table) + + fun update(table: String): UpdateBuilder = UpdateBuilder(table) + + fun deleteFrom(table: String): DeleteBuilder = DeleteBuilder(table) + + fun col(name: String): SqlExpression = ColumnExpression(name) + + fun column(name: String): SqlExpression = col(name) + + fun table(name: String, alias: String? = null): SqlTable = + NamedTable(name, alias) + + fun value(value: Any?): SqlExpression = ValueExpression(value) + + fun raw(sql: String, vararg parameters: Any?): SqlExpression = + RawExpression(sql, parameters.toList()) + + fun count(expression: SqlExpression = StarExpression, distinct: Boolean = false): SqlExpression = + FunctionExpression("COUNT", listOf(expression), distinct) + + fun countDistinct(expression: SqlExpression): SqlExpression = + count(expression, distinct = true) + + fun exists(query: SelectBuilder): SqlExpression = ExistsExpression(query) + + fun and(vararg conditions: SqlExpression): SqlExpression = + CompoundExpression("AND", conditions.toList()).normalized() + + fun or(vararg conditions: SqlExpression): SqlExpression = + CompoundExpression("OR", conditions.toList()).normalized() +} + +enum class SortDirection(private val sql: String) { + ASC("ASC"), + DESC("DESC"); + + override fun toString(): String = sql +} + +enum class NullsOrder(private val sql: String) { + FIRST("FIRST"), + LAST("LAST"); + + override fun toString(): String = sql +} + +data class OrderTerm( + val expression: SqlExpression, + val direction: SortDirection = SortDirection.ASC, + val nulls: NullsOrder? = null, +) : SqlFragment { + override fun render(ctx: SqlRenderContext, out: StringBuilder) { + expression.render(ctx, out) + out.append(' ').append(direction) + if (nulls != null) { + out.append(" NULLS ").append(nulls) + } + } +} + +fun SqlExpression.asAlias(alias: String): SqlExpression = + AliasedExpression(this, alias) + +fun SqlExpression.asc(nulls: NullsOrder? = null): OrderTerm = + OrderTerm(this, SortDirection.ASC, nulls) + +fun SqlExpression.desc(nulls: NullsOrder? = null): OrderTerm = + OrderTerm(this, SortDirection.DESC, nulls) + +fun SqlExpression.eq(value: Any?): SqlExpression = + if (value == null) isNull() else BinaryExpression(this, "=", value.toExpression()) + +fun SqlExpression.eq(other: SqlExpression): SqlExpression = + BinaryExpression(this, "=", other) + +fun SqlExpression.ne(value: Any?): SqlExpression = + if (value == null) isNotNull() else BinaryExpression(this, "<>", value.toExpression()) + +fun SqlExpression.ne(other: SqlExpression): SqlExpression = + BinaryExpression(this, "<>", other) + +fun SqlExpression.gt(value: Any?): SqlExpression = + BinaryExpression(this, ">", value.toExpression()) + +fun SqlExpression.gt(other: SqlExpression): SqlExpression = + BinaryExpression(this, ">", other) + +fun SqlExpression.gte(value: Any?): SqlExpression = + BinaryExpression(this, ">=", value.toExpression()) + +fun SqlExpression.gte(other: SqlExpression): SqlExpression = + BinaryExpression(this, ">=", other) + +fun SqlExpression.lt(value: Any?): SqlExpression = + BinaryExpression(this, "<", value.toExpression()) + +fun SqlExpression.lt(other: SqlExpression): SqlExpression = + BinaryExpression(this, "<", other) + +fun SqlExpression.lte(value: Any?): SqlExpression = + BinaryExpression(this, "<=", value.toExpression()) + +fun SqlExpression.lte(other: SqlExpression): SqlExpression = + BinaryExpression(this, "<=", other) + +fun SqlExpression.like(pattern: Any?): SqlExpression = + BinaryExpression(this, "LIKE", pattern.toExpression()) + +fun SqlExpression.isNull(): SqlExpression = + UnaryPostfixExpression(this, "IS NULL") + +fun SqlExpression.isNotNull(): SqlExpression = + UnaryPostfixExpression(this, "IS NOT NULL") + +fun SqlExpression.between(start: Any?, end: Any?): SqlExpression = + BetweenExpression(this, start.toExpression(), end.toExpression()) + +fun SqlExpression.inValues(vararg values: Any?): SqlExpression = + InValuesExpression(this, values.toList(), negated = false) + +fun SqlExpression.notInValues(vararg values: Any?): SqlExpression = + InValuesExpression(this, values.toList(), negated = true) + +fun SqlExpression.inSubquery(query: SelectBuilder): SqlExpression = + InSubqueryExpression(this, query, negated = false) + +fun SqlExpression.notInSubquery(query: SelectBuilder): SqlExpression = + InSubqueryExpression(this, query, negated = true) + +fun SqlExpression.and(other: SqlExpression): SqlExpression = + CompoundExpression("AND", listOf(this, other)).normalized() + +fun SqlExpression.or(other: SqlExpression): SqlExpression = + CompoundExpression("OR", listOf(this, other)).normalized() + +fun SqlExpression.not(): SqlExpression = + PrefixExpression("NOT", this) + +class SelectBuilder : SqlStatementBuilder { + private val ctes = mutableListOf() + private val selections = mutableListOf() + private val joins = mutableListOf() + private val groupBy = mutableListOf() + private val orderBy = mutableListOf() + private val unions = mutableListOf() + private var distinct = false + private var from: SqlTable? = null + private var where: SqlExpression? = null + private var having: SqlExpression? = null + private var limit: Int? = null + private var offset: Int? = null + + fun with(name: String, query: SqlStatementBuilder, vararg columns: String): SelectBuilder = + apply { + ctes += CommonTableExpression(name, columns.toList(), query) + } + + fun distinct(): SelectBuilder = + apply { + distinct = true + } + + fun select(vararg columns: String): SelectBuilder = + apply { + selections += columns.map { Sql.col(it) } + } + + fun select(vararg expressions: SqlExpression): SelectBuilder = + apply { + selections += expressions + } + + fun from(table: String, alias: String? = null): SelectBuilder = + from(Sql.table(table, alias)) + + fun from(table: SqlTable): SelectBuilder = + apply { + from = table + } + + fun from(query: SelectBuilder, alias: String): SelectBuilder = + from(SubqueryTable(query, alias)) + + fun join(table: String, on: SqlExpression): SelectBuilder = + join(Sql.table(table), JoinType.INNER, on) + + fun join(table: String, alias: String, on: SqlExpression): SelectBuilder = + join(Sql.table(table, alias), JoinType.INNER, on) + + fun leftJoin(table: String, on: SqlExpression): SelectBuilder = + join(Sql.table(table), JoinType.LEFT, on) + + fun leftJoin(table: String, alias: String, on: SqlExpression): SelectBuilder = + join(Sql.table(table, alias), JoinType.LEFT, on) + + fun rightJoin(table: String, on: SqlExpression): SelectBuilder = + join(Sql.table(table), JoinType.RIGHT, on) + + fun rightJoin(table: String, alias: String, on: SqlExpression): SelectBuilder = + join(Sql.table(table, alias), JoinType.RIGHT, on) + + fun fullJoin(table: String, on: SqlExpression): SelectBuilder = + join(Sql.table(table), JoinType.FULL, on) + + fun fullJoin(table: String, alias: String, on: SqlExpression): SelectBuilder = + join(Sql.table(table, alias), JoinType.FULL, on) + + fun crossJoin(table: String, alias: String? = null): SelectBuilder = + join(Sql.table(table, alias), JoinType.CROSS, null) + + fun join(table: SqlTable, type: JoinType = JoinType.INNER, on: SqlExpression? = null): SelectBuilder = + apply { + require(type == JoinType.CROSS || on != null) { "${type.sql} requires an ON condition" } + joins += JoinClause(type, table, on) + } + + fun where(condition: SqlExpression): SelectBuilder = + apply { + where = condition + } + + fun andWhere(condition: SqlExpression): SelectBuilder = + apply { + where = where?.and(condition) ?: condition + } + + fun orWhere(condition: SqlExpression): SelectBuilder = + apply { + where = where?.or(condition) ?: condition + } + + fun groupBy(vararg columns: String): SelectBuilder = + apply { + groupBy += columns.map { Sql.col(it) } + } + + fun groupBy(vararg expressions: SqlExpression): SelectBuilder = + apply { + groupBy += expressions + } + + fun having(condition: SqlExpression): SelectBuilder = + apply { + having = condition + } + + fun andHaving(condition: SqlExpression): SelectBuilder = + apply { + having = having?.and(condition) ?: condition + } + + fun orderBy(vararg columns: String): SelectBuilder = + apply { + orderBy += columns.map { Sql.col(it).asc() } + } + + fun orderBy(vararg terms: OrderTerm): SelectBuilder = + apply { + orderBy += terms + } + + fun limit(limit: Int): SelectBuilder = + apply { + require(limit >= 0) { "LIMIT must be greater than or equal to zero" } + this.limit = limit + } + + fun offset(offset: Int): SelectBuilder = + apply { + require(offset >= 0) { "OFFSET must be greater than or equal to zero" } + this.offset = offset + } + + fun union(query: SelectBuilder): SelectBuilder = + apply { + unions += UnionClause(all = false, query) + } + + fun unionAll(query: SelectBuilder): SelectBuilder = + apply { + unions += UnionClause(all = true, query) + } + + override fun render(ctx: SqlRenderContext, out: StringBuilder) { + renderSingle(ctx, out) + unions.forEach { union -> + out.append(if (union.all) " UNION ALL " else " UNION ") + union.query.renderSingle(ctx, out) + } + } + + private fun renderSingle(ctx: SqlRenderContext, out: StringBuilder) { + renderCtes(ctes, ctx, out) + out.append("SELECT ") + if (distinct) { + out.append("DISTINCT ") + } + + if (selections.isEmpty()) { + StarExpression.render(ctx, out) + } else { + renderCommaSeparated(selections, ctx, out) + } + + if (from != null) { + out.append(" FROM ") + from?.render(ctx, out) + } + + if (joins.isNotEmpty()) { + require(from != null) { "JOIN clauses require a FROM clause" } + joins.forEach { join -> + out.append(' ') + join.render(ctx, out) + } + } + + if (where != null) { + out.append(" WHERE ") + where?.render(ctx, out) + } + + if (groupBy.isNotEmpty()) { + out.append(" GROUP BY ") + renderCommaSeparated(groupBy, ctx, out) + } + + if (having != null) { + out.append(" HAVING ") + having?.render(ctx, out) + } + + if (orderBy.isNotEmpty()) { + out.append(" ORDER BY ") + renderCommaSeparated(orderBy, ctx, out) + } + + if (limit != null) { + out.append(" LIMIT ").append(limit) + } + + if (offset != null) { + out.append(" OFFSET ").append(offset) + } + } +} + +class InsertBuilder( + private val table: String, +) : SqlStatementBuilder { + private val ctes = mutableListOf() + private val columns = mutableListOf() + private val rows = mutableListOf>() + private val returning = mutableListOf() + private var defaultValues = false + private var selectSource: SelectBuilder? = null + + fun with(name: String, query: SqlStatementBuilder, vararg columns: String): InsertBuilder = + apply { + ctes += CommonTableExpression(name, columns.toList(), query) + } + + fun columns(vararg columns: String): InsertBuilder = + apply { + require(this.columns.isEmpty()) { "INSERT columns are already set" } + this.columns += columns + this.columns.forEach(::identifierParts) + } + + fun values(vararg values: Any?): InsertBuilder = + apply { + require(!defaultValues) { "Cannot add VALUES after DEFAULT VALUES" } + require(selectSource == null) { "Cannot add VALUES after INSERT SELECT" } + require(columns.isNotEmpty()) { "INSERT columns must be set before values" } + require(values.size == columns.size) { + "VALUES size (${values.size}) must match INSERT columns size (${columns.size})" + } + rows += values.map { it.toExpression() } + } + + fun set(values: Map): InsertBuilder = + apply { + require(values.isNotEmpty()) { "INSERT values cannot be empty" } + if (columns.isEmpty()) { + columns += values.keys + columns.forEach(::identifierParts) + } else { + require(values.keys.toList() == columns) { + "INSERT map keys must match existing columns and order" + } + } + rows += values.values.map { it.toExpression() } + } + + fun defaultValues(): InsertBuilder = + apply { + require(columns.isEmpty() && rows.isEmpty() && selectSource == null) { + "DEFAULT VALUES cannot be combined with explicit columns, values, or INSERT SELECT" + } + defaultValues = true + } + + fun fromSelect(columns: List, query: SelectBuilder): InsertBuilder = + apply { + require(!defaultValues && rows.isEmpty()) { + "INSERT SELECT cannot be combined with DEFAULT VALUES or VALUES" + } + require(columns.isNotEmpty()) { "INSERT SELECT requires at least one column" } + this.columns.clear() + this.columns += columns + this.columns.forEach(::identifierParts) + selectSource = query + } + + fun returning(vararg columns: String): InsertBuilder = + apply { + returning += columns.map { Sql.col(it) } + } + + fun returning(vararg expressions: SqlExpression): InsertBuilder = + apply { + returning += expressions + } + + override fun render(ctx: SqlRenderContext, out: StringBuilder) { + renderCtes(ctes, ctx, out) + out.append("INSERT INTO ").append(ctx.quoteIdentifier(table)) + + if (columns.isNotEmpty()) { + out.append(" (") + out.append(columns.joinToString(", ") { ctx.quoteIdentifier(it) }) + out.append(')') + } + + when { + defaultValues -> out.append(" DEFAULT VALUES") + selectSource != null -> { + out.append(' ') + selectSource?.render(ctx, out) + } + rows.isNotEmpty() -> { + out.append(" VALUES ") + rows.forEachIndexed { rowIndex, row -> + if (rowIndex > 0) { + out.append(", ") + } + out.append('(') + renderCommaSeparated(row, ctx, out) + out.append(')') + } + } + else -> error("INSERT requires values, DEFAULT VALUES, or INSERT SELECT") + } + + renderReturning(returning, ctx, out) + } +} + +class UpdateBuilder( + private val table: String, +) : SqlStatementBuilder { + private val ctes = mutableListOf() + private val assignments = mutableListOf>() + private val returning = mutableListOf() + private var where: SqlExpression? = null + + fun with(name: String, query: SqlStatementBuilder, vararg columns: String): UpdateBuilder = + apply { + ctes += CommonTableExpression(name, columns.toList(), query) + } + + fun set(column: String, value: Any?): UpdateBuilder = + apply { + identifierParts(column) + assignments += column to value.toExpression() + } + + fun where(condition: SqlExpression): UpdateBuilder = + apply { + where = condition + } + + fun andWhere(condition: SqlExpression): UpdateBuilder = + apply { + where = where?.and(condition) ?: condition + } + + fun returning(vararg columns: String): UpdateBuilder = + apply { + returning += columns.map { Sql.col(it) } + } + + fun returning(vararg expressions: SqlExpression): UpdateBuilder = + apply { + returning += expressions + } + + override fun render(ctx: SqlRenderContext, out: StringBuilder) { + require(assignments.isNotEmpty()) { "UPDATE requires at least one assignment" } + + renderCtes(ctes, ctx, out) + out.append("UPDATE ").append(ctx.quoteIdentifier(table)).append(" SET ") + assignments.forEachIndexed { index, assignment -> + if (index > 0) { + out.append(", ") + } + out.append(ctx.quoteIdentifier(assignment.first)).append(" = ") + assignment.second.render(ctx, out) + } + + if (where != null) { + out.append(" WHERE ") + where?.render(ctx, out) + } + + renderReturning(returning, ctx, out) + } +} + +class DeleteBuilder( + private val table: String, +) : SqlStatementBuilder { + private val ctes = mutableListOf() + private val returning = mutableListOf() + private var where: SqlExpression? = null + + fun with(name: String, query: SqlStatementBuilder, vararg columns: String): DeleteBuilder = + apply { + ctes += CommonTableExpression(name, columns.toList(), query) + } + + fun where(condition: SqlExpression): DeleteBuilder = + apply { + where = condition + } + + fun andWhere(condition: SqlExpression): DeleteBuilder = + apply { + where = where?.and(condition) ?: condition + } + + fun returning(vararg columns: String): DeleteBuilder = + apply { + returning += columns.map { Sql.col(it) } + } + + fun returning(vararg expressions: SqlExpression): DeleteBuilder = + apply { + returning += expressions + } + + override fun render(ctx: SqlRenderContext, out: StringBuilder) { + renderCtes(ctes, ctx, out) + out.append("DELETE FROM ").append(ctx.quoteIdentifier(table)) + + if (where != null) { + out.append(" WHERE ") + where?.render(ctx, out) + } + + renderReturning(returning, ctx, out) + } +} + +enum class JoinType( + val sql: String, +) { + INNER("JOIN"), + LEFT("LEFT JOIN"), + RIGHT("RIGHT JOIN"), + FULL("FULL JOIN"), + CROSS("CROSS JOIN"), +} + +private data class JoinClause( + val type: JoinType, + val table: SqlTable, + val on: SqlExpression?, +) : SqlFragment { + override fun render(ctx: SqlRenderContext, out: StringBuilder) { + out.append(type.sql).append(' ') + table.render(ctx, out) + if (on != null) { + out.append(" ON ") + on.render(ctx, out) + } + } +} + +private data class UnionClause( + val all: Boolean, + val query: SelectBuilder, +) + +private data class CommonTableExpression( + val name: String, + val columns: List, + val query: SqlStatementBuilder, +) : SqlFragment { + init { + identifierParts(name) + columns.forEach(::identifierParts) + } + + override fun render(ctx: SqlRenderContext, out: StringBuilder) { + out.append(ctx.quoteIdentifier(name)) + if (columns.isNotEmpty()) { + out.append(" (") + out.append(columns.joinToString(", ") { ctx.quoteIdentifier(it) }) + out.append(')') + } + out.append(" AS (") + query.render(ctx, out) + out.append(')') + } +} + +private data class NamedTable( + val name: String, + val alias: String?, +) : SqlTable { + init { + identifierParts(name) + if (alias != null) { + identifierParts(alias) + } + } + + override fun render(ctx: SqlRenderContext, out: StringBuilder) { + out.append(ctx.quoteIdentifier(name)) + if (alias != null) { + out.append(" AS ").append(ctx.quoteIdentifier(alias)) + } + } +} + +private data class SubqueryTable( + val query: SelectBuilder, + val alias: String, +) : SqlTable { + init { + identifierParts(alias) + } + + override fun render(ctx: SqlRenderContext, out: StringBuilder) { + out.append('(') + query.render(ctx, out) + out.append(") AS ").append(ctx.quoteIdentifier(alias)) + } +} + +private data object StarExpression : SqlExpression { + override fun render(ctx: SqlRenderContext, out: StringBuilder) { + out.append('*') + } +} + +private data class ColumnExpression( + val name: String, +) : SqlExpression { + init { + identifierParts(name) + } + + override fun render(ctx: SqlRenderContext, out: StringBuilder) { + out.append(ctx.quoteIdentifier(name)) + } +} + +private data class ValueExpression( + val value: Any?, +) : SqlExpression { + override fun render(ctx: SqlRenderContext, out: StringBuilder) { + out.append(ctx.bind(value)) + } +} + +private data class RawExpression( + val sql: String, + val parameters: List, +) : SqlExpression { + override fun render(ctx: SqlRenderContext, out: StringBuilder) { + if (parameters.isEmpty()) { + out.append(sql) + return + } + + val parts = sql.split('?') + require(parts.size == parameters.size + 1) { + "Raw SQL parameter count must match '?' placeholder count" + } + + parts.forEachIndexed { index, part -> + out.append(part) + if (index < parameters.size) { + out.append(ctx.bind(parameters[index])) + } + } + } +} + +private data class FunctionExpression( + val name: String, + val arguments: List, + val distinct: Boolean, +) : SqlExpression { + init { + functionNameParts(name) + } + + override fun render(ctx: SqlRenderContext, out: StringBuilder) { + out.append(name).append('(') + if (distinct) { + out.append("DISTINCT ") + } + renderCommaSeparated(arguments, ctx, out) + out.append(')') + } +} + +private data class AliasedExpression( + val expression: SqlExpression, + val alias: String, +) : SqlExpression { + init { + identifierParts(alias) + } + + override fun render(ctx: SqlRenderContext, out: StringBuilder) { + expression.render(ctx, out) + out.append(" AS ").append(ctx.quoteIdentifier(alias)) + } +} + +private data class BinaryExpression( + val left: SqlExpression, + val operator: String, + val right: SqlExpression, +) : SqlExpression { + override fun render(ctx: SqlRenderContext, out: StringBuilder) { + out.append('(') + left.render(ctx, out) + out.append(' ').append(operator).append(' ') + right.render(ctx, out) + out.append(')') + } +} + +private data class UnaryPostfixExpression( + val expression: SqlExpression, + val operator: String, +) : SqlExpression { + override fun render(ctx: SqlRenderContext, out: StringBuilder) { + out.append('(') + expression.render(ctx, out) + out.append(' ').append(operator).append(')') + } +} + +private data class PrefixExpression( + val operator: String, + val expression: SqlExpression, +) : SqlExpression { + override fun render(ctx: SqlRenderContext, out: StringBuilder) { + out.append('(').append(operator).append(' ') + expression.render(ctx, out) + out.append(')') + } +} + +private data class BetweenExpression( + val expression: SqlExpression, + val start: SqlExpression, + val end: SqlExpression, +) : SqlExpression { + override fun render(ctx: SqlRenderContext, out: StringBuilder) { + out.append('(') + expression.render(ctx, out) + out.append(" BETWEEN ") + start.render(ctx, out) + out.append(" AND ") + end.render(ctx, out) + out.append(')') + } +} + +private data class InValuesExpression( + val expression: SqlExpression, + val values: List, + val negated: Boolean, +) : SqlExpression { + override fun render(ctx: SqlRenderContext, out: StringBuilder) { + if (values.isEmpty()) { + out.append(if (negated) "(1 = 1)" else "(1 = 0)") + return + } + + out.append('(') + expression.render(ctx, out) + out.append(if (negated) " NOT IN (" else " IN (") + values.forEachIndexed { index, value -> + if (index > 0) { + out.append(", ") + } + value.toExpression().render(ctx, out) + } + out.append("))") + } +} + +private data class InSubqueryExpression( + val expression: SqlExpression, + val query: SelectBuilder, + val negated: Boolean, +) : SqlExpression { + override fun render(ctx: SqlRenderContext, out: StringBuilder) { + out.append('(') + expression.render(ctx, out) + out.append(if (negated) " NOT IN (" else " IN (") + query.render(ctx, out) + out.append("))") + } +} + +private data class ExistsExpression( + val query: SelectBuilder, +) : SqlExpression { + override fun render(ctx: SqlRenderContext, out: StringBuilder) { + out.append("EXISTS (") + query.render(ctx, out) + out.append(')') + } +} + +private data class CompoundExpression( + val operator: String, + val expressions: List, +) : SqlExpression { + fun normalized(): SqlExpression { + val flattened = expressions.flatMap { expression -> + if (expression is CompoundExpression && expression.operator == operator) { + expression.expressions + } else { + listOf(expression) + } + } + + return when (flattened.size) { + 0 -> RawExpression(if (operator == "AND") "1 = 1" else "1 = 0", emptyList()) + 1 -> flattened.single() + else -> CompoundExpression(operator, flattened) + } + } + + override fun render(ctx: SqlRenderContext, out: StringBuilder) { + out.append('(') + expressions.forEachIndexed { index, expression -> + if (index > 0) { + out.append(' ').append(operator).append(' ') + } + expression.render(ctx, out) + } + out.append(')') + } +} + +private fun Any?.toExpression(): SqlExpression = + if (this is SqlExpression) this else ValueExpression(this) + +private fun renderCtes( + ctes: List, + ctx: SqlRenderContext, + out: StringBuilder, +) { + if (ctes.isEmpty()) { + return + } + + out.append("WITH ") + renderCommaSeparated(ctes, ctx, out) + out.append(' ') +} + +private fun renderReturning( + returning: List, + ctx: SqlRenderContext, + out: StringBuilder, +) { + if (returning.isEmpty()) { + return + } + + require(ctx.dialect.supportsReturning) { + "${ctx.dialect.name} dialect does not support RETURNING" + } + + out.append(" RETURNING ") + renderCommaSeparated(returning, ctx, out) +} + +private fun renderCommaSeparated( + expressions: List, + ctx: SqlRenderContext, + out: StringBuilder, +) { + expressions.forEachIndexed { index, expression -> + if (index > 0) { + out.append(", ") + } + expression.render(ctx, out) + } +} + +private val identifierPattern = Regex("[A-Za-z_][A-Za-z0-9_]*|\\*") +private val functionNamePattern = Regex("[A-Za-z_][A-Za-z0-9_]*") + +private fun identifierParts(identifier: String): List { + require(identifier.isNotBlank()) { "Identifier cannot be blank" } + + val parts = identifier.split('.') + require(parts.all { identifierPattern.matches(it) }) { + "Invalid identifier '$identifier'. Use Sql.raw(...) for SQL fragments." + } + require(parts.dropLast(1).none { it == "*" }) { + "Wildcard '*' can only be the final identifier part" + } + + return parts +} + +private fun functionNameParts(name: String): List { + require(name.isNotBlank()) { "Function name cannot be blank" } + + val parts = name.split('.') + require(parts.all { functionNamePattern.matches(it) }) { + "Invalid function name '$name'" + } + + return parts +} diff --git a/src/test/kotlin/dev/mduchene/sql/SqlBuilderTest.kt b/src/test/kotlin/dev/mduchene/sql/SqlBuilderTest.kt new file mode 100644 index 0000000..06961dc --- /dev/null +++ b/src/test/kotlin/dev/mduchene/sql/SqlBuilderTest.kt @@ -0,0 +1,135 @@ +package dev.mduchene.sql + +import kotlin.test.Test +import kotlin.test.assertEquals +import kotlin.test.assertFailsWith + +class SqlBuilderTest { + @Test + fun `builds select with joins count grouping and paging`() { + val statement = Sql + .select( + Sql.col("u.id"), + Sql.col("u.email"), + Sql.count(Sql.col("p.id")).asAlias("post_count"), + ) + .from("users", "u") + .leftJoin("posts", "p", Sql.col("p.user_id").eq(Sql.col("u.id"))) + .where(Sql.col("u.status").eq("active")) + .andWhere(Sql.col("u.deleted_at").isNull()) + .groupBy("u.id", "u.email") + .having(Sql.count(Sql.col("p.id")).gt(0)) + .orderBy(Sql.col("u.created_at").desc(NullsOrder.LAST)) + .limit(25) + .offset(50) + .toSql() + + assertEquals( + """SELECT "u"."id", "u"."email", COUNT("p"."id") AS "post_count" FROM "users" AS "u" LEFT JOIN "posts" AS "p" ON ("p"."user_id" = "u"."id") WHERE (("u"."status" = $1) AND ("u"."deleted_at" IS NULL)) GROUP BY "u"."id", "u"."email" HAVING (COUNT("p"."id") > $2) ORDER BY "u"."created_at" DESC NULLS LAST LIMIT 25 OFFSET 50""", + statement.sql, + ) + assertEquals(listOf("active", 0), statement.parameters) + } + + @Test + fun `builds insert with multiple rows and returning`() { + val statement = Sql + .insertInto("users") + .columns("email", "display_name") + .values("sakura@example.com", "Sakura") + .values("naruto@example.com", "Naruto") + .returning("id") + .toSql() + + assertEquals( + """INSERT INTO "users" ("email", "display_name") VALUES ($1, $2), ($3, $4) RETURNING "id"""", + statement.sql, + ) + assertEquals( + listOf("sakura@example.com", "Sakura", "naruto@example.com", "Naruto"), + statement.parameters, + ) + } + + @Test + fun `builds update with raw expression and where clause`() { + val statement = Sql + .update("users") + .set("display_name", "Hinata") + .set("updated_at", Sql.raw("CURRENT_TIMESTAMP")) + .where(Sql.col("id").eq(42)) + .returning("id", "updated_at") + .toSql() + + assertEquals( + """UPDATE "users" SET "display_name" = $1, "updated_at" = CURRENT_TIMESTAMP WHERE ("id" = $2) RETURNING "id", "updated_at"""", + statement.sql, + ) + assertEquals(listOf("Hinata", 42), statement.parameters) + } + + @Test + fun `builds delete with cte and subquery`() { + val inactiveUsers = Sql + .select("id") + .from("users") + .where(Sql.col("last_seen_at").lt(Sql.raw("CURRENT_DATE - INTERVAL '1 year'"))) + + val statement = Sql + .deleteFrom("sessions") + .with("inactive_users", inactiveUsers) + .where( + Sql.col("user_id").inSubquery( + Sql.select("id").from("inactive_users"), + ), + ) + .toSql() + + assertEquals( + """WITH "inactive_users" AS (SELECT "id" FROM "users" WHERE ("last_seen_at" < CURRENT_DATE - INTERVAL '1 year')) DELETE FROM "sessions" WHERE ("user_id" IN (SELECT "id" FROM "inactive_users"))""", + statement.sql, + ) + assertEquals(emptyList(), statement.parameters) + } + + @Test + fun `can render ansi placeholders from the same builder`() { + val builder = Sql + .select("id") + .from("users") + .where(Sql.col("email").eq("sakura@example.com")) + .andWhere(Sql.col("status").eq("active")) + + assertEquals( + """SELECT "id" FROM "users" WHERE (("email" = $1) AND ("status" = $2))""", + builder.toSql(PostgresDialect).sql, + ) + + val ansi = builder.toSql(AnsiDialect) + assertEquals( + """SELECT "id" FROM "users" WHERE (("email" = ?) AND ("status" = ?))""", + ansi.sql, + ) + assertEquals(listOf("sakura@example.com", "active"), ansi.parameters) + } + + @Test + fun `rejects returning for dialects that do not support it`() { + val builder = Sql + .insertInto("users") + .columns("email") + .values("sakura@example.com") + .returning("id") + + assertFailsWith { + builder.toSql(AnsiDialect) + } + } + + @Test + fun `rejects unsafe identifiers unless raw sql is explicit`() { + assertFailsWith { + Sql.select("id; DROP TABLE users").from("users").toSql() + } + } +}