Expand Gotlin language and tooling

This commit is contained in:
pavel 2026-08-27 01:46:57 +02:00
commit de1262b4cf
41 changed files with 6059 additions and 379 deletions

View file

@ -4,9 +4,24 @@ type Program struct {
PackagePath string
Imports []ImportDecl
Interfaces []InterfaceDecl
Enums []EnumDecl
Classes []ClassDecl
Workers []WorkerDecl
Functions []FunctionDecl
Embeds []EmbedDecl
}
type EmbedDecl struct{ Path, Name, Type string }
type EnumDecl struct {
Name string
Variants []EnumVariant
String bool
}
type EnumVariant struct {
Name string
PayloadTypes []string
StringValue string
}
type ImportDecl struct {
@ -20,10 +35,13 @@ type InterfaceDecl struct {
}
type ClassDecl struct {
Name string
Fields []FieldDecl
Parents []string
Methods []FunctionDecl
Name string
Data bool
JSONNaming string
Table string
Fields []FieldDecl
Parents []string
Methods []FunctionDecl
}
type WorkerDecl struct {
@ -40,9 +58,13 @@ type WorkerFieldDecl struct {
}
type FieldDecl struct {
Mutable bool
Name string
Type string
Mutable bool
Private bool
Name string
Type string
Column string
ID bool
Generated bool
}
type FunctionSignature struct {
@ -90,6 +112,7 @@ func (MultiVarDecl) stmtNode() {}
type AssignStmt struct {
Name string
Pos int
Value Expr
}
@ -97,14 +120,16 @@ func (AssignStmt) stmtNode() {}
type AddAssignStmt struct {
Name string
Pos int
Value Expr
}
func (AddAssignStmt) stmtNode() {}
type MultiAssignStmt struct {
Names []string
Value Expr
Names []string
Positions []int
Value Expr
}
func (MultiAssignStmt) stmtNode() {}
@ -127,6 +152,12 @@ type GoStmt struct {
func (GoStmt) stmtNode() {}
type DeferStmt struct {
Value Expr
}
func (DeferStmt) stmtNode() {}
type ExprStmt struct {
Value Expr
}
@ -148,6 +179,14 @@ type WhileStmt struct {
func (WhileStmt) stmtNode() {}
type ForEachStmt struct {
Name string
Source Expr
Body []Stmt
}
func (ForEachStmt) stmtNode() {}
type TryCatchStmt struct {
TryBody []Stmt
CatchName string
@ -163,6 +202,19 @@ type SelectStmt struct {
func (SelectStmt) stmtNode() {}
type MatchStmt struct {
Value Expr
Cases []MatchCase
}
func (MatchStmt) stmtNode() {}
type MatchCase struct {
EnumName, VariantName string
Bindings []string
Body []Stmt
}
type SelectCase struct {
Source Expr
Body []Stmt
@ -180,6 +232,12 @@ type IntExpr struct {
func (IntExpr) exprNode() {}
type FloatExpr struct {
Value string
}
func (FloatExpr) exprNode() {}
type StringExpr struct {
Value string
}
@ -212,9 +270,15 @@ type BinaryExpr struct {
func (BinaryExpr) exprNode() {}
type CallExpr struct {
Callee Expr
Args []Expr
TypeArgs []string
Callee Expr
Args []Expr
TypeArgs []string
NamedArgs []NamedArg
}
type NamedArg struct {
Name string
Value Expr
}
func (CallExpr) exprNode() {}
@ -226,6 +290,20 @@ type SelectorExpr struct {
func (SelectorExpr) exprNode() {}
type IndexExpr struct {
Receiver Expr
Index Expr
}
func (IndexExpr) exprNode() {}
type EnumVariantExpr struct {
EnumName, VariantName string
Values []Expr
}
func (EnumVariantExpr) exprNode() {}
type LambdaExpr struct {
Params []Param
ImplicitIt bool

View file

@ -867,6 +867,183 @@ worker Counter {
}
}
func TestGenerateGoRejectsValReassignment(t *testing.T) {
src := `
package demo
fun main() {
val count = 0
count = 1
}
`
prog, err := Parse(src)
if err != nil {
t.Fatalf("parse failed: %v", err)
}
_, err = GenerateGo(prog)
if err == nil {
t.Fatal("expected val reassignment error")
}
if !strings.Contains(err.Error(), "cannot reassign immutable name count") {
t.Fatalf("unexpected error: %v", err)
}
}
func TestGenerateGoDefer(t *testing.T) {
src := `
package demo
import database.sql
fun main() {
val rows = sql.Open("postgres", "dsn")
defer rows.Close()
}
`
prog, err := Parse(src)
if err != nil {
t.Fatalf("parse failed: %v", err)
}
out, err := GenerateGo(prog)
if err != nil {
t.Fatalf("go generation failed: %v", err)
}
if !strings.Contains(string(out), "defer rows.Close()") {
t.Fatalf("generated Go missing defer:\n%s", out)
}
}
func TestGenerateGoDecimalLiteral(t *testing.T) {
src := `
package demo
fun interest(): Double {
return 0.025
}
`
prog, err := Parse(src)
if err != nil {
t.Fatalf("parse failed: %v", err)
}
out, err := GenerateGo(prog)
if err != nil {
t.Fatalf("go generation failed: %v", err)
}
for _, want := range []string{"func interest() float64", "return 0.025"} {
if !strings.Contains(string(out), want) {
t.Fatalf("generated Go missing %q:\n%s", want, out)
}
}
}
func TestGenerateGoRejectsMultiValReassignment(t *testing.T) {
src := `
package demo
import database.sql
fun main() {
val db, err = sql.Open("postgres", "postgres://localhost/postgres?sslmode=disable")
err = nil
}
`
prog, err := Parse(src)
if err != nil {
t.Fatalf("parse failed: %v", err)
}
_, err = GenerateGo(prog)
if err == nil {
t.Fatal("expected multi-val reassignment error")
}
if !strings.Contains(err.Error(), "cannot reassign immutable name err") {
t.Fatalf("unexpected error: %v", err)
}
}
func TestGenerateGoAllowsVarReassignment(t *testing.T) {
src := `
package demo
fun main() {
var count = 0
count = 1
count += 2
println(count)
}
`
prog, err := Parse(src)
if err != nil {
t.Fatalf("parse failed: %v", err)
}
out, err := GenerateGo(prog)
if err != nil {
t.Fatalf("go generation failed: %v", err)
}
code := string(out)
for _, want := range []string{
`count := 0`,
`count = 1`,
`count += 2`,
} {
if !strings.Contains(code, want) {
t.Fatalf("generated Go missing %q:\n%s", want, code)
}
}
}
func TestGenerateGoRejectsReassignmentToImmutableClassField(t *testing.T) {
src := `
package demo
class Counter(val count: Int, var total: Int) {
fun freeze() {
count = 1
}
}
`
prog, err := Parse(src)
if err != nil {
t.Fatalf("parse failed: %v", err)
}
_, err = GenerateGo(prog)
if err == nil {
t.Fatal("expected immutable class field reassignment error")
}
if !strings.Contains(err.Error(), "cannot reassign immutable name count") {
t.Fatalf("unexpected error: %v", err)
}
}
func TestGenerateGoAllowsReassignmentToMutableClassField(t *testing.T) {
src := `
package demo
class Counter(var count: Int) {
fun inc() {
count += 1
}
}
`
prog, err := Parse(src)
if err != nil {
t.Fatalf("parse failed: %v", err)
}
out, err := GenerateGo(prog)
if err != nil {
t.Fatalf("go generation failed: %v", err)
}
if !strings.Contains(string(out), `self.count += 1`) {
t.Fatalf("generated Go missing mutable field assignment:\n%s", string(out))
}
}
func TestGenerateGoWorkerSyncPanicPropagationScaffolding(t *testing.T) {
src := `
package demo

View file

@ -0,0 +1,72 @@
package lang
import (
"strings"
"testing"
)
func TestGenerateGoDataClassUsesSnakeCaseJSONTags(t *testing.T) {
prog, err := Parse(`
package demo
data class KeycloakToken(var accessToken: String, var refreshToken: String)
`)
if err != nil {
t.Fatalf("parse failed: %v", err)
}
out, err := GenerateGo(prog)
if err != nil {
t.Fatalf("generation failed: %v", err)
}
for _, want := range []string{"AccessToken", "`json:\"access_token\"`", "RefreshToken", "`json:\"refresh_token\"`"} {
if !strings.Contains(string(out), want) {
t.Fatalf("generated Go missing %q:\n%s", want, out)
}
}
}
func TestGenerateGoDataClassJSONNamingOverride(t *testing.T) {
prog, err := Parse(`
package demo
@jsonNaming(camelCase)
data class Account(var availableBalance: Double)
`)
if err != nil {
t.Fatalf("parse failed: %v", err)
}
out, err := GenerateGo(prog)
if err != nil {
t.Fatalf("generation failed: %v", err)
}
if !strings.Contains(string(out), "`json:\"availableBalance\"`") {
t.Fatalf("generated Go missing camelCase tag:\n%s", out)
}
}
func TestGenerateGoDataClassFieldVisibilityAndSelector(t *testing.T) {
prog, err := Parse(`
package demo
import json encoding.json
data class Request(var email: String, private val traceId: String)
fun email(body: ByteSlice): String {
val request = json.decode<Request>(body)
return request.email
}
`)
if err != nil {
t.Fatalf("parse failed: %v", err)
}
out, err := GenerateGo(prog)
if err != nil {
t.Fatalf("generation failed: %v", err)
}
for _, want := range []string{"Email", "`json:\"email\"`", "traceId string `json:\"-\"`", "return request.Email"} {
if !strings.Contains(string(out), want) {
t.Fatalf("generated Go missing %q:\n%s", want, out)
}
}
}

View file

@ -0,0 +1,79 @@
package lang
import (
"strings"
"testing"
)
func TestGenerateRustStyleEnumAndExhaustiveMatch(t *testing.T) {
prog, err := Parse(`
package demo
enum PaymentResult { Accepted(String), Rejected(String), Pending }
fun describe(result: PaymentResult): String {
var description = ""
match (result) {
PaymentResult::Accepted(id) -> { description = id }
PaymentResult::Rejected(reason) -> { description = reason }
PaymentResult::Pending -> { description = "pending" }
}
return description
}
fun main() { println(describe(PaymentResult::Accepted("p1"))) }
`)
if err != nil {
t.Fatal(err)
}
out, err := GenerateGo(prog)
if err != nil {
t.Fatal(err)
}
for _, want := range []string{"type PaymentResult interface", "type PaymentResultAccepted struct", "&PaymentResultAccepted{Value0: \"p1\"}", "case *PaymentResultRejected:", "reason := gotlinMatch1.Value0"} {
if !strings.Contains(string(out), want) {
t.Fatalf("missing %q:\n%s", want, out)
}
}
}
func TestRejectNonExhaustiveEnumMatch(t *testing.T) {
prog, err := Parse(`package demo
enum Result { Ok, Error(String) }
fun use(result: Result) { match (result) { Result::Ok -> { println("ok") } } }`)
if err != nil {
t.Fatal(err)
}
_, err = GenerateGo(prog)
if err == nil || !strings.Contains(err.Error(), "missing Error") {
t.Fatalf("expected exhaustive-match error, got %v", err)
}
}
func TestRejectWrongVariantPayloadCount(t *testing.T) {
prog, err := Parse(`package demo
enum Result { Ok(String) }
fun main() { val result = Result::Ok() }`)
if err != nil {
t.Fatal(err)
}
_, err = GenerateGo(prog)
if err == nil || !strings.Contains(err.Error(), "expects 1 values") {
t.Fatalf("expected payload error, got %v", err)
}
}
func TestPayloadlessEnumUsesExactVariantStrings(t *testing.T) {
prog, err := Parse(`package demo
enum Status { PendingReservation, Initiated }
fun main() { println(Status::PendingReservation) }`)
if err != nil {
t.Fatal(err)
}
out, err := GenerateGo(prog)
if err != nil {
t.Fatal(err)
}
for _, want := range []string{`type Status string`, `StatusPendingReservation`, `"PendingReservation"`, `StatusInitiated`, `"Initiated"`} {
if !strings.Contains(string(out), want) {
t.Fatalf("missing %q:\n%s", want, out)
}
}
}

View file

@ -0,0 +1,62 @@
package lang
import (
"strings"
"testing"
)
func TestGenerateExternalGoStructConstruction(t *testing.T) {
prog, err := Parse(`package demo
import platform "example/platform"
fun create(issuer: String, audience: String, jwksUrl: String): *platform.Authorizer {
return platform.Authorizer(issuer = issuer, audience = audience, jwksUrl = jwksUrl)
}`)
if err != nil {
t.Fatal(err)
}
out, err := GenerateGo(prog)
if err != nil {
t.Fatal(err)
}
want := `&platform.Authorizer{Issuer: issuer, Audience: audience, JWKSURL: jwksUrl}`
if !strings.Contains(string(out), want) {
t.Fatalf("missing %q:\n%s", want, out)
}
}
func TestGenerateEmbeddedValue(t *testing.T) {
prog, err := Parse(`package demo
import embed
@embed("assets/*") val assets: embed.FS
fun main() { println(assets) }`)
if err != nil {
t.Fatal(err)
}
out, err := GenerateGo(prog)
if err != nil {
t.Fatal(err)
}
for _, want := range []string{`"embed"`, `//go:embed assets/*`, `var assets embed.FS`} {
if !strings.Contains(string(out), want) {
t.Fatalf("missing %q:\n%s", want, out)
}
}
}
func TestGenerateEmbeddedStringAddsBlankImport(t *testing.T) {
prog, err := Parse(`package demo
@embed("schema.sql") val schema: String
fun main() { println(schema) }`)
if err != nil {
t.Fatal(err)
}
out, err := GenerateGo(prog)
if err != nil {
t.Fatal(err)
}
for _, want := range []string{`_ "embed"`, `//go:embed schema.sql`, `var schema string`} {
if !strings.Contains(string(out), want) {
t.Fatalf("missing %q:\n%s", want, out)
}
}
}

28
internal/lang/foreach.go Normal file
View file

@ -0,0 +1,28 @@
package lang
import "fmt"
func (p *parser) parseForEach() (Stmt, error) {
if _, err := p.expect(tokenLParen, "expected '(' after 'for'"); err != nil {
return nil, err
}
name, err := p.expect(tokenIdent, "expected iteration variable")
if err != nil {
return nil, err
}
if _, err := p.expect(tokenIn, "expected 'in' after iteration variable"); err != nil {
return nil, err
}
source, err := p.parseExpr(0)
if err != nil {
return nil, err
}
if _, err := p.expect(tokenRParen, "expected ')' after iteration source"); err != nil {
return nil, err
}
body, err := p.parseBlock()
if err != nil {
return nil, fmt.Errorf("invalid for body: %w", err)
}
return ForEachStmt{Name: name.lexeme, Source: source, Body: body}, nil
}

View file

@ -0,0 +1,31 @@
package lang
import (
"strings"
"testing"
)
func TestGenerateGoForEach(t *testing.T) {
prog, err := Parse(`
package demo
fun main() {
val accounts: List<String> = listOf("one", "two")
for (account in accounts) {
println(account)
}
}
`)
if err != nil {
t.Fatalf("parse failed: %v", err)
}
out, err := GenerateGo(prog)
if err != nil {
t.Fatalf("generation failed: %v", err)
}
for _, want := range []string{"for _, account := range accounts {", "fmt.Println(account)"} {
if !strings.Contains(string(out), want) {
t.Fatalf("generated Go missing %q:\n%s", want, out)
}
}
}

File diff suppressed because it is too large Load diff

View file

@ -0,0 +1,32 @@
package lang
import (
"strings"
"testing"
)
func TestGenerateGoLowercaseImportedSelectors(t *testing.T) {
prog, err := Parse(`
package demo
import http net.http
import json encoding.json
fun main() {
http.handleFunc("/health", handler)
json.unmarshal(body, &target)
}
`)
if err != nil {
t.Fatalf("parse failed: %v", err)
}
out, err := GenerateGo(prog)
if err != nil {
t.Fatalf("generation failed: %v", err)
}
for _, want := range []string{"http.HandleFunc(\"/health\", handler)", "json.Unmarshal(body, &target)"} {
if !strings.Contains(string(out), want) {
t.Fatalf("generated Go missing %q:\n%s", want, out)
}
}
}

View file

@ -0,0 +1,26 @@
package lang
import (
"strings"
"testing"
)
func TestGenerateGoIndexedDataClassSelector(t *testing.T) {
prog, err := Parse(`
package demo
data class User(var id: String)
fun first(users: List<User>): String { return users[0].id }
`)
if err != nil {
t.Fatalf("parse failed: %v", err)
}
out, err := GenerateGo(prog)
if err != nil {
t.Fatalf("generation failed: %v", err)
}
if !strings.Contains(string(out), "return users[0].Id") {
t.Fatalf("generated Go missing indexed selector:\n%s", out)
}
}

View file

@ -0,0 +1,31 @@
package lang
import (
"strings"
"testing"
)
func TestGenerateGoJSONDecodeReturnsTypedValue(t *testing.T) {
prog, err := Parse(`
package demo
import json encoding.json
fun decode(body: ByteSlice): List<Account> {
val accounts = json.decode<List<Account>>(body)
return accounts
}
`)
if err != nil {
t.Fatalf("parse failed: %v", err)
}
out, err := GenerateGo(prog)
if err != nil {
t.Fatalf("generation failed: %v", err)
}
for _, want := range []string{"accounts := gotlinAutoThrow(gotlinJSONDecode[[]Account](body))", "json.Unmarshal(body, &value)"} {
if !strings.Contains(string(out), want) {
t.Fatalf("generated Go missing %q:\n%s", want, out)
}
}
}

View file

@ -49,6 +49,13 @@ func (l *lexer) next() (token, error) {
for l.pos < len(l.src) && unicode.IsDigit(l.src[l.pos]) {
l.pos++
}
if l.pos+1 < len(l.src) && l.src[l.pos] == '.' && unicode.IsDigit(l.src[l.pos+1]) {
l.pos++
for l.pos < len(l.src) && unicode.IsDigit(l.src[l.pos]) {
l.pos++
}
return token{kind: tokenFloat, lexeme: string(l.src[start:l.pos]), pos: start}, nil
}
return token{kind: tokenInt, lexeme: string(l.src[start:l.pos]), pos: start}, nil
case ch == '"':
l.pos++
@ -77,11 +84,18 @@ func (l *lexer) next() (token, error) {
return token{kind: tokenLBrace, lexeme: "{", pos: start}, nil
case '}':
return token{kind: tokenRBrace, lexeme: "}", pos: start}, nil
case '[':
return token{kind: tokenLBracket, lexeme: "[", pos: start}, nil
case ']':
return token{kind: tokenRBracket, lexeme: "]", pos: start}, nil
case ',':
return token{kind: tokenComma, lexeme: ",", pos: start}, nil
case '.':
return token{kind: tokenDot, lexeme: ".", pos: start}, nil
case ':':
if l.match(':') {
return token{kind: tokenDoubleColon, lexeme: "::", pos: start}, nil
}
return token{kind: tokenColon, lexeme: ":", pos: start}, nil
case ';':
return token{kind: tokenSemicolon, lexeme: ";", pos: start}, nil
@ -131,6 +145,11 @@ func (l *lexer) next() (token, error) {
if l.match('&') {
return token{kind: tokenAnd, lexeme: "&&", pos: start}, nil
}
return token{kind: tokenAmp, lexeme: "&", pos: start}, nil
case '@':
return token{kind: tokenAt, lexeme: "@", pos: start}, nil
case '?':
return token{kind: tokenQuestion, lexeme: "?", pos: start}, nil
case '|':
if l.match('|') {
return token{kind: tokenOr, lexeme: "||", pos: start}, nil

269
internal/lang/mapping.go Normal file
View file

@ -0,0 +1,269 @@
package lang
import (
"fmt"
"strconv"
"strings"
)
type mappingPair struct{ source, target, function string }
type mappingState struct {
functions map[string]string
pairs []mappingPair
}
func (g *goGenerator) mappingTopLevelTarget(target string) string {
if _, ok := g.classForType(target); ok && !strings.HasPrefix(target, "*") {
return "*" + target
}
return target
}
func (g *goGenerator) ensureMapping(source, target, path string) (string, error) {
key := source + "->" + target
if function, ok := g.mappings.functions[key]; ok {
return function, nil
}
if err := g.validateMapping(source, target, path, map[string]bool{}); err != nil {
return "", err
}
function := fmt.Sprintf("gotlinMap%d", len(g.mappings.pairs)+1)
g.mappings.functions[key] = function
g.mappings.pairs = append(g.mappings.pairs, mappingPair{source: source, target: target, function: function})
return function, nil
}
func (g *goGenerator) validateMapping(source, target, path string, seen map[string]bool) error {
if source == target {
return nil
}
key := source + "->" + target
if seen[key] {
return nil
}
seen[key] = true
if strings.HasSuffix(source, "?") || strings.HasSuffix(target, "?") {
if strings.HasSuffix(source, "?") && !strings.HasSuffix(target, "?") {
return mappingError(path, source, target)
}
return g.validateMapping(strings.TrimSuffix(source, "?"), strings.TrimSuffix(target, "?"), path, seen)
}
if sourceBase, sourceArgs, ok := parseGenericType(source); ok {
targetBase, targetArgs, targetOK := parseGenericType(target)
if !targetOK || sourceBase != targetBase || len(sourceArgs) != len(targetArgs) {
return mappingError(path, source, target)
}
for i := range sourceArgs {
if sourceBase == "Map" || sourceBase == "MutableMap" {
if i == 0 && sourceArgs[i] != targetArgs[i] {
return mappingError(path+".<key>", sourceArgs[i], targetArgs[i])
}
}
if err := g.validateMapping(sourceArgs[i], targetArgs[i], path+"[]", seen); err != nil {
return err
}
}
return nil
}
sourceClass, sourceClassOK := g.classForType(source)
targetClass, targetClassOK := g.classForType(target)
if sourceClassOK || targetClassOK {
if !sourceClassOK || !targetClassOK {
return mappingError(path, source, target)
}
for _, targetField := range targetClass.Fields {
sourceField, ok := classFieldByName(sourceClass, targetField.Name)
fieldPath := path + "." + targetField.Name
if !ok {
return fmt.Errorf("cannot map %s: source field is missing for %s.%s", fieldPath, targetClass.Name, targetField.Name)
}
if err := g.validateMapping(sourceField.Type, targetField.Type, fieldPath, seen); err != nil {
return err
}
}
return nil
}
sourceEnum, sourceEnumOK := g.enums[strings.TrimPrefix(source, "*")]
targetEnum, targetEnumOK := g.enums[strings.TrimPrefix(target, "*")]
if sourceEnumOK || targetEnumOK {
if sourceEnumOK && enumIsString(sourceEnum) && target == "String" {
return nil
}
if targetEnumOK && enumIsString(targetEnum) && source == "String" {
return nil
}
if !sourceEnumOK || !targetEnumOK {
return mappingError(path, source, target)
}
for _, sourceVariant := range sourceEnum.Variants {
targetVariant := enumVariant(targetEnum, sourceVariant.Name)
variantPath := path + "::" + sourceVariant.Name
if targetVariant == nil {
return fmt.Errorf("cannot map %s: target enum %s has no compatible variant", variantPath, targetEnum.Name)
}
if len(sourceVariant.PayloadTypes) != len(targetVariant.PayloadTypes) {
return fmt.Errorf("cannot map %s: payload count %d is incompatible with %d", variantPath, len(sourceVariant.PayloadTypes), len(targetVariant.PayloadTypes))
}
for i := range sourceVariant.PayloadTypes {
if err := g.validateMapping(sourceVariant.PayloadTypes[i], targetVariant.PayloadTypes[i], fmt.Sprintf("%s[%d]", variantPath, i), seen); err != nil {
return err
}
}
}
return nil
}
return mappingError(path, source, target)
}
func mappingError(path, source, target string) error {
return fmt.Errorf("cannot map %s: %s is incompatible with %s", path, source, target)
}
func classFieldByName(class ClassDecl, name string) (FieldDecl, bool) {
for _, field := range class.Fields {
if field.Name == name {
return field, true
}
}
return FieldDecl{}, false
}
func mappingFieldName(class ClassDecl, field FieldDecl) string {
if class.Data && !field.Private {
return exportedGoName(field.Name)
}
return field.Name
}
func (g *goGenerator) emitMapping(pair mappingPair) error {
g.line(fmt.Sprintf("func %s(source %s) %s {", pair.function, mapGoType(pair.source), mapGoType(pair.target)))
g.indentLevel++
if sourceEnum, ok := g.enums[strings.TrimPrefix(pair.source, "*")]; ok {
if enumIsString(sourceEnum) && pair.target == "String" {
g.line("return string(source)")
g.indentLevel--
g.line("}")
return nil
}
targetEnum := g.enums[strings.TrimPrefix(pair.target, "*")]
if enumIsString(sourceEnum) && enumIsString(targetEnum) {
g.line("return " + targetEnum.Name + "(source)")
g.indentLevel--
g.line("}")
return nil
}
g.line("switch value := source.(type) {")
g.indentLevel++
for _, sourceVariant := range sourceEnum.Variants {
targetVariant := enumVariant(targetEnum, sourceVariant.Name)
g.line("case *" + sourceEnum.Name + sourceVariant.Name + ":")
g.indentLevel++
fields := make([]string, 0, len(sourceVariant.PayloadTypes))
for i := range sourceVariant.PayloadTypes {
expr, err := g.mappingExpr(fmt.Sprintf("value.Value%d", i), sourceVariant.PayloadTypes[i], targetVariant.PayloadTypes[i], sourceEnum.Name+"::"+sourceVariant.Name)
if err != nil {
return err
}
fields = append(fields, fmt.Sprintf("Value%d: %s", i, expr))
}
g.line("return &" + targetEnum.Name + targetVariant.Name + "{" + strings.Join(fields, ", ") + "}")
g.indentLevel--
}
g.indentLevel--
g.line("}")
g.line(`panic("unreachable enum mapping")`)
} else if targetEnum, ok := g.enums[strings.TrimPrefix(pair.target, "*")]; ok && pair.source == "String" && enumIsString(targetEnum) {
g.line("switch source {")
g.indentLevel++
for _, variant := range targetEnum.Variants {
g.line("case " + strconv.Quote(enumStringValue(variant)) + ":")
g.indentLevel++
g.line("return " + targetEnum.Name + variant.Name)
g.indentLevel--
}
g.indentLevel--
g.line("}")
g.line(`panic("unknown enum string: " + source)`)
} else {
expr, err := g.mappingExpr("source", pair.source, pair.target, strings.TrimPrefix(pair.source, "*"))
if err != nil {
return err
}
g.line("return " + expr)
}
g.indentLevel--
g.line("}")
return nil
}
func (g *goGenerator) mappingExpr(expr, source, target, path string) (string, error) {
if source == target {
return expr, nil
}
if strings.HasSuffix(source, "?") || strings.HasSuffix(target, "?") {
sourceInner := strings.TrimSuffix(source, "?")
targetInner := strings.TrimSuffix(target, "?")
inner, err := g.mappingExpr("*value", sourceInner, targetInner, path)
if err != nil {
return "", err
}
if strings.HasSuffix(source, "?") {
return fmt.Sprintf("func(value %s) %s { if value == nil { return nil }; mapped := %s; return &mapped }(%s)", mapGoType(source), mapGoType(target), inner, expr), nil
}
inner, err = g.mappingExpr("value", sourceInner, targetInner, path)
if err != nil {
return "", err
}
return fmt.Sprintf("func(value %s) %s { mapped := %s; return &mapped }(%s)", mapGoType(source), mapGoType(target), inner, expr), nil
}
if sourceBase, sourceArgs, ok := parseGenericType(source); ok {
_, targetArgs, _ := parseGenericType(target)
if sourceBase == "List" || sourceBase == "MutableList" {
item, err := g.mappingExpr("item", sourceArgs[0], targetArgs[0], path+"[]")
if err != nil {
return "", err
}
return fmt.Sprintf("func(values %s) %s { var result %s; for _, item := range values { result = append(result, %s) }; return result }(%s)", mapGoType(source), mapGoType(target), mapGoType(target), item, expr), nil
}
if sourceBase == "Map" || sourceBase == "MutableMap" {
value, err := g.mappingExpr("item", sourceArgs[1], targetArgs[1], path+"[]")
if err != nil {
return "", err
}
return fmt.Sprintf("func(values %s) %s { result := make(%s, len(values)); for key, item := range values { result[key] = %s }; return result }(%s)", mapGoType(source), mapGoType(target), mapGoType(target), value, expr), nil
}
}
if sourceClass, ok := g.classForType(source); ok {
targetClass, _ := g.classForType(target)
fields := make([]string, 0, len(targetClass.Fields))
for _, targetField := range targetClass.Fields {
sourceField, _ := classFieldByName(sourceClass, targetField.Name)
mapped, err := g.mappingExpr(expr+"."+mappingFieldName(sourceClass, sourceField), sourceField.Type, targetField.Type, path+"."+targetField.Name)
if err != nil {
return "", err
}
fields = append(fields, mappingFieldName(targetClass, targetField)+": "+mapped)
}
literal := targetClass.Name + "{" + strings.Join(fields, ", ") + "}"
if strings.HasPrefix(target, "*") {
literal = "&" + literal
}
if strings.HasPrefix(source, "*") {
return fmt.Sprintf("func(value %s) %s { if value == nil { return nil }; return %s }(%s)", mapGoType(source), mapGoType(target), strings.ReplaceAll(literal, expr+".", "value."), expr), nil
}
return literal, nil
}
if _, ok := g.enums[strings.TrimPrefix(source, "*")]; ok {
function, err := g.ensureMapping(source, target, path)
if err != nil {
return "", err
}
return function + "(" + expr + ")", nil
}
if targetEnum, ok := g.enums[strings.TrimPrefix(target, "*")]; ok && source == "String" && enumIsString(targetEnum) {
function, err := g.ensureMapping(source, target, path)
if err != nil {
return "", err
}
return function + "(" + expr + ")", nil
}
return "", mappingError(path, source, target)
}

View file

@ -0,0 +1,159 @@
package lang
import (
"strings"
"testing"
)
func TestGenerateRecursiveClassMapping(t *testing.T) {
prog, err := Parse(`
package demo
data class SourceAddress(var city: String)
data class TargetAddress(var city: String)
data class Source(var id: String, var address: *SourceAddress)
data class Target(var address: *TargetAddress, var id: String)
fun convert(source: *Source): *Target { return source.mapTo<Target>() }
`)
if err != nil {
t.Fatal(err)
}
out, err := GenerateGo(prog)
if err != nil {
t.Fatal(err)
}
code := string(out)
for _, want := range []string{"gotlinMap1(source)", "Address:", "Id: value.Id", "TargetAddress{City: value.City}"} {
if !strings.Contains(code, want) {
t.Fatalf("missing %q:\n%s", want, code)
}
}
}
func TestGenerateListAndMapMapping(t *testing.T) {
prog, err := Parse(`
package demo
data class Source(var id: String)
data class Target(var id: String)
fun list(values: List<*Source>): List<*Target> { return values.mapTo<List<*Target>>() }
fun mapping(values: Map<String, *Source>): Map<String, *Target> { return values.mapTo<Map<String, *Target>>() }
`)
if err != nil {
t.Fatal(err)
}
out, err := GenerateGo(prog)
if err != nil {
t.Fatal(err)
}
for _, want := range []string{"for _, item := range values", "for key, item := range values", "result[key]"} {
if !strings.Contains(string(out), want) {
t.Fatalf("missing %q:\n%s", want, out)
}
}
}
func TestGenerateRecursiveEnumMapping(t *testing.T) {
prog, err := Parse(`
package demo
enum Source { Ready(String), Failed(String) }
enum Target { Ready(String), Failed(String), Pending }
fun convert(value: Source): Target { return value.mapTo<Target>() }
`)
if err != nil {
t.Fatal(err)
}
out, err := GenerateGo(prog)
if err != nil {
t.Fatal(err)
}
for _, want := range []string{"case *SourceReady:", "return &TargetReady{Value0: value.Value0}", "case *SourceFailed:"} {
if !strings.Contains(string(out), want) {
t.Fatalf("missing %q:\n%s", want, out)
}
}
}
func TestMappingReportsNestedFieldPath(t *testing.T) {
prog, err := Parse(`
package demo
data class SourceAddress(var zip: String)
data class TargetAddress(var zip: Int)
data class Source(var address: *SourceAddress)
data class Target(var address: *TargetAddress)
fun convert(value: *Source): *Target { return value.mapTo<Target>() }
`)
if err != nil {
t.Fatal(err)
}
_, err = GenerateGo(prog)
if err == nil || !strings.Contains(err.Error(), "Source.address.zip") || !strings.Contains(err.Error(), "String is incompatible with Int") {
t.Fatalf("unexpected mapping error: %v", err)
}
}
func TestMappingRejectsMissingFieldAndEnumVariant(t *testing.T) {
for _, source := range []string{
`package demo data class Source(var id: String) data class Target(var id: String, var name: String) fun convert(value: *Source): *Target { return value.mapTo<Target>() }`,
`package demo enum Source { Ready, Failed } enum Target { Ready } fun convert(value: Source): Target { return value.mapTo<Target>() }`,
} {
prog, err := Parse(source)
if err != nil {
t.Fatal(err)
}
if _, err = GenerateGo(prog); err == nil {
t.Fatalf("expected mapping error for %s", source)
}
}
}
func TestMapStringBackedEnumToAndFromString(t *testing.T) {
prog, err := Parse(`package demo
enum Status { PendingReservation, Initiated }
data class Domain(var status: Status)
data class Row(var status: String)
fun toRow(value: *Domain): *Row { return value.mapTo<Row>() }
fun toDomain(value: *Row): *Domain { return value.mapTo<Domain>() }`)
if err != nil {
t.Fatal(err)
}
out, err := GenerateGo(prog)
if err != nil {
t.Fatal(err)
}
for _, want := range []string{`return string(source)`, `case "PendingReservation":`, `return StatusPendingReservation`} {
if !strings.Contains(string(out), want) {
t.Fatalf("missing %q:\n%s", want, out)
}
}
}
func TestMapToInfersExpectedTargetType(t *testing.T) {
prog, err := Parse(`package demo
data class Source(var id: String)
data class Target(var id: String)
data class Wrapper(var target: *Target)
fun returned(value: *Source): *Target { return value.mapTo() }
fun wrapped(value: *Source): *Wrapper { return Wrapper(value.mapTo()) }
fun local(value: *Source): *Target { val target: *Target = value.mapTo(); return target }`)
if err != nil {
t.Fatal(err)
}
out, err := GenerateGo(prog)
if err != nil {
t.Fatal(err)
}
if strings.Count(string(out), "gotlinMap1(value)") < 3 {
t.Fatalf("expected inferred mapping calls:\n%s", out)
}
}
func TestMapToWithoutTargetContextHasHelpfulError(t *testing.T) {
prog, err := Parse(`package demo
fun convert(value: *Source) { val target = value.mapTo() }`)
if err != nil {
t.Fatal(err)
}
_, err = GenerateGo(prog)
if err == nil || !strings.Contains(err.Error(), "target type cannot be inferred") {
t.Fatalf("unexpected error: %v", err)
}
}

View file

@ -0,0 +1,53 @@
package lang
import (
"strings"
"testing"
)
func TestClassFieldTypeResolvesInternalMethodWithoutAlias(t *testing.T) {
prog, err := Parse(`
package demo
class Repository {
fun healthy(): Boolean { return true }
}
class Service(val repository: *Repository) {
fun healthy(): Boolean { return repository.healthy() }
}
`)
if err != nil {
t.Fatal(err)
}
out, err := GenerateGo(prog)
if err != nil {
t.Fatal(err)
}
if !strings.Contains(string(out), "self.repository.healthy()") {
t.Fatalf("internal method was treated as foreign:\n%s", out)
}
}
func TestInternalMethodReturnTypeFlowsThroughSelectors(t *testing.T) {
prog, err := Parse(`
package demo
class Transaction { fun commit(): Boolean { return true } }
class Result(val transaction: *Transaction)
class Repository { fun begin(): *Result { return Result(Transaction()) } }
class Service(val repository: *Repository) {
fun run(): Boolean { return repository.begin().transaction.commit() }
}
`)
if err != nil {
t.Fatal(err)
}
out, err := GenerateGo(prog)
if err != nil {
t.Fatal(err)
}
if !strings.Contains(string(out), "self.repository.begin().transaction.commit()") {
t.Fatalf("method return type did not flow through selector:\n%s", out)
}
}

View file

@ -2,6 +2,7 @@ package lang
import (
"fmt"
"strconv"
"strings"
)
@ -44,12 +45,92 @@ func (p *parser) parseProgram() (*Program, error) {
return nil, err
}
prog.Interfaces = append(prog.Interfaces, decl)
case p.check(tokenClass):
case p.check(tokenEnum):
decl, err := p.parseEnum()
if err != nil {
return nil, err
}
prog.Enums = append(prog.Enums, decl)
case p.check(tokenClass) || p.check(tokenData):
decl, err := p.parseClass()
if err != nil {
return nil, err
}
prog.Classes = append(prog.Classes, decl)
case p.match(tokenAt):
annotation, err := p.expect(tokenIdent, "expected annotation name")
if err != nil {
return nil, err
}
if annotation.lexeme == "embed" {
if _, err := p.expect(tokenLParen, "expected '(' after embed"); err != nil {
return nil, err
}
path, err := p.expect(tokenString, "expected embed path")
if err != nil {
return nil, err
}
if _, err := p.expect(tokenRParen, "expected ')' after embed path"); err != nil {
return nil, err
}
if _, err := p.expect(tokenVal, "expected 'val' after embed annotation"); err != nil {
return nil, err
}
name, err := p.expect(tokenIdent, "expected embedded value name")
if err != nil {
return nil, err
}
if _, err := p.expect(tokenColon, "expected ':' after embedded value name"); err != nil {
return nil, err
}
typ, err := p.parseTypeRef()
if err != nil {
return nil, err
}
prog.Embeds = append(prog.Embeds, EmbedDecl{Path: strings.Trim(path.lexeme, "\""), Name: name.lexeme, Type: typ})
continue
}
switch annotation.lexeme {
case "jsonNaming":
if _, err := p.expect(tokenLParen, "expected '(' after jsonNaming"); err != nil {
return nil, err
}
policy, err := p.expect(tokenIdent, "expected JSON naming policy")
if err != nil {
return nil, err
}
if _, err := p.expect(tokenRParen, "expected ')' after JSON naming policy"); err != nil {
return nil, err
}
if !p.check(tokenData) {
return nil, fmt.Errorf("jsonNaming is only valid on data classes")
}
decl, err := p.parseClass()
if err != nil {
return nil, err
}
decl.JSONNaming = policy.lexeme
prog.Classes = append(prog.Classes, decl)
case "table":
table, err := p.parseStringAnnotationArgument("table")
if err != nil {
return nil, err
}
if !validSQLName(table) {
return nil, fmt.Errorf("invalid SQL table name %q", table)
}
if !p.check(tokenData) {
return nil, fmt.Errorf("table is only valid on data classes")
}
decl, err := p.parseClass()
if err != nil {
return nil, err
}
decl.Table = table
prog.Classes = append(prog.Classes, decl)
default:
return nil, fmt.Errorf("unsupported annotation %q", annotation.lexeme)
}
case p.check(tokenWorker):
decl, err := p.parseWorker()
if err != nil {
@ -70,6 +151,64 @@ func (p *parser) parseProgram() (*Program, error) {
return prog, nil
}
func (p *parser) parseEnum() (EnumDecl, error) {
if _, err := p.expect(tokenEnum, "expected 'enum'"); err != nil {
return EnumDecl{}, err
}
name, err := p.expect(tokenIdent, "expected enum name")
if err != nil {
return EnumDecl{}, err
}
if _, err := p.expect(tokenLBrace, "expected '{' after enum name"); err != nil {
return EnumDecl{}, err
}
var variants []EnumVariant
seen := map[string]bool{}
for !p.check(tokenRBrace) && !p.check(tokenEOF) {
variant, err := p.expect(tokenIdent, "expected enum variant")
if err != nil {
return EnumDecl{}, err
}
if seen[variant.lexeme] {
return EnumDecl{}, fmt.Errorf("duplicate enum variant %s", variant.lexeme)
}
seen[variant.lexeme] = true
var payload []string
if p.match(tokenLParen) {
if !p.check(tokenRParen) {
for {
typ, err := p.parseTypeRef()
if err != nil {
return EnumDecl{}, err
}
payload = append(payload, typ)
if !p.match(tokenComma) {
break
}
}
}
if _, err := p.expect(tokenRParen, "expected ')' after enum payload"); err != nil {
return EnumDecl{}, err
}
}
stringValue := ""
if p.match(tokenAssign) {
value, err := p.expect(tokenString, "expected string enum value")
if err != nil {
return EnumDecl{}, err
}
stringValue = strings.Trim(value.lexeme, "\"")
}
variants = append(variants, EnumVariant{Name: variant.lexeme, PayloadTypes: payload, StringValue: stringValue})
p.match(tokenComma)
p.match(tokenSemicolon)
}
if _, err := p.expect(tokenRBrace, "expected '}' after enum"); err != nil {
return EnumDecl{}, err
}
return EnumDecl{Name: name.lexeme, Variants: variants}, nil
}
func (p *parser) parsePackageDecl() (string, error) {
path, err := p.parseImportPath()
if err != nil {
@ -151,6 +290,7 @@ func (p *parser) parseInterface() (InterfaceDecl, error) {
}
func (p *parser) parseClass() (ClassDecl, error) {
data := p.match(tokenData)
if _, err := p.expect(tokenClass, "expected 'class'"); err != nil {
return ClassDecl{}, err
}
@ -174,7 +314,7 @@ func (p *parser) parseClass() (ClassDecl, error) {
return ClassDecl{}, err
}
if !p.match(tokenLBrace) {
return ClassDecl{Name: name.lexeme, Fields: fields, Parents: parents, Methods: nil}, nil
return ClassDecl{Name: name.lexeme, Data: data, Fields: fields, Parents: parents, Methods: nil}, nil
}
var methods []FunctionDecl
@ -195,7 +335,7 @@ func (p *parser) parseClass() (ClassDecl, error) {
return ClassDecl{}, err
}
return ClassDecl{Name: name.lexeme, Fields: fields, Parents: parents, Methods: methods}, nil
return ClassDecl{Name: name.lexeme, Data: data, Fields: fields, Parents: parents, Methods: methods}, nil
}
func (p *parser) parseWorker() (WorkerDecl, error) {
@ -284,6 +424,41 @@ func (p *parser) parseClassFields() ([]FieldDecl, error) {
return fields, nil
}
for {
column := ""
id := false
generated := false
for p.match(tokenAt) {
annotation, err := p.expect(tokenIdent, "expected field annotation name")
if err != nil {
return nil, err
}
switch annotation.lexeme {
case "id":
if id {
return nil, fmt.Errorf("duplicate id annotation")
}
id = true
case "generated":
if generated {
return nil, fmt.Errorf("duplicate generated annotation")
}
generated = true
case "column":
if column != "" {
return nil, fmt.Errorf("duplicate column annotation")
}
column, err = p.parseStringAnnotationArgument("column")
if err != nil {
return nil, err
}
if !validSQLName(column) {
return nil, fmt.Errorf("invalid SQL column name %q", column)
}
default:
return nil, fmt.Errorf("unsupported field annotation %q", annotation.lexeme)
}
}
private := p.match(tokenPrivate)
mutable := false
switch {
case p.match(tokenVal):
@ -305,13 +480,49 @@ func (p *parser) parseClassFields() ([]FieldDecl, error) {
if err != nil {
return nil, err
}
fields = append(fields, FieldDecl{Mutable: mutable, Name: name.lexeme, Type: typ})
fields = append(fields, FieldDecl{Mutable: mutable, Private: private, Name: name.lexeme, Type: typ, Column: column, ID: id, Generated: generated})
if !p.match(tokenComma) {
return fields, nil
}
}
}
func (p *parser) parseStringAnnotationArgument(name string) (string, error) {
if _, err := p.expect(tokenLParen, "expected '(' after "+name); err != nil {
return "", err
}
value, err := p.expect(tokenString, "expected string argument for "+name)
if err != nil {
return "", err
}
if _, err := p.expect(tokenRParen, "expected ')' after "+name+" argument"); err != nil {
return "", err
}
decoded, err := strconv.Unquote(value.lexeme)
if err != nil {
return "", fmt.Errorf("invalid string argument for %s: %w", name, err)
}
return decoded, nil
}
func validSQLName(name string) bool {
if name == "" {
return false
}
for i, r := range name {
if i == 0 {
if !isIdentStart(r) {
return false
}
continue
}
if !isIdentPart(r) {
return false
}
}
return true
}
func (p *parser) parseFunction() (FunctionDecl, error) {
signature, err := p.parseFunctionSignature()
if err != nil {
@ -448,12 +659,25 @@ func (p *parser) parseStmt() (Stmt, error) {
return nil, fmt.Errorf("'go' expects a function call expression")
}
return GoStmt{Value: expr}, nil
case p.match(tokenDefer):
expr, err := p.parseExpr(0)
if err != nil {
return nil, err
}
if _, ok := expr.(CallExpr); !ok {
return nil, fmt.Errorf("'defer' expects a function call expression")
}
return DeferStmt{Value: expr}, nil
case p.match(tokenIf):
return p.parseIf()
case p.match(tokenWhile):
return p.parseWhile()
case p.match(tokenFor):
return p.parseForEach()
case p.match(tokenSelect):
return p.parseSelect()
case p.match(tokenMatch):
return p.parseMatch()
case p.match(tokenTry):
return p.parseTryCatch()
case p.check(tokenIdent) && (p.peekN(1).kind == tokenAssign || p.peekN(1).kind == tokenComma || p.peekN(1).kind == tokenPlusAssign):
@ -466,7 +690,7 @@ func (p *parser) parseStmt() (Stmt, error) {
if err != nil {
return nil, err
}
return AddAssignStmt{Name: name.lexeme, Value: value}, nil
return AddAssignStmt{Name: name.lexeme, Pos: name.pos, Value: value}, nil
}
names, err := p.parseNameList()
if err != nil {
@ -480,9 +704,18 @@ func (p *parser) parseStmt() (Stmt, error) {
return nil, err
}
if len(names) == 1 {
return AssignStmt{Name: names[0], Value: value}, nil
return AssignStmt{Name: names[0].lexeme, Pos: names[0].pos, Value: value}, nil
}
return MultiAssignStmt{Names: names, Value: value}, nil
assign := MultiAssignStmt{
Names: make([]string, 0, len(names)),
Positions: make([]int, 0, len(names)),
Value: value,
}
for _, name := range names {
assign.Names = append(assign.Names, name.lexeme)
assign.Positions = append(assign.Positions, name.pos)
}
return assign, nil
default:
expr, err := p.parseExpr(0)
if err != nil {
@ -492,6 +725,68 @@ func (p *parser) parseStmt() (Stmt, error) {
}
}
func (p *parser) parseMatch() (Stmt, error) {
if _, err := p.expect(tokenLParen, "expected '(' after match"); err != nil {
return nil, err
}
value, err := p.parseExpr(0)
if err != nil {
return nil, err
}
if _, err := p.expect(tokenRParen, "expected ')' after match value"); err != nil {
return nil, err
}
if _, err := p.expect(tokenLBrace, "expected '{' after match value"); err != nil {
return nil, err
}
var cases []MatchCase
for !p.check(tokenRBrace) && !p.check(tokenEOF) {
enumName, err := p.expect(tokenIdent, "expected enum name in match case")
if err != nil {
return nil, err
}
if _, err := p.expect(tokenDoubleColon, "expected '::' in match case"); err != nil {
return nil, err
}
variant, err := p.expect(tokenIdent, "expected variant name")
if err != nil {
return nil, err
}
var bindings []string
if p.match(tokenLParen) {
if !p.check(tokenRParen) {
for {
binding, err := p.expect(tokenIdent, "expected variant binding")
if err != nil {
return nil, err
}
bindings = append(bindings, binding.lexeme)
if !p.match(tokenComma) {
break
}
}
}
if _, err := p.expect(tokenRParen, "expected ')' after variant bindings"); err != nil {
return nil, err
}
}
if _, err := p.expect(tokenArrow, "expected '->' after match pattern"); err != nil {
return nil, err
}
body, err := p.parseBlock()
if err != nil {
return nil, err
}
cases = append(cases, MatchCase{EnumName: enumName.lexeme, VariantName: variant.lexeme, Bindings: bindings, Body: body})
p.match(tokenComma)
p.match(tokenSemicolon)
}
if _, err := p.expect(tokenRBrace, "expected '}' after match"); err != nil {
return nil, err
}
return MatchStmt{Value: value, Cases: cases}, nil
}
func (p *parser) parseSelect() (Stmt, error) {
if _, err := p.expect(tokenLBrace, "expected '{' after select"); err != nil {
return nil, err
@ -570,7 +865,7 @@ func (p *parser) parseVarDecl(mutable bool) (Stmt, error) {
if err != nil {
return nil, err
}
name := names[0]
name := names[0].lexeme
var typ string
if p.match(tokenColon) {
if len(names) > 1 {
@ -592,7 +887,11 @@ func (p *parser) parseVarDecl(mutable bool) (Stmt, error) {
if len(names) == 1 {
return VarDecl{Mutable: mutable, Name: name, Type: typ, Value: value}, nil
}
return MultiVarDecl{Mutable: mutable, Names: names, Value: value}, nil
decl := MultiVarDecl{Mutable: mutable, Names: make([]string, 0, len(names)), Value: value}
for _, name := range names {
decl.Names = append(decl.Names, name.lexeme)
}
return decl, nil
}
func (p *parser) parseIf() (Stmt, error) {
@ -662,9 +961,36 @@ func (p *parser) parsePrefix() (Expr, error) {
tok := p.advance()
switch tok.kind {
case tokenIdent:
if p.match(tokenDoubleColon) {
variant, err := p.expect(tokenIdent, "expected enum variant")
if err != nil {
return nil, err
}
var values []Expr
if p.match(tokenLParen) {
if !p.check(tokenRParen) {
for {
value, err := p.parseExpr(0)
if err != nil {
return nil, err
}
values = append(values, value)
if !p.match(tokenComma) {
break
}
}
}
if _, err := p.expect(tokenRParen, "expected ')' after variant values"); err != nil {
return nil, err
}
}
return p.parsePostfix(EnumVariantExpr{EnumName: tok.lexeme, VariantName: variant.lexeme, Values: values})
}
return p.parsePostfix(IdentExpr{Name: tok.lexeme})
case tokenInt:
return IntExpr{Value: tok.lexeme}, nil
case tokenFloat:
return FloatExpr{Value: tok.lexeme}, nil
case tokenString:
return StringExpr{Value: tok.lexeme}, nil
case tokenTrue:
@ -684,7 +1010,7 @@ func (p *parser) parsePrefix() (Expr, error) {
return p.parsePostfix(expr)
case tokenLBrace:
return p.parseLambdaExpr()
case tokenBang, tokenMinus:
case tokenBang, tokenMinus, tokenAmp, tokenStar:
value, err := p.parseExpr(7)
if err != nil {
return nil, err
@ -695,18 +1021,18 @@ func (p *parser) parsePrefix() (Expr, error) {
}
}
func (p *parser) parseNameList() ([]string, error) {
func (p *parser) parseNameList() ([]token, error) {
name, err := p.expect(tokenIdent, "expected variable name")
if err != nil {
return nil, err
}
names := []string{name.lexeme}
names := []token{name}
for p.match(tokenComma) {
next, err := p.expect(tokenIdent, "expected variable name")
if err != nil {
return nil, err
}
names = append(names, next.lexeme)
names = append(names, next)
}
return names, nil
}
@ -774,11 +1100,20 @@ func (p *parser) parsePostfix(expr Expr) (Expr, error) {
for {
switch {
case p.match(tokenDot):
name, err := p.expect(tokenIdent, "expected selector name")
name := p.advance()
if !selectorName(name.lexeme) {
return nil, fmt.Errorf("expected selector name at %d, found %q", name.pos, name.lexeme)
}
expr = SelectorExpr{Receiver: expr, Name: name.lexeme}
case p.match(tokenLBracket):
index, err := p.parseExpr(0)
if err != nil {
return nil, err
}
expr = SelectorExpr{Receiver: expr, Name: name.lexeme}
if _, err := p.expect(tokenRBracket, "expected ']' after index"); err != nil {
return nil, err
}
expr = IndexExpr{Receiver: expr, Index: index}
case p.check(tokenLt):
typeArgs, hasTypeArgs, err := p.tryParseCallTypeArgs()
if err != nil {
@ -790,44 +1125,31 @@ func (p *parser) parsePostfix(expr Expr) (Expr, error) {
if _, err := p.expect(tokenLParen, "expected '(' after generic type arguments"); err != nil {
return nil, err
}
var args []Expr
if !p.check(tokenRParen) {
for {
arg, err := p.parseExpr(0)
if err != nil {
return nil, err
}
args = append(args, arg)
if !p.match(tokenComma) {
break
}
}
args, namedArgs, err := p.parseCallArguments()
if err != nil {
return nil, err
}
if _, err := p.expect(tokenRParen, "expected ')' after arguments"); err != nil {
return nil, err
}
expr = CallExpr{Callee: expr, Args: args, TypeArgs: typeArgs}
expr = CallExpr{Callee: expr, Args: args, NamedArgs: namedArgs, TypeArgs: typeArgs}
case p.match(tokenLParen):
var args []Expr
if !p.check(tokenRParen) {
for {
arg, err := p.parseExpr(0)
if err != nil {
return nil, err
}
args = append(args, arg)
if !p.match(tokenComma) {
break
}
}
args, namedArgs, err := p.parseCallArguments()
if err != nil {
return nil, err
}
if _, err := p.expect(tokenRParen, "expected ')' after arguments"); err != nil {
return nil, err
}
expr = CallExpr{Callee: expr, Args: args}
expr = CallExpr{Callee: expr, Args: args, NamedArgs: namedArgs}
case p.check(tokenLBrace):
call, ok := expr.(CallExpr)
if !ok {
var call CallExpr
switch current := expr.(type) {
case CallExpr:
call = current
case SelectorExpr:
call = CallExpr{Callee: current}
default:
return expr, nil
}
lambda, err := p.parsePrefix()
@ -842,6 +1164,54 @@ func (p *parser) parsePostfix(expr Expr) (Expr, error) {
}
}
func selectorName(value string) bool {
runes := []rune(value)
if len(runes) == 0 || !isIdentStart(runes[0]) {
return false
}
for _, r := range runes[1:] {
if !isIdentPart(r) {
return false
}
}
return true
}
func (p *parser) parseCallArguments() ([]Expr, []NamedArg, error) {
var args []Expr
var named []NamedArg
if p.check(tokenRParen) {
return args, named, nil
}
for {
if p.check(tokenIdent) && p.peekN(1).kind == tokenAssign {
if len(args) > 0 {
return nil, nil, fmt.Errorf("cannot mix positional and named arguments")
}
name := p.advance().lexeme
p.advance()
value, err := p.parseExpr(0)
if err != nil {
return nil, nil, err
}
named = append(named, NamedArg{Name: name, Value: value})
} else {
if len(named) > 0 {
return nil, nil, fmt.Errorf("cannot mix named and positional arguments")
}
value, err := p.parseExpr(0)
if err != nil {
return nil, nil, err
}
args = append(args, value)
}
if !p.match(tokenComma) {
break
}
}
return args, named, nil
}
func (p *parser) tryParseCallTypeArgs() ([]string, bool, error) {
if !p.check(tokenLt) {
return nil, false, nil
@ -937,6 +1307,9 @@ func (p *parser) parseTypeRef() (string, error) {
}
b += "<" + strings.Join(args, ", ") + ">"
}
if p.match(tokenQuestion) {
b += "?"
}
return b, nil
}

View file

@ -0,0 +1,28 @@
package lang
import (
"strings"
"testing"
)
func TestGenerateGoAddressOfExpression(t *testing.T) {
prog, err := Parse(`
package demo
import encoding.json
fun decode(body: ByteSlice, target: Account) {
json.Unmarshal(body, &target)
}
`)
if err != nil {
t.Fatalf("parse failed: %v", err)
}
out, err := GenerateGo(prog)
if err != nil {
t.Fatalf("generation failed: %v", err)
}
if !strings.Contains(string(out), "json.Unmarshal(body, &target)") {
t.Fatalf("generated Go missing address-of expression:\n%s", out)
}
}

282
internal/lang/semantics.go Normal file
View file

@ -0,0 +1,282 @@
package lang
import "fmt"
func validateMutability(program *Program) error {
checker := mutabilityChecker{}
for _, fn := range program.Functions {
if err := checker.checkFunction(fn, nil, nil); err != nil {
return err
}
}
for _, class := range program.Classes {
fields := make(map[string]bool, len(class.Fields))
for _, field := range class.Fields {
fields[field.Name] = field.Mutable
}
for _, method := range class.Methods {
if err := checker.checkFunction(method, fields, nil); err != nil {
return err
}
}
}
for _, worker := range program.Workers {
fields := make(map[string]bool, len(worker.Fields))
for _, field := range worker.Fields {
fields[field.Name] = field.Mutable
}
for _, method := range worker.Methods {
if err := checker.checkFunction(method, nil, fields); err != nil {
return err
}
}
}
return nil
}
type mutabilityChecker struct {
scopes []map[string]bool
classFields map[string]bool
workerFields map[string]bool
}
func (c *mutabilityChecker) checkFunction(fn FunctionDecl, classFields map[string]bool, workerFields map[string]bool) error {
c.scopes = nil
c.classFields = classFields
c.workerFields = workerFields
c.pushScope()
defer c.popScope()
for _, param := range fn.Params {
c.define(param.Name, false)
}
return c.checkStmts(fn.Body)
}
func (c *mutabilityChecker) checkStmts(stmts []Stmt) error {
for _, stmt := range stmts {
switch s := stmt.(type) {
case VarDecl:
if err := c.checkExpr(s.Value); err != nil {
return err
}
c.define(s.Name, s.Mutable)
case MultiVarDecl:
if err := c.checkExpr(s.Value); err != nil {
return err
}
for _, name := range s.Names {
c.define(name, s.Mutable)
}
case AssignStmt:
if err := c.requireMutable(s.Name, s.Pos); err != nil {
return err
}
if err := c.checkExpr(s.Value); err != nil {
return err
}
case AddAssignStmt:
if err := c.requireMutable(s.Name, s.Pos); err != nil {
return err
}
if err := c.checkExpr(s.Value); err != nil {
return err
}
case MultiAssignStmt:
for i, name := range s.Names {
pos := 0
if i < len(s.Positions) {
pos = s.Positions[i]
}
if err := c.requireMutable(name, pos); err != nil {
return err
}
}
if err := c.checkExpr(s.Value); err != nil {
return err
}
case ReturnStmt:
if s.Value != nil {
if err := c.checkExpr(s.Value); err != nil {
return err
}
}
case ThrowStmt:
if err := c.checkExpr(s.Value); err != nil {
return err
}
case GoStmt:
if err := c.checkExpr(s.Value); err != nil {
return err
}
case DeferStmt:
if err := c.checkExpr(s.Value); err != nil {
return err
}
case ExprStmt:
if err := c.checkExpr(s.Value); err != nil {
return err
}
case IfStmt:
if err := c.checkExpr(s.Cond); err != nil {
return err
}
if err := c.checkBlock(s.Then, nil); err != nil {
return err
}
if err := c.checkBlock(s.Else, nil); err != nil {
return err
}
case WhileStmt:
if err := c.checkExpr(s.Cond); err != nil {
return err
}
if err := c.checkBlock(s.Body, nil); err != nil {
return err
}
case ForEachStmt:
if err := c.checkExpr(s.Source); err != nil {
return err
}
if err := c.checkBlock(s.Body, map[string]bool{s.Name: false}); err != nil {
return err
}
case SelectStmt:
for _, sc := range s.Cases {
if err := c.checkExpr(sc.Source); err != nil {
return err
}
if err := c.checkBlock(sc.Body, map[string]bool{"it": false}); err != nil {
return err
}
}
case MatchStmt:
if err := c.checkExpr(s.Value); err != nil {
return err
}
for _, matchCase := range s.Cases {
bindings := map[string]bool{}
for _, binding := range matchCase.Bindings {
bindings[binding] = false
}
if err := c.checkBlock(matchCase.Body, bindings); err != nil {
return err
}
}
case TryCatchStmt:
if err := c.checkBlock(s.TryBody, nil); err != nil {
return err
}
if err := c.checkBlock(s.CatchBody, map[string]bool{s.CatchName: false}); err != nil {
return err
}
}
}
return nil
}
func (c *mutabilityChecker) checkBlock(stmts []Stmt, bindings map[string]bool) error {
c.pushScope()
defer c.popScope()
for name, mutable := range bindings {
c.define(name, mutable)
}
return c.checkStmts(stmts)
}
func (c *mutabilityChecker) checkExpr(expr Expr) error {
switch e := expr.(type) {
case UnaryExpr:
return c.checkExpr(e.Value)
case BinaryExpr:
if err := c.checkExpr(e.Left); err != nil {
return err
}
return c.checkExpr(e.Right)
case CallExpr:
if err := c.checkExpr(e.Callee); err != nil {
return err
}
for _, arg := range e.Args {
if err := c.checkExpr(arg); err != nil {
return err
}
}
for _, arg := range e.NamedArgs {
if err := c.checkExpr(arg.Value); err != nil {
return err
}
}
case SelectorExpr:
return c.checkExpr(e.Receiver)
case IndexExpr:
if err := c.checkExpr(e.Receiver); err != nil {
return err
}
return c.checkExpr(e.Index)
case EnumVariantExpr:
for _, value := range e.Values {
if err := c.checkExpr(value); err != nil {
return err
}
}
case LambdaExpr:
bindings := map[string]bool{}
if e.ImplicitIt {
bindings["it"] = false
}
for _, param := range e.Params {
bindings[param.Name] = false
}
return c.checkBlock(e.Body, bindings)
}
return nil
}
func (c *mutabilityChecker) pushScope() {
c.scopes = append(c.scopes, map[string]bool{})
}
func (c *mutabilityChecker) popScope() {
if len(c.scopes) == 0 {
return
}
c.scopes = c.scopes[:len(c.scopes)-1]
}
func (c *mutabilityChecker) define(name string, mutable bool) {
if len(c.scopes) == 0 {
c.pushScope()
}
c.scopes[len(c.scopes)-1][name] = mutable
}
func (c *mutabilityChecker) requireMutable(name string, pos int) error {
for i := len(c.scopes) - 1; i >= 0; i-- {
if mutable, ok := c.scopes[i][name]; ok {
if mutable {
return nil
}
return immutableAssignmentError(name, pos)
}
}
if mutable, ok := c.classFields[name]; ok {
if mutable {
return nil
}
return immutableAssignmentError(name, pos)
}
if mutable, ok := c.workerFields[name]; ok {
if mutable {
return nil
}
return immutableAssignmentError(name, pos)
}
return nil
}
func immutableAssignmentError(name string, pos int) error {
if pos > 0 {
return fmt.Errorf("cannot reassign immutable name %s at %d", name, pos)
}
return fmt.Errorf("cannot reassign immutable name %s", name)
}

1244
internal/lang/sql.go Normal file

File diff suppressed because it is too large Load diff

View file

@ -0,0 +1,277 @@
package lang
import (
"strings"
"testing"
)
const eventSQLSource = `
import time
@table("outbox_events")
data class EventRow(
@generated @id var id: String,
var payload: String,
var attempts: Int,
var publishedAt: time.Time?,
var claimedUntil: time.Time?,
var createdAt: time.Time
)
data class EventProjection(var id: String, var payload: String)
data class WrongProjection(var id: Int)
`
func TestSQLGeneratedNullableAndLockingMetadata(t *testing.T) {
program, err := Parse(eventSQLSource)
if err != nil {
t.Fatalf("parse failed: %v", err)
}
row := program.Classes[0]
if !row.Fields[0].Generated {
t.Fatal("id field is missing @generated metadata")
}
if row.Fields[3].Type != "time.Time?" {
t.Fatalf("publishedAt type = %q, want time.Time?", row.Fields[3].Type)
}
if got := mapGoType(row.Fields[3].Type); got != "*time.Time" {
t.Fatalf("mapped nullable timestamp = %q, want *time.Time", got)
}
code := compileSQL(t, eventSQLSource+`
fun claimQuery(nowLimit: Int): GotlinSQLQuery {
return sql.from<EventRow>()
.where { it.publishedAt == null && (it.claimedUntil == null || it.claimedUntil < now()) }
.orderByDescending { it.createdAt }
.limit(nowLimit)
.forUpdate()
.skipLocked()
.build()
}
`)
for _, want := range []string{
`SQL: "SELECT id, payload, attempts, published_at, claimed_until, created_at FROM outbox_events WHERE (published_at IS NULL AND (claimed_until IS NULL OR claimed_until < CURRENT_TIMESTAMP)) ORDER BY created_at DESC LIMIT $1 FOR UPDATE SKIP LOCKED"`,
`Args: []any{nowLimit}`,
} {
if !strings.Contains(code, want) {
t.Fatalf("generated Go missing %q:\n%s", want, code)
}
}
}
func TestSQLTypedProjectionExecution(t *testing.T) {
code := compileSQL(t, eventSQLSource+`
import context
import pgxpool "github.com/jackc/pgx/v5/pgxpool"
fun events(pool: *pgxpool.Pool, ctx: context.Context): List<*EventProjection> {
return sql.from<EventRow>()
.select { row -> EventProjection(row.id, row.payload) }
.orderBy { it.createdAt }
.fetch(pool, ctx)
}
fun eventStream(pool: *pgxpool.Pool, ctx: context.Context): GotlinSQLIterator<EventProjection> {
return sql.from<EventRow>()
.select { row -> EventProjection(row.id, row.payload) }
.iterator(pool, ctx)
}
`)
for _, want := range []string{
`return gotlinSQLFetch[EventProjection](pool, ctx, GotlinSQLQuery{SQL: "SELECT id, payload FROM outbox_events ORDER BY created_at", Args: []any{}}, gotlinSQLScanEventProjection)`,
`return gotlinSQLIterate[EventProjection](pool, ctx, GotlinSQLQuery{SQL: "SELECT id, payload FROM outbox_events", Args: []any{}}, gotlinSQLScanEventProjection)`,
`func gotlinSQLScanEventProjection(row gotlinSQLRow) (*EventProjection, error)`,
`err := row.Scan(&value.Id, &value.Payload)`,
} {
if !strings.Contains(code, want) {
t.Fatalf("generated Go missing %q:\n%s", want, code)
}
}
}
func TestSQLInsertOmitsGeneratedAndReturnsProjection(t *testing.T) {
code := compileSQL(t, eventSQLSource+`
import context
import pgxpool "github.com/jackc/pgx/v5/pgxpool"
fun create(row: EventRow, pool: *pgxpool.Pool, ctx: context.Context): *EventProjection {
return sql.insert<EventRow>(row)
.returning { value -> EventProjection(value.id, value.payload) }
.single(pool, ctx)
}
`)
for _, want := range []string{
`SQL: "INSERT INTO outbox_events (payload, attempts, published_at, claimed_until, created_at) VALUES ($1, $2, $3, $4, $5) RETURNING id, payload"`,
`Args: []any{row.Payload, row.Attempts, row.PublishedAt, row.ClaimedUntil, row.CreatedAt}`,
`gotlinSQLSingle[EventProjection]`,
} {
if !strings.Contains(code, want) {
t.Fatalf("generated Go missing %q:\n%s", want, code)
}
}
}
func TestSQLGeneratedOnlyInsertUsesDefaultValues(t *testing.T) {
code := compileSQL(t, `
@table("tokens") data class TokenRow(@generated @id var id: String)
fun insert(row: TokenRow): GotlinSQLQuery { return sql.insert<TokenRow>(row).build() }
`)
if want := `SQL: "INSERT INTO tokens DEFAULT VALUES"`; !strings.Contains(code, want) {
t.Fatalf("generated Go missing %q:\n%s", want, code)
}
}
func TestSQLTypedUpdateAndReturning(t *testing.T) {
code := compileSQL(t, eventSQLSource+`
import context
import pgxpool "github.com/jackc/pgx/v5/pgxpool"
fun claim(id: String, payload: String, pool: *pgxpool.Pool, ctx: context.Context): *EventProjection {
return sql.update<EventRow>()
.set { row ->
set(row.payload, payload)
set(row.attempts, row.attempts + 1)
set(row.claimedUntil, now())
}
.where { it.id == id && it.publishedAt == null }
.returning { row -> EventProjection(row.id, row.payload) }
.single(pool, ctx)
}
`)
for _, want := range []string{
`SQL: "UPDATE outbox_events SET payload = $1, attempts = (attempts + $2), claimed_until = CURRENT_TIMESTAMP WHERE (id = $3 AND published_at IS NULL) RETURNING id, payload"`,
`Args: []any{payload, 1, id}`,
`gotlinSQLSingle[EventProjection]`,
} {
if !strings.Contains(code, want) {
t.Fatalf("generated Go missing %q:\n%s", want, code)
}
}
}
func TestSQLDeleteReturningFullRow(t *testing.T) {
code := compileSQL(t, eventSQLSource+`
import context
import pgxpool "github.com/jackc/pgx/v5/pgxpool"
fun remove(id: String, pool: *pgxpool.Pool, ctx: context.Context): *EventRow {
return sql.delete<EventRow>()
.where { it.id == id }
.returning { it }
.single(pool, ctx)
}
`)
for _, want := range []string{
`SQL: "DELETE FROM outbox_events WHERE id = $1 RETURNING id, payload, attempts, published_at, claimed_until, created_at"`,
`gotlinSQLSingle[EventRow]`,
`err := row.Scan(&value.Id, &value.Payload, &value.Attempts, &value.PublishedAt, &value.ClaimedUntil, &value.CreatedAt)`,
} {
if !strings.Contains(code, want) {
t.Fatalf("generated Go missing %q:\n%s", want, code)
}
}
}
func TestRejectExpandedInvalidSQLQueries(t *testing.T) {
tests := []struct {
name string
src string
want string
}{
{
name: "skip locked without for update",
src: eventSQLSource + `fun query(): GotlinSQLQuery { return sql.from<EventRow>().skipLocked().build() }`,
want: "requires a preceding forUpdate",
},
{
name: "lock before limit",
src: eventSQLSource + `fun query(): GotlinSQLQuery { return sql.from<EventRow>().forUpdate().limit(1).build() }`,
want: "limit() may appear once",
},
{
name: "non nullable null comparison",
src: eventSQLSource + `fun query(): GotlinSQLQuery { return sql.from<EventRow>().where { it.payload == null }.build() }`,
want: "null comparison requires a nullable operand",
},
{
name: "wrong limit type",
src: eventSQLSource + `fun query(limit: String): GotlinSQLQuery { return sql.from<EventRow>().limit(limit).build() }`,
want: "argument must have type Int",
},
{
name: "negative literal limit",
src: eventSQLSource + `fun query(): GotlinSQLQuery { return sql.from<EventRow>().limit(-1).build() }`,
want: "requires a non-negative Int",
},
{
name: "projection type mismatch",
src: eventSQLSource + `fun query(): GotlinSQLQuery { return sql.from<EventRow>().select { WrongProjection(it.id) }.build() }`,
want: "has type Int but row field id has type String",
},
{
name: "projection arity mismatch",
src: eventSQLSource + `fun query(): GotlinSQLQuery { return sql.from<EventRow>().select { EventProjection(it.id) }.build() }`,
want: "expects 2 fields, got 1",
},
{
name: "projection after where",
src: eventSQLSource + `fun query(): GotlinSQLQuery { return sql.from<EventRow>().where { it.id == "x" }.select { EventProjection(it.id, it.payload) }.build() }`,
want: "must be the first sql.from method",
},
{
name: "update missing set",
src: eventSQLSource + `fun query(): GotlinSQLQuery { return sql.update<EventRow>().where { it.id == "x" }.build() }`,
want: "requires set() as its first method",
},
{
name: "update incompatible value",
src: eventSQLSource + `fun query(): GotlinSQLQuery { return sql.update<EventRow>().set { set(it.attempts, "bad") }.build() }`,
want: "has type Int but value has type String",
},
{
name: "update null non nullable",
src: eventSQLSource + `fun query(): GotlinSQLQuery { return sql.update<EventRow>().set { set(it.payload, null) }.build() }`,
want: "has non-nullable type String",
},
{
name: "update nullable into non nullable",
src: eventSQLSource + `fun query(value: String?): GotlinSQLQuery { return sql.update<EventRow>().set { set(it.payload, value) }.build() }`,
want: "has type String but value has type String?",
},
{
name: "duplicate update target",
src: eventSQLSource + `fun query(): GotlinSQLQuery { return sql.update<EventRow>().set { set(it.payload, "a"); set(it.payload, "b") }.build() }`,
want: "duplicate set target payload",
},
{
name: "write execution without returning",
src: eventSQLSource + `fun query(pool: Any, ctx: Any): *EventRow { return sql.delete<EventRow>().single(pool, ctx) }`,
want: "requires returning()",
},
{
name: "returning out of order",
src: eventSQLSource + `fun query(): GotlinSQLQuery { return sql.delete<EventRow>().returning { it }.where { it.id == "x" }.build() }`,
want: "invalid method order",
},
}
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 TestRejectDuplicateGeneratedAnnotation(t *testing.T) {
_, err := Parse(`data class Row(@generated @generated var id: String)`)
if err == nil || !strings.Contains(err.Error(), "duplicate generated annotation") {
t.Fatalf("Parse() error = %v", err)
}
}

383
internal/lang/sql_test.go Normal file
View file

@ -0,0 +1,383 @@
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<AccountRow>()
.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<AccountRow>().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"
fun execute(pool: *pgxpool.Pool, ctx: Context, customerId: String) {
val query = sql.from<AccountRow>().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<*AccountRow> {
return sql.from<AccountRow>()
.where { it.customerId == customerId }
.fetch(pool, ctx)
}
`)
for _, want := range []string{
`"github.com/jackc/pgx/v5"`,
`func accounts(pool *pgxpool.Pool, ctx context.Context, customerId string) []*AccountRow`,
`return 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, gotlinAutoThrow(scan(rows)))`,
`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<AccountRow>().where { it.id == id }.single(pool, ctx)
return account.balance
}
`)
for _, want := range []string{
`account := gotlinSQLSingle[AccountRow](pool, ctx, GotlinSQLQuery{SQL: "SELECT id, customer_id, kind, balance FROM accounts WHERE id = $1", Args: []any{id}}, gotlinSQLScanAccountRow)`,
`return account.Balance`,
`gotlinAutoThrow(gotlinSQLError("SQL single() expected exactly one row, got zero"))`,
`gotlinAutoThrow(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<AccountRow>().iterator(pool, ctx)
defer rows.close()
while (rows.next()) {
val account = rows.value()
println(account.customerId)
}
val checked = rows.err()
}
`)
for _, want := range []string{
`rows := gotlinSQLIterate[AccountRow](pool, ctx, GotlinSQLQuery{SQL: "SELECT id, customer_id, kind, balance FROM accounts", Args: []any{}}, gotlinSQLScanAccountRow)`,
`defer rows.close()`,
`for rows.next()`,
`account := gotlinAutoThrow(rows.value())`,
`fmt.Println(account.CustomerId)`,
`_ = gotlinAutoThrow(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<AccountRow>().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<AccountRow>(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<AccountRow>(row)
.onConflict { it.id }
.doUpdate { excluded -> set(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<BalanceRow>(row)
.onConflict { listOf(it.tenantId, it.id) }
.doUpdate { excluded ->
set(BalanceRow.amount, excluded.amount)
set(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 TestRejectInvalidSQLQueries(t *testing.T) {
tests := []struct {
name string
src string
want string
}{
{
name: "unknown row class",
src: `fun query(): GotlinSQLQuery { return sql.from<MissingRow>().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<Row>().build() }`,
want: "requires @table",
},
{
name: "unknown predicate field",
src: accountRowSource + `fun query(): GotlinSQLQuery { return sql.from<AccountRow>().where { it.missing == "x" }.build() }`,
want: "has no field missing",
},
{
name: "incompatible predicate operands",
src: accountRowSource + `fun query(customerId: String): GotlinSQLQuery { return sql.from<AccountRow>().where { it.balance == customerId }.build() }`,
want: "incompatible types Double and String",
},
{
name: "non numeric ordering predicate",
src: accountRowSource + `fun query(): GotlinSQLQuery { return sql.from<AccountRow>().where { it.customerId > "a" }.build() }`,
want: "ordering comparison requires numeric operands",
},
{
name: "unknown order field",
src: accountRowSource + `fun query(): GotlinSQLQuery { return sql.from<AccountRow>().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<AccountRow>(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<AccountRow>(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<AccountRow>(row).onConflict { it.id }.doUpdate { excluded -> set(AccountRow.missing, excluded.balance) }.build() }`,
want: "has no field missing",
},
{
name: "incompatible update fields",
src: accountRowSource + `fun query(row: AccountRow): GotlinSQLQuery { return sql.insert<AccountRow>(row).onConflict { it.id }.doUpdate { excluded -> set(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<AccountRow>(row).onConflict { it.id }.doNothing().build() }
`,
want: "cannot insert value of type OtherRow",
},
{
name: "incomplete chain",
src: accountRowSource + `fun query(): GotlinSQLQuery { return sql.from<AccountRow>() }`,
want: "must end with build()",
},
{
name: "fetch missing arguments",
src: accountRowSource + `fun query(): List<*AccountRow> { return sql.from<AccountRow>().fetch() }`,
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<AccountRow>().single(pool, ctx, ctx) }`,
want: "single() expects exactly pool and ctx positional arguments",
},
{
name: "iterator named arguments",
src: accountRowSource + `fun query(pool: Any, ctx: Any): GotlinSQLIterator<AccountRow> { return sql.from<AccountRow>().iterator(pool = pool, ctx = ctx) }`,
want: "iterator() expects exactly pool and ctx positional arguments",
},
{
name: "fetch type arguments",
src: accountRowSource + `fun query(pool: Any, ctx: Any): List<*AccountRow> { return sql.from<AccountRow>().fetch<String>(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<*AccountRow> { return sql.from<AccountRow>().fetch(pool, ctx) }`,
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<AccountRow>().single(pool, ctx) }`,
want: "ctx argument has non-context type Int",
},
{
name: "insert execution terminal",
src: accountRowSource + `fun query(row: AccountRow, pool: Any, ctx: Any): List<*AccountRow> { return sql.insert<AccountRow>(row).fetch(pool, ctx) }`,
want: "fetch() is only supported for sql.from",
},
}
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)
}

View file

@ -3,56 +3,70 @@ package lang
type tokenKind string
const (
tokenEOF tokenKind = "EOF"
tokenIdent tokenKind = "IDENT"
tokenInt tokenKind = "INT"
tokenString tokenKind = "STRING"
tokenTrue tokenKind = "TRUE"
tokenFalse tokenKind = "FALSE"
tokenNull tokenKind = "NULL"
tokenImport tokenKind = "IMPORT"
tokenPackage tokenKind = "PACKAGE"
tokenClass tokenKind = "CLASS"
tokenWorker tokenKind = "WORKER"
tokenInterface tokenKind = "INTERFACE"
tokenFun tokenKind = "FUN"
tokenOverride tokenKind = "OVERRIDE"
tokenVal tokenKind = "VAL"
tokenVar tokenKind = "VAR"
tokenIf tokenKind = "IF"
tokenElse tokenKind = "ELSE"
tokenWhile tokenKind = "WHILE"
tokenSelect tokenKind = "SELECT"
tokenReturn tokenKind = "RETURN"
tokenGo tokenKind = "GO"
tokenTry tokenKind = "TRY"
tokenCatch tokenKind = "CATCH"
tokenThrow tokenKind = "THROW"
tokenLParen tokenKind = "("
tokenRParen tokenKind = ")"
tokenLBrace tokenKind = "{"
tokenRBrace tokenKind = "}"
tokenComma tokenKind = ","
tokenDot tokenKind = "."
tokenColon tokenKind = ":"
tokenSemicolon tokenKind = ";"
tokenPlus tokenKind = "+"
tokenMinus tokenKind = "-"
tokenStar tokenKind = "*"
tokenSlash tokenKind = "/"
tokenPercent tokenKind = "%"
tokenBang tokenKind = "!"
tokenAssign tokenKind = "="
tokenPlusAssign tokenKind = "+="
tokenEq tokenKind = "=="
tokenNeq tokenKind = "!="
tokenLt tokenKind = "<"
tokenLte tokenKind = "<="
tokenGt tokenKind = ">"
tokenGte tokenKind = ">="
tokenAnd tokenKind = "&&"
tokenOr tokenKind = "||"
tokenArrow tokenKind = "->"
tokenEOF tokenKind = "EOF"
tokenIdent tokenKind = "IDENT"
tokenInt tokenKind = "INT"
tokenFloat tokenKind = "FLOAT"
tokenString tokenKind = "STRING"
tokenTrue tokenKind = "TRUE"
tokenFalse tokenKind = "FALSE"
tokenNull tokenKind = "NULL"
tokenImport tokenKind = "IMPORT"
tokenPackage tokenKind = "PACKAGE"
tokenClass tokenKind = "CLASS"
tokenData tokenKind = "DATA"
tokenWorker tokenKind = "WORKER"
tokenInterface tokenKind = "INTERFACE"
tokenEnum tokenKind = "ENUM"
tokenMatch tokenKind = "MATCH"
tokenFun tokenKind = "FUN"
tokenOverride tokenKind = "OVERRIDE"
tokenPrivate tokenKind = "PRIVATE"
tokenVal tokenKind = "VAL"
tokenVar tokenKind = "VAR"
tokenIf tokenKind = "IF"
tokenElse tokenKind = "ELSE"
tokenWhile tokenKind = "WHILE"
tokenFor tokenKind = "FOR"
tokenIn tokenKind = "IN"
tokenSelect tokenKind = "SELECT"
tokenReturn tokenKind = "RETURN"
tokenGo tokenKind = "GO"
tokenDefer tokenKind = "DEFER"
tokenTry tokenKind = "TRY"
tokenCatch tokenKind = "CATCH"
tokenThrow tokenKind = "THROW"
tokenLParen tokenKind = "("
tokenRParen tokenKind = ")"
tokenLBrace tokenKind = "{"
tokenRBrace tokenKind = "}"
tokenLBracket tokenKind = "["
tokenRBracket tokenKind = "]"
tokenComma tokenKind = ","
tokenDot tokenKind = "."
tokenColon tokenKind = ":"
tokenDoubleColon tokenKind = "::"
tokenSemicolon tokenKind = ";"
tokenPlus tokenKind = "+"
tokenMinus tokenKind = "-"
tokenStar tokenKind = "*"
tokenSlash tokenKind = "/"
tokenPercent tokenKind = "%"
tokenBang tokenKind = "!"
tokenAssign tokenKind = "="
tokenPlusAssign tokenKind = "+="
tokenEq tokenKind = "=="
tokenNeq tokenKind = "!="
tokenLt tokenKind = "<"
tokenLte tokenKind = "<="
tokenGt tokenKind = ">"
tokenGte tokenKind = ">="
tokenAnd tokenKind = "&&"
tokenAmp tokenKind = "&"
tokenAt tokenKind = "@"
tokenQuestion tokenKind = "?"
tokenOr tokenKind = "||"
tokenArrow tokenKind = "->"
)
var keywords = map[string]tokenKind{
@ -60,17 +74,24 @@ var keywords = map[string]tokenKind{
"import": tokenImport,
"package": tokenPackage,
"class": tokenClass,
"data": tokenData,
"worker": tokenWorker,
"interface": tokenInterface,
"enum": tokenEnum,
"match": tokenMatch,
"val": tokenVal,
"var": tokenVar,
"override": tokenOverride,
"private": tokenPrivate,
"if": tokenIf,
"else": tokenElse,
"while": tokenWhile,
"for": tokenFor,
"in": tokenIn,
"select": tokenSelect,
"return": tokenReturn,
"go": tokenGo,
"defer": tokenDefer,
"try": tokenTry,
"catch": tokenCatch,
"throw": tokenThrow,