package lang import ( "strings" "testing" ) const accountRowSource = ` @table("accounts") data class AccountRow( @id var id: String, var customerId: String, @column("kind") var accountType: String, var balance: Double ) ` func TestParseSQLMetadata(t *testing.T) { program, err := Parse(accountRowSource) if err != nil { t.Fatalf("parse failed: %v", err) } if len(program.Classes) != 1 { t.Fatalf("expected one class, got %d", len(program.Classes)) } class := program.Classes[0] if class.Table != "accounts" { t.Fatalf("table = %q, want accounts", class.Table) } if !class.Fields[0].ID { t.Fatal("id field is missing @id metadata") } if class.Fields[2].Column != "kind" { t.Fatalf("accountType column = %q, want kind", class.Fields[2].Column) } if got := sqlColumn(class.Fields[1]); got != "customer_id" { t.Fatalf("default customerId column = %q, want customer_id", got) } } func TestGenerateSQLSelect(t *testing.T) { code := compileSQL(t, accountRowSource+` fun accountQuery(customerId: String): GotlinSQLQuery { return sql.from() .where { row -> row.customerId == customerId && row.balance > 0.0 } .orderBy { row -> row.accountType } .build() } `) for _, want := range []string{ "type GotlinSQLQuery struct", "SQL string", "Args []any", `SQL: "SELECT id, customer_id, kind, balance FROM accounts WHERE (customer_id = $1 AND balance > $2) ORDER BY kind"`, `Args: []any{customerId, 0.0}`, } { if !strings.Contains(code, want) { t.Fatalf("generated Go missing %q:\n%s", want, code) } } if strings.Contains(code, `github.com/jackc/pgx/v5`) { t.Fatalf("build-only SQL unexpectedly emitted pgx runtime:\n%s", code) } } func TestGenerateSQLSelectUsesTypedLocalAsArgument(t *testing.T) { code := compileSQL(t, accountRowSource+` fun accountQuery(): GotlinSQLQuery { val customerId = "customer-1" return sql.from().where { it.customerId == customerId }.build() } `) if want := `Args: []any{customerId}`; !strings.Contains(code, want) { t.Fatalf("generated Go missing %q:\n%s", want, code) } } func TestGeneratedQueryCanBePassedToVariadicPGXCall(t *testing.T) { code := compileSQL(t, accountRowSource+` import pgxpool "github.com/jackc/pgx/v5/pgxpool" import context fun execute(pool: *pgxpool.Pool, ctx: context.Context, customerId: String) { val query = sql.from().where { it.customerId == customerId }.build() pool.query(ctx, query.sql, *query.args) } `) if want := `pool.Query(ctx, query.SQL, query.Args...)`; !strings.Contains(code, want) { t.Fatalf("generated Go missing %q:\n%s", want, code) } } func TestGenerateSQLFetchTerminal(t *testing.T) { code := compileSQL(t, accountRowSource+` import context import pgxpool "github.com/jackc/pgx/v5/pgxpool" fun accounts(pool: *pgxpool.Pool, ctx: context.Context, customerId: String): List { return sql.from() .where { it.customerId == customerId } .fetch(pool, ctx).unwrap() } `) for _, want := range []string{ `"github.com/jackc/pgx/v5"`, `func accounts(pool *pgxpool.Pool, ctx context.Context, customerId string) []*AccountRow`, `return gotlinResultUnwrap(gotlinSQLFetch[AccountRow](pool, ctx, GotlinSQLQuery{SQL: "SELECT id, customer_id, kind, balance FROM accounts WHERE customer_id = $1", Args: []any{customerId}}, gotlinSQLScanAccountRow))`, `defer rows.Close()`, `values = append(values, value)`, `err := row.Scan(&value.Id, &value.CustomerId, &value.AccountType, &value.Balance)`, } { if !strings.Contains(code, want) { t.Fatalf("generated Go missing %q:\n%s", want, code) } } } func TestGenerateSQLSingleInfersRowSelectors(t *testing.T) { code := compileSQL(t, accountRowSource+` import context import pgxpool "github.com/jackc/pgx/v5/pgxpool" fun balance(pool: *pgxpool.Pool, ctx: context.Context, id: String): Double { val account = sql.from().where { it.id == id }.single(pool, ctx).unwrap() return account.balance } `) for _, want := range []string{ `account := gotlinResultUnwrap(gotlinSQLSingle[AccountRow](pool, ctx, GotlinSQLQuery{SQL: "SELECT id, customer_id, kind, balance FROM accounts WHERE id = $1", Args: []any{id}}, gotlinSQLScanAccountRow))`, `return account.Balance`, `GotlinResult[*T]{Err: gotlinSQLError("SQL single() expected exactly one row, got zero")}`, `GotlinResult[*T]{Err: gotlinSQLError("SQL single() expected exactly one row, got more than one")}`, } { if !strings.Contains(code, want) { t.Fatalf("generated Go missing %q:\n%s", want, code) } } } func TestGenerateSQLIteratorTerminalAndTypedValue(t *testing.T) { code := compileSQL(t, accountRowSource+` import context import pgxpool "github.com/jackc/pgx/v5/pgxpool" fun printAccounts(pool: *pgxpool.Pool, ctx: context.Context) { val rows = sql.from().iterator(pool, ctx).unwrap() defer rows.close() while (rows.next()) { val account = rows.value() println(account.customerId) } val checked = rows.err() } `) for _, want := range []string{ `rows := gotlinResultUnwrap(gotlinSQLIterate[AccountRow](pool, ctx, GotlinSQLQuery{SQL: "SELECT id, customer_id, kind, balance FROM accounts", Args: []any{}}, gotlinSQLScanAccountRow))`, `defer rows.close()`, `for rows.next()`, `account := rows.value()`, `fmt.Println(account.CustomerId)`, `_ = rows.err()`, `type GotlinSQLIterator[T any] struct`, `func (iterator *GotlinSQLIterator[T]) next() bool`, `func (iterator *GotlinSQLIterator[T]) value() *T`, `func (iterator *GotlinSQLIterator[T]) close()`, `func (iterator *GotlinSQLIterator[T]) err() error`, } { if !strings.Contains(code, want) { t.Fatalf("generated Go missing %q:\n%s", want, code) } } } func TestGenerateSQLSelectInClassMethodUsesTypedParameter(t *testing.T) { code := compileSQL(t, accountRowSource+` class AccountQueries { fun byCustomer(customerId: String): GotlinSQLQuery { return sql.from().where { it.customerId == customerId }.build() } } `) if want := `Args: []any{customerId}`; !strings.Contains(code, want) { t.Fatalf("generated Go missing %q:\n%s", want, code) } } func TestGenerateSQLInsertDoNothing(t *testing.T) { code := compileSQL(t, accountRowSource+` fun insertAccount(row: AccountRow): GotlinSQLQuery { return sql.insert(row).onConflict { it.id }.doNothing().build() } `) for _, want := range []string{ `SQL: "INSERT INTO accounts (id, customer_id, kind, balance) VALUES ($1, $2, $3, $4) ON CONFLICT (id) DO NOTHING"`, `Args: []any{row.Id, row.CustomerId, row.AccountType, row.Balance}`, } { if !strings.Contains(code, want) { t.Fatalf("generated Go missing %q:\n%s", want, code) } } } func TestGenerateSQLInsertDoUpdate(t *testing.T) { code := compileSQL(t, accountRowSource+` fun upsertAccount(row: AccountRow): GotlinSQLQuery { return sql.insert(row) .onConflict { it.id } .doUpdate { excluded -> AccountRow.balance = excluded.balance } .build() } `) if want := `SQL: "INSERT INTO accounts (id, customer_id, kind, balance) VALUES ($1, $2, $3, $4) ON CONFLICT (id) DO UPDATE SET balance = EXCLUDED.balance"`; !strings.Contains(code, want) { t.Fatalf("generated Go missing %q:\n%s", want, code) } } func TestGenerateSQLCompositeConflictAndMultipleUpdates(t *testing.T) { code := compileSQL(t, ` @table("balances") data class BalanceRow(@id var tenantId: String, @id var id: String, var amount: Double, var pending: Double) fun upsert(row: BalanceRow): GotlinSQLQuery { return sql.insert(row) .onConflict { listOf(it.tenantId, it.id) } .doUpdate { excluded -> BalanceRow.amount = excluded.amount BalanceRow.pending = excluded.pending } .build() } `) if want := `ON CONFLICT (tenant_id, id) DO UPDATE SET amount = EXCLUDED.amount, pending = EXCLUDED.pending`; !strings.Contains(code, want) { t.Fatalf("generated Go missing %q:\n%s", want, code) } } func TestGenerateSQLJoinsAliasesGroupingOffsetAndAggregates(t *testing.T) { code := compileSQL(t, accountRowSource+` @table("customers") data class CustomerRow(@id var id: String, var name: String) data class AccountSummary(var customerId: String, var customerName: String, var total: Double, var entries: Long) fun summaries(minimum: Double): GotlinSQLQuery { return sql.from() .alias("account") .leftJoin("customer") { account, customer -> account.customerId == customer.id } .select { account, customer -> AccountSummary(account.customerId, customer.name, sum(account.balance), count()) } .where { account, customer -> account.balance > minimum } .groupBy { account, customer -> listOf(account.customerId, customer.name) } .having { account, customer -> sum(account.balance) > minimum } .orderByDescending { account, customer -> sum(account.balance) } .limit(10) .offset(20) .build() } `) want := `SQL: "SELECT account.customer_id, customer.name, SUM(account.balance), COUNT(*) FROM accounts AS account LEFT JOIN customers AS customer ON account.customer_id = customer.id WHERE account.balance > $1 GROUP BY account.customer_id, customer.name HAVING SUM(account.balance) > $2 ORDER BY SUM(account.balance) DESC LIMIT 10 OFFSET 20"` if !strings.Contains(code, want) { t.Fatalf("generated Go missing %q:\n%s", want, code) } if !strings.Contains(code, `Args: []any{minimum, minimum}`) { t.Fatalf("aggregate arguments missing:\n%s", code) } } func TestGenerateSQLBulkInsert(t *testing.T) { code := compileSQL(t, accountRowSource+` fun insertAccounts(rows: List): GotlinSQLQuery = sql.insertAll(rows).build() `) for _, want := range []string{ `gotlinSQLBulkInsert[AccountRow](rows, "accounts", []string{"id", "customer_id", "kind", "balance"}`, `func(row *AccountRow) []any { return []any{row.Id, row.CustomerId, row.AccountType, row.Balance} }`, `func gotlinSQLBulkInsert[T any]`, `panic("bulk insert requires at least one row")`, } { if !strings.Contains(code, want) { t.Fatalf("generated Go missing %q:\n%s", want, code) } } } func TestRejectInvalidSQLQueries(t *testing.T) { tests := []struct { name string src string want string }{ { name: "unknown row class", src: `fun query(): GotlinSQLQuery { return sql.from().build() }`, want: `SQL row class "MissingRow" does not exist`, }, { name: "missing table metadata", src: `data class Row(var id: String) fun query(): GotlinSQLQuery { return sql.from().build() }`, want: "requires @table", }, { name: "unknown predicate field", src: accountRowSource + `fun query(): GotlinSQLQuery { return sql.from().where { it.missing == "x" }.build() }`, want: "has no field missing", }, { name: "incompatible predicate operands", src: accountRowSource + `fun query(customerId: String): GotlinSQLQuery { return sql.from().where { it.balance == customerId }.build() }`, want: "incompatible types Double and String", }, { name: "non numeric ordering predicate", src: accountRowSource + `fun query(): GotlinSQLQuery { return sql.from().where { it.customerId > "a" }.build() }`, want: "ordering comparison requires numeric operands", }, { name: "unknown order field", src: accountRowSource + `fun query(): GotlinSQLQuery { return sql.from().orderBy { it.missing }.build() }`, want: "has no field missing", }, { name: "conflict field must be id", src: accountRowSource + `fun query(row: AccountRow): GotlinSQLQuery { return sql.insert(row).onConflict { it.customerId }.doNothing().build() }`, want: "must be annotated @id", }, { name: "unknown conflict field", src: accountRowSource + `fun query(row: AccountRow): GotlinSQLQuery { return sql.insert(row).onConflict { it.missing }.doNothing().build() }`, want: "has no field missing", }, { name: "unknown update target", src: accountRowSource + `fun query(row: AccountRow): GotlinSQLQuery { return sql.insert(row).onConflict { it.id }.doUpdate { excluded -> AccountRow.missing = excluded.balance }.build() }`, want: "has no field missing", }, { name: "incompatible update fields", src: accountRowSource + `fun query(row: AccountRow): GotlinSQLQuery { return sql.insert(row).onConflict { it.id }.doUpdate { excluded -> AccountRow.balance = excluded.accountType }.build() }`, want: "has type Double", }, { name: "wrong insert row type", src: accountRowSource + ` @table("other") data class OtherRow(@id var id: String) fun query(row: OtherRow): GotlinSQLQuery { return sql.insert(row).onConflict { it.id }.doNothing().build() } `, want: "cannot insert value of type OtherRow", }, { name: "incomplete chain", src: accountRowSource + `fun query(): GotlinSQLQuery { return sql.from() }`, want: "must end with build()", }, { name: "fetch missing arguments", src: accountRowSource + `fun query(): List { return sql.from().fetch().unwrap() }`, want: "fetch() expects exactly pool and ctx positional arguments", }, { name: "single extra argument", src: accountRowSource + `fun query(pool: Any, ctx: Any): AccountRow { return sql.from().single(pool, ctx, ctx).unwrap() }`, want: "single() expects exactly pool and ctx positional arguments", }, { name: "iterator named arguments", src: accountRowSource + `fun query(pool: Any, ctx: Any): GotlinSQLIterator { return sql.from().iterator(pool = pool, ctx = ctx).unwrap() }`, want: "iterator() expects exactly pool and ctx positional arguments", }, { name: "fetch type arguments", src: accountRowSource + `fun query(pool: Any, ctx: Any): List { return sql.from().fetch(pool, ctx) }`, want: "fetch() expects exactly pool and ctx positional arguments", }, { name: "fetch invalid pool type", src: accountRowSource + `fun query(pool: String, ctx: Any): List { return sql.from().fetch(pool, ctx).unwrap() }`, want: "pool argument has non-query type String", }, { name: "single invalid context type", src: accountRowSource + `fun query(pool: Any, ctx: Int): AccountRow { return sql.from().single(pool, ctx).unwrap() }`, want: "ctx argument has non-context type Int", }, { name: "insert execution terminal", src: accountRowSource + `fun query(row: AccountRow, pool: Any, ctx: Any): List { return sql.insert(row).fetch(pool, ctx).unwrap() }`, want: "fetch() is only supported for sql.from", }, { name: "duplicate join alias", src: accountRowSource + ` @table("customers") data class CustomerRow(@id var id: String) fun query(): GotlinSQLQuery = sql.from().alias("row").join("row") { account, customer -> account.customerId == customer.id }.build() `, want: `duplicate SQL alias "row"`, }, { name: "having without grouping", src: accountRowSource + `fun query(): GotlinSQLQuery = sql.from().having { count() > 0 }.build()`, want: "after groupBy", }, { name: "negative offset", src: accountRowSource + `fun query(): GotlinSQLQuery = sql.from().offset(-1).build()`, want: "offset() requires a non-negative Int", }, { name: "bulk insert wrong element type", src: accountRowSource + ` @table("other") data class OtherRow(@id var id: String) fun query(rows: List): GotlinSQLQuery = sql.insertAll(rows).build() `, want: "requires List", }, } for _, test := range tests { t.Run(test.name, func(t *testing.T) { program, err := Parse(test.src) if err != nil { t.Fatalf("parse failed: %v", err) } _, err = GenerateGo(program) if err == nil || !strings.Contains(err.Error(), test.want) { t.Fatalf("error = %v, want substring %q", err, test.want) } }) } } func TestRejectInvalidSQLAnnotations(t *testing.T) { for _, test := range []struct { src string want string }{ {`@table("accounts") class Row(var id: String)`, "table is only valid on data classes"}, {`@table("bad-name") data class Row(var id: String)`, "invalid SQL table name"}, {`data class Row(@column("bad-name") var id: String)`, "invalid SQL column name"}, } { _, err := Parse(test.src) if err == nil || !strings.Contains(err.Error(), test.want) { t.Fatalf("Parse() error = %v, want substring %q", err, test.want) } } } func compileSQL(t *testing.T, src string) string { t.Helper() program, err := Parse(src) if err != nil { t.Fatalf("parse failed: %v", err) } out, err := GenerateGo(program) if err != nil { t.Fatalf("Go generation failed: %v", err) } return string(out) }