Introduce typed semantic analysis pipeline

This commit is contained in:
pavel 2026-08-27 19:28:24 +02:00
commit f4cd4f4458
30 changed files with 2079 additions and 2381 deletions

View file

@ -8,6 +8,11 @@ This is the practical boundary of the prototype:
- It targets the Go toolchain by generating valid Go source and building through `go build`. - It targets the Go toolchain by generating valid Go source and building through `go build`.
- It is not a direct integration into Go's internal `cmd/compile` backend APIs. - It is not a direct integration into Go's internal `cmd/compile` backend APIs.
The compiler keeps source spelling in its syntax AST, then builds lexical
symbols, structural semantic types, resolved expression meanings, and typed
HIR before Go emission. Class reference semantics live in `ClassType`; only
the semantic Type-to-Go mapping turns a class such as `User` into `*User`.
## Supported language slice ## Supported language slice
- `fun` declarations - `fun` declarations
@ -170,8 +175,8 @@ destructuring, or `.unwrap()`.
Gotlin-defined `class` and `data class` values are references automatically, Gotlin-defined `class` and `data class` values are references automatically,
including nested generic types such as `List<User>`. Explicit pointer syntax is including nested generic types such as `List<User>`. Explicit pointer syntax is
still supported for compatibility. Go interop remains explicit, for example reserved for Go interop, for example `*http.Request` and `*pgxpool.Pool`;
`*http.Request` and `*pgxpool.Pool`. applying `*` to a Gotlin class is a compile error.
Enums whose variants carry no payload are represented as string-backed values. Enums whose variants carry no payload are represented as string-backed values.
The exact variant identifier is used for JSON and PostgreSQL text values: The exact variant identifier is used for JSON and PostgreSQL text values:
@ -251,8 +256,9 @@ go run ./cmd/gotlinc run ./examples/hello.gt
Gotlin coroutines use Go goroutines underneath, but expose only structured Gotlin coroutines use Go goroutines underneath, but expose only structured
scopes. A scope waits for its children, propagates child failures, and cancels scopes. A scope waits for its children, propagates child failures, and cancels
sibling coroutine contexts. The removed `worker` and bare `go` forms are not sibling coroutine contexts. The removed `worker`, bare `go`, and channel
valid Gotlin syntax. `select` forms are not valid Gotlin syntax; use coroutine scopes, `delay`, and
explicit channel `read()`/`send()` operations.
```kotlin ```kotlin
suspend fun load(): Int { suspend fun load(): Int {

View file

@ -306,6 +306,9 @@ func buildDocumentState(text string) documentState {
return state return state
} }
state.program = program state.program = program
if _, err := lang.Analyze(program); err != nil {
state.diagnostics = append(state.diagnostics, diagnosticFromError(text, err))
}
state.symbols = indexSymbols(text, program) state.symbols = indexSymbols(text, program)
state.diagnostics = semanticDiagnostics(text, program, state.symbols) state.diagnostics = semanticDiagnostics(text, program, state.symbols)
if _, err := lang.GenerateGo(program); err != nil { if _, err := lang.GenerateGo(program); err != nil {
@ -681,38 +684,6 @@ func indexSymbols(text string, program *lang.Program) []symbol {
}) })
} }
} }
workerRanges := findAllLineMatches(lines, regexp.MustCompile(`^\s*worker\s+([A-Za-z_][A-Za-z0-9_]*)`))
for i, decl := range program.Workers {
r := rng{}
if i < len(workerRanges) {
r = workerRanges[i]
}
symbols = append(symbols, symbol{
Name: decl.Name,
Kind: symbolKindClass,
Detail: renderWorkerSignature(decl),
Range: r,
})
for _, field := range decl.Fields {
symbols = append(symbols, symbol{
Name: field.Name,
Kind: symbolKindField,
Detail: renderWorkerFieldSignature(decl.Name, field),
Range: r,
Targets: []string{decl.Name},
})
}
for _, method := range decl.Methods {
symbols = append(symbols, symbol{
Name: method.Name,
Kind: symbolKindMethod,
Detail: renderMethodSignature(decl.Name, method),
Range: r,
Targets: []string{decl.Name},
})
}
}
funcRanges := findAllLineMatches(lines, regexp.MustCompile(`^\s*fun\s+([A-Za-z_][A-Za-z0-9_]*)`)) funcRanges := findAllLineMatches(lines, regexp.MustCompile(`^\s*fun\s+([A-Za-z_][A-Za-z0-9_]*)`))
for i, fn := range program.Functions { for i, fn := range program.Functions {
detail := renderFunctionSignature(fn) detail := renderFunctionSignature(fn)
@ -809,9 +780,6 @@ func semanticDiagnosticsWithPackage(text string, program *lang.Program, symbols
for _, decl := range program.Classes { for _, decl := range program.Classes {
types[decl.Name] = true types[decl.Name] = true
} }
for _, decl := range program.Workers {
types[decl.Name] = true
}
for _, decl := range program.Enums { for _, decl := range program.Enums {
types[decl.Name] = true types[decl.Name] = true
} }
@ -825,9 +793,6 @@ func semanticDiagnosticsWithPackage(text string, program *lang.Program, symbols
for _, decl := range sibling.Classes { for _, decl := range sibling.Classes {
types[decl.Name] = true types[decl.Name] = true
} }
for _, decl := range sibling.Workers {
types[decl.Name] = true
}
for _, decl := range sibling.Enums { for _, decl := range sibling.Enums {
types[decl.Name] = true types[decl.Name] = true
} }
@ -846,17 +811,6 @@ func semanticDiagnosticsWithPackage(text string, program *lang.Program, symbols
diagnostics = append(diagnostics, functionSemanticDiagnostics(text, method, functions, imports, types, fields)...) diagnostics = append(diagnostics, functionSemanticDiagnostics(text, method, functions, imports, types, fields)...)
} }
} }
for _, worker := range program.Workers {
fields := map[string]bool{
"this": true,
}
for _, field := range worker.Fields {
fields[field.Name] = true
}
for _, method := range worker.Methods {
diagnostics = append(diagnostics, functionSemanticDiagnostics(text, method, functions, imports, types, fields)...)
}
}
sort.SliceStable(diagnostics, func(i, j int) bool { sort.SliceStable(diagnostics, func(i, j int) bool {
if diagnostics[i].Range.Start.Line != diagnostics[j].Range.Start.Line { if diagnostics[i].Range.Start.Line != diagnostics[j].Range.Start.Line {
@ -913,8 +867,6 @@ func functionSemanticDiagnostics(text string, fn lang.FunctionDecl, functions ma
} }
case lang.ThrowStmt: case lang.ThrowStmt:
walkExpr(s.Value, scope) walkExpr(s.Value, scope)
case lang.GoStmt:
walkExpr(s.Value, scope)
case lang.DeferStmt: case lang.DeferStmt:
walkExpr(s.Value, scope) walkExpr(s.Value, scope)
case lang.ExprStmt: case lang.ExprStmt:
@ -934,13 +886,6 @@ func functionSemanticDiagnostics(text string, fn lang.FunctionDecl, functions ma
bodyScope := copyScope(scope) bodyScope := copyScope(scope)
bodyScope[s.Name] = true bodyScope[s.Name] = true
walkStmts(s.Body, bodyScope) walkStmts(s.Body, bodyScope)
case lang.SelectStmt:
for _, c := range s.Cases {
walkExpr(c.Source, scope)
caseScope := copyScope(scope)
caseScope["it"] = true
walkStmts(c.Body, caseScope)
}
case lang.MatchStmt: case lang.MatchStmt:
walkExpr(s.Value, scope) walkExpr(s.Value, scope)
for _, matchCase := range s.Cases { for _, matchCase := range s.Cases {
@ -973,25 +918,6 @@ func functionSemanticDiagnostics(text string, fn lang.FunctionDecl, functions ma
walkExpr(e.Left, scope) walkExpr(e.Left, scope)
walkExpr(e.Right, scope) walkExpr(e.Right, scope)
case lang.CallExpr: case lang.CallExpr:
if ident, ok := e.Callee.(lang.IdentExpr); ok {
switch ident.Name {
case "after", "every":
if len(e.Args) != 1 {
diagnostics = append(diagnostics, undefinedNameDiagnostic(text, ident.Name, ident.Name+"(ms) expects exactly one Int argument"))
}
if len(e.Args) == 1 {
switch e.Args[0].(type) {
case lang.StringExpr, lang.BoolExpr, lang.NullExpr:
diagnostics = append(diagnostics, undefinedNameDiagnostic(text, ident.Name, ident.Name+"(ms) expects an Int argument"))
}
}
if ident.Name == "every" && len(e.Args) == 1 {
if value, ok := staticIntExprValue(e.Args[0]); ok && value <= 0 {
diagnostics = append(diagnostics, undefinedNameDiagnostic(text, ident.Name, "every(ms) requires ms > 0"))
}
}
}
}
walkExpr(e.Callee, scope) walkExpr(e.Callee, scope)
for _, arg := range e.Args { for _, arg := range e.Args {
walkExpr(arg, scope) walkExpr(arg, scope)
@ -1339,10 +1265,6 @@ func renderClassSignature(decl lang.ClassDecl) string {
return b.String() return b.String()
} }
func renderWorkerSignature(decl lang.WorkerDecl) string {
return "worker " + decl.Name
}
func renderMethodSignature(className string, fn lang.FunctionDecl) string { func renderMethodSignature(className string, fn lang.FunctionDecl) string {
return className + "." + renderFunctionSignature(fn) return className + "." + renderFunctionSignature(fn)
} }
@ -1355,21 +1277,6 @@ func renderFieldSignature(className string, field lang.FieldDecl) string {
return className + "." + keyword + " " + field.Name + ": " + field.Type return className + "." + keyword + " " + field.Name + ": " + field.Type
} }
func renderWorkerFieldSignature(workerName string, field lang.WorkerFieldDecl) string {
keyword := "val"
if field.Mutable {
keyword = "var"
}
typ := field.Type
if typ == "" {
typ = inferWorkerFieldType(field)
}
if typ != "" {
return workerName + "." + keyword + " " + field.Name + ": " + typ
}
return workerName + "." + keyword + " " + field.Name
}
func renderVariableSignature(name string, mutable bool, typ string) string { func renderVariableSignature(name string, mutable bool, typ string) string {
keyword := "val" keyword := "val"
if mutable { if mutable {
@ -1390,10 +1297,6 @@ func collectVariableDeclInfos(program *lang.Program) []variableDeclInfo {
for _, class := range program.Classes { for _, class := range program.Classes {
classes[class.Name] = class classes[class.Name] = class
} }
workers := map[string]lang.WorkerDecl{}
for _, worker := range program.Workers {
workers[worker.Name] = worker
}
var out []variableDeclInfo var out []variableDeclInfo
var walkStmts func(stmts []lang.Stmt, scope map[string]string) var walkStmts func(stmts []lang.Stmt, scope map[string]string)
@ -1409,14 +1312,9 @@ func collectVariableDeclInfos(program *lang.Program) []variableDeclInfo {
if _, ok := classes[callee.Name]; ok { if _, ok := classes[callee.Name]; ok {
return []string{callee.Name} return []string{callee.Name}
} }
if _, ok := workers[callee.Name]; ok {
return []string{callee.Name}
}
switch callee.Name { switch callee.Name {
case "println": case "println":
return []string{"Unit"} return []string{"Unit"}
case "after", "every":
return []string{"Channel<time.Time>"}
case "Channel": case "Channel":
if len(call.TypeArgs) == 1 { if len(call.TypeArgs) == 1 {
return []string{"Channel<" + call.TypeArgs[0] + ">"} return []string{"Channel<" + call.TypeArgs[0] + ">"}
@ -1562,16 +1460,6 @@ func collectVariableDeclInfos(program *lang.Program) []variableDeclInfo {
} }
} }
} }
if worker, ok := workers[trimmed]; ok {
for _, field := range worker.Fields {
if field.Name == e.Name {
if field.Type != "" {
return field.Type
}
return inferWorkerFieldType(field)
}
}
}
return "" return ""
default: default:
return "" return ""
@ -1584,7 +1472,12 @@ func collectVariableDeclInfos(program *lang.Program) []variableDeclInfo {
case lang.VarDecl: case lang.VarDecl:
typ := s.Type typ := s.Type
if typ == "" { if typ == "" {
typ = inferExprType(s.Value, scope) resolved := lang.ResolvedType(s.Value)
if resolved.String() != "<unknown>" {
typ = resolved.String()
} else {
typ = inferExprType(s.Value, scope)
}
} }
out = append(out, variableDeclInfo{Name: s.Name, Mutable: s.Mutable, Type: typ}) out = append(out, variableDeclInfo{Name: s.Name, Mutable: s.Mutable, Type: typ})
scope[s.Name] = typ scope[s.Name] = typ
@ -1613,16 +1506,6 @@ func collectVariableDeclInfos(program *lang.Program) []variableDeclInfo {
case lang.WhileStmt: case lang.WhileStmt:
bodyScope := copyTypeScope(scope) bodyScope := copyTypeScope(scope)
walkStmts(s.Body, bodyScope) walkStmts(s.Body, bodyScope)
case lang.SelectStmt:
for _, c := range s.Cases {
caseScope := copyTypeScope(scope)
if elemType, ok := channelElementType(inferExprType(c.Source, scope)); ok {
caseScope["it"] = elemType
} else {
caseScope["it"] = "any"
}
walkStmts(c.Body, caseScope)
}
case lang.TryCatchStmt: case lang.TryCatchStmt:
tryScope := copyTypeScope(scope) tryScope := copyTypeScope(scope)
walkStmts(s.TryBody, tryScope) walkStmts(s.TryBody, tryScope)
@ -1653,23 +1536,6 @@ func collectVariableDeclInfos(program *lang.Program) []variableDeclInfo {
walkStmts(method.Body, scope) walkStmts(method.Body, scope)
} }
} }
for _, worker := range program.Workers {
fieldScope := map[string]string{}
for _, field := range worker.Fields {
typ := field.Type
if typ == "" {
typ = inferWorkerFieldType(field)
}
fieldScope[field.Name] = typ
}
for _, method := range worker.Methods {
scope := copyTypeScope(fieldScope)
for _, param := range method.Params {
scope[param.Name] = param.Type
}
walkStmts(method.Body, scope)
}
}
return out return out
} }
@ -1693,44 +1559,6 @@ func channelElementType(typ string) (string, bool) {
return inner, true return inner, true
} }
func staticIntExprValue(expr lang.Expr) (int, bool) {
switch e := expr.(type) {
case lang.IntExpr:
v, err := strconv.Atoi(e.Value)
if err != nil {
return 0, false
}
return v, true
case lang.UnaryExpr:
if e.Op != "-" {
return 0, false
}
v, ok := staticIntExprValue(e.Value)
if !ok {
return 0, false
}
return -v, true
default:
return 0, false
}
}
func inferWorkerFieldType(field lang.WorkerFieldDecl) string {
if field.Type != "" {
return field.Type
}
switch field.Value.(type) {
case lang.IntExpr:
return "Int"
case lang.StringExpr:
return "String"
case lang.BoolExpr:
return "Boolean"
default:
return ""
}
}
func selectorPathLang(expr lang.Expr) (string, bool) { func selectorPathLang(expr lang.Expr) (string, bool) {
switch e := expr.(type) { switch e := expr.(type) {
case lang.IdentExpr: case lang.IdentExpr:
@ -2199,8 +2027,6 @@ var builtinDetails = map[string]string{
"println": "fun println(value: Any): Unit", "println": "fun println(value: Any): Unit",
"runCatching": "fun runCatching(block: () -> Unit): Result", "runCatching": "fun runCatching(block: () -> Unit): Result",
"Channel": "fun Channel<T>(capacity: Int = 0): Channel<T>", "Channel": "fun Channel<T>(capacity: Int = 0): Channel<T>",
"after": "fun after(ms: Int): Channel<time.Time>",
"every": "fun every(ms: Int): Channel<time.Time>",
"listOf": "fun listOf<T>(values: T...): List<T>", "listOf": "fun listOf<T>(values: T...): List<T>",
"mutableListOf": "fun mutableListOf<T>(values: T...): MutableList<T>", "mutableListOf": "fun mutableListOf<T>(values: T...): MutableList<T>",
"mapOf": "fun mapOf<K, V>(pairs: Any...): Map<K, V>", "mapOf": "fun mapOf<K, V>(pairs: Any...): Map<K, V>",

View file

@ -707,7 +707,7 @@ fun main() {
} }
} }
func TestBuildDocumentStateChannelsAndSelectSemantics(t *testing.T) { func TestBuildDocumentStateChannelSemantics(t *testing.T) {
text := strings.TrimSpace(` text := strings.TrimSpace(`
package demo package demo
@ -716,13 +716,8 @@ fun writer(ch: Channel<Int>) {
} }
fun main() { fun main() {
val ch = Channel<Int>() val ch = Channel<Int>(1)
runBlocking { writer(ch)
launch { writer(ch) }
select {
ch -> println(it)
}
}
val v = ch.read() val v = ch.read()
println(v) println(v)
} }
@ -787,138 +782,3 @@ fun main() {
t.Fatal("expected removed worker syntax diagnostic") t.Fatal("expected removed worker syntax diagnostic")
} }
} }
func TestBuildDocumentStateTimerBuiltinsSemantics(t *testing.T) {
text := strings.TrimSpace(`
package demo
fun main() {
val once = after(1000)
val repeat = every(250)
select {
after(1000) -> println(it)
every(250) -> println("tick")
}
println(once)
println(repeat)
}
`)
state := buildDocumentState(text)
if state.program == nil {
t.Fatal("expected parsed program")
}
if len(state.diagnostics) != 0 {
t.Fatalf("expected no diagnostics, got %+v", state.diagnostics)
}
var onceDetail string
var repeatDetail string
for _, sym := range state.symbols {
if sym.Kind != symbolKindVariable {
continue
}
if sym.Name == "once" {
onceDetail = sym.Detail
}
if sym.Name == "repeat" {
repeatDetail = sym.Detail
}
}
if onceDetail != "val once: Channel<time.Time>" {
t.Fatalf("unexpected once detail: %q", onceDetail)
}
if repeatDetail != "val repeat: Channel<time.Time>" {
t.Fatalf("unexpected repeat detail: %q", repeatDetail)
}
}
func TestBuildDocumentStateEveryNonPositiveDiagnostics(t *testing.T) {
text := strings.TrimSpace(`
package demo
fun main() {
select {
every(0) -> println("x")
}
}
`)
state := buildDocumentState(text)
if state.program == nil {
t.Fatal("expected parsed program")
}
if len(state.diagnostics) == 0 {
t.Fatal("expected diagnostics")
}
found := false
for _, d := range state.diagnostics {
if strings.Contains(d.Message, "every(ms) requires ms > 0") {
found = true
break
}
}
if !found {
t.Fatalf("expected every(ms) diagnostic, got %+v", state.diagnostics)
}
}
func TestBuildDocumentStateAfterIntArgumentDiagnostics(t *testing.T) {
text := strings.TrimSpace(`
package demo
fun main() {
select {
after("1s") -> println("x")
}
}
`)
state := buildDocumentState(text)
if state.program == nil {
t.Fatal("expected parsed program")
}
found := false
for _, d := range state.diagnostics {
if strings.Contains(d.Message, "after(ms) expects an Int argument") {
found = true
break
}
}
if !found {
t.Fatalf("expected after int argument diagnostic, got %+v", state.diagnostics)
}
}
func TestHoverBuiltinEverySignature(t *testing.T) {
text := strings.TrimSpace(`
package demo
fun main() {
every(100)
}
`)
uri := "file:///tmp/demo.gt"
s := &server{
docs: map[string]documentState{
uri: buildDocumentState(text),
},
}
result := s.hover(uri, position{Line: 3, Character: 5})
if result == nil {
t.Fatal("expected hover result")
}
m, ok := result.(map[string]any)
if !ok {
t.Fatalf("unexpected hover type: %T", result)
}
contents, ok := m["contents"].(map[string]any)
if !ok {
t.Fatalf("unexpected hover contents: %+v", m)
}
value, _ := contents["value"].(string)
if !strings.Contains(value, "fun every(ms: Int): Channel<time.Time>") {
t.Fatalf("unexpected hover value: %q", value)
}
}

View file

@ -140,7 +140,6 @@ func compileFiles(inputPaths []string, forceMain bool) string {
program.Interfaces = append(program.Interfaces, parsed.Interfaces...) program.Interfaces = append(program.Interfaces, parsed.Interfaces...)
program.Enums = append(program.Enums, parsed.Enums...) program.Enums = append(program.Enums, parsed.Enums...)
program.Classes = append(program.Classes, parsed.Classes...) program.Classes = append(program.Classes, parsed.Classes...)
program.Workers = append(program.Workers, parsed.Workers...)
program.Functions = append(program.Functions, parsed.Functions...) program.Functions = append(program.Functions, parsed.Functions...)
program.Embeds = append(program.Embeds, parsed.Embeds...) program.Embeds = append(program.Embeds, parsed.Embeds...)
if firstSource == "" { if firstSource == "" {

View file

@ -36,9 +36,8 @@ fun main() {
delay(120) delay(120)
ready.send("timer fired") ready.send("timer fired")
} }
select { val message = ready.read()
ready -> println("channel says: " + it) println("channel says: " + message)
}
val answer = async<Int> { val answer = async<Int> {
delay(50) delay(50)

View file

@ -6,7 +6,6 @@ type Program struct {
Interfaces []InterfaceDecl Interfaces []InterfaceDecl
Enums []EnumDecl Enums []EnumDecl
Classes []ClassDecl Classes []ClassDecl
Workers []WorkerDecl
Functions []FunctionDecl Functions []FunctionDecl
Embeds []EmbedDecl Embeds []EmbedDecl
} }
@ -44,19 +43,6 @@ type ClassDecl struct {
Methods []FunctionDecl Methods []FunctionDecl
} }
type WorkerDecl struct {
Name string
Fields []WorkerFieldDecl
Methods []FunctionDecl
}
type WorkerFieldDecl struct {
Mutable bool
Name string
Type string
Value Expr
}
type FieldDecl struct { type FieldDecl struct {
Mutable bool Mutable bool
Private bool Private bool
@ -148,12 +134,6 @@ type ThrowStmt struct {
func (ThrowStmt) stmtNode() {} func (ThrowStmt) stmtNode() {}
type GoStmt struct {
Value Expr
}
func (GoStmt) stmtNode() {}
type DeferStmt struct { type DeferStmt struct {
Value Expr Value Expr
} }
@ -198,12 +178,6 @@ type TryCatchStmt struct {
func (TryCatchStmt) stmtNode() {} func (TryCatchStmt) stmtNode() {}
type SelectStmt struct {
Cases []SelectCase
}
func (SelectStmt) stmtNode() {}
type MatchStmt struct { type MatchStmt struct {
Value Expr Value Expr
Cases []MatchCase Cases []MatchCase
@ -218,6 +192,7 @@ type MatchCase struct {
} }
type MatchExpr struct { type MatchExpr struct {
Meta ExprMeta
Value Expr Value Expr
Cases []MatchExprCase Cases []MatchExprCase
} }
@ -230,46 +205,47 @@ type MatchExprCase struct {
Value Expr Value Expr
} }
type SelectCase struct {
Source Expr
Body []Stmt
}
type IdentExpr struct { type IdentExpr struct {
Meta ExprMeta
Name string Name string
} }
func (IdentExpr) exprNode() {} func (IdentExpr) exprNode() {}
type IntExpr struct { type IntExpr struct {
Meta ExprMeta
Value string Value string
} }
func (IntExpr) exprNode() {} func (IntExpr) exprNode() {}
type FloatExpr struct { type FloatExpr struct {
Meta ExprMeta
Value string Value string
} }
func (FloatExpr) exprNode() {} func (FloatExpr) exprNode() {}
type StringExpr struct { type StringExpr struct {
Meta ExprMeta
Value string Value string
} }
func (StringExpr) exprNode() {} func (StringExpr) exprNode() {}
type BoolExpr struct { type BoolExpr struct {
Meta ExprMeta
Value bool Value bool
} }
func (BoolExpr) exprNode() {} func (BoolExpr) exprNode() {}
type NullExpr struct{} type NullExpr struct{ Meta ExprMeta }
func (NullExpr) exprNode() {} func (NullExpr) exprNode() {}
type UnaryExpr struct { type UnaryExpr struct {
Meta ExprMeta
Op string Op string
Value Expr Value Expr
} }
@ -277,6 +253,7 @@ type UnaryExpr struct {
func (UnaryExpr) exprNode() {} func (UnaryExpr) exprNode() {}
type BinaryExpr struct { type BinaryExpr struct {
Meta ExprMeta
Left Expr Left Expr
Op string Op string
Right Expr Right Expr
@ -285,6 +262,7 @@ type BinaryExpr struct {
func (BinaryExpr) exprNode() {} func (BinaryExpr) exprNode() {}
type CallExpr struct { type CallExpr struct {
Meta ExprMeta
Callee Expr Callee Expr
Args []Expr Args []Expr
TypeArgs []string TypeArgs []string
@ -299,6 +277,7 @@ type NamedArg struct {
func (CallExpr) exprNode() {} func (CallExpr) exprNode() {}
type SelectorExpr struct { type SelectorExpr struct {
Meta ExprMeta
Receiver Expr Receiver Expr
Name string Name string
} }
@ -306,21 +285,29 @@ type SelectorExpr struct {
func (SelectorExpr) exprNode() {} func (SelectorExpr) exprNode() {}
type SafeSelectorExpr struct { type SafeSelectorExpr struct {
Meta ExprMeta
Receiver Expr Receiver Expr
Name string Name string
} }
func (SafeSelectorExpr) exprNode() {} func (SafeSelectorExpr) exprNode() {}
type NonNullExpr struct{ Value Expr } type NonNullExpr struct {
Meta ExprMeta
Value Expr
}
func (NonNullExpr) exprNode() {} func (NonNullExpr) exprNode() {}
type TryExpr struct{ Value Expr } type TryExpr struct {
Meta ExprMeta
Value Expr
}
func (TryExpr) exprNode() {} func (TryExpr) exprNode() {}
type IndexExpr struct { type IndexExpr struct {
Meta ExprMeta
Receiver Expr Receiver Expr
Index Expr Index Expr
} }
@ -328,6 +315,7 @@ type IndexExpr struct {
func (IndexExpr) exprNode() {} func (IndexExpr) exprNode() {}
type EnumVariantExpr struct { type EnumVariantExpr struct {
Meta ExprMeta
EnumName, VariantName string EnumName, VariantName string
Values []Expr Values []Expr
} }
@ -335,6 +323,7 @@ type EnumVariantExpr struct {
func (EnumVariantExpr) exprNode() {} func (EnumVariantExpr) exprNode() {}
type LambdaExpr struct { type LambdaExpr struct {
Meta ExprMeta
Params []Param Params []Param
ImplicitIt bool ImplicitIt bool
Body []Stmt Body []Stmt

View file

@ -1096,7 +1096,7 @@ fun main() {
} }
} }
func TestGenerateGoChannelsAndSelect(t *testing.T) { func TestGenerateGoChannelSendAndRead(t *testing.T) {
src := ` src := `
package demo package demo
@ -1105,13 +1105,8 @@ fun writer(ch: Channel<Int>) {
} }
fun main() { fun main() {
val ch = Channel<Int>() val ch = Channel<Int>(1)
runBlocking { writer(ch)
launch { writer(ch) }
select {
ch -> println(it)
}
}
val v = ch.read() val v = ch.read()
println(v) println(v)
} }
@ -1131,12 +1126,8 @@ fun main() {
for _, want := range []string{ for _, want := range []string{
`func writer(ch chan int)`, `func writer(ch chan int)`,
`ch <- 7`, `ch <- 7`,
`ch := make(chan int)`, `ch := make(chan int, 1)`,
`gotlinScope.Launch`,
`writer(ch)`, `writer(ch)`,
`select {`,
`case it := <-ch:`,
`fmt.Println(it)`,
`v := <-ch`, `v := <-ch`,
`fmt.Println(v)`, `fmt.Println(v)`,
} { } {
@ -1146,138 +1137,6 @@ fun main() {
} }
} }
func TestGenerateGoSelectAfterAndEvery(t *testing.T) {
src := `
package demo
fun main() {
select {
after(1000) -> println(it)
every(250) -> println("tick")
}
}
`
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{
`"time"`,
`func gotlinEveryMs(ms int) <-chan time.Time {`,
`case it := <-time.After(time.Duration(1000) * time.Millisecond):`,
`fmt.Println(it)`,
`case <-gotlinEveryMs(250):`,
`fmt.Println("tick")`,
} {
if !strings.Contains(code, want) {
t.Fatalf("generated Go missing %q:\n%s", want, code)
}
}
}
func TestGenerateGoEveryRejectsNonPositiveMilliseconds(t *testing.T) {
cases := []string{
`every(0)`,
`every(-1)`,
}
for _, timerCall := range cases {
src := `
package demo
fun main() {
select {
` + timerCall + ` -> println("x")
}
}
`
prog, err := Parse(src)
if err != nil {
t.Fatalf("parse failed: %v", err)
}
_, err = GenerateGo(prog)
if err == nil {
t.Fatalf("expected go generation error for %s", timerCall)
}
if !strings.Contains(err.Error(), "every(ms) requires ms > 0") {
t.Fatalf("unexpected error for %s: %v", timerCall, err)
}
}
}
func TestGenerateGoAfterEveryArityValidation(t *testing.T) {
cases := []struct {
call string
want string
}{
{`after()`, "after(ms) expects exactly one Int argument"},
{`after(1, 2)`, "after(ms) expects exactly one Int argument"},
{`every()`, "every(ms) expects exactly one Int argument"},
{`every(1, 2)`, "every(ms) expects exactly one Int argument"},
}
for _, tc := range cases {
src := `
package demo
fun main() {
select {
` + tc.call + ` -> println("x")
}
}
`
prog, err := Parse(src)
if err != nil {
t.Fatalf("parse failed: %v", err)
}
_, err = GenerateGo(prog)
if err == nil {
t.Fatalf("expected go generation error for %s", tc.call)
}
if !strings.Contains(err.Error(), tc.want) {
t.Fatalf("unexpected error for %s: %v", tc.call, err)
}
}
}
func TestGenerateGoAfterEveryIntValidation(t *testing.T) {
cases := []struct {
call string
want string
}{
{`after("1s")`, "after(ms) expects an Int argument"},
{`every(true)`, "every(ms) expects an Int argument"},
}
for _, tc := range cases {
src := `
package demo
fun main() {
select {
` + tc.call + ` -> println("x")
}
}
`
prog, err := Parse(src)
if err != nil {
t.Fatalf("parse failed: %v", err)
}
_, err = GenerateGo(prog)
if err == nil {
t.Fatalf("expected go generation error for %s", tc.call)
}
if !strings.Contains(err.Error(), tc.want) {
t.Fatalf("unexpected error for %s: %v", tc.call, err)
}
}
}
func TestGenerateGoNoTimeImportWhenUnused(t *testing.T) { func TestGenerateGoNoTimeImportWhenUnused(t *testing.T) {
src := ` src := `
package demo package demo
@ -1300,36 +1159,6 @@ fun main() {
} }
} }
func TestGenerateGoNoTickerShorthandRewrite(t *testing.T) {
src := `
package demo
import time
fun main() {
val tick = time.NewTicker(time.Second)
select {
tick -> println("x")
}
}
`
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)
if strings.Contains(code, "<-tick.C") {
t.Fatalf("ticker shorthand rewrite should be removed:\n%s", code)
}
if !strings.Contains(code, "case <-tick:") {
t.Fatalf("expected raw source receive in generated code:\n%s", code)
}
}
func TestGenerateGoExplicitErrorsForBunStyleCalls(t *testing.T) { func TestGenerateGoExplicitErrorsForBunStyleCalls(t *testing.T) {
src := ` src := `
package demo package demo

View file

@ -104,14 +104,14 @@ func (g *goGenerator) emitCoroutineSupport() {
func (g *goGenerator) coroutineLambda(lambda LambdaExpr, returnType string) (string, error) { func (g *goGenerator) coroutineLambda(lambda LambdaExpr, returnType string) (string, error) {
var b strings.Builder var b strings.Builder
b.WriteString("func(gotlinScope *GotlinCoroutineScope)") b.WriteString("func(gotlinScope *GotlinCoroutineScope)")
if mapped := mapGoType(returnType); mapped != "" { if mapped := g.goType(returnType); mapped != "" {
b.WriteString(" ") b.WriteString(" ")
b.WriteString(mapped) b.WriteString(mapped)
} }
b.WriteString(" {\n") b.WriteString(" {\n")
sub := goGenerator{indentLevel: 1, needsFmt: g.needsFmt, needsTime: g.needsTime, needsCoroutines: true, functions: g.functions, classes: g.classes, workers: g.workers, enums: g.enums, imports: g.imports, currentFunc: FunctionDecl{ReturnType: returnType}, currentClass: g.currentClass, currentWorker: g.currentWorker, currentCoroutineScope: "gotlinScope", mappings: g.mappings} sub := goGenerator{semantic: g.semantic, indentLevel: 1, needsFmt: g.needsFmt, needsTime: g.needsTime, needsCoroutines: true, currentFunc: FunctionDecl{ReturnType: returnType}, currentClass: g.currentClass, currentCoroutineScope: "gotlinScope"}
sub.scopes = g.cloneScopes() sub.scopes = g.cloneScopes()
sub.typeScopes = g.cloneTypeScopes() sub.semanticScope = g.semanticScope
sub.pushScope() sub.pushScope()
if err := sub.block(lambda.Body); err != nil { if err := sub.block(lambda.Body); err != nil {
return "", err return "", err

View file

@ -44,8 +44,8 @@ fun main() { println(load()) }`)
} }
} }
func TestWorkerAndBareGoAreRemoved(t *testing.T) { func TestLegacyConcurrencySyntaxIsRemoved(t *testing.T) {
for _, source := range []string{`package demo worker Counter { var count = 0 }`, `package demo fun main() { go println("x") }`} { for _, source := range []string{`package demo worker Counter { var count = 0 }`, `package demo fun main() { go println("x") }`, `package demo fun main() { select { channel -> println(it) } }`} {
if _, err := Parse(source); err == nil { if _, err := Parse(source); err == nil {
t.Fatalf("deprecated concurrency syntax parsed: %s", source) t.Fatalf("deprecated concurrency syntax parsed: %s", source)
} }

File diff suppressed because it is too large Load diff

170
internal/lang/go_types.go Normal file
View file

@ -0,0 +1,170 @@
package lang
import (
"go/importer"
gotypes "go/types"
"strings"
)
func importGoPackage(decl ImportDecl) (*gotypes.Package, error) {
path := strings.Trim(decl.Path, `"`)
if !strings.Contains(path, "/") {
path = importPathToGoPath(path)
}
return importer.Default().Import(path)
}
func (semantic *SemanticProgram) goSelectorType(alias, name string) Type {
pkg := semantic.GoPackages[alias]
if pkg == nil {
return UnknownType{}
}
object := pkg.Scope().Lookup(exportedGoName(name))
if object == nil {
return UnknownType{}
}
if function, ok := object.(*gotypes.Func); ok {
if signature, ok := function.Type().(*gotypes.Signature); ok {
return semanticTypeFromGoSignature(signature)
}
}
return semanticTypeFromGo(object.Type())
}
func (semantic *SemanticProgram) goMethodType(receiver Type, name string) Type {
text := receiver.String()
pointer := strings.HasPrefix(text, "*")
text = strings.TrimPrefix(text, "*")
parts := strings.Split(text, ".")
if len(parts) != 2 {
return UnknownType{}
}
for _, pkg := range semantic.GoPackages {
if pkg.Name() != parts[0] {
continue
}
object, ok := pkg.Scope().Lookup(parts[1]).(*gotypes.TypeName)
if !ok {
return UnknownType{}
}
var typ gotypes.Type = object.Type()
if pointer {
typ = gotypes.NewPointer(typ)
}
selection := gotypes.NewMethodSet(typ).Lookup(pkg, exportedGoName(name))
if selection == nil {
return UnknownType{}
}
if signature, ok := selection.Obj().Type().(*gotypes.Signature); ok {
return semanticTypeFromGoSignature(signature)
}
}
return UnknownType{}
}
func (semantic *SemanticProgram) goFieldType(receiver Type, name string) Type {
text := strings.TrimPrefix(receiver.String(), "*")
parts := strings.Split(text, ".")
if len(parts) != 2 {
return UnknownType{}
}
for _, pkg := range semantic.GoPackages {
if pkg.Name() != parts[0] {
continue
}
object, ok := pkg.Scope().Lookup(parts[1]).(*gotypes.TypeName)
if !ok {
return UnknownType{}
}
underlying, ok := object.Type().Underlying().(*gotypes.Struct)
if !ok {
return UnknownType{}
}
for index := 0; index < underlying.NumFields(); index++ {
field := underlying.Field(index)
if field.Name() == exportedGoName(name) {
return semanticTypeFromGo(field.Type())
}
}
}
return UnknownType{}
}
func semanticTypeFromGoSignature(signature *gotypes.Signature) Type {
params := make([]Type, signature.Params().Len())
for index := range params {
params[index] = semanticTypeFromGo(signature.Params().At(index).Type())
}
result := semanticTypeFromGoResults(signature.Results())
return FunctionType{Params: params, Result: result}
}
func semanticTypeFromGoResults(results *gotypes.Tuple) Type {
if results == nil || results.Len() == 0 {
return NamedType{Name: "Unit"}
}
last := semanticTypeFromGo(results.At(results.Len() - 1).Type())
if last.String() == "Error" {
if results.Len() == 1 {
return NullableType{Element: NamedType{Name: "Error"}}
}
value := Type(NamedType{Name: "Unit"})
if results.Len() == 2 {
value = semanticTypeFromGo(results.At(0).Type())
} else if results.Len() > 2 {
return UnknownType{}
}
return GenericType{Base: NamedType{Name: "Result"}, Args: []Type{value, NamedType{Name: "Error"}}}
}
if results.Len() == 1 {
return last
}
return UnknownType{}
}
func semanticTypeFromGo(typ gotypes.Type) Type {
switch value := typ.(type) {
case *gotypes.Basic:
switch value.Kind() {
case gotypes.Bool:
return NamedType{Name: "Boolean"}
case gotypes.String:
return NamedType{Name: "String"}
case gotypes.Int64, gotypes.Uint64:
return NamedType{Name: "Long"}
case gotypes.Float32, gotypes.Float64:
return NamedType{Name: "Double"}
case gotypes.Int, gotypes.Int8, gotypes.Int16, gotypes.Int32, gotypes.Uint, gotypes.Uint8, gotypes.Uint16, gotypes.Uint32:
return NamedType{Name: "Int"}
case gotypes.UnsafePointer:
return NamedType{Name: "Any"}
}
case *gotypes.Pointer:
return GoPointerType{Element: semanticTypeFromGo(value.Elem())}
case *gotypes.Slice:
if basic, ok := value.Elem().(*gotypes.Basic); ok && basic.Kind() == gotypes.Byte {
return NamedType{Name: "ByteSlice"}
}
return GenericType{Base: NamedType{Name: "List"}, Args: []Type{semanticTypeFromGo(value.Elem())}}
case *gotypes.Map:
return GenericType{Base: NamedType{Name: "Map"}, Args: []Type{semanticTypeFromGo(value.Key()), semanticTypeFromGo(value.Elem())}}
case *gotypes.Chan:
return GenericType{Base: NamedType{Name: "Channel"}, Args: []Type{semanticTypeFromGo(value.Elem())}}
case *gotypes.Signature:
return semanticTypeFromGoSignature(value)
case *gotypes.Named:
if value.Obj().Pkg() == nil && value.Obj().Name() == "error" {
return NamedType{Name: "Error"}
}
if value.Obj().Pkg() != nil {
return NamedType{Name: value.Obj().Pkg().Name() + "." + value.Obj().Name()}
}
return NamedType{Name: value.Obj().Name()}
case *gotypes.Interface:
if value.String() == "error" {
return NamedType{Name: "Error"}
}
return NamedType{Name: "Any"}
}
return UnknownType{}
}

765
internal/lang/hir.go Normal file
View file

@ -0,0 +1,765 @@
package lang
import "fmt"
type ExprMeaning int
const (
UnresolvedExpr ExprMeaning = iota
LiteralExpr
LocalReferenceExpr
FunctionReferenceExpr
ClassReferenceExpr
EnumReferenceExpr
ImportReferenceExpr
GotlinCallExpr
GoCallExpr
ClassConstructionExpr
EnumConstructionExpr
FieldAccessExpr
MethodCallExpr
PropagateResultExpr
MatchValueExpr
SQLExpression
MappingExpression
CoroutineExpression
)
type ExprMeta struct{ Semantic *HIRExpr }
type HIRExpr struct {
Type Type
Meaning ExprMeaning
Symbol *Symbol
Node HIRNode
}
type HIRNode interface{ hirNode() }
type HIRLiteral struct{}
func (HIRLiteral) hirNode() {}
type HIRReference struct{ Target *Symbol }
func (HIRReference) hirNode() {}
type HIRGoCall struct {
Callee *HIRExpr
Result Type
}
func (HIRGoCall) hirNode() {}
type HIRGotlinCall struct {
Target *Symbol
Result Type
}
func (HIRGotlinCall) hirNode() {}
type HIRClassConstruction struct{ Class *ClassSymbol }
func (HIRClassConstruction) hirNode() {}
type HIREnumConstruction struct {
EnumName, VariantName string
}
func (HIREnumConstruction) hirNode() {}
type HIRPropagateResult struct {
Value *HIRExpr
Type Type
}
func (HIRPropagateResult) hirNode() {}
type HIRMatch struct {
Value *HIRExpr
Cases []HIRMatchCase
Type Type
}
func (HIRMatch) hirNode() {}
type HIRMatchCase struct {
EnumName, VariantName string
Bindings []string
Value *HIRExpr
}
type HIRCoroutine struct{ Operation string }
func (HIRCoroutine) hirNode() {}
type HIRSQL struct{ Type Type }
func (HIRSQL) hirNode() {}
type HIRMapping struct {
Source Type
Target Type
}
func (HIRMapping) hirNode() {}
type HIRFunction struct {
Symbol *Symbol
Decl *FunctionDecl
Scope *Scope
}
type HIRProgram struct {
Functions []*HIRFunction
Methods []*HIRFunction
}
func exprMeta(expr Expr) *HIRExpr {
switch value := expr.(type) {
case IdentExpr:
return value.Meta.Semantic
case IntExpr:
return value.Meta.Semantic
case FloatExpr:
return value.Meta.Semantic
case StringExpr:
return value.Meta.Semantic
case BoolExpr:
return value.Meta.Semantic
case NullExpr:
return value.Meta.Semantic
case UnaryExpr:
return value.Meta.Semantic
case BinaryExpr:
return value.Meta.Semantic
case CallExpr:
return value.Meta.Semantic
case SelectorExpr:
return value.Meta.Semantic
case SafeSelectorExpr:
return value.Meta.Semantic
case NonNullExpr:
return value.Meta.Semantic
case TryExpr:
return value.Meta.Semantic
case IndexExpr:
return value.Meta.Semantic
case EnumVariantExpr:
return value.Meta.Semantic
case MatchExpr:
return value.Meta.Semantic
case LambdaExpr:
return value.Meta.Semantic
default:
return nil
}
}
func ResolvedType(expr Expr) Type {
if semantic := exprMeta(expr); semantic != nil {
return semantic.Type
}
return UnknownType{}
}
func ResolvedMeaning(expr Expr) ExprMeaning {
if semantic := exprMeta(expr); semantic != nil {
return semantic.Meaning
}
return UnresolvedExpr
}
func withExprMeta(expr Expr, semantic *HIRExpr) Expr {
switch value := expr.(type) {
case IdentExpr:
value.Meta.Semantic = semantic
return value
case IntExpr:
value.Meta.Semantic = semantic
return value
case FloatExpr:
value.Meta.Semantic = semantic
return value
case StringExpr:
value.Meta.Semantic = semantic
return value
case BoolExpr:
value.Meta.Semantic = semantic
return value
case NullExpr:
value.Meta.Semantic = semantic
return value
case UnaryExpr:
value.Meta.Semantic = semantic
return value
case BinaryExpr:
value.Meta.Semantic = semantic
return value
case CallExpr:
value.Meta.Semantic = semantic
return value
case SelectorExpr:
value.Meta.Semantic = semantic
return value
case SafeSelectorExpr:
value.Meta.Semantic = semantic
return value
case NonNullExpr:
value.Meta.Semantic = semantic
return value
case TryExpr:
value.Meta.Semantic = semantic
return value
case IndexExpr:
value.Meta.Semantic = semantic
return value
case EnumVariantExpr:
value.Meta.Semantic = semantic
return value
case MatchExpr:
value.Meta.Semantic = semantic
return value
case LambdaExpr:
value.Meta.Semantic = semantic
return value
default:
return expr
}
}
type semanticResolver struct {
program *SemanticProgram
err error
}
func (resolver *semanticResolver) resolve() error {
resolver.program.HIR = &HIRProgram{}
for index := range resolver.program.Syntax.Functions {
decl := &resolver.program.Syntax.Functions[index]
symbol, _ := resolver.program.Global.Lookup(decl.Name)
scope := NewScope(resolver.program.Global)
for _, param := range decl.Params {
typ, _ := resolver.program.ResolveType(param.Type)
_ = scope.Define(&Symbol{Name: param.Name, Kind: VariableSymbol, Type: typ})
}
result, _ := resolver.program.ResolveType(decl.ReturnType)
resolver.resolveStmts(decl.Body, scope, nil, result)
resolver.program.HIR.Functions = append(resolver.program.HIR.Functions, &HIRFunction{Symbol: symbol, Decl: decl, Scope: scope})
}
for classIndex := range resolver.program.Syntax.Classes {
decl := &resolver.program.Syntax.Classes[classIndex]
class := resolver.program.ClassInfo[decl.Name]
for methodIndex := range decl.Methods {
method := &decl.Methods[methodIndex]
scope := NewScope(resolver.program.Global)
_ = scope.Define(&Symbol{Name: "this", Kind: VariableSymbol, Type: ClassType{Class: class}})
for _, param := range method.Params {
typ, _ := resolver.program.ResolveType(param.Type)
_ = scope.Define(&Symbol{Name: param.Name, Kind: VariableSymbol, Type: typ})
}
result, _ := resolver.program.ResolveType(method.ReturnType)
resolver.resolveStmts(method.Body, scope, class, result)
resolver.program.HIR.Methods = append(resolver.program.HIR.Methods, &HIRFunction{Symbol: class.Methods[method.Name], Decl: method, Scope: scope})
}
}
return resolver.err
}
func (resolver *semanticResolver) resolveStmts(stmts []Stmt, scope *Scope, class *ClassSymbol, returnType Type) {
for index, stmt := range stmts {
switch value := stmt.(type) {
case VarDecl:
expected := Type(UnknownType{})
if value.Type != "" {
expected, _ = resolver.program.ResolveType(value.Type)
}
value.Value, _ = resolver.resolveExpr(value.Value, scope, class, expected)
if isUnknownType(expected) {
expected = exprMeta(value.Value).Type
}
_ = scope.Define(&Symbol{Name: value.Name, Kind: VariableSymbol, Type: expected, Mutable: value.Mutable, Decl: &value})
stmts[index] = value
case MultiVarDecl:
value.Value, _ = resolver.resolveExpr(value.Value, scope, class, UnknownType{})
valueTypes := make([]Type, len(value.Names))
for index := range valueTypes {
valueTypes[index] = UnknownType{}
}
if result, ok := ResolvedType(value.Value).(GenericType); ok && result.Base.String() == "Result" && len(result.Args) == 2 && len(valueTypes) == 2 {
valueTypes[0] = result.Args[0]
valueTypes[1] = NullableType{Element: result.Args[1]}
}
for index, name := range value.Names {
_ = scope.Define(&Symbol{Name: name, Kind: VariableSymbol, Type: valueTypes[index], Mutable: value.Mutable})
}
stmts[index] = value
case AssignStmt:
expected := Type(UnknownType{})
if symbol, ok := scope.Lookup(value.Name); ok {
expected = symbol.Type
}
value.Value, _ = resolver.resolveExpr(value.Value, scope, class, expected)
stmts[index] = value
case AddAssignStmt:
value.Value, _ = resolver.resolveExpr(value.Value, scope, class, UnknownType{})
stmts[index] = value
case MultiAssignStmt:
value.Value, _ = resolver.resolveExpr(value.Value, scope, class, UnknownType{})
stmts[index] = value
case ReturnStmt:
if value.Value != nil {
value.Value, _ = resolver.resolveExpr(value.Value, scope, class, returnType)
stmts[index] = value
}
case ThrowStmt:
value.Value, _ = resolver.resolveExpr(value.Value, scope, class, NamedType{Name: "Error"})
stmts[index] = value
case DeferStmt:
value.Value, _ = resolver.resolveExpr(value.Value, scope, class, NamedType{Name: "Unit"})
stmts[index] = value
case ExprStmt:
value.Value, _ = resolver.resolveExpr(value.Value, scope, class, UnknownType{})
stmts[index] = value
case IfStmt:
value.Cond, _ = resolver.resolveExpr(value.Cond, scope, class, NamedType{Name: "Boolean"})
resolver.resolveStmts(value.Then, NewScope(scope), class, returnType)
resolver.resolveStmts(value.Else, NewScope(scope), class, returnType)
stmts[index] = value
case WhileStmt:
value.Cond, _ = resolver.resolveExpr(value.Cond, scope, class, NamedType{Name: "Boolean"})
resolver.resolveStmts(value.Body, NewScope(scope), class, returnType)
stmts[index] = value
case ForEachStmt:
value.Source, _ = resolver.resolveExpr(value.Source, scope, class, UnknownType{})
bodyScope := NewScope(scope)
element := collectionElement(exprMeta(value.Source).Type)
_ = bodyScope.Define(&Symbol{Name: value.Name, Kind: VariableSymbol, Type: element})
resolver.resolveStmts(value.Body, bodyScope, class, returnType)
stmts[index] = value
case MatchStmt:
value.Value, _ = resolver.resolveExpr(value.Value, scope, class, UnknownType{})
for caseIndex := range value.Cases {
matchCase := &value.Cases[caseIndex]
caseScope := NewScope(scope)
resolver.defineMatchBindings(caseScope, matchCase.EnumName, matchCase.VariantName, matchCase.Bindings)
resolver.resolveStmts(matchCase.Body, caseScope, class, returnType)
}
resolver.validateStatementMatch(value)
stmts[index] = value
case TryCatchStmt:
resolver.resolveStmts(value.TryBody, NewScope(scope), class, returnType)
catchScope := NewScope(scope)
catchType, _ := resolver.program.ResolveType(value.CatchType)
_ = catchScope.Define(&Symbol{Name: value.CatchName, Kind: VariableSymbol, Type: catchType})
resolver.resolveStmts(value.CatchBody, catchScope, class, returnType)
}
}
}
func (resolver *semanticResolver) resolveExpr(expr Expr, scope *Scope, class *ClassSymbol, expected Type) (Expr, Type) {
environment := TypeEnvironment{Scope: scope, Class: class}
switch value := expr.(type) {
case UnaryExpr:
value.Value, _ = resolver.resolveExpr(value.Value, scope, class, UnknownType{})
expr = value
case BinaryExpr:
value.Left, _ = resolver.resolveExpr(value.Left, scope, class, UnknownType{})
value.Right, _ = resolver.resolveExpr(value.Right, scope, class, UnknownType{})
expr = value
case CallExpr:
value.Callee, _ = resolver.resolveExpr(value.Callee, scope, class, UnknownType{})
for index := range value.Args {
argumentExpected := Type(UnknownType{})
if selector, ok := value.Callee.(SelectorExpr); ok {
if receiver, ok := selector.Receiver.(IdentExpr); ok && receiver.Name == "Result" {
if selector.Name == "Err" {
argumentExpected = NamedType{Name: "Error"}
} else if selector.Name == "Ok" {
if result, ok := expected.(GenericType); ok && result.Base.String() == "Result" && len(result.Args) == 2 {
argumentExpected = result.Args[0]
}
}
}
}
value.Args[index], _ = resolver.resolveExpr(value.Args[index], scope, class, argumentExpected)
}
for index := range value.NamedArgs {
value.NamedArgs[index].Value, _ = resolver.resolveExpr(value.NamedArgs[index].Value, scope, class, UnknownType{})
}
expr = value
case SelectorExpr:
value.Receiver, _ = resolver.resolveExpr(value.Receiver, scope, class, UnknownType{})
expr = value
case SafeSelectorExpr:
value.Receiver, _ = resolver.resolveExpr(value.Receiver, scope, class, UnknownType{})
expr = value
case NonNullExpr:
value.Value, _ = resolver.resolveExpr(value.Value, scope, class, UnknownType{})
expr = value
case TryExpr:
value.Value, _ = resolver.resolveExpr(value.Value, scope, class, UnknownType{})
expr = value
case IndexExpr:
value.Receiver, _ = resolver.resolveExpr(value.Receiver, scope, class, UnknownType{})
value.Index, _ = resolver.resolveExpr(value.Index, scope, class, NamedType{Name: "Int"})
expr = value
case EnumVariantExpr:
for index := range value.Values {
value.Values[index], _ = resolver.resolveExpr(value.Values[index], scope, class, UnknownType{})
}
expr = value
case MatchExpr:
value.Value, _ = resolver.resolveExpr(value.Value, scope, class, UnknownType{})
for index := range value.Cases {
matchCase := &value.Cases[index]
caseScope := NewScope(scope)
resolver.defineMatchBindings(caseScope, matchCase.EnumName, matchCase.VariantName, matchCase.Bindings)
matchCase.Value, _ = resolver.resolveExpr(matchCase.Value, caseScope, class, expected)
}
resolver.validateValueMatch(value)
expr = value
case LambdaExpr:
lambdaScope := NewScope(scope)
for _, param := range value.Params {
typ, _ := resolver.program.ResolveType(param.Type)
_ = lambdaScope.Define(&Symbol{Name: param.Name, Kind: VariableSymbol, Type: typ})
}
resolver.resolveStmts(value.Body, lambdaScope, class, functionResult(expected))
expr = value
}
typ := resolver.program.TypeOf(expr, environment)
if isUnknownType(typ) && !isUnknownType(expected) {
typ = expected
}
if binary, ok := expr.(BinaryExpr); ok {
switch binary.Op {
case "==", "!=", "<", "<=", ">", ">=", "&&", "||":
typ = NamedType{Name: "Boolean"}
default:
if meta := exprMeta(binary.Left); meta != nil {
typ = meta.Type
}
}
}
meaning, symbol := resolver.meaning(expr, scope)
resolver.validateEnumExpression(expr, meaning)
if meaning == GoCallExpr && isResultType(expected) {
typ = expected
}
if nullable, ok := typ.(NullableType); ok && typeEqual(nullable.Element, expected) {
typ = expected
}
semantic := &HIRExpr{Type: typ, Meaning: meaning, Symbol: symbol}
semantic.Node = resolver.hirNode(expr, semantic)
return withExprMeta(expr, semantic), typ
}
func (resolver *semanticResolver) hirNode(expr Expr, semantic *HIRExpr) HIRNode {
switch semantic.Meaning {
case LiteralExpr:
return HIRLiteral{}
case LocalReferenceExpr, FunctionReferenceExpr, ClassReferenceExpr, EnumReferenceExpr, ImportReferenceExpr, FieldAccessExpr:
return HIRReference{Target: semantic.Symbol}
case GoCallExpr:
if call, ok := expr.(CallExpr); ok {
return HIRGoCall{Callee: exprMeta(call.Callee), Result: semantic.Type}
}
case GotlinCallExpr, MethodCallExpr:
return HIRGotlinCall{Target: semantic.Symbol, Result: semantic.Type}
case ClassConstructionExpr:
if class := classTypeOf(semantic.Type); class != nil {
return HIRClassConstruction{Class: class}
}
case EnumConstructionExpr:
enumName, variantName := enumExpressionName(expr)
return HIREnumConstruction{EnumName: enumName, VariantName: variantName}
case PropagateResultExpr:
if attempt, ok := expr.(TryExpr); ok {
return HIRPropagateResult{Value: exprMeta(attempt.Value), Type: semantic.Type}
}
case MatchValueExpr:
if match, ok := expr.(MatchExpr); ok {
cases := make([]HIRMatchCase, len(match.Cases))
for index, matchCase := range match.Cases {
cases[index] = HIRMatchCase{EnumName: matchCase.EnumName, VariantName: matchCase.VariantName, Bindings: matchCase.Bindings, Value: exprMeta(matchCase.Value)}
}
return HIRMatch{Value: exprMeta(match.Value), Cases: cases, Type: semantic.Type}
}
case CoroutineExpression:
if call, ok := expr.(CallExpr); ok {
if ident, ok := call.Callee.(IdentExpr); ok {
return HIRCoroutine{Operation: ident.Name}
}
}
case SQLExpression:
return HIRSQL{Type: semantic.Type}
case MappingExpression:
if call, ok := expr.(CallExpr); ok {
if selector, ok := call.Callee.(SelectorExpr); ok {
target := semantic.Type
return HIRMapping{Source: ResolvedType(selector.Receiver), Target: target}
}
}
}
return nil
}
func enumExpressionName(expr Expr) (string, string) {
switch value := expr.(type) {
case CallExpr:
if selector, ok := value.Callee.(SelectorExpr); ok {
if receiver, ok := selector.Receiver.(IdentExpr); ok {
return receiver.Name, selector.Name
}
}
case SelectorExpr:
if receiver, ok := value.Receiver.(IdentExpr); ok {
return receiver.Name, value.Name
}
case EnumVariantExpr:
return value.EnumName, value.VariantName
}
return "", ""
}
func (resolver *semanticResolver) meaning(expr Expr, scope *Scope) (ExprMeaning, *Symbol) {
switch value := expr.(type) {
case IntExpr, FloatExpr, StringExpr, BoolExpr, NullExpr:
return LiteralExpr, nil
case IdentExpr:
if symbol, ok := scope.Lookup(value.Name); ok {
switch symbol.Kind {
case FunctionSymbolKind:
return FunctionReferenceExpr, symbol
case ClassSymbolKind:
return ClassReferenceExpr, symbol
case EnumSymbolKind:
return EnumReferenceExpr, symbol
case ImportSymbolKind:
return ImportReferenceExpr, symbol
default:
return LocalReferenceExpr, symbol
}
}
case CallExpr:
if _, _, _, ok := splitSQLChain(value); ok {
return SQLExpression, nil
}
if _, ok := value.Callee.(SelectorExpr); ok {
selector := value.Callee.(SelectorExpr)
if receiver, ok := selector.Receiver.(IdentExpr); ok && receiver.Name == "json" && selector.Name == "decode" {
return GotlinCallExpr, nil
}
if selector.Name == "mapTo" {
return MappingExpression, nil
}
if receiver, ok := selector.Receiver.(IdentExpr); ok {
if receiver.Name == "Result" {
return EnumConstructionExpr, nil
}
if _, ok := resolver.program.Enums[receiver.Name]; ok {
return EnumConstructionExpr, nil
}
}
if root, ok := selectorRootAlias(selector); ok && resolver.program.Imports[root] {
return GoCallExpr, nil
}
receiverType := ResolvedType(selector.Receiver)
if !isUnknownType(resolver.program.goMethodType(receiverType, selector.Name)) {
return GoCallExpr, nil
}
return MethodCallExpr, nil
}
if ident, ok := value.Callee.(IdentExpr); ok {
if coroutineBuiltins[ident.Name] {
return CoroutineExpression, nil
}
if symbol, ok := resolver.program.Global.Lookup(ident.Name); ok {
if symbol.Kind == ClassSymbolKind {
return ClassConstructionExpr, symbol
}
if symbol.Kind == FunctionSymbolKind {
return GotlinCallExpr, symbol
}
}
}
return GoCallExpr, nil
case SelectorExpr:
if receiver, ok := value.Receiver.(IdentExpr); ok {
if _, ok := resolver.program.Enums[receiver.Name]; ok {
return EnumReferenceExpr, nil
}
}
return FieldAccessExpr, nil
case SafeSelectorExpr:
return FieldAccessExpr, nil
case EnumVariantExpr:
return EnumConstructionExpr, nil
case TryExpr:
return PropagateResultExpr, nil
case MatchExpr:
return MatchValueExpr, nil
}
return UnresolvedExpr, nil
}
func (resolver *semanticResolver) defineMatchBindings(scope *Scope, enumName, variantName string, bindings []string) {
decl, ok := resolver.program.Enums[enumName]
if !ok {
return
}
variant := enumVariant(decl, variantName)
if variant == nil {
return
}
for index, binding := range bindings {
if index >= len(variant.PayloadTypes) {
break
}
typ, _ := resolver.program.ResolveType(variant.PayloadTypes[index])
_ = scope.Define(&Symbol{Name: binding, Kind: VariableSymbol, Type: typ})
}
}
func (resolver *semanticResolver) validateStatementMatch(match MatchStmt) {
patterns := make([]matchPattern, len(match.Cases))
for index, matchCase := range match.Cases {
patterns[index] = matchPattern{enumName: matchCase.EnumName, variantName: matchCase.VariantName, bindings: matchCase.Bindings}
}
resolver.validateMatchPatterns(patterns)
}
func (resolver *semanticResolver) validateValueMatch(match MatchExpr) {
patterns := make([]matchPattern, len(match.Cases))
var result Type = UnknownType{}
for index, matchCase := range match.Cases {
patterns[index] = matchPattern{enumName: matchCase.EnumName, variantName: matchCase.VariantName, bindings: matchCase.Bindings}
armType := ResolvedType(matchCase.Value)
if isUnknownType(result) {
result = armType
} else if !isUnknownType(armType) && !typeEqual(result, armType) {
resolver.fail(fmt.Errorf("match expression arm %s.%s has type %s, expected %s", matchCase.EnumName, matchCase.VariantName, armType.String(), result.String()))
}
}
if isUnknownType(result) {
resolver.fail(fmt.Errorf("match expression result type cannot be inferred"))
}
resolver.validateMatchPatterns(patterns)
}
type matchPattern struct {
enumName, variantName string
bindings []string
}
func (resolver *semanticResolver) validateMatchPatterns(patterns []matchPattern) {
if len(patterns) == 0 {
resolver.fail(fmt.Errorf("match requires at least one case"))
return
}
enumName := patterns[0].enumName
decl, ok := resolver.program.Enums[enumName]
if !ok {
resolver.fail(fmt.Errorf("match value is not a known enum"))
return
}
seen := map[string]bool{}
for _, pattern := range patterns {
if pattern.enumName != enumName {
resolver.fail(fmt.Errorf("match case %s.%s does not match enum %s", pattern.enumName, pattern.variantName, enumName))
return
}
if seen[pattern.variantName] {
resolver.fail(fmt.Errorf("duplicate match case %s.%s", enumName, pattern.variantName))
return
}
seen[pattern.variantName] = true
variant := enumVariant(decl, pattern.variantName)
if variant == nil {
resolver.fail(fmt.Errorf("unknown variant %s.%s", enumName, pattern.variantName))
return
}
if len(pattern.bindings) != len(variant.PayloadTypes) {
resolver.fail(fmt.Errorf("match case %s.%s expects %d bindings", enumName, pattern.variantName, len(variant.PayloadTypes)))
return
}
}
for _, variant := range decl.Variants {
if !seen[variant.Name] {
resolver.fail(fmt.Errorf("non-exhaustive match for %s: missing %s", enumName, variant.Name))
return
}
}
}
func (resolver *semanticResolver) fail(err error) {
if resolver.err == nil {
resolver.err = err
}
}
func (resolver *semanticResolver) validateEnumExpression(expr Expr, meaning ExprMeaning) {
if meaning != EnumConstructionExpr {
return
}
var enumName, variantName string
valueCount := 0
switch value := expr.(type) {
case CallExpr:
selector, ok := value.Callee.(SelectorExpr)
if !ok {
return
}
receiver, ok := selector.Receiver.(IdentExpr)
if !ok {
return
}
enumName, variantName, valueCount = receiver.Name, selector.Name, len(value.Args)
case SelectorExpr:
receiver, ok := value.Receiver.(IdentExpr)
if !ok {
return
}
enumName, variantName = receiver.Name, value.Name
default:
return
}
if enumName == "Result" {
if valueCount != 1 {
resolver.fail(fmt.Errorf("Result.%s expects one value", variantName))
}
return
}
decl, ok := resolver.program.Enums[enumName]
if !ok {
return
}
variant := enumVariant(decl, variantName)
if variant == nil {
resolver.fail(fmt.Errorf("unknown variant %s.%s", enumName, variantName))
return
}
if len(variant.PayloadTypes) != valueCount {
resolver.fail(fmt.Errorf("variant %s.%s expects %d values", enumName, variantName, len(variant.PayloadTypes)))
}
}
func collectionElement(typ Type) Type {
if generic, ok := typ.(GenericType); ok && len(generic.Args) > 0 {
return generic.Args[len(generic.Args)-1]
}
return UnknownType{}
}
func functionResult(typ Type) Type {
if function, ok := typ.(FunctionType); ok {
return function.Result
}
return UnknownType{}
}
func isResultType(typ Type) bool {
result, ok := typ.(GenericType)
return ok && result.Base.String() == "Result" && len(result.Args) == 2
}

View file

@ -13,23 +13,20 @@ type mappingState struct {
} }
func (g *goGenerator) mappingTopLevelTarget(target string) string { func (g *goGenerator) mappingTopLevelTarget(target string) string {
if _, ok := g.classForType(target); ok && !strings.HasPrefix(target, "*") {
return "*" + target
}
return target return target
} }
func (g *goGenerator) ensureMapping(source, target, path string) (string, error) { func (g *goGenerator) ensureMapping(source, target, path string) (string, error) {
key := source + "->" + target key := source + "->" + target
if function, ok := g.mappings.functions[key]; ok { if function, ok := g.semantic.Mappings.functions[key]; ok {
return function, nil return function, nil
} }
if err := g.validateMapping(source, target, path, map[string]bool{}); err != nil { if err := g.validateMapping(source, target, path, map[string]bool{}); err != nil {
return "", err return "", err
} }
function := fmt.Sprintf("gotlinMap%d", len(g.mappings.pairs)+1) function := fmt.Sprintf("gotlinMap%d", len(g.semantic.Mappings.pairs)+1)
g.mappings.functions[key] = function g.semantic.Mappings.functions[key] = function
g.mappings.pairs = append(g.mappings.pairs, mappingPair{source: source, target: target, function: function}) g.semantic.Mappings.pairs = append(g.semantic.Mappings.pairs, mappingPair{source: source, target: target, function: function})
return function, nil return function, nil
} }
@ -83,8 +80,8 @@ func (g *goGenerator) validateMapping(source, target, path string, seen map[stri
} }
return nil return nil
} }
sourceEnum, sourceEnumOK := g.enums[strings.TrimPrefix(source, "*")] sourceEnum, sourceEnumOK := g.semantic.Enums[strings.TrimPrefix(source, "*")]
targetEnum, targetEnumOK := g.enums[strings.TrimPrefix(target, "*")] targetEnum, targetEnumOK := g.semantic.Enums[strings.TrimPrefix(target, "*")]
if sourceEnumOK || targetEnumOK { if sourceEnumOK || targetEnumOK {
if sourceEnumOK && enumIsString(sourceEnum) && target == "String" { if sourceEnumOK && enumIsString(sourceEnum) && target == "String" {
return nil return nil
@ -134,16 +131,16 @@ func mappingFieldName(class ClassDecl, field FieldDecl) string {
} }
func (g *goGenerator) emitMapping(pair mappingPair) error { 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.line(fmt.Sprintf("func %s(source %s) %s {", pair.function, g.goType(pair.source), g.goType(pair.target)))
g.indentLevel++ g.indentLevel++
if sourceEnum, ok := g.enums[strings.TrimPrefix(pair.source, "*")]; ok { if sourceEnum, ok := g.semantic.Enums[strings.TrimPrefix(pair.source, "*")]; ok {
if enumIsString(sourceEnum) && pair.target == "String" { if enumIsString(sourceEnum) && pair.target == "String" {
g.line("return string(source)") g.line("return string(source)")
g.indentLevel-- g.indentLevel--
g.line("}") g.line("}")
return nil return nil
} }
targetEnum := g.enums[strings.TrimPrefix(pair.target, "*")] targetEnum := g.semantic.Enums[strings.TrimPrefix(pair.target, "*")]
if enumIsString(sourceEnum) && enumIsString(targetEnum) { if enumIsString(sourceEnum) && enumIsString(targetEnum) {
g.line("return " + targetEnum.Name + "(source)") g.line("return " + targetEnum.Name + "(source)")
g.indentLevel-- g.indentLevel--
@ -170,7 +167,7 @@ func (g *goGenerator) emitMapping(pair mappingPair) error {
g.indentLevel-- g.indentLevel--
g.line("}") g.line("}")
g.line(`panic("unreachable enum mapping")`) g.line(`panic("unreachable enum mapping")`)
} else if targetEnum, ok := g.enums[strings.TrimPrefix(pair.target, "*")]; ok && pair.source == "String" && enumIsString(targetEnum) { } else if targetEnum, ok := g.semantic.Enums[strings.TrimPrefix(pair.target, "*")]; ok && pair.source == "String" && enumIsString(targetEnum) {
g.line("switch source {") g.line("switch source {")
g.indentLevel++ g.indentLevel++
for _, variant := range targetEnum.Variants { for _, variant := range targetEnum.Variants {
@ -201,18 +198,22 @@ func (g *goGenerator) mappingExpr(expr, source, target, path string) (string, er
if strings.HasSuffix(source, "?") || strings.HasSuffix(target, "?") { if strings.HasSuffix(source, "?") || strings.HasSuffix(target, "?") {
sourceInner := strings.TrimSuffix(source, "?") sourceInner := strings.TrimSuffix(source, "?")
targetInner := strings.TrimSuffix(target, "?") targetInner := strings.TrimSuffix(target, "?")
inner, err := g.mappingExpr("*value", sourceInner, targetInner, path) innerExpr := "*value"
if _, class := g.classForType(sourceInner); class {
innerExpr = "value"
}
inner, err := g.mappingExpr(innerExpr, sourceInner, targetInner, path)
if err != nil { if err != nil {
return "", err return "", err
} }
if strings.HasSuffix(source, "?") { 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 return fmt.Sprintf("func(value %s) %s { if value == nil { return nil }; return %s }(%s)", g.goType(source), g.goType(target), inner, expr), nil
} }
inner, err = g.mappingExpr("value", sourceInner, targetInner, path) inner, err = g.mappingExpr("value", sourceInner, targetInner, path)
if err != nil { if err != nil {
return "", err return "", err
} }
return fmt.Sprintf("func(value %s) %s { mapped := %s; return &mapped }(%s)", mapGoType(source), mapGoType(target), inner, expr), nil return fmt.Sprintf("func(value %s) %s { return %s }(%s)", g.goType(source), g.goType(target), inner, expr), nil
} }
if sourceBase, sourceArgs, ok := parseGenericType(source); ok { if sourceBase, sourceArgs, ok := parseGenericType(source); ok {
_, targetArgs, _ := parseGenericType(target) _, targetArgs, _ := parseGenericType(target)
@ -221,14 +222,14 @@ func (g *goGenerator) mappingExpr(expr, source, target, path string) (string, er
if err != nil { if err != nil {
return "", err 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 return fmt.Sprintf("func(values %s) %s { var result %s; for _, item := range values { result = append(result, %s) }; return result }(%s)", g.goType(source), g.goType(target), g.goType(target), item, expr), nil
} }
if sourceBase == "Map" || sourceBase == "MutableMap" { if sourceBase == "Map" || sourceBase == "MutableMap" {
value, err := g.mappingExpr("item", sourceArgs[1], targetArgs[1], path+"[]") value, err := g.mappingExpr("item", sourceArgs[1], targetArgs[1], path+"[]")
if err != nil { if err != nil {
return "", err 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 return fmt.Sprintf("func(values %s) %s { result := make(%s, len(values)); for key, item := range values { result[key] = %s }; return result }(%s)", g.goType(source), g.goType(target), g.goType(target), value, expr), nil
} }
} }
if sourceClass, ok := g.classForType(source); ok { if sourceClass, ok := g.classForType(source); ok {
@ -242,23 +243,16 @@ func (g *goGenerator) mappingExpr(expr, source, target, path string) (string, er
} }
fields = append(fields, mappingFieldName(targetClass, targetField)+": "+mapped) fields = append(fields, mappingFieldName(targetClass, targetField)+": "+mapped)
} }
literal := targetClass.Name + "{" + strings.Join(fields, ", ") + "}" return "&" + targetClass.Name + "{" + strings.Join(fields, ", ") + "}", nil
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 { if _, ok := g.semantic.Enums[strings.TrimPrefix(source, "*")]; ok {
function, err := g.ensureMapping(source, target, path) function, err := g.ensureMapping(source, target, path)
if err != nil { if err != nil {
return "", err return "", err
} }
return function + "(" + expr + ")", nil return function + "(" + expr + ")", nil
} }
if targetEnum, ok := g.enums[strings.TrimPrefix(target, "*")]; ok && source == "String" && enumIsString(targetEnum) { if targetEnum, ok := g.semantic.Enums[strings.TrimPrefix(target, "*")]; ok && source == "String" && enumIsString(targetEnum) {
function, err := g.ensureMapping(source, target, path) function, err := g.ensureMapping(source, target, path)
if err != nil { if err != nil {
return "", err return "", err

View file

@ -10,9 +10,9 @@ func TestGenerateRecursiveClassMapping(t *testing.T) {
package demo package demo
data class SourceAddress(var city: String) data class SourceAddress(var city: String)
data class TargetAddress(var city: String) data class TargetAddress(var city: String)
data class Source(var id: String, var address: *SourceAddress) data class Source(var id: String, var address: SourceAddress)
data class Target(var address: *TargetAddress, var id: String) data class Target(var address: TargetAddress, var id: String)
fun convert(source: *Source): *Target { return source.mapTo<Target>() } fun convert(source: Source): Target { return source.mapTo<Target>() }
`) `)
if err != nil { if err != nil {
t.Fatal(err) t.Fatal(err)
@ -22,7 +22,7 @@ fun convert(source: *Source): *Target { return source.mapTo<Target>() }
t.Fatal(err) t.Fatal(err)
} }
code := string(out) code := string(out)
for _, want := range []string{"gotlinMap1(source)", "Address:", "Id: value.Id", "TargetAddress{City: value.City}"} { for _, want := range []string{"gotlinMap1(source)", "Address:", "Id: source.Id", "&TargetAddress{City: source.Address.City}"} {
if !strings.Contains(code, want) { if !strings.Contains(code, want) {
t.Fatalf("missing %q:\n%s", want, code) t.Fatalf("missing %q:\n%s", want, code)
} }
@ -34,8 +34,8 @@ func TestGenerateListAndMapMapping(t *testing.T) {
package demo package demo
data class Source(var id: String) data class Source(var id: String)
data class Target(var id: String) data class Target(var id: String)
fun list(values: List<*Source>): List<*Target> { return values.mapTo<List<*Target>>() } 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>>() } fun mapping(values: Map<String, Source>): Map<String, Target> { return values.mapTo<Map<String, Target>>() }
`) `)
if err != nil { if err != nil {
t.Fatal(err) t.Fatal(err)
@ -77,9 +77,9 @@ func TestMappingReportsNestedFieldPath(t *testing.T) {
package demo package demo
data class SourceAddress(var zip: String) data class SourceAddress(var zip: String)
data class TargetAddress(var zip: Int) data class TargetAddress(var zip: Int)
data class Source(var address: *SourceAddress) data class Source(var address: SourceAddress)
data class Target(var address: *TargetAddress) data class Target(var address: TargetAddress)
fun convert(value: *Source): *Target { return value.mapTo<Target>() } fun convert(value: Source): Target { return value.mapTo<Target>() }
`) `)
if err != nil { if err != nil {
t.Fatal(err) t.Fatal(err)
@ -92,7 +92,7 @@ fun convert(value: *Source): *Target { return value.mapTo<Target>() }
func TestMappingRejectsMissingFieldAndEnumVariant(t *testing.T) { func TestMappingRejectsMissingFieldAndEnumVariant(t *testing.T) {
for _, source := range []string{ 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 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>() }`, `package demo enum Source { Ready, Failed } enum Target { Ready } fun convert(value: Source): Target { return value.mapTo<Target>() }`,
} { } {
prog, err := Parse(source) prog, err := Parse(source)
@ -110,8 +110,8 @@ func TestMapStringBackedEnumToAndFromString(t *testing.T) {
enum Status { PendingReservation, Initiated } enum Status { PendingReservation, Initiated }
data class Domain(var status: Status) data class Domain(var status: Status)
data class Row(var status: String) data class Row(var status: String)
fun toRow(value: *Domain): *Row { return value.mapTo<Row>() } fun toRow(value: Domain): Row { return value.mapTo<Row>() }
fun toDomain(value: *Row): *Domain { return value.mapTo<Domain>() }`) fun toDomain(value: Row): Domain { return value.mapTo<Domain>() }`)
if err != nil { if err != nil {
t.Fatal(err) t.Fatal(err)
} }
@ -130,10 +130,10 @@ func TestMapToInfersExpectedTargetType(t *testing.T) {
prog, err := Parse(`package demo prog, err := Parse(`package demo
data class Source(var id: String) data class Source(var id: String)
data class Target(var id: String) data class Target(var id: String)
data class Wrapper(var target: *Target) data class Wrapper(var target: Target)
fun returned(value: *Source): *Target { return value.mapTo() } fun returned(value: Source): Target { return value.mapTo() }
fun wrapped(value: *Source): *Wrapper { return Wrapper(value.mapTo()) } fun wrapped(value: Source): Wrapper { return Wrapper(value.mapTo()) }
fun local(value: *Source): *Target { val target: *Target = value.mapTo(); return target }`) fun local(value: Source): Target { val target: Target = value.mapTo(); return target }`)
if err != nil { if err != nil {
t.Fatal(err) t.Fatal(err)
} }
@ -148,7 +148,7 @@ fun local(value: *Source): *Target { val target: *Target = value.mapTo(); return
func TestMapToWithoutTargetContextHasHelpfulError(t *testing.T) { func TestMapToWithoutTargetContextHasHelpfulError(t *testing.T) {
prog, err := Parse(`package demo prog, err := Parse(`package demo
fun convert(value: *Source) { val target = value.mapTo() }`) fun convert(value: Source) { val target = value.mapTo() }`)
if err != nil { if err != nil {
t.Fatal(err) t.Fatal(err)
} }

View file

@ -13,7 +13,7 @@ class Repository {
fun healthy(): Boolean { return true } fun healthy(): Boolean { return true }
} }
class Service(val repository: *Repository) { class Service(val repository: Repository) {
fun healthy(): Boolean { return repository.healthy() } fun healthy(): Boolean { return repository.healthy() }
} }
@ -34,9 +34,9 @@ func TestInternalMethodReturnTypeFlowsThroughSelectors(t *testing.T) {
prog, err := Parse(` prog, err := Parse(`
package demo package demo
class Transaction { fun commit(): Boolean { return true } } class Transaction { fun commit(): Boolean { return true } }
class Result(val transaction: *Transaction) class Result(val transaction: Transaction)
class Repository { fun begin(): *Result { return Result(Transaction()) } } class Repository { fun begin(): Result { return Result(Transaction()) } }
class Service(val repository: *Repository) { class Service(val repository: Repository) {
fun run(): Boolean { return repository.begin().transaction.commit() } fun run(): Boolean { return repository.begin().transaction.commit() }
} }
`) `)

View file

@ -8,8 +8,8 @@ import (
func TestSafeAccessAndNonNullAssertion(t *testing.T) { func TestSafeAccessAndNonNullAssertion(t *testing.T) {
prog, err := Parse(`package demo prog, err := Parse(`package demo
data class User(var email: String) data class User(var email: String)
fun safe(user: *User?): String? { return user?.email } fun safe(user: User?): String? { return user?.email }
fun required(user: *User?): String { return user!!.email }`) fun required(user: User?): String { return user!!.email }`)
if err != nil { if err != nil {
t.Fatal(err) t.Fatal(err)
} }
@ -26,7 +26,7 @@ fun required(user: *User?): String { return user!!.email }`)
func TestRejectNullableDereference(t *testing.T) { func TestRejectNullableDereference(t *testing.T) {
prog, err := Parse(`package demo prog, err := Parse(`package demo
fun unsafe(user: *User?): String { return user.email }`) fun unsafe(user: User?): String { return user.email }`)
if err != nil { if err != nil {
t.Fatal(err) t.Fatal(err)
} }
@ -38,8 +38,8 @@ fun unsafe(user: *User?): String { return user.email }`)
func TestRejectNullForNonNullableTypes(t *testing.T) { func TestRejectNullForNonNullableTypes(t *testing.T) {
for _, source := range []string{ for _, source := range []string{
`package demo data class User(var email: String) fun main() { val user: *User = null }`, `package demo data class User(var email: String) fun main() { val user: User = null }`,
`package demo data class User(var email: String) fun use(user: *User) {} fun main() { use(null) }`, `package demo data class User(var email: String) fun use(user: User) {} fun main() { use(null) }`,
`package demo fun name(): String { return null }`, `package demo fun name(): String { return null }`,
`package demo fun main() { val value = null }`, `package demo fun main() { val value = null }`,
} { } {
@ -56,8 +56,8 @@ func TestRejectNullForNonNullableTypes(t *testing.T) {
func TestNullableSmartCasts(t *testing.T) { func TestNullableSmartCasts(t *testing.T) {
prog, err := Parse(`package demo prog, err := Parse(`package demo
data class User(var email: String) data class User(var email: String)
fun guarded(user: *User?): String { if (user == null) { return "missing" }; return user.email } fun guarded(user: User?): String { if (user == null) { return "missing" }; return user.email }
fun branched(user: *User?): String { if (user != null) { return user.email } else { return "missing" } }`) fun branched(user: User?): String { if (user != null) { return user.email } else { return "missing" } }`)
if err != nil { if err != nil {
t.Fatal(err) t.Fatal(err)
} }

View file

@ -131,12 +131,6 @@ func (p *parser) parseProgram() (*Program, error) {
default: default:
return nil, fmt.Errorf("unsupported annotation %q", annotation.lexeme) return nil, fmt.Errorf("unsupported annotation %q", annotation.lexeme)
} }
case p.check(tokenWorker):
decl, err := p.parseWorker()
if err != nil {
return nil, err
}
prog.Workers = append(prog.Workers, decl)
case p.check(tokenFun) || p.check(tokenSuspend): case p.check(tokenFun) || p.check(tokenSuspend):
fn, err := p.parseFunction() fn, err := p.parseFunction()
if err != nil { if err != nil {
@ -338,68 +332,6 @@ func (p *parser) parseClass() (ClassDecl, error) {
return ClassDecl{Name: name.lexeme, Data: data, 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) {
if _, err := p.expect(tokenWorker, "expected 'worker'"); err != nil {
return WorkerDecl{}, err
}
name, err := p.expect(tokenIdent, "expected worker name")
if err != nil {
return WorkerDecl{}, err
}
if _, err := p.expect(tokenLBrace, "expected '{'"); err != nil {
return WorkerDecl{}, err
}
var fields []WorkerFieldDecl
var methods []FunctionDecl
for !p.check(tokenRBrace) && !p.check(tokenEOF) {
switch {
case p.match(tokenVal), p.match(tokenVar):
mutable := p.tokens[p.pos-1].kind == tokenVar
fieldName, err := p.expect(tokenIdent, "expected field name")
if err != nil {
return WorkerDecl{}, err
}
typ := ""
if p.match(tokenColon) {
parsed, err := p.parseTypeRef()
if err != nil {
return WorkerDecl{}, err
}
typ = parsed
}
if _, err := p.expect(tokenAssign, "expected '=' after field declaration"); err != nil {
return WorkerDecl{}, err
}
value, err := p.parseExpr(0)
if err != nil {
return WorkerDecl{}, err
}
fields = append(fields, WorkerFieldDecl{
Mutable: mutable,
Name: fieldName.lexeme,
Type: typ,
Value: value,
})
p.match(tokenSemicolon)
case p.check(tokenFun):
method, err := p.parseFunction()
if err != nil {
return WorkerDecl{}, err
}
methods = append(methods, method)
p.match(tokenSemicolon)
default:
tok := p.peek()
return WorkerDecl{}, fmt.Errorf("expected worker member at %d, found %q", tok.pos, tok.lexeme)
}
}
if _, err := p.expect(tokenRBrace, "expected '}'"); err != nil {
return WorkerDecl{}, err
}
return WorkerDecl{Name: name.lexeme, Fields: fields, Methods: methods}, nil
}
func (p *parser) parseClassParents() ([]string, error) { func (p *parser) parseClassParents() ([]string, error) {
if !p.match(tokenColon) { if !p.match(tokenColon) {
return nil, nil return nil, nil
@ -626,6 +558,9 @@ func (p *parser) parseStmt() (Stmt, error) {
if p.check(tokenIdent) && p.peek().lexeme == "go" { if p.check(tokenIdent) && p.peek().lexeme == "go" {
return nil, fmt.Errorf("bare go is removed; use launch inside a coroutine scope") return nil, fmt.Errorf("bare go is removed; use launch inside a coroutine scope")
} }
if p.check(tokenIdent) && p.peek().lexeme == "select" {
return nil, fmt.Errorf("select is removed; use structured coroutines and channel read()")
}
switch { switch {
case p.match(tokenVal): case p.match(tokenVal):
return p.parseVarDecl(false) return p.parseVarDecl(false)
@ -646,25 +581,6 @@ func (p *parser) parseStmt() (Stmt, error) {
return nil, err return nil, err
} }
return ThrowStmt{Value: expr}, nil return ThrowStmt{Value: expr}, nil
case p.match(tokenGo):
var expr Expr
var err error
if p.check(tokenLBrace) {
body, err := p.parseBlock()
if err != nil {
return nil, err
}
expr = CallExpr{Callee: LambdaExpr{Body: body}}
} else {
expr, err = p.parseExpr(0)
if err != nil {
return nil, err
}
}
if _, ok := expr.(CallExpr); !ok {
return nil, fmt.Errorf("'go' expects a function call expression")
}
return GoStmt{Value: expr}, nil
case p.match(tokenDefer): case p.match(tokenDefer):
expr, err := p.parseExpr(0) expr, err := p.parseExpr(0)
if err != nil { if err != nil {
@ -680,8 +596,6 @@ func (p *parser) parseStmt() (Stmt, error) {
return p.parseWhile() return p.parseWhile()
case p.match(tokenFor): case p.match(tokenFor):
return p.parseForEach() return p.parseForEach()
case p.match(tokenSelect):
return p.parseSelect()
case p.match(tokenMatch): case p.match(tokenMatch):
return p.parseMatch() return p.parseMatch()
case p.match(tokenTry): case p.match(tokenTry):
@ -855,41 +769,6 @@ func (p *parser) parseMatchExpr() (Expr, error) {
return MatchExpr{Value: value, Cases: cases}, nil return MatchExpr{Value: value, Cases: cases}, nil
} }
func (p *parser) parseSelect() (Stmt, error) {
if _, err := p.expect(tokenLBrace, "expected '{' after select"); err != nil {
return nil, err
}
var cases []SelectCase
for !p.check(tokenRBrace) && !p.check(tokenEOF) {
source, err := p.parseExpr(0)
if err != nil {
return nil, err
}
if _, err := p.expect(tokenArrow, "expected '->' in select case"); err != nil {
return nil, err
}
var body []Stmt
if p.check(tokenLBrace) {
body, err = p.parseBlock()
if err != nil {
return nil, err
}
} else {
stmt, err := p.parseStmt()
if err != nil {
return nil, err
}
body = []Stmt{stmt}
p.match(tokenSemicolon)
}
cases = append(cases, SelectCase{Source: source, Body: body})
}
if _, err := p.expect(tokenRBrace, "expected '}' after select"); err != nil {
return nil, err
}
return SelectStmt{Cases: cases}, nil
}
func (p *parser) parseTryCatch() (Stmt, error) { func (p *parser) parseTryCatch() (Stmt, error) {
tryBody, err := p.parseBlock() tryBody, err := p.parseBlock()
if err != nil { if err != nil {

View file

@ -1,197 +0,0 @@
package lang
import "strings"
func normalizeClassReferences(program *Program) {
classes := map[string]bool{}
for _, class := range program.Classes {
classes[class.Name] = true
}
normalize := func(value string) string { return normalizeReferenceType(value, classes) }
for i := range program.Interfaces {
for j := range program.Interfaces[i].Methods {
normalizeSignature(&program.Interfaces[i].Methods[j], normalize)
}
}
for i := range program.Enums {
for j := range program.Enums[i].Variants {
for k := range program.Enums[i].Variants[j].PayloadTypes {
program.Enums[i].Variants[j].PayloadTypes[k] = normalize(program.Enums[i].Variants[j].PayloadTypes[k])
}
}
}
for i := range program.Classes {
for j := range program.Classes[i].Fields {
program.Classes[i].Fields[j].Type = normalize(program.Classes[i].Fields[j].Type)
}
for j := range program.Classes[i].Methods {
normalizeFunction(&program.Classes[i].Methods[j], normalize)
}
}
for i := range program.Workers {
for j := range program.Workers[i].Fields {
program.Workers[i].Fields[j].Type = normalize(program.Workers[i].Fields[j].Type)
normalizeExprTypes(program.Workers[i].Fields[j].Value, normalize)
}
for j := range program.Workers[i].Methods {
normalizeFunction(&program.Workers[i].Methods[j], normalize)
}
}
for i := range program.Functions {
normalizeFunction(&program.Functions[i], normalize)
}
}
func normalizeSignature(signature *FunctionSignature, normalize func(string) string) {
for i := range signature.Params {
signature.Params[i].Type = normalize(signature.Params[i].Type)
}
signature.ReturnType = normalize(signature.ReturnType)
}
func normalizeFunction(function *FunctionDecl, normalize func(string) string) {
for i := range function.Params {
function.Params[i].Type = normalize(function.Params[i].Type)
}
function.ReturnType = normalize(function.ReturnType)
normalizeStmtTypes(function.Body, normalize)
}
func normalizeStmtTypes(statements []Stmt, normalize func(string) string) {
for index, statement := range statements {
switch value := statement.(type) {
case VarDecl:
value.Type = normalize(value.Type)
normalizeExprTypes(value.Value, normalize)
statements[index] = value
case MultiVarDecl:
normalizeExprTypes(value.Value, normalize)
case AssignStmt:
normalizeExprTypes(value.Value, normalize)
case AddAssignStmt:
normalizeExprTypes(value.Value, normalize)
case MultiAssignStmt:
normalizeExprTypes(value.Value, normalize)
case ReturnStmt:
if value.Value != nil {
normalizeExprTypes(value.Value, normalize)
}
case ThrowStmt:
normalizeExprTypes(value.Value, normalize)
case GoStmt:
normalizeExprTypes(value.Value, normalize)
case DeferStmt:
normalizeExprTypes(value.Value, normalize)
case ExprStmt:
normalizeExprTypes(value.Value, normalize)
case IfStmt:
normalizeExprTypes(value.Cond, normalize)
normalizeStmtTypes(value.Then, normalize)
normalizeStmtTypes(value.Else, normalize)
case WhileStmt:
normalizeExprTypes(value.Cond, normalize)
normalizeStmtTypes(value.Body, normalize)
case ForEachStmt:
normalizeExprTypes(value.Source, normalize)
normalizeStmtTypes(value.Body, normalize)
case SelectStmt:
for _, c := range value.Cases {
normalizeExprTypes(c.Source, normalize)
normalizeStmtTypes(c.Body, normalize)
}
case MatchStmt:
normalizeExprTypes(value.Value, normalize)
for _, c := range value.Cases {
normalizeStmtTypes(c.Body, normalize)
}
case TryCatchStmt:
value.CatchType = normalize(value.CatchType)
normalizeStmtTypes(value.TryBody, normalize)
normalizeStmtTypes(value.CatchBody, normalize)
statements[index] = value
}
}
}
func normalizeExprTypes(expression Expr, normalize func(string) string) {
switch value := expression.(type) {
case UnaryExpr:
normalizeExprTypes(value.Value, normalize)
case NonNullExpr:
normalizeExprTypes(value.Value, normalize)
case BinaryExpr:
normalizeExprTypes(value.Left, normalize)
normalizeExprTypes(value.Right, normalize)
case SelectorExpr:
normalizeExprTypes(value.Receiver, normalize)
case SafeSelectorExpr:
normalizeExprTypes(value.Receiver, normalize)
case IndexExpr:
normalizeExprTypes(value.Receiver, normalize)
normalizeExprTypes(value.Index, normalize)
case EnumVariantExpr:
for _, item := range value.Values {
normalizeExprTypes(item, normalize)
}
case MatchExpr:
normalizeExprTypes(value.Value, normalize)
for _, item := range value.Cases {
normalizeExprTypes(item.Value, normalize)
}
case LambdaExpr:
for i := range value.Params {
value.Params[i].Type = normalize(value.Params[i].Type)
}
normalizeStmtTypes(value.Body, normalize)
case CallExpr:
skipTypeArgs := false
if selector, ok := value.Callee.(SelectorExpr); ok {
if root, ok := selector.Receiver.(IdentExpr); ok && root.Name == "sql" {
skipTypeArgs = true
}
}
if !skipTypeArgs {
for i := range value.TypeArgs {
value.TypeArgs[i] = normalize(value.TypeArgs[i])
}
}
normalizeExprTypes(value.Callee, normalize)
for _, item := range value.Args {
normalizeExprTypes(item, normalize)
}
for _, item := range value.NamedArgs {
normalizeExprTypes(item.Value, normalize)
}
}
}
func normalizeReferenceType(value string, classes map[string]bool) string {
if value == "" {
return value
}
nullable := strings.HasSuffix(value, "?")
if nullable {
value = strings.TrimSuffix(value, "?")
}
explicitPointer := strings.HasPrefix(value, "*")
if explicitPointer {
value = strings.TrimPrefix(value, "*")
}
if params, result, ok := parseFunctionType(value); ok {
for i := range params {
params[i] = normalizeReferenceType(params[i], classes)
}
value = "(" + strings.Join(params, ", ") + ") -> " + normalizeReferenceType(result, classes)
} else if base, args, ok := parseGenericType(value); ok {
for i := range args {
args[i] = normalizeReferenceType(args[i], classes)
}
value = base + "<" + strings.Join(args, ", ") + ">"
} else if classes[value] || explicitPointer {
value = "*" + value
}
if nullable {
value += "?"
}
return value
}

View file

@ -26,18 +26,15 @@ fun optional(): Repository? { return null }`)
} }
} }
func TestExplicitClassPointerRemainsCompatible(t *testing.T) { func TestExplicitClassPointerIsRejected(t *testing.T) {
prog, err := Parse(`package demo prog, err := Parse(`package demo
class Repository class Repository
fun use(repository: *Repository): *Repository { return repository }`) fun use(repository: *Repository): *Repository { return repository }`)
if err != nil { if err != nil {
t.Fatal(err) t.Fatal(err)
} }
out, err := GenerateGo(prog) _, err = GenerateGo(prog)
if err != nil { if err == nil || !strings.Contains(err.Error(), "already reference-valued") {
t.Fatal(err) t.Fatalf("unexpected explicit pointer result: %v", err)
}
if strings.Contains(string(out), "**Repository") || !strings.Contains(string(out), "repository *Repository") {
t.Fatalf("unexpected explicit pointer output:\n%s", out)
} }
} }

View file

@ -2,31 +2,17 @@ package lang
import "fmt" import "fmt"
func validateMutability(program *Program) error { func validateMutability(semantic *SemanticProgram) error {
checker := mutabilityChecker{} checker := mutabilityChecker{semantic: semantic}
for _, fn := range program.Functions { for index := range semantic.Syntax.Functions {
if err := checker.checkFunction(fn, nil, nil); err != nil { if err := checker.checkFunction(&semantic.Syntax.Functions[index], nil); err != nil {
return err return err
} }
} }
for _, class := range program.Classes { for index := range semantic.Syntax.Classes {
fields := make(map[string]bool, len(class.Fields)) class := semantic.ClassInfo[semantic.Syntax.Classes[index].Name]
for _, field := range class.Fields { for methodIndex := range semantic.Syntax.Classes[index].Methods {
fields[field.Name] = field.Mutable if err := checker.checkFunction(&semantic.Syntax.Classes[index].Methods[methodIndex], class); err != nil {
}
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 err
} }
} }
@ -35,138 +21,123 @@ func validateMutability(program *Program) error {
} }
type mutabilityChecker struct { type mutabilityChecker struct {
scopes []map[string]bool semantic *SemanticProgram
classFields map[string]bool scope *Scope
workerFields map[string]bool class *ClassSymbol
} }
func (c *mutabilityChecker) checkFunction(fn FunctionDecl, classFields map[string]bool, workerFields map[string]bool) error { func (checker *mutabilityChecker) checkFunction(function *FunctionDecl, class *ClassSymbol) error {
c.scopes = nil checker.class = class
c.classFields = classFields checker.scope = NewScope(checker.semantic.Global)
c.workerFields = workerFields if class != nil {
c.pushScope() _ = checker.scope.Define(&Symbol{Name: "this", Kind: VariableSymbol, Type: ClassType{Class: class}})
defer c.popScope()
for _, param := range fn.Params {
c.define(param.Name, false)
} }
return c.checkStmts(fn.Body) for _, param := range function.Params {
typ, _ := checker.semantic.ResolveType(param.Type)
_ = checker.scope.Define(&Symbol{Name: param.Name, Kind: VariableSymbol, Type: typ})
}
return checker.checkStmts(function.Body)
} }
func (c *mutabilityChecker) checkStmts(stmts []Stmt) error { func (checker *mutabilityChecker) checkStmts(statements []Stmt) error {
for _, stmt := range stmts { for _, statement := range statements {
switch s := stmt.(type) { switch value := statement.(type) {
case VarDecl: case VarDecl:
if err := c.checkExpr(s.Value); err != nil { if err := checker.checkExpr(value.Value); err != nil {
return err return err
} }
c.define(s.Name, s.Mutable) typ, _ := checker.semantic.ResolveType(value.Type)
_ = checker.scope.Define(&Symbol{Name: value.Name, Kind: VariableSymbol, Type: typ, Mutable: value.Mutable})
case MultiVarDecl: case MultiVarDecl:
if err := c.checkExpr(s.Value); err != nil { if err := checker.checkExpr(value.Value); err != nil {
return err return err
} }
for _, name := range s.Names { for _, name := range value.Names {
c.define(name, s.Mutable) _ = checker.scope.Define(&Symbol{Name: name, Kind: VariableSymbol, Type: UnknownType{}, Mutable: value.Mutable})
} }
case AssignStmt: case AssignStmt:
if err := c.requireMutable(s.Name, s.Pos); err != nil { if err := checker.requireMutable(value.Name, value.Pos); err != nil {
return err return err
} }
if err := c.checkExpr(s.Value); err != nil { if err := checker.checkExpr(value.Value); err != nil {
return err return err
} }
case AddAssignStmt: case AddAssignStmt:
if err := c.requireMutable(s.Name, s.Pos); err != nil { if err := checker.requireMutable(value.Name, value.Pos); err != nil {
return err return err
} }
if err := c.checkExpr(s.Value); err != nil { if err := checker.checkExpr(value.Value); err != nil {
return err return err
} }
case MultiAssignStmt: case MultiAssignStmt:
for i, name := range s.Names { for index, name := range value.Names {
pos := 0 position := 0
if i < len(s.Positions) { if index < len(value.Positions) {
pos = s.Positions[i] position = value.Positions[index]
} }
if err := c.requireMutable(name, pos); err != nil { if err := checker.requireMutable(name, position); err != nil {
return err return err
} }
} }
if err := c.checkExpr(s.Value); err != nil { if err := checker.checkExpr(value.Value); err != nil {
return err return err
} }
case ReturnStmt: case ReturnStmt:
if s.Value != nil { if value.Value != nil {
if err := c.checkExpr(s.Value); err != nil { if err := checker.checkExpr(value.Value); err != nil {
return err return err
} }
} }
case ThrowStmt: case ThrowStmt:
if err := c.checkExpr(s.Value); err != nil { if err := checker.checkExpr(value.Value); err != nil {
return err
}
case GoStmt:
if err := c.checkExpr(s.Value); err != nil {
return err return err
} }
case DeferStmt: case DeferStmt:
if err := c.checkExpr(s.Value); err != nil { if err := checker.checkExpr(value.Value); err != nil {
return err return err
} }
case ExprStmt: case ExprStmt:
if err := c.checkExpr(s.Value); err != nil { if err := checker.checkExpr(value.Value); err != nil {
return err return err
} }
case IfStmt: case IfStmt:
if err := c.checkExpr(s.Cond); err != nil { if err := checker.checkExpr(value.Cond); err != nil {
return err return err
} }
if err := c.checkBlock(s.Then, nil); err != nil { if err := checker.checkBlock(value.Then, nil); err != nil {
return err return err
} }
if err := c.checkBlock(s.Else, nil); err != nil { if err := checker.checkBlock(value.Else, nil); err != nil {
return err return err
} }
case WhileStmt: case WhileStmt:
if err := c.checkExpr(s.Cond); err != nil { if err := checker.checkExpr(value.Cond); err != nil {
return err return err
} }
if err := c.checkBlock(s.Body, nil); err != nil { if err := checker.checkBlock(value.Body, nil); err != nil {
return err return err
} }
case ForEachStmt: case ForEachStmt:
if err := c.checkExpr(s.Source); err != nil { if err := checker.checkExpr(value.Source); err != nil {
return err return err
} }
if err := c.checkBlock(s.Body, map[string]bool{s.Name: false}); err != nil { if err := checker.checkBlock(value.Body, []string{value.Name}); err != nil {
return err 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: case MatchStmt:
if err := c.checkExpr(s.Value); err != nil { if err := checker.checkExpr(value.Value); err != nil {
return err return err
} }
for _, matchCase := range s.Cases { for _, matchCase := range value.Cases {
bindings := map[string]bool{} if err := checker.checkBlock(matchCase.Body, matchCase.Bindings); err != nil {
for _, binding := range matchCase.Bindings {
bindings[binding] = false
}
if err := c.checkBlock(matchCase.Body, bindings); err != nil {
return err return err
} }
} }
case TryCatchStmt: case TryCatchStmt:
if err := c.checkBlock(s.TryBody, nil); err != nil { if err := checker.checkBlock(value.TryBody, nil); err != nil {
return err return err
} }
if err := c.checkBlock(s.CatchBody, map[string]bool{s.CatchName: false}); err != nil { if err := checker.checkBlock(value.CatchBody, []string{value.CatchName}); err != nil {
return err return err
} }
} }
@ -174,124 +145,104 @@ func (c *mutabilityChecker) checkStmts(stmts []Stmt) error {
return nil return nil
} }
func (c *mutabilityChecker) checkBlock(stmts []Stmt, bindings map[string]bool) error { func (checker *mutabilityChecker) checkBlock(statements []Stmt, bindings []string) error {
c.pushScope() previous := checker.scope
defer c.popScope() checker.scope = NewScope(previous)
for name, mutable := range bindings { defer func() { checker.scope = previous }()
c.define(name, mutable) for _, name := range bindings {
_ = checker.scope.Define(&Symbol{Name: name, Kind: VariableSymbol, Type: UnknownType{}})
} }
return c.checkStmts(stmts) return checker.checkStmts(statements)
} }
func (c *mutabilityChecker) checkExpr(expr Expr) error { func (checker *mutabilityChecker) checkExpr(expr Expr) error {
switch e := expr.(type) { switch value := expr.(type) {
case UnaryExpr: case UnaryExpr:
return c.checkExpr(e.Value) return checker.checkExpr(value.Value)
case BinaryExpr: case BinaryExpr:
if err := c.checkExpr(e.Left); err != nil { if err := checker.checkExpr(value.Left); err != nil {
return err return err
} }
return c.checkExpr(e.Right) return checker.checkExpr(value.Right)
case CallExpr: case CallExpr:
if err := c.checkExpr(e.Callee); err != nil { if err := checker.checkExpr(value.Callee); err != nil {
return err return err
} }
for _, arg := range e.Args { for _, argument := range value.Args {
if err := c.checkExpr(arg); err != nil { if err := checker.checkExpr(argument); err != nil {
return err return err
} }
} }
for _, arg := range e.NamedArgs { for _, argument := range value.NamedArgs {
if err := c.checkExpr(arg.Value); err != nil { if err := checker.checkExpr(argument.Value); err != nil {
return err return err
} }
} }
case SelectorExpr: case SelectorExpr:
return c.checkExpr(e.Receiver) return checker.checkExpr(value.Receiver)
case SafeSelectorExpr: case SafeSelectorExpr:
return c.checkExpr(e.Receiver) return checker.checkExpr(value.Receiver)
case NonNullExpr: case NonNullExpr:
return c.checkExpr(e.Value) return checker.checkExpr(value.Value)
case TryExpr: case TryExpr:
return c.checkExpr(e.Value) return checker.checkExpr(value.Value)
case IndexExpr: case IndexExpr:
if err := c.checkExpr(e.Receiver); err != nil { if err := checker.checkExpr(value.Receiver); err != nil {
return err return err
} }
return c.checkExpr(e.Index) return checker.checkExpr(value.Index)
case EnumVariantExpr: case EnumVariantExpr:
for _, value := range e.Values { for _, item := range value.Values {
if err := c.checkExpr(value); err != nil { if err := checker.checkExpr(item); err != nil {
return err return err
} }
} }
case MatchExpr: case MatchExpr:
if err := c.checkExpr(e.Value); err != nil { if err := checker.checkExpr(value.Value); err != nil {
return err return err
} }
for _, matchCase := range e.Cases { for _, matchCase := range value.Cases {
if err := c.checkExpr(matchCase.Value); err != nil { if err := checker.checkExpr(matchCase.Value); err != nil {
return err return err
} }
} }
case LambdaExpr: case LambdaExpr:
bindings := map[string]bool{} previous := checker.scope
if e.ImplicitIt { checker.scope = NewScope(previous)
bindings["it"] = false defer func() { checker.scope = previous }()
if value.ImplicitIt {
_ = checker.scope.Define(&Symbol{Name: "it", Kind: VariableSymbol, Type: UnknownType{}})
} }
for _, param := range e.Params { for _, param := range value.Params {
bindings[param.Name] = false typ, _ := checker.semantic.ResolveType(param.Type)
_ = checker.scope.Define(&Symbol{Name: param.Name, Kind: VariableSymbol, Type: typ})
} }
return c.checkBlock(e.Body, bindings) return checker.checkStmts(value.Body)
} }
return nil return nil
} }
func (c *mutabilityChecker) pushScope() { func (checker *mutabilityChecker) requireMutable(name string, position int) error {
c.scopes = append(c.scopes, map[string]bool{}) if symbol, ok := checker.scope.Lookup(name); ok && symbol.Kind == VariableSymbol {
} if symbol.Mutable {
return nil
func (c *mutabilityChecker) popScope() { }
if len(c.scopes) == 0 { return immutableAssignmentError(name, position)
return
} }
c.scopes = c.scopes[:len(c.scopes)-1] if checker.class != nil {
} if field, ok := checker.class.Fields[name]; ok {
if field.Mutable {
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 nil
} }
return immutableAssignmentError(name, pos) return immutableAssignmentError(name, position)
} }
} }
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 return nil
} }
func immutableAssignmentError(name string, pos int) error { func immutableAssignmentError(name string, position int) error {
if pos > 0 { if position > 0 {
return fmt.Errorf("cannot reassign immutable name %s at %d", name, pos) return fmt.Errorf("cannot reassign immutable name %s at %d", name, position)
} }
return fmt.Errorf("cannot reassign immutable name %s", name) return fmt.Errorf("cannot reassign immutable name %s", name)
} }

View file

@ -171,9 +171,9 @@ func sqlChainResultType(expr Expr) (string, bool) {
} }
switch terminal { switch terminal {
case "fetch": case "fetch":
return "Result<List<*" + resultType + ">, Error>", true return "Result<List<" + resultType + ">, Error>", true
case "single": case "single":
return "Result<*" + resultType + ", Error>", true return "Result<" + resultType + ", Error>", true
case "iterator": case "iterator":
return "Result<GotlinSQLIterator<" + resultType + ">, Error>", true return "Result<GotlinSQLIterator<" + resultType + ">, Error>", true
default: default:
@ -215,7 +215,7 @@ func (g *goGenerator) sqlClass(root CallExpr, operation string) (ClassDecl, erro
if len(root.TypeArgs) != 1 { if len(root.TypeArgs) != 1 {
return ClassDecl{}, fmt.Errorf("sql.%s expects exactly one row type", operation) return ClassDecl{}, fmt.Errorf("sql.%s expects exactly one row type", operation)
} }
class, ok := g.classes[root.TypeArgs[0]] class, ok := g.semantic.Classes[root.TypeArgs[0]]
if !ok { if !ok {
return ClassDecl{}, fmt.Errorf("SQL row class %q does not exist", root.TypeArgs[0]) return ClassDecl{}, fmt.Errorf("SQL row class %q does not exist", root.TypeArgs[0])
} }
@ -561,7 +561,7 @@ func (g *goGenerator) sqlProjection(call CallExpr, rowClass ClassDecl, method st
if !ok { if !ok {
return ClassDecl{}, nil, fmt.Errorf("%s projection must construct a local data class", method) return ClassDecl{}, nil, fmt.Errorf("%s projection must construct a local data class", method)
} }
projection, ok := g.classes[callee.Name] projection, ok := g.semantic.Classes[callee.Name]
if !ok || !projection.Data { if !ok || !projection.Data {
return ClassDecl{}, nil, fmt.Errorf("%s projection type %s must be a data class", method, callee.Name) return ClassDecl{}, nil, fmt.Errorf("%s projection type %s must be a data class", method, callee.Name)
} }
@ -1190,10 +1190,6 @@ func stmtsMatch(stmts []Stmt, match func(Expr) bool) bool {
if exprMatches(s.Value, match) { if exprMatches(s.Value, match) {
return true return true
} }
case GoStmt:
if exprMatches(s.Value, match) {
return true
}
case DeferStmt: case DeferStmt:
if exprMatches(s.Value, match) { if exprMatches(s.Value, match) {
return true return true
@ -1214,12 +1210,6 @@ func stmtsMatch(stmts []Stmt, match func(Expr) bool) bool {
if exprMatches(s.Source, match) || stmtsMatch(s.Body, match) { if exprMatches(s.Source, match) || stmtsMatch(s.Body, match) {
return true return true
} }
case SelectStmt:
for _, c := range s.Cases {
if exprMatches(c.Source, match) || stmtsMatch(c.Body, match) {
return true
}
}
case TryCatchStmt: case TryCatchStmt:
if stmtsMatch(s.TryBody, match) || stmtsMatch(s.CatchBody, match) { if stmtsMatch(s.TryBody, match) || stmtsMatch(s.CatchBody, match) {
return true return true
@ -1242,17 +1232,5 @@ func programExprMatches(program *Program, match func(Expr) bool) bool {
} }
} }
} }
for _, worker := range program.Workers {
for _, field := range worker.Fields {
if exprMatches(field.Value, match) {
return true
}
}
for _, fn := range worker.Methods {
if stmtsMatch(fn.Body, match) {
return true
}
}
}
return false return false
} }

View file

@ -34,7 +34,11 @@ func TestSQLGeneratedNullableAndLockingMetadata(t *testing.T) {
if row.Fields[3].Type != "time.Time?" { if row.Fields[3].Type != "time.Time?" {
t.Fatalf("publishedAt type = %q, want time.Time?", row.Fields[3].Type) t.Fatalf("publishedAt type = %q, want time.Time?", row.Fields[3].Type)
} }
if got := mapGoType(row.Fields[3].Type); got != "*time.Time" { typ, err := ParseType(row.Fields[3].Type)
if err != nil {
t.Fatal(err)
}
if got := renderGoType(typ); got != "*time.Time" {
t.Fatalf("mapped nullable timestamp = %q, want *time.Time", got) t.Fatalf("mapped nullable timestamp = %q, want *time.Time", got)
} }
@ -64,7 +68,7 @@ func TestSQLTypedProjectionExecution(t *testing.T) {
import context import context
import pgxpool "github.com/jackc/pgx/v5/pgxpool" import pgxpool "github.com/jackc/pgx/v5/pgxpool"
fun events(pool: *pgxpool.Pool, ctx: context.Context): List<*EventProjection> { fun events(pool: *pgxpool.Pool, ctx: context.Context): List<EventProjection> {
return sql.from<EventRow>() return sql.from<EventRow>()
.select { row -> EventProjection(row.id, row.payload) } .select { row -> EventProjection(row.id, row.payload) }
.orderBy { it.createdAt } .orderBy { it.createdAt }
@ -94,7 +98,7 @@ func TestSQLInsertOmitsGeneratedAndReturnsProjection(t *testing.T) {
import context import context
import pgxpool "github.com/jackc/pgx/v5/pgxpool" import pgxpool "github.com/jackc/pgx/v5/pgxpool"
fun create(row: EventRow, pool: *pgxpool.Pool, ctx: context.Context): *EventProjection { fun create(row: EventRow, pool: *pgxpool.Pool, ctx: context.Context): EventProjection {
return sql.insert<EventRow>(row) return sql.insert<EventRow>(row)
.returning { value -> EventProjection(value.id, value.payload) } .returning { value -> EventProjection(value.id, value.payload) }
.single(pool, ctx).unwrap() .single(pool, ctx).unwrap()
@ -126,7 +130,7 @@ func TestSQLTypedUpdateAndReturning(t *testing.T) {
import context import context
import pgxpool "github.com/jackc/pgx/v5/pgxpool" import pgxpool "github.com/jackc/pgx/v5/pgxpool"
fun claim(id: String, payload: String, pool: *pgxpool.Pool, ctx: context.Context): *EventProjection { fun claim(id: String, payload: String, pool: *pgxpool.Pool, ctx: context.Context): EventProjection {
return sql.update<EventRow>() return sql.update<EventRow>()
.set { row -> .set { row ->
set(row.payload, payload) set(row.payload, payload)
@ -154,7 +158,7 @@ func TestSQLDeleteReturningFullRow(t *testing.T) {
import context import context
import pgxpool "github.com/jackc/pgx/v5/pgxpool" import pgxpool "github.com/jackc/pgx/v5/pgxpool"
fun remove(id: String, pool: *pgxpool.Pool, ctx: context.Context): *EventRow { fun remove(id: String, pool: *pgxpool.Pool, ctx: context.Context): EventRow {
return sql.delete<EventRow>() return sql.delete<EventRow>()
.where { it.id == id } .where { it.id == id }
.returning { it } .returning { it }
@ -245,7 +249,7 @@ func TestRejectExpandedInvalidSQLQueries(t *testing.T) {
}, },
{ {
name: "write execution without returning", name: "write execution without returning",
src: eventSQLSource + `fun query(pool: Any, ctx: Any): *EventRow { return sql.delete<EventRow>().single(pool, ctx).unwrap() }`, src: eventSQLSource + `fun query(pool: Any, ctx: Any): EventRow { return sql.delete<EventRow>().single(pool, ctx).unwrap() }`,
want: "requires returning()", want: "requires returning()",
}, },
{ {

View file

@ -94,7 +94,7 @@ func TestGenerateSQLFetchTerminal(t *testing.T) {
import context import context
import pgxpool "github.com/jackc/pgx/v5/pgxpool" import pgxpool "github.com/jackc/pgx/v5/pgxpool"
fun accounts(pool: *pgxpool.Pool, ctx: context.Context, customerId: String): List<*AccountRow> { fun accounts(pool: *pgxpool.Pool, ctx: context.Context, customerId: String): List<AccountRow> {
return sql.from<AccountRow>() return sql.from<AccountRow>()
.where { it.customerId == customerId } .where { it.customerId == customerId }
.fetch(pool, ctx).unwrap() .fetch(pool, ctx).unwrap()
@ -304,12 +304,12 @@ fun query(row: OtherRow): GotlinSQLQuery { return sql.insert<AccountRow>(row).on
}, },
{ {
name: "fetch missing arguments", name: "fetch missing arguments",
src: accountRowSource + `fun query(): List<*AccountRow> { return sql.from<AccountRow>().fetch().unwrap() }`, src: accountRowSource + `fun query(): List<AccountRow> { return sql.from<AccountRow>().fetch().unwrap() }`,
want: "fetch() expects exactly pool and ctx positional arguments", want: "fetch() expects exactly pool and ctx positional arguments",
}, },
{ {
name: "single extra argument", name: "single extra argument",
src: accountRowSource + `fun query(pool: Any, ctx: Any): *AccountRow { return sql.from<AccountRow>().single(pool, ctx, ctx).unwrap() }`, src: accountRowSource + `fun query(pool: Any, ctx: Any): AccountRow { return sql.from<AccountRow>().single(pool, ctx, ctx).unwrap() }`,
want: "single() expects exactly pool and ctx positional arguments", want: "single() expects exactly pool and ctx positional arguments",
}, },
{ {
@ -319,22 +319,22 @@ fun query(row: OtherRow): GotlinSQLQuery { return sql.insert<AccountRow>(row).on
}, },
{ {
name: "fetch type arguments", name: "fetch type arguments",
src: accountRowSource + `fun query(pool: Any, ctx: Any): List<*AccountRow> { return sql.from<AccountRow>().fetch<String>(pool, ctx) }`, 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", want: "fetch() expects exactly pool and ctx positional arguments",
}, },
{ {
name: "fetch invalid pool type", name: "fetch invalid pool type",
src: accountRowSource + `fun query(pool: String, ctx: Any): List<*AccountRow> { return sql.from<AccountRow>().fetch(pool, ctx).unwrap() }`, src: accountRowSource + `fun query(pool: String, ctx: Any): List<AccountRow> { return sql.from<AccountRow>().fetch(pool, ctx).unwrap() }`,
want: "pool argument has non-query type String", want: "pool argument has non-query type String",
}, },
{ {
name: "single invalid context type", name: "single invalid context type",
src: accountRowSource + `fun query(pool: Any, ctx: Int): *AccountRow { return sql.from<AccountRow>().single(pool, ctx).unwrap() }`, src: accountRowSource + `fun query(pool: Any, ctx: Int): AccountRow { return sql.from<AccountRow>().single(pool, ctx).unwrap() }`,
want: "ctx argument has non-context type Int", want: "ctx argument has non-context type Int",
}, },
{ {
name: "insert execution terminal", name: "insert execution terminal",
src: accountRowSource + `fun query(row: AccountRow, pool: Any, ctx: Any): List<*AccountRow> { return sql.insert<AccountRow>(row).fetch(pool, ctx).unwrap() }`, src: accountRowSource + `fun query(row: AccountRow, pool: Any, ctx: Any): List<AccountRow> { return sql.insert<AccountRow>(row).fetch(pool, ctx).unwrap() }`,
want: "fetch() is only supported for sql.from", want: "fetch() is only supported for sql.from",
}, },
} }

215
internal/lang/symbols.go Normal file
View file

@ -0,0 +1,215 @@
package lang
import (
"fmt"
gotypes "go/types"
)
type SymbolKind int
const (
VariableSymbol SymbolKind = iota
FunctionSymbolKind
ClassSymbolKind
EnumSymbolKind
ImportSymbolKind
)
type Symbol struct {
Name string
Kind SymbolKind
Type Type
Mutable bool
Decl any
}
type Scope struct {
Parent *Scope
Symbols map[string]*Symbol
}
func NewScope(parent *Scope) *Scope { return &Scope{Parent: parent, Symbols: map[string]*Symbol{}} }
func (scope *Scope) Define(symbol *Symbol) error {
if _, exists := scope.Symbols[symbol.Name]; exists {
return fmt.Errorf("duplicate symbol %s", symbol.Name)
}
scope.Symbols[symbol.Name] = symbol
return nil
}
func (scope *Scope) Lookup(name string) (*Symbol, bool) {
for current := scope; current != nil; current = current.Parent {
if symbol, ok := current.Symbols[name]; ok {
return symbol, true
}
}
return nil, false
}
type ClassSymbol struct {
Name string
Decl *ClassDecl
Fields map[string]*Symbol
Methods map[string]*Symbol
}
type SemanticProgram struct {
Syntax *Program
Global *Scope
Classes map[string]ClassDecl
ClassInfo map[string]*ClassSymbol
Functions map[string]FunctionDecl
Enums map[string]EnumDecl
Imports map[string]bool
GoPackages map[string]*gotypes.Package
HIR *HIRProgram
Mappings *mappingState
}
func Analyze(program *Program) (*SemanticProgram, error) {
semantic := &SemanticProgram{
Syntax: program,
Global: NewScope(nil),
Classes: map[string]ClassDecl{},
ClassInfo: map[string]*ClassSymbol{},
Functions: map[string]FunctionDecl{},
Enums: map[string]EnumDecl{},
Imports: map[string]bool{},
GoPackages: map[string]*gotypes.Package{},
Mappings: &mappingState{functions: map[string]string{}},
}
for index := range program.Classes {
decl := &program.Classes[index]
class := &ClassSymbol{Name: decl.Name, Decl: decl, Fields: map[string]*Symbol{}, Methods: map[string]*Symbol{}}
semantic.Classes[decl.Name] = *decl
semantic.ClassInfo[decl.Name] = class
if err := semantic.Global.Define(&Symbol{Name: decl.Name, Kind: ClassSymbolKind, Type: ClassType{Class: class}, Decl: decl}); err != nil {
return nil, err
}
}
for index := range program.Enums {
decl := &program.Enums[index]
semantic.Enums[decl.Name] = *decl
if err := semantic.Global.Define(&Symbol{Name: decl.Name, Kind: EnumSymbolKind, Type: NamedType{Name: decl.Name}, Decl: decl}); err != nil {
return nil, err
}
}
for _, imported := range program.Imports {
name := imported.Alias
if name == "" {
name = defaultImportAlias(imported)
}
semantic.Imports[name] = true
if importedPackage, err := importGoPackage(imported); err == nil {
semantic.GoPackages[name] = importedPackage
}
if existing, ok := semantic.Global.Lookup(name); ok && existing.Kind == ImportSymbolKind {
continue
}
if err := semantic.Global.Define(&Symbol{Name: name, Kind: ImportSymbolKind, Type: NamedType{Name: name}, Decl: imported}); err != nil {
return nil, err
}
}
for index := range program.Functions {
decl := &program.Functions[index]
semantic.Functions[decl.Name] = *decl
typ, err := semantic.functionType(decl.Params, decl.ReturnType)
if err != nil {
return nil, fmt.Errorf("function %s: %w", decl.Name, err)
}
if err := semantic.Global.Define(&Symbol{Name: decl.Name, Kind: FunctionSymbolKind, Type: typ, Decl: decl}); err != nil {
return nil, err
}
}
for index := range program.Classes {
decl := &program.Classes[index]
class := semantic.ClassInfo[decl.Name]
for fieldIndex := range decl.Fields {
field := &decl.Fields[fieldIndex]
typ, err := semantic.ResolveType(field.Type)
if err != nil {
return nil, fmt.Errorf("field %s.%s: %w", decl.Name, field.Name, err)
}
class.Fields[field.Name] = &Symbol{Name: field.Name, Kind: VariableSymbol, Type: typ, Mutable: field.Mutable, Decl: field}
}
for methodIndex := range decl.Methods {
method := &decl.Methods[methodIndex]
typ, err := semantic.functionType(method.Params, method.ReturnType)
if err != nil {
return nil, fmt.Errorf("method %s.%s: %w", decl.Name, method.Name, err)
}
class.Methods[method.Name] = &Symbol{Name: method.Name, Kind: FunctionSymbolKind, Type: typ, Decl: method}
}
}
if err := validateMutability(semantic); err != nil {
return nil, err
}
resolver := semanticResolver{program: semantic}
if err := resolver.resolve(); err != nil {
return nil, err
}
return semantic, nil
}
func (semantic *SemanticProgram) ResolveType(text string) (Type, error) {
typ, err := ParseType(text)
if err != nil {
return nil, err
}
resolved := resolveClassTypes(typ, semantic.ClassInfo)
if err := validateNoClassPointer(resolved); err != nil {
return nil, err
}
return resolved, nil
}
func validateNoClassPointer(typ Type) error {
switch value := typ.(type) {
case GoPointerType:
if class, ok := value.Element.(ClassType); ok {
return fmt.Errorf("Gotlin class %s is already reference-valued; remove '*'", class.Class.Name)
}
return validateNoClassPointer(value.Element)
case NullableType:
return validateNoClassPointer(value.Element)
case GenericType:
for _, arg := range value.Args {
if err := validateNoClassPointer(arg); err != nil {
return err
}
}
case FunctionType:
for _, param := range value.Params {
if err := validateNoClassPointer(param); err != nil {
return err
}
}
return validateNoClassPointer(value.Result)
}
return nil
}
func (semantic *SemanticProgram) functionType(params []Param, result string) (Type, error) {
paramTypes := make([]Type, len(params))
for index, param := range params {
typ, err := semantic.ResolveType(param.Type)
if err != nil {
return nil, err
}
paramTypes[index] = typ
}
resultType, err := semantic.ResolveType(result)
if err != nil {
return nil, err
}
return FunctionType{Params: paramTypes, Result: resultType}, nil
}
func (semantic *SemanticProgram) GoType(text string) string {
typ, err := semantic.ResolveType(text)
if err != nil {
return ""
}
return renderGoType(typ)
}

View file

@ -15,7 +15,6 @@ const (
tokenPackage tokenKind = "PACKAGE" tokenPackage tokenKind = "PACKAGE"
tokenClass tokenKind = "CLASS" tokenClass tokenKind = "CLASS"
tokenData tokenKind = "DATA" tokenData tokenKind = "DATA"
tokenWorker tokenKind = "WORKER"
tokenInterface tokenKind = "INTERFACE" tokenInterface tokenKind = "INTERFACE"
tokenEnum tokenKind = "ENUM" tokenEnum tokenKind = "ENUM"
tokenMatch tokenKind = "MATCH" tokenMatch tokenKind = "MATCH"
@ -30,9 +29,7 @@ const (
tokenWhile tokenKind = "WHILE" tokenWhile tokenKind = "WHILE"
tokenFor tokenKind = "FOR" tokenFor tokenKind = "FOR"
tokenIn tokenKind = "IN" tokenIn tokenKind = "IN"
tokenSelect tokenKind = "SELECT"
tokenReturn tokenKind = "RETURN" tokenReturn tokenKind = "RETURN"
tokenGo tokenKind = "GO"
tokenDefer tokenKind = "DEFER" tokenDefer tokenKind = "DEFER"
tokenTry tokenKind = "TRY" tokenTry tokenKind = "TRY"
tokenCatch tokenKind = "CATCH" tokenCatch tokenKind = "CATCH"
@ -90,7 +87,6 @@ var keywords = map[string]tokenKind{
"while": tokenWhile, "while": tokenWhile,
"for": tokenFor, "for": tokenFor,
"in": tokenIn, "in": tokenIn,
"select": tokenSelect,
"return": tokenReturn, "return": tokenReturn,
"defer": tokenDefer, "defer": tokenDefer,
"try": tokenTry, "try": tokenTry,

View file

@ -0,0 +1,191 @@
package lang
type TypeEnvironment struct {
Scope *Scope
Class *ClassSymbol
}
func (semantic *SemanticProgram) TypeOf(expr Expr, environment TypeEnvironment) Type {
resolve := func(text string) Type {
typ, err := semantic.ResolveType(text)
if err != nil {
return UnknownType{}
}
return typ
}
switch value := expr.(type) {
case IdentExpr:
if environment.Scope != nil {
if symbol, ok := environment.Scope.Lookup(value.Name); ok {
return symbol.Type
}
}
if environment.Class != nil {
if value.Name == "this" {
return ClassType{Class: environment.Class}
}
if field, ok := environment.Class.Fields[value.Name]; ok {
return field.Type
}
}
if symbol, ok := semantic.Global.Lookup(value.Name); ok {
return symbol.Type
}
case IntExpr:
return NamedType{Name: "Int"}
case FloatExpr:
return NamedType{Name: "Double"}
case StringExpr:
return NamedType{Name: "String"}
case BoolExpr:
return NamedType{Name: "Boolean"}
case NullExpr:
return NullableType{Element: UnknownType{}}
case CallExpr:
if selector, ok := value.Callee.(SelectorExpr); ok {
if receiver, ok := selector.Receiver.(IdentExpr); ok {
if semantic.Imports[receiver.Name] {
if function, ok := semantic.goSelectorType(receiver.Name, selector.Name).(FunctionType); ok {
return function.Result
}
}
if _, ok := semantic.Enums[receiver.Name]; ok {
return NamedType{Name: receiver.Name}
}
if receiver.Name == "json" && selector.Name == "decode" && len(value.TypeArgs) == 1 {
decoded := resolve(value.TypeArgs[0])
return GenericType{Base: NamedType{Name: "Result"}, Args: []Type{decoded, NamedType{Name: "Error"}}}
}
}
if selector.Name == "unwrap" || selector.Name == "unwrapOr" {
if result, ok := semantic.TypeOf(selector.Receiver, environment).(GenericType); ok && result.Base.String() == "Result" && len(result.Args) == 2 {
return result.Args[0]
}
}
if selector.Name == "mapTo" && len(value.TypeArgs) == 1 {
return resolve(value.TypeArgs[0])
}
receiverType := semantic.TypeOf(selector.Receiver, environment)
if iterator, ok := receiverType.(GenericType); ok && iterator.Base.String() == "GotlinSQLIterator" && len(iterator.Args) == 1 {
switch selector.Name {
case "next":
return NamedType{Name: "Boolean"}
case "value":
return iterator.Args[0]
case "err":
return NamedType{Name: "Error"}
case "close":
return NamedType{Name: "Unit"}
}
}
if class := classTypeOf(receiverType); class != nil {
if method, ok := class.Methods[selector.Name]; ok {
if function, ok := method.Type.(FunctionType); ok {
return function.Result
}
}
}
if function, ok := semantic.goMethodType(receiverType, selector.Name).(FunctionType); ok {
return function.Result
}
}
if ident, ok := value.Callee.(IdentExpr); ok {
if ident.Name == "keys" && len(value.Args) == 1 {
if mapping, ok := semantic.TypeOf(value.Args[0], environment).(GenericType); ok && (mapping.Base.String() == "Map" || mapping.Base.String() == "MutableMap") && len(mapping.Args) == 2 {
return GenericType{Base: NamedType{Name: "List"}, Args: []Type{mapping.Args[0]}}
}
}
if function, ok := semantic.Functions[ident.Name]; ok {
return resolve(function.ReturnType)
}
if class, ok := semantic.ClassInfo[ident.Name]; ok {
return ClassType{Class: class}
}
}
if sqlType, ok := sqlChainResultType(value); ok {
return resolve(sqlType)
}
if len(value.NamedArgs) > 0 {
if selector, ok := value.Callee.(SelectorExpr); ok {
if receiver, ok := selector.Receiver.(IdentExpr); ok {
return NamedType{Name: receiver.Name + "." + selector.Name}
}
}
}
case SelectorExpr:
if receiver, ok := value.Receiver.(IdentExpr); ok {
if semantic.Imports[receiver.Name] {
return semantic.goSelectorType(receiver.Name, value.Name)
}
if _, ok := semantic.Enums[receiver.Name]; ok {
return NamedType{Name: receiver.Name}
}
}
if class := classTypeOf(semantic.TypeOf(value.Receiver, environment)); class != nil {
if field, ok := class.Fields[value.Name]; ok {
return field.Type
}
}
if field := semantic.goFieldType(semantic.TypeOf(value.Receiver, environment), value.Name); !isUnknownType(field) {
return field
}
case IndexExpr:
if generic, ok := semantic.TypeOf(value.Receiver, environment).(GenericType); ok && len(generic.Args) > 0 {
return generic.Args[len(generic.Args)-1]
}
case MatchExpr:
for _, matchCase := range value.Cases {
if typ := semantic.TypeOf(matchCase.Value, environment); !isUnknownType(typ) {
return typ
}
}
case EnumVariantExpr:
return resolve(value.EnumName)
case NonNullExpr:
if nullable, ok := semantic.TypeOf(value.Value, environment).(NullableType); ok {
return nullable.Element
}
case TryExpr:
if result, ok := semantic.TypeOf(value.Value, environment).(GenericType); ok && result.Base.String() == "Result" && len(result.Args) == 2 {
return result.Args[0]
}
case SafeSelectorExpr:
receiver := semantic.TypeOf(value.Receiver, environment)
if nullable, ok := receiver.(NullableType); ok {
receiver = nullable.Element
}
if class := classTypeOf(receiver); class != nil {
if field, ok := class.Fields[value.Name]; ok {
return nullableSemanticType(field.Type)
}
if method, ok := class.Methods[value.Name]; ok {
if function, ok := method.Type.(FunctionType); ok {
return nullableSemanticType(function.Result)
}
}
}
}
return UnknownType{}
}
func classTypeOf(typ Type) *ClassSymbol {
switch value := typ.(type) {
case ClassType:
return value.Class
case NullableType:
return classTypeOf(value.Element)
}
return nil
}
func nullableSemanticType(typ Type) Type {
if _, ok := typ.(NullableType); ok {
return typ
}
return NullableType{Element: typ}
}
func isUnknownType(typ Type) bool {
_, ok := typ.(UnknownType)
return ok
}

231
internal/lang/types.go Normal file
View file

@ -0,0 +1,231 @@
package lang
import (
"fmt"
"strings"
)
type Type interface {
typeNode()
String() string
}
type UnknownType struct{}
func (UnknownType) typeNode() {}
func (UnknownType) String() string { return "<unknown>" }
type NamedType struct{ Name string }
func (NamedType) typeNode() {}
func (t NamedType) String() string { return t.Name }
type ClassType struct{ Class *ClassSymbol }
func (ClassType) typeNode() {}
func (t ClassType) String() string { return t.Class.Name }
type NullableType struct{ Element Type }
func (NullableType) typeNode() {}
func (t NullableType) String() string { return t.Element.String() + "?" }
type GoPointerType struct{ Element Type }
func (GoPointerType) typeNode() {}
func (t GoPointerType) String() string { return "*" + t.Element.String() }
type FunctionType struct {
Params []Type
Result Type
}
func (FunctionType) typeNode() {}
func (t FunctionType) String() string {
params := make([]string, 0, len(t.Params))
for _, param := range t.Params {
params = append(params, param.String())
}
return "(" + strings.Join(params, ", ") + ") -> " + t.Result.String()
}
type GenericType struct {
Base Type
Args []Type
}
func (GenericType) typeNode() {}
func (t GenericType) String() string {
args := make([]string, 0, len(t.Args))
for _, arg := range t.Args {
args = append(args, arg.String())
}
return t.Base.String() + "<" + strings.Join(args, ", ") + ">"
}
func ParseType(text string) (Type, error) {
text = strings.TrimSpace(text)
if text == "" {
return UnknownType{}, nil
}
if strings.HasSuffix(text, "?") {
element, err := ParseType(strings.TrimSpace(strings.TrimSuffix(text, "?")))
if err != nil {
return nil, err
}
return NullableType{Element: element}, nil
}
if strings.HasPrefix(text, "*") {
element, err := ParseType(strings.TrimSpace(strings.TrimPrefix(text, "*")))
if err != nil {
return nil, err
}
return GoPointerType{Element: element}, nil
}
if params, result, ok := parseFunctionType(text); ok {
resolvedParams := make([]Type, 0, len(params))
for _, param := range params {
resolved, err := ParseType(param)
if err != nil {
return nil, err
}
resolvedParams = append(resolvedParams, resolved)
}
resolvedResult, err := ParseType(result)
if err != nil {
return nil, err
}
return FunctionType{Params: resolvedParams, Result: resolvedResult}, nil
}
if base, args, ok := parseGenericType(text); ok {
resolvedBase, err := ParseType(base)
if err != nil {
return nil, err
}
resolvedArgs := make([]Type, 0, len(args))
for _, arg := range args {
resolved, err := ParseType(arg)
if err != nil {
return nil, err
}
resolvedArgs = append(resolvedArgs, resolved)
}
return GenericType{Base: resolvedBase, Args: resolvedArgs}, nil
}
if strings.ContainsAny(text, "<>?()") {
return nil, fmt.Errorf("invalid type %q", text)
}
return NamedType{Name: text}, nil
}
func resolveClassTypes(typ Type, classes map[string]*ClassSymbol) Type {
switch value := typ.(type) {
case NamedType:
if class, ok := classes[value.Name]; ok {
return ClassType{Class: class}
}
return value
case NullableType:
return NullableType{Element: resolveClassTypes(value.Element, classes)}
case GoPointerType:
return GoPointerType{Element: resolveClassTypes(value.Element, classes)}
case FunctionType:
params := make([]Type, len(value.Params))
for i, param := range value.Params {
params[i] = resolveClassTypes(param, classes)
}
return FunctionType{Params: params, Result: resolveClassTypes(value.Result, classes)}
case GenericType:
args := make([]Type, len(value.Args))
for i, arg := range value.Args {
args[i] = resolveClassTypes(arg, classes)
}
return GenericType{Base: resolveClassTypes(value.Base, classes), Args: args}
default:
return typ
}
}
func typeEqual(left, right Type) bool { return left.String() == right.String() }
func renderGoType(typ Type) string {
switch value := typ.(type) {
case UnknownType:
return ""
case ClassType:
return "*" + value.Class.Name
case NullableType:
element := renderGoType(value.Element)
switch value.Element.(type) {
case ClassType, GoPointerType:
return element
}
if element == "error" || element == "any" {
return element
}
return "*" + element
case GoPointerType:
return "*" + renderGoType(value.Element)
case FunctionType:
params := make([]string, len(value.Params))
for i, param := range value.Params {
params[i] = renderGoType(param)
}
result := renderGoType(value.Result)
if result == "" {
return "func(" + strings.Join(params, ", ") + ")"
}
return "func(" + strings.Join(params, ", ") + ") " + result
case GenericType:
base := value.Base.String()
args := make([]string, len(value.Args))
for i, arg := range value.Args {
args[i] = renderGoType(arg)
}
switch base {
case "List", "MutableList":
return "[]" + args[0]
case "Map", "MutableMap":
return "map[" + args[0] + "]" + args[1]
case "Channel":
return "chan " + args[0]
case "Result":
if args[0] == "" {
args[0] = "struct{}"
}
return "GotlinResult[" + args[0] + "]"
case "GotlinSQLIterator":
argument := args[0]
if class, ok := value.Args[0].(ClassType); ok {
argument = class.Class.Name
}
return "*GotlinSQLIterator[" + argument + "]"
}
return base + "[" + strings.Join(args, ", ") + "]"
case NamedType:
switch value.Name {
case "Int":
return "int"
case "Long":
return "int64"
case "Float", "Double":
return "float64"
case "String":
return "string"
case "Any":
return "any"
case "ByteSlice":
return "[]byte"
case "Boolean":
return "bool"
case "Unit":
return ""
case "Error":
return "error"
default:
return value.Name
}
default:
return ""
}
}

View file

@ -0,0 +1,85 @@
package lang
import "testing"
func TestStructuralTypeParsing(t *testing.T) {
typ, err := ParseType("(String?, List<User>) -> Result<User, Error>")
if err != nil {
t.Fatal(err)
}
function, ok := typ.(FunctionType)
if !ok || len(function.Params) != 2 {
t.Fatalf("unexpected function type: %#v", typ)
}
if _, ok := function.Params[0].(NullableType); !ok {
t.Fatalf("expected nullable parameter: %#v", function.Params[0])
}
result, ok := function.Result.(GenericType)
if !ok || result.Base.String() != "Result" || len(result.Args) != 2 {
t.Fatalf("unexpected result type: %#v", function.Result)
}
}
func TestClassReferenceSemanticsStayOutOfSyntaxAST(t *testing.T) {
program, err := Parse(`package demo
class Repository
class Service(val repository: Repository)
fun create(): Service { return Service(Repository()) }`)
if err != nil {
t.Fatal(err)
}
semantic, err := Analyze(program)
if err != nil {
t.Fatal(err)
}
if program.Classes[1].Fields[0].Type != "Repository" || program.Functions[0].ReturnType != "Service" {
t.Fatalf("syntax types were rewritten: %#v %#v", program.Classes[1].Fields[0], program.Functions[0])
}
if _, err := GenerateGo(program); err != nil {
t.Fatal(err)
}
if program.Classes[1].Fields[0].Type != "Repository" || program.Functions[0].ReturnType != "Service" {
t.Fatal("Go generation mutated syntax type spelling")
}
fieldType := semantic.ClassInfo["Service"].Fields["repository"].Type
if _, ok := fieldType.(ClassType); !ok {
t.Fatalf("expected semantic class type, got %#v", fieldType)
}
if got := semantic.GoType("List<Repository>"); got != "[]*Repository" {
t.Fatalf("unexpected Go type: %s", got)
}
}
func TestResolverBuildsTypedHIR(t *testing.T) {
program, err := Parse(`package demo
import strconv
fun parse(value: String): Result<Int, Error> {
val parsed = strconv.atoi(value)?
return Result.Ok(parsed)
}`)
if err != nil {
t.Fatal(err)
}
semantic, err := Analyze(program)
if err != nil {
t.Fatal(err)
}
if semantic.HIR == nil || len(semantic.HIR.Functions) != 1 {
t.Fatal("typed HIR function was not created")
}
decl := program.Functions[0].Body[0].(VarDecl)
attempt := decl.Value.(TryExpr)
if attempt.Meta.Semantic == nil || attempt.Meta.Semantic.Meaning != PropagateResultExpr || attempt.Meta.Semantic.Type.String() != "Int" {
t.Fatalf("unexpected propagation HIR: %#v", attempt.Meta.Semantic)
}
if _, ok := attempt.Meta.Semantic.Node.(HIRPropagateResult); !ok {
t.Fatalf("propagation did not lower to HIRPropagateResult: %#v", attempt.Meta.Semantic.Node)
}
call := attempt.Value.(CallExpr)
if call.Meta.Semantic == nil || call.Meta.Semantic.Meaning != GoCallExpr || call.Meta.Semantic.Type.String() != "Result<Int, Error>" {
t.Fatalf("external call was not resolved: %#v", call.Meta.Semantic)
}
if _, ok := call.Meta.Semantic.Node.(HIRGoCall); !ok {
t.Fatalf("external call did not lower to HIRGoCall: %#v", call.Meta.Semantic.Node)
}
}

View file

@ -43,7 +43,7 @@ assert(language.folding?.markers?.start && language.indentationRules?.increaseIn
const grammarSource = JSON.stringify(grammar); const grammarSource = JSON.stringify(grammar);
const expectedTokens = [ const expectedTokens = [
"data", "class", "suspend", "private", "override", "val", "var", "if", "else", "data", "class", "suspend", "private", "override", "val", "var", "if", "else",
"while", "for", "in", "select", "return", "defer", "try", "catch", "while", "for", "in", "return", "defer", "try", "catch",
"throw", "enum", "match", "package", "import", "jsonNaming", "embed", "table", "column", "id", "throw", "enum", "match", "package", "import", "jsonNaming", "embed", "table", "column", "id",
"generated", "Int", "Long", "String", "Boolean", "Unit", "Double", "Float", "Any", "generated", "Int", "Long", "String", "Boolean", "Unit", "Double", "Float", "Any",
"ByteSlice", "Error", "Result", "List", "MutableList", "Map", "MutableMap", "Channel", "GotlinSQLQuery", "ByteSlice", "Error", "Result", "List", "MutableList", "Map", "MutableMap", "Channel", "GotlinSQLQuery",

View file

@ -274,7 +274,7 @@
}, },
{ {
"name": "keyword.control.concurrency.gotlin", "name": "keyword.control.concurrency.gotlin",
"match": "\\b(select|defer|runBlocking|coroutineScope|launch|async|await|delay|withTimeout|isActive)\\b" "match": "\\b(defer|runBlocking|coroutineScope|launch|async|await|delay|withTimeout|isActive)\\b"
}, },
{ {
"name": "keyword.control.exception.gotlin", "name": "keyword.control.exception.gotlin",