Add value-returning match expressions

This commit is contained in:
pavel 2026-08-27 17:56:13 +02:00
commit ab783c33ed
12 changed files with 384 additions and 14 deletions

View file

@ -108,18 +108,18 @@ enum PaymentResult {
}
fun describe(result: PaymentResult): String {
var description = ""
match (result) {
PaymentResult::Accepted(id) -> { description = id }
PaymentResult::Rejected(reason) -> { description = reason }
PaymentResult::Pending -> { description = "pending" }
return match (result) {
PaymentResult::Accepted(id) -> id
PaymentResult::Rejected(reason) -> reason
PaymentResult::Pending -> "pending"
}
return description
}
```
Enum matches must contain each variant exactly once. Variant payload arity is
checked during Gotlin compilation.
checked during Gotlin compilation. A match used as an expression also requires
every arm to return the same type. Block-style statement matches remain
available for side effects.
## Null safety

View file

@ -7,13 +7,11 @@ enum PaymentResult {
}
fun describe(result: PaymentResult): String {
var description = ""
match (result) {
PaymentResult::Accepted(id) -> { description = "accepted " + id }
PaymentResult::Rejected(reason) -> { description = "rejected " + reason }
PaymentResult::Pending -> { description = "pending" }
return match (result) {
PaymentResult::Accepted(id) -> "accepted " + id
PaymentResult::Rejected(reason) -> "rejected " + reason
PaymentResult::Pending -> "pending"
}
return description
}
fun main() {

View file

@ -217,6 +217,19 @@ type MatchCase struct {
Body []Stmt
}
type MatchExpr struct {
Value Expr
Cases []MatchExprCase
}
func (MatchExpr) exprNode() {}
type MatchExprCase struct {
EnumName, VariantName string
Bindings []string
Value Expr
}
type SelectCase struct {
Source Expr
Body []Stmt

View file

@ -66,6 +66,15 @@ func expressionUsesCoroutines(expr Expr) bool {
}
case LambdaExpr:
return statementsUseCoroutines(e.Body)
case MatchExpr:
if expressionUsesCoroutines(e.Value) {
return true
}
for _, matchCase := range e.Cases {
if expressionUsesCoroutines(matchCase.Value) {
return true
}
}
case SelectorExpr:
return expressionUsesCoroutines(e.Receiver)
case BinaryExpr:

View file

@ -77,3 +77,68 @@ fun main() { println(Status::PendingReservation) }`)
}
}
}
func TestGenerateExhaustiveMatchExpression(t *testing.T) {
prog, err := Parse(`package demo
enum AccountType { BASIC, SAVINGS }
fun interestRate(accountType: AccountType): Double {
return match (accountType) {
AccountType::BASIC -> 0.0
AccountType::SAVINGS -> 0.02
}
}`)
if err != nil {
t.Fatal(err)
}
out, err := GenerateGo(prog)
if err != nil {
t.Fatal(err)
}
for _, want := range []string{"func() float64", "case AccountTypeBASIC:", "return 0.0", "case AccountTypeSAVINGS:", "return 0.02", `panic("unreachable exhaustive match")`} {
if !strings.Contains(string(out), want) {
t.Fatalf("missing %q:\n%s", want, out)
}
}
}
func TestGeneratePayloadMatchExpression(t *testing.T) {
prog, err := Parse(`package demo
enum Outcome { Success(String), Failure(String) }
fun message(outcome: Outcome): String {
return match (outcome) {
Outcome::Success(value) -> value
Outcome::Failure(reason) -> reason
}
}`)
if err != nil {
t.Fatal(err)
}
out, err := GenerateGo(prog)
if err != nil {
t.Fatal(err)
}
for _, want := range []string{"switch gotlinMatch", "value := gotlinMatch", "return value", "return reason"} {
if !strings.Contains(string(out), want) {
t.Fatalf("missing %q:\n%s", want, out)
}
}
}
func TestRejectInvalidMatchExpression(t *testing.T) {
for _, test := range []struct {
source string
want string
}{
{`package demo enum State { On, Off } fun value(state: State): Int { return match (state) { State::On -> 1 } }`, "missing Off"},
{`package demo enum State { On, Off } fun value(state: State): Int { return match (state) { State::On -> 1 State::Off -> "off" } }`, "has type String, expected Int"},
} {
prog, err := Parse(test.source)
if err != nil {
t.Fatal(err)
}
_, err = GenerateGo(prog)
if err == nil || !strings.Contains(err.Error(), test.want) {
t.Fatalf("expected %q, got %v", test.want, err)
}
}
}

View file

@ -1679,6 +1679,8 @@ func (g *goGenerator) expr(expr Expr, expectedType string) (string, error) {
return "", err
}
return receiver + "[" + index + "]", nil
case MatchExpr:
return g.valueMatch(e, expectedType)
case EnumVariantExpr:
if e.EnumName == "Result" {
base, args, ok := parseGenericType(expectedType)
@ -1743,6 +1745,128 @@ func (g *goGenerator) expr(expr Expr, expectedType string) (string, error) {
}
}
func (g *goGenerator) valueMatch(match MatchExpr, expectedType string) (string, error) {
enumName := strings.TrimPrefix(g.exprType(match.Value), "*")
if enumName == "" && len(match.Cases) > 0 {
enumName = match.Cases[0].EnumName
}
decl, ok := g.enums[enumName]
if !ok {
return "", fmt.Errorf("match value is not a known enum")
}
seen := map[string]bool{}
for _, c := range match.Cases {
if c.EnumName != enumName {
return "", fmt.Errorf("match case %s::%s does not match enum %s", c.EnumName, c.VariantName, enumName)
}
if seen[c.VariantName] {
return "", fmt.Errorf("duplicate match case %s::%s", enumName, c.VariantName)
}
seen[c.VariantName] = true
variant := enumVariant(decl, c.VariantName)
if variant == nil {
return "", fmt.Errorf("unknown variant %s::%s", enumName, c.VariantName)
}
if len(c.Bindings) != len(variant.PayloadTypes) {
return "", fmt.Errorf("match case %s::%s expects %d bindings", enumName, c.VariantName, len(variant.PayloadTypes))
}
}
for _, variant := range decl.Variants {
if !seen[variant.Name] {
return "", fmt.Errorf("non-exhaustive match for %s: missing %s", enumName, variant.Name)
}
}
resultType := expectedType
if resultType == "" || resultType == "Any" {
for _, c := range match.Cases {
variant := enumVariant(decl, c.VariantName)
g.pushScope()
for i, binding := range c.Bindings {
g.defineType(binding, variant.PayloadTypes[i])
}
armType := g.exprType(c.Value)
g.popScope()
if armType != "" {
resultType = armType
break
}
}
}
if resultType == "" || resultType == "Unit" {
return "", fmt.Errorf("match expression result type cannot be inferred")
}
for _, c := range match.Cases {
variant := enumVariant(decl, c.VariantName)
g.pushScope()
for i, binding := range c.Bindings {
g.defineType(binding, variant.PayloadTypes[i])
}
armType := g.exprType(c.Value)
g.popScope()
if armType != "" && armType != resultType {
return "", fmt.Errorf("match expression arm %s::%s has type %s, expected %s", enumName, c.VariantName, armType, resultType)
}
}
value, err := g.expr(match.Value, enumName)
if err != nil {
return "", err
}
g.matchCounter++
matchName := fmt.Sprintf("gotlinMatch%d", g.matchCounter)
var out strings.Builder
out.WriteString("func() ")
out.WriteString(mapGoType(resultType))
out.WriteString(" { ")
if enumIsString(decl) {
out.WriteString("switch ")
out.WriteString(value)
out.WriteString(" { ")
for _, c := range match.Cases {
arm, err := g.expr(c.Value, resultType)
if err != nil {
return "", err
}
out.WriteString("case ")
out.WriteString(enumName + c.VariantName)
out.WriteString(": return ")
out.WriteString(arm)
out.WriteString("; ")
}
out.WriteString("}; ")
} else {
out.WriteString("switch ")
out.WriteString(matchName)
out.WriteString(" := ")
out.WriteString(value)
out.WriteString(".(type) { ")
for _, c := range match.Cases {
variant := enumVariant(decl, c.VariantName)
g.pushScope()
out.WriteString("case *")
out.WriteString(enumName + c.VariantName)
out.WriteString(": ")
for i, binding := range c.Bindings {
g.defineType(binding, variant.PayloadTypes[i])
out.WriteString(binding)
out.WriteString(" := ")
out.WriteString(matchName)
out.WriteString(fmt.Sprintf(".Value%d; ", i))
}
arm, err := g.expr(c.Value, resultType)
g.popScope()
if err != nil {
return "", err
}
out.WriteString("return ")
out.WriteString(arm)
out.WriteString("; ")
}
out.WriteString("}; ")
}
out.WriteString(`panic("unreachable exhaustive match") }()`)
return out.String(), nil
}
func (g *goGenerator) emitJSONDecodeSupport() {
g.line("func gotlinJSONDecode[T any](body []byte) GotlinResult[T] {")
g.indentLevel++
@ -2883,6 +3007,22 @@ func exprUsesName(expr Expr, name string, shadowed bool) bool {
}
}
return false
case MatchExpr:
if exprUsesName(e.Value, name, shadowed) {
return true
}
for _, matchCase := range e.Cases {
caseShadowed := shadowed
for _, binding := range matchCase.Bindings {
if binding == name {
caseShadowed = true
}
}
if exprUsesName(matchCase.Value, name, caseShadowed) {
return true
}
}
return false
case LambdaExpr:
lambdaShadowed := shadowed
if e.ImplicitIt && name == "it" {
@ -3044,6 +3184,12 @@ func (g *goGenerator) exprType(expr Expr) string {
if _, args, ok := parseGenericType(g.exprType(e.Receiver)); ok && len(args) > 0 {
return args[len(args)-1]
}
case MatchExpr:
for _, c := range e.Cases {
if typ := g.exprType(c.Value); typ != "" {
return typ
}
}
case EnumVariantExpr:
return e.EnumName
case NonNullExpr:
@ -3280,6 +3426,15 @@ func exprUsesPrintln(expr Expr) bool {
return exprUsesPrintln(e.Left) || exprUsesPrintln(e.Right)
case SelectorExpr:
return exprUsesPrintln(e.Receiver)
case MatchExpr:
if exprUsesPrintln(e.Value) {
return true
}
for _, matchCase := range e.Cases {
if exprUsesPrintln(matchCase.Value) {
return true
}
}
case LambdaExpr:
return usesPrintln(e.Body)
}
@ -3306,6 +3461,15 @@ func exprUsesRunCatching(expr Expr) bool {
return exprUsesRunCatching(e.Left) || exprUsesRunCatching(e.Right)
case SelectorExpr:
return exprUsesRunCatching(e.Receiver)
case MatchExpr:
if exprUsesRunCatching(e.Value) {
return true
}
for _, matchCase := range e.Cases {
if exprUsesRunCatching(matchCase.Value) {
return true
}
}
case LambdaExpr:
return usesRunCatching(e.Body)
}
@ -3332,6 +3496,15 @@ func exprUsesTimerBuiltins(expr Expr) bool {
return exprUsesTimerBuiltins(e.Left) || exprUsesTimerBuiltins(e.Right)
case SelectorExpr:
return exprUsesTimerBuiltins(e.Receiver)
case MatchExpr:
if exprUsesTimerBuiltins(e.Value) {
return true
}
for _, matchCase := range e.Cases {
if exprUsesTimerBuiltins(matchCase.Value) {
return true
}
}
case LambdaExpr:
return usesTimerBuiltins(e.Body)
}
@ -3358,6 +3531,15 @@ func exprUsesEveryBuiltin(expr Expr) bool {
return exprUsesEveryBuiltin(e.Left) || exprUsesEveryBuiltin(e.Right)
case SelectorExpr:
return exprUsesEveryBuiltin(e.Receiver)
case MatchExpr:
if exprUsesEveryBuiltin(e.Value) {
return true
}
for _, matchCase := range e.Cases {
if exprUsesEveryBuiltin(matchCase.Value) {
return true
}
}
case LambdaExpr:
return usesEveryBuiltin(e.Body)
}
@ -3485,6 +3667,11 @@ func usedImportAliases(program *Program) map[string]bool {
for _, value := range e.Values {
walkExpr(value)
}
case MatchExpr:
walkExpr(e.Value)
for _, matchCase := range e.Cases {
walkExpr(matchCase.Value)
}
case LambdaExpr:
for _, param := range e.Params {
markType(param.Type)

View file

@ -793,6 +793,68 @@ func (p *parser) parseMatch() (Stmt, error) {
return MatchStmt{Value: value, Cases: cases}, nil
}
func (p *parser) parseMatchExpr() (Expr, error) {
if _, err := p.expect(tokenLParen, "expected '(' after match"); err != nil {
return nil, err
}
value, err := p.parseExpr(0)
if err != nil {
return nil, err
}
if _, err := p.expect(tokenRParen, "expected ')' after match value"); err != nil {
return nil, err
}
if _, err := p.expect(tokenLBrace, "expected '{' after match value"); err != nil {
return nil, err
}
var cases []MatchExprCase
for !p.check(tokenRBrace) && !p.check(tokenEOF) {
enumName, err := p.expect(tokenIdent, "expected enum name in match case")
if err != nil {
return nil, err
}
if _, err := p.expect(tokenDoubleColon, "expected '::' in match case"); err != nil {
return nil, err
}
variant, err := p.expect(tokenIdent, "expected variant name")
if err != nil {
return nil, err
}
var bindings []string
if p.match(tokenLParen) {
if !p.check(tokenRParen) {
for {
binding, err := p.expect(tokenIdent, "expected variant binding")
if err != nil {
return nil, err
}
bindings = append(bindings, binding.lexeme)
if !p.match(tokenComma) {
break
}
}
}
if _, err := p.expect(tokenRParen, "expected ')' after variant bindings"); err != nil {
return nil, err
}
}
if _, err := p.expect(tokenArrow, "expected '->' after match pattern"); err != nil {
return nil, err
}
result, err := p.parseExpr(0)
if err != nil {
return nil, err
}
cases = append(cases, MatchExprCase{EnumName: enumName.lexeme, VariantName: variant.lexeme, Bindings: bindings, Value: result})
p.match(tokenComma)
p.match(tokenSemicolon)
}
if _, err := p.expect(tokenRBrace, "expected '}' after match"); err != nil {
return nil, err
}
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
@ -1016,6 +1078,8 @@ func (p *parser) parsePrefix() (Expr, error) {
return p.parsePostfix(expr)
case tokenLBrace:
return p.parseLambdaExpr()
case tokenMatch:
return p.parseMatchExpr()
case tokenBang, tokenMinus, tokenAmp, tokenStar:
value, err := p.parseExpr(7)
if err != nil {

View file

@ -133,6 +133,11 @@ func normalizeExprTypes(expression Expr, normalize func(string) string) {
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)

View file

@ -225,6 +225,15 @@ func (c *mutabilityChecker) checkExpr(expr Expr) error {
return err
}
}
case MatchExpr:
if err := c.checkExpr(e.Value); err != nil {
return err
}
for _, matchCase := range e.Cases {
if err := c.checkExpr(matchCase.Value); err != nil {
return err
}
}
case LambdaExpr:
bindings := map[string]bool{}
if e.ImplicitIt {

View file

@ -1144,6 +1144,15 @@ func exprMatches(expr Expr, match func(Expr) bool) bool {
return exprMatches(e.Receiver, match)
case IndexExpr:
return exprMatches(e.Receiver, match) || exprMatches(e.Index, match)
case MatchExpr:
if exprMatches(e.Value, match) {
return true
}
for _, matchCase := range e.Cases {
if exprMatches(matchCase.Value, match) {
return true
}
}
case LambdaExpr:
return stmtsMatch(e.Body, match)
}

View file

@ -82,7 +82,7 @@ for (const declaration of [
const prefixes = new Set(Object.values(snippets).map((snippet) => snippet.prefix));
for (const prefix of [
"dataclass", "tablerow", "embed", "coroutinescope", "launch", "async", "enum", "match", "mapto", "safe", "nonnull", "resultfun", "resultmatch", "defer", "foreach", "sqlfetch",
"dataclass", "tablerow", "embed", "coroutinescope", "launch", "async", "enum", "match", "matchvalue", "mapto", "safe", "nonnull", "resultfun", "resultmatch", "defer", "foreach", "sqlfetch",
"sqlsingle", "sqliterator", "sqlinsertnothing", "sqlinsertupdate", "sqlupdatereturning"
]) {
assert(prefixes.has(prefix), `snippets are missing prefix ${prefix}`);

View file

@ -80,6 +80,17 @@
],
"description": "Exhaustive match over enum variants"
},
"Exhaustive Match Expression": {
"prefix": "matchvalue",
"body": [
"val ${1:value} = match (${2:result}) {",
" ${3:Result}::${4:Success}(${5:item}) -> ${5:item}",
" ${3:Result}::${6:Failure}(${7:reason}) -> ${7:reason}",
" $0",
"}"
],
"description": "Return a typed value from an exhaustive enum match"
},
"Recursive Structural Mapping": {
"prefix": "mapto",
"body": ["val ${1:target} = ${2:source}.mapTo<${3:Target}>()"],