gotlin/internal/lang/sql.go

1242 lines
37 KiB
Go

package lang
import (
"fmt"
"strconv"
"strings"
)
type sqlCallStep struct {
name string
call CallExpr
}
type sqlLowered struct {
value string
result ClassDecl
hasResult bool
}
func splitSQLChain(expr Expr) (CallExpr, string, []sqlCallStep, bool) {
var reversed []sqlCallStep
current := expr
for {
call, ok := current.(CallExpr)
if !ok {
return CallExpr{}, "", nil, false
}
selector, ok := call.Callee.(SelectorExpr)
if !ok {
return CallExpr{}, "", nil, false
}
if receiver, ok := selector.Receiver.(IdentExpr); ok && receiver.Name == "sql" {
switch selector.Name {
case "from", "insert", "update", "delete":
default:
return CallExpr{}, "", nil, false
}
steps := make([]sqlCallStep, len(reversed))
for i := range reversed {
steps[len(reversed)-1-i] = reversed[i]
}
return call, selector.Name, steps, true
}
reversed = append(reversed, sqlCallStep{name: selector.Name, call: call})
current = selector.Receiver
}
}
func (g *goGenerator) lowerSQLQuery(expr Expr) (string, bool, error) {
if call, ok := expr.(CallExpr); ok {
if selector, ok := call.Callee.(SelectorExpr); ok && (selector.Name == "unwrap" || selector.Name == "unwrapOr") {
return "", false, nil
}
}
root, operation, steps, ok := splitSQLChain(expr)
if !ok {
return "", false, nil
}
if len(steps) == 0 {
return "", true, sqlTerminalError(operation)
}
terminal := steps[len(steps)-1]
switch terminal.name {
case "build":
if len(terminal.call.Args) != 0 || len(terminal.call.NamedArgs) != 0 || len(terminal.call.TypeArgs) != 0 {
return "", true, fmt.Errorf("SQL build() does not accept arguments")
}
case "fetch", "single", "iterator":
if err := g.validateSQLExecutionTerminal(terminal); err != nil {
return "", true, err
}
default:
return "", true, sqlTerminalError(operation)
}
steps = steps[:len(steps)-1]
var lowered sqlLowered
var err error
switch operation {
case "from":
lowered, err = g.lowerSQLSelect(root, steps)
case "insert":
lowered, err = g.lowerSQLInsert(root, steps)
case "update":
lowered, err = g.lowerSQLUpdate(root, steps)
case "delete":
lowered, err = g.lowerSQLDelete(root, steps)
default:
err = fmt.Errorf("unsupported sql operation %q", operation)
}
if err != nil || terminal.name == "build" {
return lowered.value, true, err
}
if !lowered.hasResult {
return "", true, fmt.Errorf("SQL %s() is only supported for sql.from or writes with returning(); sql.%s requires returning()", terminal.name, operation)
}
pool, err := g.expr(terminal.call.Args[0], "")
if err != nil {
return "", true, err
}
ctx, err := g.expr(terminal.call.Args[1], "")
if err != nil {
return "", true, err
}
helper := map[string]string{
"fetch": "gotlinSQLFetch",
"single": "gotlinSQLSingle",
"iterator": "gotlinSQLIterate",
}[terminal.name]
return fmt.Sprintf("%s[%s](%s, %s, %s, gotlinSQLScan%s)", helper, lowered.result.Name, pool, ctx, lowered.value, lowered.result.Name), true, nil
}
func sqlTerminalError(operation string) error {
return fmt.Errorf("sql.%s query must end with build(), fetch(pool, ctx), single(pool, ctx), or iterator(pool, ctx)", operation)
}
func (g *goGenerator) validateSQLExecutionTerminal(terminal sqlCallStep) error {
call := terminal.call
if len(call.TypeArgs) != 0 || len(call.NamedArgs) != 0 || len(call.Args) != 2 {
return fmt.Errorf("SQL %s() expects exactly pool and ctx positional arguments", terminal.name)
}
if typ := g.exprType(call.Args[0]); sqlInvalidExecutionArgType(typ) {
return fmt.Errorf("SQL %s() pool argument has non-query type %s", terminal.name, typ)
}
if typ := g.exprType(call.Args[1]); sqlInvalidExecutionArgType(typ) {
return fmt.Errorf("SQL %s() ctx argument has non-context type %s", terminal.name, typ)
}
return nil
}
func sqlInvalidExecutionArgType(typ string) bool {
switch sqlBaseType(typ) {
case "Int", "Float", "Double", "String", "Boolean", "Unit", "ByteSlice":
return true
default:
return false
}
}
func sqlChainResultType(expr Expr) (string, bool) {
root, operation, steps, ok := splitSQLChain(expr)
if !ok || len(root.TypeArgs) != 1 || len(steps) == 0 {
return "", false
}
terminal := steps[len(steps)-1].name
if terminal == "build" {
return "GotlinSQLQuery", true
}
if terminal != "fetch" && terminal != "single" && terminal != "iterator" {
return "", false
}
resultType := root.TypeArgs[0]
if operation == "from" {
for _, step := range steps[:len(steps)-1] {
if step.name == "select" {
if typ, ok := sqlProjectionType(step.call, resultType); ok {
resultType = typ
}
}
}
} else {
found := false
for _, step := range steps[:len(steps)-1] {
if step.name == "returning" {
resultType, found = sqlProjectionType(step.call, resultType)
}
}
if !found {
return "", false
}
}
switch terminal {
case "fetch":
return "Result<List<" + resultType + ">, Error>", true
case "single":
return "Result<" + resultType + ", Error>", true
case "iterator":
return "Result<GotlinSQLIterator<" + resultType + ">, Error>", true
default:
return "", false
}
}
func sqlProjectionType(call CallExpr, rowType string) (string, bool) {
lambda, err := sqlLambdaArg(call, "projection")
if err != nil {
return "", false
}
rowName := "it"
if !lambda.ImplicitIt {
if len(lambda.Params) != 1 {
return "", false
}
rowName = lambda.Params[0].Name
}
if len(lambda.Body) != 1 {
return "", false
}
stmt, ok := lambda.Body[0].(ExprStmt)
if !ok {
return "", false
}
if ident, ok := stmt.Value.(IdentExpr); ok && ident.Name == rowName {
return rowType, true
}
constructor, ok := stmt.Value.(CallExpr)
if !ok || len(constructor.NamedArgs) != 0 || len(constructor.TypeArgs) != 0 {
return "", false
}
ident, ok := constructor.Callee.(IdentExpr)
return ident.Name, ok
}
func (g *goGenerator) sqlClass(root CallExpr, operation string) (ClassDecl, error) {
if len(root.TypeArgs) != 1 {
return ClassDecl{}, fmt.Errorf("sql.%s expects exactly one row type", operation)
}
class, ok := g.semantic.Classes[root.TypeArgs[0]]
if !ok {
return ClassDecl{}, fmt.Errorf("SQL row class %q does not exist", root.TypeArgs[0])
}
if !class.Data {
return ClassDecl{}, fmt.Errorf("SQL row class %s must be a data class", class.Name)
}
if class.Table == "" {
return ClassDecl{}, fmt.Errorf("SQL row class %s requires @table", class.Name)
}
columns := map[string]string{}
for _, field := range class.Fields {
column := sqlColumn(field)
if previous, exists := columns[column]; exists {
return ClassDecl{}, fmt.Errorf("SQL fields %s and %s map to duplicate column %q", previous, field.Name, column)
}
columns[column] = field.Name
}
return class, nil
}
func (g *goGenerator) lowerSQLSelect(root CallExpr, steps []sqlCallStep) (sqlLowered, error) {
if len(root.Args) != 0 || len(root.NamedArgs) != 0 {
return sqlLowered{}, fmt.Errorf("sql.from does not accept value arguments")
}
class, err := g.sqlClass(root, "from")
if err != nil {
return sqlLowered{}, err
}
columns := sqlAllColumns(class)
result := class
var where, order, limit string
var args []string
stage := 0
seenSelect := false
seenWhere := false
seenOrder := false
seenLimit := false
seenForUpdate := false
seenSkipLocked := false
for _, step := range steps {
switch step.name {
case "select":
if seenSelect || stage != 0 {
return sqlLowered{}, fmt.Errorf("select() may appear once and must be the first sql.from method")
}
result, columns, err = g.sqlProjection(step.call, class, "select")
if err != nil {
return sqlLowered{}, err
}
seenSelect = true
stage = 1
case "where":
if seenWhere || stage > 1 {
return sqlLowered{}, fmt.Errorf("where() may appear once after select() and before ordering")
}
where, args, err = g.sqlWhere(step.call, class, len(args)+1)
if err != nil {
return sqlLowered{}, err
}
seenWhere = true
stage = 2
case "orderBy", "orderByDescending":
if seenOrder || stage > 2 {
return sqlLowered{}, fmt.Errorf("orderBy()/orderByDescending() may appear once before limit()")
}
order, err = sqlOrder(step, class)
if err != nil {
return sqlLowered{}, err
}
seenOrder = true
stage = 3
case "limit":
if seenLimit || stage > 3 {
return sqlLowered{}, fmt.Errorf("limit() may appear once after ordering and before forUpdate()")
}
var limitArgs []string
limit, limitArgs, err = g.sqlLimit(step.call)
if err != nil {
return sqlLowered{}, err
}
if limit == "?" {
limit = "$" + strconv.Itoa(len(args)+1)
}
args = append(args, limitArgs...)
seenLimit = true
stage = 4
case "forUpdate":
if seenForUpdate || stage > 4 || !sqlNoArgs(step.call) {
return sqlLowered{}, fmt.Errorf("forUpdate() accepts no arguments and must appear once after limit()")
}
seenForUpdate = true
stage = 5
case "skipLocked":
if seenSkipLocked || !seenForUpdate || stage != 5 || !sqlNoArgs(step.call) {
return sqlLowered{}, fmt.Errorf("skipLocked() accepts no arguments and requires a preceding forUpdate()")
}
seenSkipLocked = true
stage = 6
default:
return sqlLowered{}, fmt.Errorf("unsupported sql.from method %q", step.name)
}
}
query := "SELECT " + strings.Join(columns, ", ") + " FROM " + class.Table
if where != "" {
query += " WHERE " + where
}
if order != "" {
query += " ORDER BY " + order
}
if limit != "" {
query += " LIMIT " + limit
}
if seenForUpdate {
query += " FOR UPDATE"
}
if seenSkipLocked {
query += " SKIP LOCKED"
}
return sqlLowered{value: sqlQueryValue(query, args), result: result, hasResult: true}, nil
}
func (g *goGenerator) lowerSQLInsert(root CallExpr, steps []sqlCallStep) (sqlLowered, error) {
if len(root.Args) != 1 || len(root.NamedArgs) != 0 {
return sqlLowered{}, fmt.Errorf("sql.insert expects exactly one row argument")
}
class, err := g.sqlClass(root, "insert")
if err != nil {
return sqlLowered{}, err
}
if typ := strings.TrimPrefix(g.exprType(root.Args[0]), "*"); typ == "" {
return sqlLowered{}, fmt.Errorf("sql.insert row argument requires a known type")
} else if typ != class.Name {
return sqlLowered{}, fmt.Errorf("sql.insert<%s> cannot insert value of type %s", class.Name, typ)
}
row, err := g.expr(root.Args[0], class.Name)
if err != nil {
return sqlLowered{}, err
}
var columns, placeholders, args []string
for _, field := range class.Fields {
if field.Generated {
continue
}
columns = append(columns, sqlColumn(field))
placeholders = append(placeholders, "$"+strconv.Itoa(len(args)+1))
args = append(args, sqlGoField(row, root.Args[0], class, field))
}
query := "INSERT INTO " + class.Table
if len(columns) == 0 {
query += " DEFAULT VALUES"
} else {
query += " (" + strings.Join(columns, ", ") + ") VALUES (" + strings.Join(placeholders, ", ") + ")"
}
index := 0
if index < len(steps) && steps[index].name == "onConflict" {
conflictLambda, err := sqlLambdaArg(steps[index].call, "onConflict")
if err != nil {
return sqlLowered{}, err
}
rowName, body, err := sqlLambdaExpr(conflictLambda, class, "onConflict")
if err != nil {
return sqlLowered{}, err
}
conflictFields, err := sqlConflictFields(body, class, rowName)
if err != nil {
return sqlLowered{}, err
}
var conflictColumns []string
for _, field := range conflictFields {
if !field.ID {
return sqlLowered{}, fmt.Errorf("onConflict field %s must be annotated @id", field.Name)
}
conflictColumns = append(conflictColumns, sqlColumn(field))
}
query += " ON CONFLICT (" + strings.Join(conflictColumns, ", ") + ")"
index++
if index >= len(steps) {
return sqlLowered{}, fmt.Errorf("onConflict() requires doNothing() or doUpdate()")
}
action := steps[index]
switch action.name {
case "doNothing":
if !sqlNoArgs(action.call) {
return sqlLowered{}, fmt.Errorf("doNothing() does not accept arguments")
}
query += " DO NOTHING"
case "doUpdate":
lambda, err := sqlLambdaArg(action.call, "doUpdate")
if err != nil {
return sqlLowered{}, err
}
assignments, err := sqlConflictUpdateAssignments(lambda, class)
if err != nil {
return sqlLowered{}, err
}
query += " DO UPDATE SET " + strings.Join(assignments, ", ")
default:
return sqlLowered{}, fmt.Errorf("onConflict() requires doNothing() or doUpdate(), found %s()", action.name)
}
index++
}
lowered := sqlLowered{}
if index < len(steps) && steps[index].name == "returning" {
result, returning, err := g.sqlProjection(steps[index].call, class, "returning")
if err != nil {
return sqlLowered{}, err
}
query += " RETURNING " + strings.Join(returning, ", ")
lowered.result = result
lowered.hasResult = true
index++
}
if index != len(steps) {
return sqlLowered{}, fmt.Errorf("unsupported sql.insert method %q or invalid method order", steps[index].name)
}
lowered.value = sqlQueryValue(query, args)
return lowered, nil
}
func (g *goGenerator) lowerSQLUpdate(root CallExpr, steps []sqlCallStep) (sqlLowered, error) {
if len(root.Args) != 0 || len(root.NamedArgs) != 0 {
return sqlLowered{}, fmt.Errorf("sql.update does not accept value arguments")
}
class, err := g.sqlClass(root, "update")
if err != nil {
return sqlLowered{}, err
}
if len(steps) == 0 || steps[0].name != "set" {
return sqlLowered{}, fmt.Errorf("sql.update requires set() as its first method")
}
assignments, args, err := g.sqlTypedUpdateAssignments(steps[0].call, class)
if err != nil {
return sqlLowered{}, err
}
query := "UPDATE " + class.Table + " SET " + postgresPlaceholders(strings.Join(assignments, ", "), 1)
index := 1
if index < len(steps) && steps[index].name == "where" {
where, whereArgs, err := g.sqlWhere(steps[index].call, class, len(args)+1)
if err != nil {
return sqlLowered{}, err
}
query += " WHERE " + where
args = append(args, whereArgs...)
index++
}
lowered := sqlLowered{}
if index < len(steps) && steps[index].name == "returning" {
result, returning, err := g.sqlProjection(steps[index].call, class, "returning")
if err != nil {
return sqlLowered{}, err
}
query += " RETURNING " + strings.Join(returning, ", ")
lowered.result = result
lowered.hasResult = true
index++
}
if index != len(steps) {
return sqlLowered{}, fmt.Errorf("unsupported sql.update method %q or invalid method order", steps[index].name)
}
lowered.value = sqlQueryValue(query, args)
return lowered, nil
}
func (g *goGenerator) lowerSQLDelete(root CallExpr, steps []sqlCallStep) (sqlLowered, error) {
if len(root.Args) != 0 || len(root.NamedArgs) != 0 {
return sqlLowered{}, fmt.Errorf("sql.delete does not accept value arguments")
}
class, err := g.sqlClass(root, "delete")
if err != nil {
return sqlLowered{}, err
}
query := "DELETE FROM " + class.Table
var args []string
index := 0
if index < len(steps) && steps[index].name == "where" {
where, whereArgs, err := g.sqlWhere(steps[index].call, class, 1)
if err != nil {
return sqlLowered{}, err
}
query += " WHERE " + where
args = whereArgs
index++
}
lowered := sqlLowered{}
if index < len(steps) && steps[index].name == "returning" {
result, returning, err := g.sqlProjection(steps[index].call, class, "returning")
if err != nil {
return sqlLowered{}, err
}
query += " RETURNING " + strings.Join(returning, ", ")
lowered.result = result
lowered.hasResult = true
index++
}
if index != len(steps) {
return sqlLowered{}, fmt.Errorf("unsupported sql.delete method %q or invalid method order", steps[index].name)
}
lowered.value = sqlQueryValue(query, args)
return lowered, nil
}
func sqlGoField(row string, source Expr, class ClassDecl, field FieldDecl) string {
fieldName := field.Name
if class.Data && !field.Private {
fieldName = exportedGoName(fieldName)
}
receiver := row
switch source.(type) {
case IdentExpr, SelectorExpr, IndexExpr:
default:
receiver = "(" + receiver + ")"
}
return receiver + "." + fieldName
}
func sqlAllColumns(class ClassDecl) []string {
columns := make([]string, 0, len(class.Fields))
for _, field := range class.Fields {
columns = append(columns, sqlColumn(field))
}
return columns
}
func (g *goGenerator) sqlProjection(call CallExpr, rowClass ClassDecl, method string) (ClassDecl, []string, error) {
lambda, err := sqlLambdaArg(call, method)
if err != nil {
return ClassDecl{}, nil, err
}
rowName, body, err := sqlLambdaExpr(lambda, rowClass, method)
if err != nil {
return ClassDecl{}, nil, err
}
if ident, ok := body.(IdentExpr); ok && ident.Name == rowName {
return rowClass, sqlAllColumns(rowClass), nil
}
constructor, ok := body.(CallExpr)
if !ok || len(constructor.TypeArgs) != 0 || len(constructor.NamedArgs) != 0 {
return ClassDecl{}, nil, fmt.Errorf("%s expects Projection(row.field, ...) or the row parameter", method)
}
callee, ok := constructor.Callee.(IdentExpr)
if !ok {
return ClassDecl{}, nil, fmt.Errorf("%s projection must construct a local data class", method)
}
projection, ok := g.semantic.Classes[callee.Name]
if !ok || !projection.Data {
return ClassDecl{}, nil, fmt.Errorf("%s projection type %s must be a data class", method, callee.Name)
}
if len(constructor.Args) != len(projection.Fields) {
return ClassDecl{}, nil, fmt.Errorf("%s projection %s expects %d fields, got %d", method, projection.Name, len(projection.Fields), len(constructor.Args))
}
columns := make([]string, 0, len(constructor.Args))
for i, arg := range constructor.Args {
source, err := sqlRowField(arg, rowClass, rowName)
if err != nil {
return ClassDecl{}, nil, fmt.Errorf("%s projection argument %d: %w", method, i+1, err)
}
target := projection.Fields[i]
if !sqlProjectionTypesCompatible(target.Type, source.Type) {
return ClassDecl{}, nil, fmt.Errorf("%s projection field %s has type %s but row field %s has type %s", method, target.Name, target.Type, source.Name, source.Type)
}
columns = append(columns, sqlColumn(source))
}
return projection, columns, nil
}
func (g *goGenerator) sqlWhere(call CallExpr, class ClassDecl, placeholderStart int) (string, []string, error) {
lambda, err := sqlLambdaArg(call, "where")
if err != nil {
return "", nil, err
}
rowName, body, err := sqlLambdaExpr(lambda, class, "where")
if err != nil {
return "", nil, err
}
predicate, args, err := g.sqlPredicate(body, class, rowName)
if err != nil {
return "", nil, err
}
return postgresPlaceholders(predicate, placeholderStart), args, nil
}
func sqlOrder(step sqlCallStep, class ClassDecl) (string, error) {
lambda, err := sqlLambdaArg(step.call, step.name)
if err != nil {
return "", err
}
rowName, body, err := sqlLambdaExpr(lambda, class, step.name)
if err != nil {
return "", err
}
field, err := sqlRowField(body, class, rowName)
if err != nil {
return "", fmt.Errorf("%s: %w", step.name, err)
}
direction := ""
if step.name == "orderByDescending" {
direction = " DESC"
}
return sqlColumn(field) + direction, nil
}
func (g *goGenerator) sqlLimit(call CallExpr) (string, []string, error) {
if len(call.Args) != 1 || len(call.NamedArgs) != 0 || len(call.TypeArgs) != 0 {
return "", nil, fmt.Errorf("limit() expects exactly one Int argument")
}
if literal, ok := call.Args[0].(IntExpr); ok {
value, err := strconv.Atoi(literal.Value)
if err != nil || value < 0 {
return "", nil, fmt.Errorf("limit() requires a non-negative Int")
}
return literal.Value, nil, nil
}
if unary, ok := call.Args[0].(UnaryExpr); ok && unary.Op == "-" {
if _, ok := unary.Value.(IntExpr); ok {
return "", nil, fmt.Errorf("limit() requires a non-negative Int")
}
}
if typ := g.exprType(call.Args[0]); typ != "Int" {
return "", nil, fmt.Errorf("limit() argument must have type Int, got %s", typ)
}
value, err := g.expr(call.Args[0], "Int")
if err != nil {
return "", nil, err
}
return "?", []string{value}, nil
}
func sqlNoArgs(call CallExpr) bool {
return len(call.Args) == 0 && len(call.NamedArgs) == 0 && len(call.TypeArgs) == 0
}
func postgresPlaceholders(query string, start int) string {
var out strings.Builder
next := start
for _, r := range query {
if r == '?' {
out.WriteString("$")
out.WriteString(strconv.Itoa(next))
next++
} else {
out.WriteRune(r)
}
}
return out.String()
}
func sqlLambdaArg(call CallExpr, method string) (LambdaExpr, error) {
if len(call.TypeArgs) != 0 || len(call.NamedArgs) != 0 || len(call.Args) != 1 {
return LambdaExpr{}, fmt.Errorf("%s expects exactly one lambda", method)
}
lambda, ok := call.Args[0].(LambdaExpr)
if !ok {
return LambdaExpr{}, fmt.Errorf("%s expects a lambda", method)
}
return lambda, nil
}
func sqlLambdaExpr(lambda LambdaExpr, class ClassDecl, method string) (string, Expr, error) {
rowName := "it"
if !lambda.ImplicitIt {
if len(lambda.Params) != 1 {
return "", nil, fmt.Errorf("%s lambda expects one row parameter", method)
}
if lambda.Params[0].Type != "" && lambda.Params[0].Type != class.Name {
return "", nil, fmt.Errorf("%s lambda parameter must have type %s", method, class.Name)
}
rowName = lambda.Params[0].Name
}
if len(lambda.Body) != 1 {
return "", nil, fmt.Errorf("%s lambda must contain one expression", method)
}
stmt, ok := lambda.Body[0].(ExprStmt)
if !ok {
return "", nil, fmt.Errorf("%s lambda must contain one expression", method)
}
return rowName, stmt.Value, nil
}
func (g *goGenerator) sqlPredicate(expr Expr, class ClassDecl, rowName string) (string, []string, error) {
switch e := expr.(type) {
case BinaryExpr:
if e.Op == "&&" || e.Op == "||" {
left, leftArgs, err := g.sqlPredicate(e.Left, class, rowName)
if err != nil {
return "", nil, err
}
right, rightArgs, err := g.sqlPredicate(e.Right, class, rowName)
if err != nil {
return "", nil, err
}
op := "AND"
if e.Op == "||" {
op = "OR"
}
return "(" + left + " " + op + " " + right + ")", append(leftArgs, rightArgs...), nil
}
if !isSQLComparison(e.Op) {
return "", nil, fmt.Errorf("unsupported SQL predicate operator %q", e.Op)
}
left, err := g.sqlOperand(e.Left, class, rowName)
if err != nil {
return "", nil, err
}
right, err := g.sqlOperand(e.Right, class, rowName)
if err != nil {
return "", nil, err
}
if left.typ == "Null" || right.typ == "Null" {
if e.Op != "==" && e.Op != "!=" {
return "", nil, fmt.Errorf("null only supports == and != in SQL predicates")
}
if left.typ == "Null" && right.typ == "Null" {
return "", nil, fmt.Errorf("SQL predicate cannot compare null with null")
}
value := left
if value.typ == "Null" {
value = right
}
if !sqlNullableType(value.typ) {
return "", nil, fmt.Errorf("SQL null comparison requires a nullable operand, got %s", value.typ)
}
op := "IS NULL"
if e.Op == "!=" {
op = "IS NOT NULL"
}
return value.sql + " " + op, value.args, nil
}
if !sqlTypesCompatible(left.typ, right.typ) {
return "", nil, fmt.Errorf("SQL predicate compares incompatible types %s and %s", left.typ, right.typ)
}
if e.Op == ">" || e.Op == ">=" || e.Op == "<" || e.Op == "<=" {
if !(sqlNumericType(left.typ) && sqlNumericType(right.typ)) && !(sqlTimestampType(left.typ) && sqlTimestampType(right.typ)) {
return "", nil, fmt.Errorf("SQL ordering comparison requires numeric operands or timestamp operands, got %s and %s", left.typ, right.typ)
}
}
op := map[string]string{"==": "=", "!=": "<>", ">": ">", ">=": ">=", "<": "<", "<=": "<="}[e.Op]
return left.sql + " " + op + " " + right.sql, append(left.args, right.args...), nil
case UnaryExpr:
if e.Op != "!" {
return "", nil, fmt.Errorf("unsupported SQL predicate unary operator %q", e.Op)
}
inner, args, err := g.sqlPredicate(e.Value, class, rowName)
if err != nil {
return "", nil, err
}
return "(NOT " + inner + ")", args, nil
case SelectorExpr:
field, err := sqlRowField(e, class, rowName)
if err != nil {
return "", nil, err
}
if field.Type != "Boolean" {
return "", nil, fmt.Errorf("SQL predicate field %s has type %s, not Boolean", field.Name, field.Type)
}
return sqlColumn(field), nil, nil
default:
return "", nil, fmt.Errorf("SQL where lambda must produce a Boolean predicate")
}
}
type sqlOperandValue struct {
sql string
typ string
args []string
}
func (g *goGenerator) sqlOperand(expr Expr, class ClassDecl, rowName string) (sqlOperandValue, error) {
if selector, ok := expr.(SelectorExpr); ok {
if receiver, ok := selector.Receiver.(IdentExpr); ok && receiver.Name == rowName {
field, err := sqlField(class, selector.Name)
if err != nil {
return sqlOperandValue{}, err
}
return sqlOperandValue{sql: sqlColumn(field), typ: field.Type}, nil
}
}
if call, ok := expr.(CallExpr); ok {
if ident, ok := call.Callee.(IdentExpr); ok && ident.Name == "now" {
if !sqlNoArgs(call) {
return sqlOperandValue{}, fmt.Errorf("now() does not accept arguments")
}
return sqlOperandValue{sql: "CURRENT_TIMESTAMP", typ: "time.Time"}, nil
}
}
if binary, ok := expr.(BinaryExpr); ok && (binary.Op == "+" || binary.Op == "-") {
left, err := g.sqlOperand(binary.Left, class, rowName)
if err != nil {
return sqlOperandValue{}, err
}
right, err := g.sqlOperand(binary.Right, class, rowName)
if err != nil {
return sqlOperandValue{}, err
}
if !sqlNumericType(left.typ) || !sqlNumericType(right.typ) {
return sqlOperandValue{}, fmt.Errorf("SQL arithmetic requires numeric operands, got %s and %s", left.typ, right.typ)
}
return sqlOperandValue{sql: "(" + left.sql + " " + binary.Op + " " + right.sql + ")", typ: left.typ, args: append(left.args, right.args...)}, nil
}
typ := ""
switch e := expr.(type) {
case IntExpr:
typ = "Int"
case FloatExpr:
typ = "Double"
case StringExpr:
typ = "String"
case BoolExpr:
typ = "Boolean"
case NullExpr:
return sqlOperandValue{sql: "NULL", typ: "Null"}, nil
case UnaryExpr:
if e.Op != "-" {
return sqlOperandValue{}, fmt.Errorf("unsupported SQL value expression")
}
switch e.Value.(type) {
case IntExpr:
typ = "Int"
case FloatExpr:
typ = "Double"
default:
return sqlOperandValue{}, fmt.Errorf("SQL unary parameters only support numeric literals")
}
default:
typ = g.exprType(expr)
}
if typ == "" {
return sqlOperandValue{}, fmt.Errorf("SQL parameter has unknown type")
}
value, err := g.expr(expr, typ)
if err != nil {
return sqlOperandValue{}, err
}
return sqlOperandValue{sql: "?", typ: typ, args: []string{value}}, nil
}
func (g *goGenerator) sqlTypedUpdateAssignments(call CallExpr, class ClassDecl) ([]string, []string, error) {
lambda, err := sqlLambdaArg(call, "set")
if err != nil {
return nil, nil, err
}
rowName := "it"
if !lambda.ImplicitIt {
if len(lambda.Params) != 1 {
return nil, nil, fmt.Errorf("set lambda expects one row parameter")
}
if lambda.Params[0].Type != "" && lambda.Params[0].Type != class.Name {
return nil, nil, fmt.Errorf("set lambda parameter must have type %s", class.Name)
}
rowName = lambda.Params[0].Name
}
if len(lambda.Body) == 0 {
return nil, nil, fmt.Errorf("set lambda requires at least one set(row.field, value) expression")
}
var assignments, args []string
seen := map[string]bool{}
for _, stmt := range lambda.Body {
exprStmt, ok := stmt.(ExprStmt)
if !ok {
return nil, nil, fmt.Errorf("set lambda only supports set(row.field, value) expressions")
}
setCall, ok := exprStmt.Value.(CallExpr)
if !ok {
return nil, nil, fmt.Errorf("set lambda only supports set(row.field, value) expressions")
}
callee, ok := setCall.Callee.(IdentExpr)
if !ok || callee.Name != "set" || len(setCall.Args) != 2 || len(setCall.NamedArgs) != 0 || len(setCall.TypeArgs) != 0 {
return nil, nil, fmt.Errorf("set lambda expects set(%s.field, value)", rowName)
}
target, err := sqlRowField(setCall.Args[0], class, rowName)
if err != nil {
return nil, nil, fmt.Errorf("set target: %w", err)
}
if !target.Mutable {
return nil, nil, fmt.Errorf("set target %s is immutable", target.Name)
}
if seen[target.Name] {
return nil, nil, fmt.Errorf("duplicate set target %s", target.Name)
}
value, err := g.sqlOperand(setCall.Args[1], class, rowName)
if err != nil {
return nil, nil, fmt.Errorf("set value for %s: %w", target.Name, err)
}
if value.typ == "Null" {
if !sqlNullableType(target.Type) {
return nil, nil, fmt.Errorf("set target %s has non-nullable type %s", target.Name, target.Type)
}
} else if !sqlAssignmentTypesCompatible(target.Type, value.typ) {
return nil, nil, fmt.Errorf("set target %s has type %s but value has type %s", target.Name, target.Type, value.typ)
}
seen[target.Name] = true
assignments = append(assignments, sqlColumn(target)+" = "+value.sql)
args = append(args, value.args...)
}
return assignments, args, nil
}
func sqlConflictUpdateAssignments(lambda LambdaExpr, class ClassDecl) ([]string, error) {
excludedName := "it"
if !lambda.ImplicitIt {
if len(lambda.Params) != 1 {
return nil, fmt.Errorf("doUpdate lambda expects one excluded-row parameter")
}
if lambda.Params[0].Type != "" && lambda.Params[0].Type != class.Name {
return nil, fmt.Errorf("doUpdate lambda parameter must have type %s", class.Name)
}
excludedName = lambda.Params[0].Name
}
if len(lambda.Body) == 0 {
return nil, fmt.Errorf("doUpdate lambda requires at least one set() expression")
}
assignments := make([]string, 0, len(lambda.Body))
seen := map[string]bool{}
for _, stmt := range lambda.Body {
exprStmt, ok := stmt.(ExprStmt)
if !ok {
return nil, fmt.Errorf("doUpdate only supports set() expressions")
}
call, ok := exprStmt.Value.(CallExpr)
if !ok {
return nil, fmt.Errorf("doUpdate only supports set() expressions")
}
callee, ok := call.Callee.(IdentExpr)
if !ok || callee.Name != "set" || len(call.Args) != 2 || len(call.NamedArgs) != 0 || len(call.TypeArgs) != 0 {
return nil, fmt.Errorf("doUpdate expects set(%s.field, %s.field)", class.Name, excludedName)
}
targetRef, ok := call.Args[0].(SelectorExpr)
if !ok {
return nil, fmt.Errorf("set target must be %s.field", class.Name)
}
targetType, ok := targetRef.Receiver.(IdentExpr)
if !ok || targetType.Name != class.Name {
return nil, fmt.Errorf("set target must be %s.field", class.Name)
}
target, err := sqlField(class, targetRef.Name)
if err != nil {
return nil, fmt.Errorf("set target: %w", err)
}
if !target.Mutable {
return nil, fmt.Errorf("set target %s is immutable", target.Name)
}
if seen[target.Name] {
return nil, fmt.Errorf("duplicate set target %s", target.Name)
}
sourceRef, ok := call.Args[1].(SelectorExpr)
if !ok {
return nil, fmt.Errorf("set value must be %s.field", excludedName)
}
sourceReceiver, ok := sourceRef.Receiver.(IdentExpr)
if !ok || sourceReceiver.Name != excludedName {
return nil, fmt.Errorf("set value must be %s.field", excludedName)
}
source, err := sqlField(class, sourceRef.Name)
if err != nil {
return nil, fmt.Errorf("set value: %w", err)
}
if !sqlAssignmentTypesCompatible(target.Type, source.Type) {
return nil, fmt.Errorf("set target %s has type %s but excluded.%s has type %s", target.Name, target.Type, source.Name, source.Type)
}
seen[target.Name] = true
assignments = append(assignments, sqlColumn(target)+" = EXCLUDED."+sqlColumn(source))
}
return assignments, nil
}
func sqlConflictFields(expr Expr, class ClassDecl, rowName string) ([]FieldDecl, error) {
if call, ok := expr.(CallExpr); ok {
callee, ok := call.Callee.(IdentExpr)
if !ok || (callee.Name != "listOf" && callee.Name != "mutableListOf") || len(call.Args) == 0 || len(call.TypeArgs) != 0 || len(call.NamedArgs) != 0 {
return nil, fmt.Errorf("onConflict expects a row field or listOf(row.field, ...)")
}
fields := make([]FieldDecl, 0, len(call.Args))
seen := map[string]bool{}
for _, arg := range call.Args {
field, err := sqlRowField(arg, class, rowName)
if err != nil {
return nil, fmt.Errorf("onConflict: %w", err)
}
if seen[field.Name] {
return nil, fmt.Errorf("onConflict contains duplicate field %s", field.Name)
}
seen[field.Name] = true
fields = append(fields, field)
}
return fields, nil
}
field, err := sqlRowField(expr, class, rowName)
if err != nil {
return nil, fmt.Errorf("onConflict: %w", err)
}
return []FieldDecl{field}, nil
}
func sqlRowField(expr Expr, class ClassDecl, rowName string) (FieldDecl, error) {
selector, ok := expr.(SelectorExpr)
if !ok {
return FieldDecl{}, fmt.Errorf("expected %s.field", rowName)
}
receiver, ok := selector.Receiver.(IdentExpr)
if !ok || receiver.Name != rowName {
return FieldDecl{}, fmt.Errorf("expected %s.field", rowName)
}
return sqlField(class, selector.Name)
}
func sqlField(class ClassDecl, name string) (FieldDecl, error) {
for _, field := range class.Fields {
if field.Name == name {
return field, nil
}
}
return FieldDecl{}, fmt.Errorf("SQL row class %s has no field %s", class.Name, name)
}
func sqlColumn(field FieldDecl) string {
if field.Column != "" {
return field.Column
}
return snakeCase(field.Name)
}
func sqlQueryValue(query string, args []string) string {
return "GotlinSQLQuery{SQL: " + strconv.Quote(query) + ", Args: []any{" + strings.Join(args, ", ") + "}}"
}
func isSQLComparison(op string) bool {
switch op {
case "==", "!=", ">", ">=", "<", "<=":
return true
default:
return false
}
}
func sqlBaseType(typ string) string {
return strings.TrimSuffix(typ, "?")
}
func sqlNullableType(typ string) bool {
return strings.HasSuffix(typ, "?")
}
func sqlNumericType(typ string) bool {
switch sqlBaseType(typ) {
case "Int", "Float", "Double":
return true
default:
return false
}
}
func sqlTimestampType(typ string) bool {
return sqlBaseType(typ) == "time.Time"
}
func sqlTypesCompatible(left, right string) bool {
if left == "Null" || right == "Null" {
return true
}
left = sqlBaseType(left)
right = sqlBaseType(right)
if left == right {
return true
}
return sqlNumericType(left) && sqlNumericType(right)
}
func sqlProjectionTypesCompatible(target, source string) bool {
return target == source
}
func sqlAssignmentTypesCompatible(target, source string) bool {
if sqlBaseType(target) != sqlBaseType(source) {
return false
}
return sqlNullableType(target) || !sqlNullableType(source)
}
func programContainsSQL(program *Program) bool {
return programExprMatches(program, func(expr Expr) bool {
_, _, _, ok := splitSQLChain(expr)
return ok
})
}
func programContainsSQLExecution(program *Program) bool {
return programExprMatches(program, func(expr Expr) bool {
_, _, steps, ok := splitSQLChain(expr)
if !ok || len(steps) == 0 {
return false
}
switch steps[len(steps)-1].name {
case "fetch", "single", "iterator":
return true
default:
return false
}
})
}
func exprMatches(expr Expr, match func(Expr) bool) bool {
if match(expr) {
return true
}
switch e := expr.(type) {
case UnaryExpr:
return exprMatches(e.Value, match)
case BinaryExpr:
return exprMatches(e.Left, match) || exprMatches(e.Right, match)
case CallExpr:
if exprMatches(e.Callee, match) {
return true
}
for _, arg := range e.Args {
if exprMatches(arg, match) {
return true
}
}
for _, arg := range e.NamedArgs {
if exprMatches(arg.Value, match) {
return true
}
}
case SelectorExpr:
return exprMatches(e.Receiver, match)
case IndexExpr:
return exprMatches(e.Receiver, match) || exprMatches(e.Index, match)
case MatchExpr:
if exprMatches(e.Value, match) {
return true
}
for _, matchCase := range e.Cases {
if exprMatches(matchCase.Value, match) {
return true
}
}
case LambdaExpr:
return stmtsMatch(e.Body, match)
}
return false
}
func stmtsMatch(stmts []Stmt, match func(Expr) bool) bool {
for _, stmt := range stmts {
switch s := stmt.(type) {
case VarDecl:
if exprMatches(s.Value, match) {
return true
}
case MultiVarDecl:
if exprMatches(s.Value, match) {
return true
}
case AssignStmt:
if exprMatches(s.Value, match) {
return true
}
case AddAssignStmt:
if exprMatches(s.Value, match) {
return true
}
case MultiAssignStmt:
if exprMatches(s.Value, match) {
return true
}
case ReturnStmt:
if s.Value != nil && exprMatches(s.Value, match) {
return true
}
case ThrowStmt:
if exprMatches(s.Value, match) {
return true
}
case DeferStmt:
if exprMatches(s.Value, match) {
return true
}
case ExprStmt:
if exprMatches(s.Value, match) {
return true
}
case IfStmt:
if exprMatches(s.Cond, match) || stmtsMatch(s.Then, match) || stmtsMatch(s.Else, match) {
return true
}
case WhileStmt:
if exprMatches(s.Cond, match) || stmtsMatch(s.Body, match) {
return true
}
case ForEachStmt:
if exprMatches(s.Source, match) || stmtsMatch(s.Body, match) {
return true
}
case TryCatchStmt:
if stmtsMatch(s.TryBody, match) || stmtsMatch(s.CatchBody, match) {
return true
}
}
}
return false
}
func programExprMatches(program *Program, match func(Expr) bool) bool {
for _, fn := range program.Functions {
if fn.ExpressionBody != nil && exprMatches(fn.ExpressionBody, match) {
return true
}
if stmtsMatch(fn.Body, match) {
return true
}
}
for _, class := range program.Classes {
for _, fn := range class.Methods {
if fn.ExpressionBody != nil && exprMatches(fn.ExpressionBody, match) {
return true
}
if stmtsMatch(fn.Body, match) {
return true
}
}
}
return false
}