diff --git a/src/internal/parser/expression.go b/src/internal/parser/expression.go new file mode 100644 index 0000000..80ba8a5 --- /dev/null +++ b/src/internal/parser/expression.go @@ -0,0 +1,513 @@ +package parser + +import ( + "strings" + + "github.com/puff-lang/puff/internal/ast" + "github.com/puff-lang/puff/internal/diagnostic" + "github.com/puff-lang/puff/internal/token" +) + +func (parser *parser) parseExpressionUntil(stops ...token.Type) ast.Expression { + previousStops := parser.expressionStops + parser.expressionStops = make(map[token.Type]bool, len(stops)) + for _, stop := range stops { + parser.expressionStops[stop] = true + } + diagnosticCount := len(parser.diagnostics) + expression := parser.parseRange() + if expression == nil && + len(parser.diagnostics) == diagnosticCount && + !parser.hasLexerDiagnosticAt(parser.peek().StartOffset) { + parser.reportExpected("expression", "") + } + parser.expressionStops = previousStops + return expression +} + +func (parser *parser) hasLexerDiagnosticAt(offset int) bool { + for index := len(parser.diagnostics) - 1; index >= 0; index-- { + item := parser.diagnostics[index] + if item.Phase != diagnostic.PhaseLexer || item.Span.StartOffset > offset { + continue + } + if item.Span.EndOffset >= offset { + return true + } + between := parser.file.Text[item.Span.EndOffset:offset] + trimmed := strings.TrimSpace(between) + if !strings.ContainsAny(between, "\r\n") && + (trimmed == "" || strings.HasPrefix(trimmed, "#")) { + return true + } + } + return false +} + +func (parser *parser) parseRange() ast.Expression { + left := parser.parseOr() + if left == nil || !parser.match(token.DotDot) { + return left + } + + right := parser.parseOr() + if right == nil { + parser.reportExpected("expression", "") + return left + } + expression := &ast.RangeExpr{ + NodeBase: parser.base(left.Span().StartOffset, right.Span().EndOffset), + Start: left, + End: right, + } + if parser.check(token.DotDot) { + parser.reportUnexpected(parser.peek(), "Ranges cannot be chained.") + parser.advance() + for !parser.atExpressionStop() && !parser.atEnd() { + parser.advance() + } + } + return expression +} + +func (parser *parser) parseOr() ast.Expression { + return parser.parseBinary(parser.parseAnd, token.Or) +} + +func (parser *parser) parseAnd() ast.Expression { + return parser.parseBinary(parser.parseEquality, token.And) +} + +func (parser *parser) parseEquality() ast.Expression { + return parser.parseBinary(parser.parseComparison, token.EqualEqual, token.BangEqual) +} + +func (parser *parser) parseComparison() ast.Expression { + return parser.parseBinary(parser.parseTerm, token.Greater, token.GreaterEq, token.Less, token.LessEq) +} + +func (parser *parser) parseTerm() ast.Expression { + return parser.parseBinary(parser.parseFactor, token.Plus, token.Minus) +} + +func (parser *parser) parseFactor() ast.Expression { + return parser.parseBinary(parser.parseUnary, token.Star, token.Slash, token.Percent) +} + +func (parser *parser) parseBinary(operand func() ast.Expression, operators ...token.Type) ast.Expression { + expression := operand() + for expression != nil && parser.match(operators...) { + operator := parser.previous() + right := operand() + if right == nil { + parser.reportExpected("expression", "") + return expression + } + expression = &ast.BinaryExpr{ + NodeBase: parser.base(expression.Span().StartOffset, right.Span().EndOffset), + Left: expression, + Operator: operator.Type, + Right: right, + } + } + return expression +} + +func (parser *parser) parseUnary() ast.Expression { + if parser.match(token.Not, token.Minus) { + operator := parser.previous() + operand := parser.parseUnary() + if operand == nil { + parser.reportExpected("expression", "") + return nil + } + return &ast.UnaryExpr{ + NodeBase: parser.base(operator.StartOffset, operand.Span().EndOffset), + Operator: operator.Type, + Operand: operand, + } + } + return parser.parsePrimary() +} + +func (parser *parser) parsePrimary() ast.Expression { + if parser.atExpressionStop() || parser.atEnd() { + return nil + } + + switch parser.peek().Type { + case token.Nil: + tok := parser.advance() + return &ast.NilLiteral{NodeBase: parser.base(tok.StartOffset, tok.EndOffset)} + case token.True, token.False: + tok := parser.advance() + return &ast.BoolLiteral{NodeBase: parser.base(tok.StartOffset, tok.EndOffset), Value: tok.Type == token.True} + case token.Int: + tok := parser.advance() + value, _ := tok.Value.(int) + return &ast.IntLiteral{NodeBase: parser.base(tok.StartOffset, tok.EndOffset), Value: int64(value)} + case token.Float: + tok := parser.advance() + value, _ := tok.Value.(float64) + return &ast.FloatLiteral{NodeBase: parser.base(tok.StartOffset, tok.EndOffset), Value: value} + case token.StringStart: + return parser.parseStringExpression() + case token.LParen: + return parser.parseGroup() + case token.LBracket: + return parser.parseList() + case token.LBrace: + return parser.parseMap() + case token.Dollar: + return parser.parseVariable(nil) + } + + if parser.checkName() { + return parser.parseNameExpression() + } + + parser.reportExpected("expression", "") + parser.advance() + return nil +} + +func (parser *parser) parseGroup() ast.Expression { + start := parser.advance().StartOffset + expression := parser.parseExpressionUntil(token.RParen) + if !parser.match(token.RParen) { + parser.reportExpected(")", "") + end := start + if expression != nil { + end = expression.Span().EndOffset + } + return &ast.GroupExpr{NodeBase: parser.base(start, end), Expression: expression} + } + return &ast.GroupExpr{ + NodeBase: parser.base(start, parser.previous().EndOffset), + Expression: expression, + } +} + +func (parser *parser) parseList() ast.Expression { + start := parser.advance().StartOffset + var elements []ast.Expression + for !parser.check(token.RBracket) && !parser.atEnd() { + element := parser.parseExpressionUntil(token.Comma, token.RBracket) + if element != nil { + elements = append(elements, element) + } + if parser.match(token.Comma) { + if parser.check(token.RBracket) { + break + } + continue + } + if !parser.check(token.RBracket) { + parser.reportExpected(`"," or "]"`, "") + parser.synchronizeUntil(token.Comma, token.RBracket, token.Newline) + parser.match(token.Comma) + } + } + end := parser.collectionEnd(start, token.RBracket, "]") + return &ast.ListExpr{NodeBase: parser.base(start, end), Elements: elements} +} + +func (parser *parser) parseMap() ast.Expression { + start := parser.advance().StartOffset + var entries []ast.MapEntry + for !parser.check(token.RBrace) && !parser.atEnd() { + if parser.check(token.Comma) { + parser.reportExpected("expression", "") + parser.advance() + continue + } + entryStart := parser.peek().StartOffset + key := parser.parseExpressionUntil(token.Colon) + hasColon := parser.match(token.Colon) + if !hasColon { + parser.reportExpected(":", "") + parser.synchronizeUntil(token.Comma, token.RBrace, token.Newline) + } + var value ast.Expression + if hasColon { + value = parser.parseExpressionUntil(token.Comma, token.RBrace) + } + entryEnd := entryStart + if value != nil { + entryEnd = value.Span().EndOffset + } else if key != nil { + entryEnd = key.Span().EndOffset + } + entries = append(entries, ast.MapEntry{ + NodeBase: parser.base(entryStart, entryEnd), + Key: key, + Value: value, + }) + if parser.match(token.Comma) { + if parser.check(token.RBrace) { + break + } + continue + } + if !parser.check(token.RBrace) { + parser.reportExpected(`"," or "}"`, "") + parser.synchronizeUntil(token.Comma, token.RBrace, token.Newline) + parser.match(token.Comma) + } + } + end := parser.collectionEnd(start, token.RBrace, "}") + return &ast.MapExpr{NodeBase: parser.base(start, end), Entries: entries} +} + +func (parser *parser) collectionEnd(start int, closing token.Type, spelling string) int { + if parser.match(closing) { + return parser.previous().EndOffset + } + parser.reportExpected(spelling, "") + if parser.current > 0 { + return parser.previous().EndOffset + } + return start +} + +func (parser *parser) parseStringExpression() ast.Expression { + startToken := parser.advance() + expression := &ast.StringExpr{Quote: startToken.Lexeme[0]} + for !parser.check(token.StringEnd) && !parser.check(token.Newline) && !parser.atEnd() { + switch parser.peek().Type { + case token.StringText: + tok := parser.advance() + value, _ := tok.Value.(string) + expression.Parts = append(expression.Parts, &ast.StringText{ + NodeBase: parser.base(tok.StartOffset, tok.EndOffset), + Raw: tok.Lexeme, + Value: value, + }) + case token.InterpStart: + partStart := parser.advance().StartOffset + value := parser.parseExpressionUntil(token.InterpEnd) + end := partStart + if parser.match(token.InterpEnd) { + end = parser.previous().EndOffset + } else { + parser.reportExpected("}", "") + parser.synchronizeUntil(token.InterpEnd, token.StringEnd, token.Newline) + if parser.match(token.InterpEnd) { + end = parser.previous().EndOffset + } + if value != nil { + if end == partStart { + end = value.Span().EndOffset + } + } + } + expression.Parts = append(expression.Parts, &ast.StringInterpolation{ + NodeBase: parser.base(partStart, end), + Expression: value, + }) + default: + parser.reportUnexpected(parser.peek(), "") + parser.advance() + } + } + + end := startToken.EndOffset + if parser.match(token.StringEnd) { + end = parser.previous().EndOffset + } else if len(expression.Parts) > 0 { + end = expression.Parts[len(expression.Parts)-1].Span().EndOffset + } + expression.NodeBase = parser.base(startToken.StartOffset, end) + return expression +} + +func (parser *parser) parseNameExpression() ast.Expression { + startIndex := parser.current + if parser.peekAt(1).Type == token.Dot && parser.peekAt(2).Type == token.Dollar { + qualifier := parser.parseIdentifier() + parser.advance() + return parser.parseVariable(qualifier) + } + + parts := []ast.Identifier{*parser.parseIdentifier()} + for parser.match(token.Dot) { + if !parser.checkName() { + parser.reportExpected("identifier", "") + break + } + parts = append(parts, *parser.parseIdentifier()) + } + + if parser.match(token.LParen) { + return parser.finishCall(parts, true, parser.tokens[startIndex].StartOffset) + } + + if parser.canContinuePattern() { + parser.current = startIndex + return parser.parsePatternPrimary() + } + + end := parts[len(parts)-1].Span().EndOffset + return &ast.CallExpr{ + NodeBase: parser.base(parser.tokens[startIndex].StartOffset, end), + Callee: ast.QualifiedName{ + NodeBase: parser.base(parser.tokens[startIndex].StartOffset, end), + Parts: parts, + }, + } +} + +func (parser *parser) finishCall(parts []ast.Identifier, explicit bool, start int) ast.Expression { + var arguments []ast.Expression + for !parser.check(token.RParen) && !parser.atEnd() { + argument := parser.parseExpressionUntil(token.Comma, token.RParen) + if argument != nil { + arguments = append(arguments, argument) + } + if parser.match(token.Comma) { + if argument != nil && parser.check(token.RParen) { + parser.reportExpected("expression", "") + break + } + continue + } + if !parser.check(token.RParen) { + parser.reportExpected(`"," or ")"`, "") + parser.synchronizeUntil(token.Comma, token.RParen, token.Newline) + parser.match(token.Comma) + } + } + end := start + if parser.match(token.RParen) { + end = parser.previous().EndOffset + } else { + parser.reportExpected(")", "") + if len(arguments) > 0 { + end = arguments[len(arguments)-1].Span().EndOffset + } + } + calleeEnd := parts[len(parts)-1].Span().EndOffset + return &ast.CallExpr{ + NodeBase: parser.base(start, end), + Callee: ast.QualifiedName{ + NodeBase: parser.base(start, calleeEnd), + Parts: parts, + }, + Arguments: arguments, + ExplicitParens: explicit, + } +} + +func (parser *parser) parsePatternPrimary() ast.Expression { + start := parser.current + for !parser.atEnd() && !parser.atExpressionStop() && !isBinaryOperator(parser.peek().Type) { + parser.advance() + } + tokens := append([]token.Token(nil), parser.tokens[start:parser.current]...) + if len(tokens) == 0 { + return nil + } + return &ast.PatternExpr{ + NodeBase: parser.base(tokens[0].StartOffset, tokens[len(tokens)-1].EndOffset), + Tokens: tokens, + } +} + +func (parser *parser) parseVariable(qualifier *ast.Identifier) ast.Expression { + start := parser.peek().StartOffset + if qualifier != nil { + start = qualifier.Span().StartOffset + } + parser.advance() + local := parser.match(token.Underscore) + if !parser.checkName() { + parser.reportExpected("variable name", "") + return &ast.VariableExpr{NodeBase: parser.base(start, parser.previous().EndOffset), Qualifier: qualifier, Local: local} + } + name := parser.parseIdentifier() + if qualifier != nil && local { + parser.reportExpected("global variable name", "") + } + + var accesses []ast.VariableAccess + for { + switch { + case parser.match(token.Dot): + accessStart := parser.previous().StartOffset + if !parser.checkName() { + parser.reportExpected("field name", "") + continue + } + field := parser.parseIdentifier() + accesses = append(accesses, &ast.FieldAccess{ + NodeBase: parser.base(accessStart, field.Span().EndOffset), + Field: *field, + }) + case parser.match(token.LBracket): + accessStart := parser.previous().StartOffset + if parser.match(token.RBracket) { + accesses = append(accesses, &ast.EmptyIndexAccess{NodeBase: parser.base(accessStart, parser.previous().EndOffset)}) + continue + } + index := parser.parseExpressionUntil(token.RBracket) + end := accessStart + if parser.match(token.RBracket) { + end = parser.previous().EndOffset + } else { + parser.reportExpected("]", "") + if index != nil { + end = index.Span().EndOffset + } + } + accesses = append(accesses, &ast.IndexAccess{ + NodeBase: parser.base(accessStart, end), + Index: index, + }) + default: + end := name.Span().EndOffset + if len(accesses) > 0 { + end = accesses[len(accesses)-1].Span().EndOffset + } + return &ast.VariableExpr{ + NodeBase: parser.base(start, end), + Qualifier: qualifier, + Name: *name, + Local: local, + Accesses: accesses, + } + } + } +} + +func (parser *parser) canContinuePattern() bool { + return !parser.atEnd() && + !parser.atExpressionStop() && + !isBinaryOperator(parser.peek().Type) && + parser.peek().Type != token.Newline +} + +func (parser *parser) atExpressionStop() bool { + return parser.expressionStops != nil && parser.expressionStops[parser.peek().Type] +} + +func isBinaryOperator(tokenType token.Type) bool { + switch tokenType { + case token.DotDot, + token.Or, + token.And, + token.EqualEqual, + token.BangEqual, + token.Greater, + token.GreaterEq, + token.Less, + token.LessEq, + token.Plus, + token.Minus, + token.Star, + token.Slash, + token.Percent: + return true + default: + return false + } +} diff --git a/src/internal/parser/expression_test.go b/src/internal/parser/expression_test.go new file mode 100644 index 0000000..c62ce40 --- /dev/null +++ b/src/internal/parser/expression_test.go @@ -0,0 +1,299 @@ +package parser + +import ( + "fmt" + "strings" + "testing" + + "github.com/puff-lang/puff/internal/ast" + "github.com/puff-lang/puff/internal/diagnostic" + "github.com/puff-lang/puff/internal/token" +) + +func TestParseExpressionPrecedence(t *testing.T) { + tests := []struct { + source string + want string + }{ + {source: "1 + 2 * 3", want: "(+ 1 (* 2 3))"}, + {source: "(1 + 2) * 3", want: "(* (group (+ 1 2)) 3)"}, + {source: "10 - 3 - 2", want: "(- (- 10 3) 2)"}, + {source: "-1 * 2", want: "(* (unary - 1) 2)"}, + {source: "not false or true", want: "(or (unary not false) true)"}, + {source: "1 + 2 >= 3 == true", want: "(== (>= (+ 1 2) 3) true)"}, + {source: "true or false and false", want: "(or true (and false false))"}, + {source: "20 / 5 % 3", want: "(% (/ 20 5) 3)"}, + } + + for _, test := range tests { + t.Run(test.source, func(t *testing.T) { + result := parseTestSource("expression.puff", "$value = "+test.source+"\n") + if len(result.Diagnostics) != 0 { + t.Fatalf("expected no diagnostics, got %#v", result.Diagnostics) + } + value := result.File.Declarations[0].(*ast.GlobalAssignment).Value + if got := expressionShape(value); got != test.want { + t.Fatalf("expected %s, got %s", test.want, got) + } + }) + } +} + +func TestParseVariablesCallsCollectionsAndRange(t *testing.T) { + result := parseTestSource("expressions.puff", ` +$variable = $player.stats[$index][] +$imported = shop.$tax +$call = shop.finalPrice(100) +$list = [1, 2, 3,] +$map = {"coins": 100, "kills": 5,} +$range = 1..10 +`) + + if len(result.Diagnostics) != 0 { + t.Fatalf("expected no diagnostics, got %#v", result.Diagnostics) + } + values := make([]ast.Expression, len(result.File.Declarations)) + for index, declaration := range result.File.Declarations { + values[index] = declaration.(*ast.GlobalAssignment).Value + } + + variable := values[0].(*ast.VariableExpr) + if variable.Name.Name != "player" || len(variable.Accesses) != 3 { + t.Fatalf("unexpected variable: %#v", variable) + } + if variable.Accesses[1].(*ast.IndexAccess).Index.(*ast.VariableExpr).Name.Name != "index" { + t.Fatalf("unexpected variable index: %#v", variable.Accesses[1]) + } + if _, ok := variable.Accesses[2].(*ast.EmptyIndexAccess); !ok { + t.Fatalf("expected empty index access, got %T", variable.Accesses[2]) + } + + imported := values[1].(*ast.VariableExpr) + if imported.Qualifier.Name != "shop" || imported.Name.Name != "tax" { + t.Fatalf("unexpected imported variable: %#v", imported) + } + call := values[2].(*ast.CallExpr) + if !call.ExplicitParens || len(call.Callee.Parts) != 2 || len(call.Arguments) != 1 { + t.Fatalf("unexpected call: %#v", call) + } + if len(values[3].(*ast.ListExpr).Elements) != 3 { + t.Fatalf("unexpected list: %#v", values[3]) + } + if len(values[4].(*ast.MapExpr).Entries) != 2 { + t.Fatalf("unexpected map: %#v", values[4]) + } + rangeExpression := values[5].(*ast.RangeExpr) + if rangeExpression.Start.(*ast.IntLiteral).Value != 1 || rangeExpression.End.(*ast.IntLiteral).Value != 10 { + t.Fatalf("unexpected range: %#v", rangeExpression) + } +} + +func TestParseStringInterpolationUsesExpressionParser(t *testing.T) { + result := parseTestSource("string.puff", `$message = "Total: {$coins + 10}; User: {shop.$name}"`+"\n") + if len(result.Diagnostics) != 0 { + t.Fatalf("expected no diagnostics, got %#v", result.Diagnostics) + } + + stringExpression := result.File.Declarations[0].(*ast.GlobalAssignment).Value.(*ast.StringExpr) + if len(stringExpression.Parts) != 4 { + t.Fatalf("expected four string parts, got %#v", stringExpression.Parts) + } + first := stringExpression.Parts[1].(*ast.StringInterpolation).Expression + if expressionShape(first) != "(+ $coins 10)" { + t.Fatalf("unexpected first interpolation: %s", expressionShape(first)) + } + second := stringExpression.Parts[3].(*ast.StringInterpolation).Expression.(*ast.VariableExpr) + if second.Qualifier.Name != "shop" || second.Name.Name != "name" { + t.Fatalf("unexpected imported interpolation: %#v", second) + } +} + +func TestParseExpressionSeparatorErrors(t *testing.T) { + tests := []struct { + source string + message string + }{ + {source: "$x = [1 2]\n", message: `Expected "\",\" or \"]\"".`}, + {source: "$x = call(1 2)\n", message: `Expected "\",\" or \")\"".`}, + {source: `$x = {"a" 1}` + "\n", message: `Expected ":".`}, + {source: `$x = {, "a": 1}` + "\n", message: `Expected "expression".`}, + {source: `$x = {"a": 1,, "b": 2}` + "\n", message: `Expected "expression".`}, + {source: "$x = 1 2\n", message: "Expected newline."}, + {source: "$x = [1,,2]\n", message: `Expected "expression".`}, + {source: "$x = call(,1)\n", message: `Expected "expression".`}, + {source: "$x = call(,)\n", message: `Expected "expression".`}, + {source: "$x = ()\n", message: `Expected "expression".`}, + {source: "$x = +\n", message: `Expected "expression".`}, + {source: "$x = 1..2..3\n", message: `Unexpected token: ..`}, + {source: `$x = "value: {1 2}"` + "\n", message: `Expected "}".`}, + } + + for _, test := range tests { + result := parseTestSource("separator.puff", test.source) + if len(result.Diagnostics) != 1 { + t.Fatalf("source %q: expected one diagnostic, got %#v", test.source, result.Diagnostics) + } + if result.Diagnostics[0].Message != test.message { + t.Fatalf("source %q: expected %q, got %#v", test.source, test.message, result.Diagnostics) + } + } +} + +func TestParseImportedVariableSpanIncludesQualifier(t *testing.T) { + result := parseTestSource("span.puff", "$value = shop.$tax\n") + if len(result.Diagnostics) != 0 { + t.Fatalf("expected no diagnostics, got %#v", result.Diagnostics) + } + variable := result.File.Declarations[0].(*ast.GlobalAssignment).Value.(*ast.VariableExpr) + if variable.Span().StartOffset != 9 || variable.Span().EndOffset != 18 { + t.Fatalf("expected imported variable span 9..18, got %#v", variable.Span()) + } +} + +func TestParseDoesNotDuplicateLexerExpressionDiagnostics(t *testing.T) { + tests := []struct { + source string + code diagnostic.Code + }{ + {source: `$x = "value: {}"` + "\n", code: diagnostic.CodeEmptyInterpolation}, + {source: "$x = 1abc\n", code: diagnostic.CodeInvalidNumber}, + {source: "$x = 1abc # comment\n", code: diagnostic.CodeInvalidNumber}, + {source: "$x = @ # comment\r\n", code: diagnostic.CodeInvalidCharacter}, + } + + for _, test := range tests { + result := parseTestSource("lexer-error.puff", test.source) + if len(result.Diagnostics) != 1 || result.Diagnostics[0].Code != test.code { + t.Fatalf("source %q: expected only %s, got %#v", test.source, test.code, result.Diagnostics) + } + } +} + +func TestParseInvalidVariableForms(t *testing.T) { + tests := []struct { + name string + source string + code diagnostic.Code + message string + }{ + { + name: "local at top level", + source: "$_price = 50\n", + code: diagnostic.CodeInvalidTopLevelStatement, + message: "Executable statements are not allowed at the top level.", + }, + { + name: "missing variable name", + source: "$ = 1\n", + code: diagnostic.CodeExpectedToken, + message: `Expected "variable name".`, + }, + { + name: "missing local variable name", + source: "$_ = 1\n", + code: diagnostic.CodeExpectedToken, + message: `Expected "variable name".`, + }, + { + name: "missing assignment value", + source: "fun f\n$_value =\nend\n", + code: diagnostic.CodeExpectedToken, + message: `Expected "expression".`, + }, + { + name: "imported local variable", + source: "fun f\nshop.$_price\nend\n", + code: diagnostic.CodeExpectedToken, + message: `Expected "global variable name".`, + }, + } + + for _, test := range tests { + t.Run(test.name, func(t *testing.T) { + result := parseTestSource("variable.puff", test.source) + if len(result.Diagnostics) != 1 { + t.Fatalf("expected one diagnostic, got %#v", result.Diagnostics) + } + got := result.Diagnostics[0] + if got.Code != test.code || got.Message != test.message { + t.Fatalf("expected %s %q, got %#v", test.code, test.message, got) + } + }) + } +} + +func expressionShape(expression ast.Expression) string { + switch node := expression.(type) { + case *ast.NilLiteral: + return "nil" + case *ast.BoolLiteral: + return fmt.Sprintf("%t", node.Value) + case *ast.IntLiteral: + return fmt.Sprintf("%d", node.Value) + case *ast.FloatLiteral: + return fmt.Sprintf("%g", node.Value) + case *ast.UnaryExpr: + return fmt.Sprintf("(unary %s %s)", operatorShape(node.Operator), expressionShape(node.Operand)) + case *ast.BinaryExpr: + return fmt.Sprintf("(%s %s %s)", operatorShape(node.Operator), expressionShape(node.Left), expressionShape(node.Right)) + case *ast.GroupExpr: + return "(group " + expressionShape(node.Expression) + ")" + case *ast.RangeExpr: + return "(range " + expressionShape(node.Start) + " " + expressionShape(node.End) + ")" + case *ast.VariableExpr: + var builder strings.Builder + if node.Qualifier != nil { + builder.WriteString(node.Qualifier.Name) + builder.WriteByte('.') + } + builder.WriteByte('$') + if node.Local { + builder.WriteByte('_') + } + builder.WriteString(node.Name.Name) + return builder.String() + case *ast.CallExpr: + parts := make([]string, len(node.Callee.Parts)) + for index, part := range node.Callee.Parts { + parts[index] = part.Name + } + return strings.Join(parts, ".") + default: + return fmt.Sprintf("%T", expression) + } +} + +func operatorShape(operator token.Type) string { + switch operator { + case token.Plus: + return "+" + case token.Minus: + return "-" + case token.Star: + return "*" + case token.Slash: + return "/" + case token.Percent: + return "%" + case token.EqualEqual: + return "==" + case token.BangEqual: + return "!=" + case token.Greater: + return ">" + case token.GreaterEq: + return ">=" + case token.Less: + return "<" + case token.LessEq: + return "<=" + case token.And: + return "and" + case token.Or: + return "or" + case token.Not: + return "not" + default: + return string(operator) + } +} diff --git a/src/internal/parser/full_program_test.go b/src/internal/parser/full_program_test.go new file mode 100644 index 0000000..5f3568f --- /dev/null +++ b/src/internal/parser/full_program_test.go @@ -0,0 +1,199 @@ +package parser + +import ( + "fmt" + "os" + "path/filepath" + "strconv" + "strings" + "testing" + + "github.com/puff-lang/puff/internal/ast" +) + +func TestParseFullGrammarExampleGolden(t *testing.T) { + input, err := os.ReadFile(filepath.Join("testdata", "full_program.puff")) + if err != nil { + t.Fatalf("read fixture: %v", err) + } + want, err := os.ReadFile(filepath.Join("testdata", "full_program.golden")) + if err != nil { + t.Fatalf("read golden: %v", err) + } + + result := parseTestSource("full_program.puff", string(input)) + if len(result.Diagnostics) != 0 { + t.Fatalf("expected no diagnostics, got %#v", result.Diagnostics) + } + wantText := strings.ReplaceAll(string(want), "\r\n", "\n") + if got := renderFullProgram(result.File); got != wantText { + t.Fatalf("unexpected full AST\nwant:\n%s\ngot:\n%s", wantText, got) + } +} + +func renderFullProgram(file *ast.File) string { + var builder strings.Builder + fmt.Fprintf(&builder, "requirements %d\n", len(file.Requirements)) + fmt.Fprintf(&builder, "declarations %d\n", len(file.Declarations)) + for _, declaration := range file.Declarations { + switch node := declaration.(type) { + case *ast.GlobalAssignment: + if node.Public { + builder.WriteString("pub ") + } + fmt.Fprintf(&builder, "global %s = %s\n", fullVariable(node.Target), fullExpression(node.Value)) + case *ast.FunctionDecl: + if node.Public { + builder.WriteString("pub ") + } + fmt.Fprintf(&builder, "fun %s(", node.Name.Name) + for index, parameter := range node.Parameters { + if index > 0 { + builder.WriteString(", ") + } + builder.WriteString(parameter.Name.Name) + if parameter.Type != nil { + fmt.Fprintf(&builder, ": %s", renderType(parameter.Type)) + } + } + builder.WriteByte(')') + if node.ReturnType != nil { + fmt.Fprintf(&builder, " -> %s", renderType(node.ReturnType)) + } + builder.WriteByte('\n') + renderFullBlock(&builder, node.Body, " ") + case *ast.EventDecl: + names := make([]string, len(node.Name)) + for index, name := range node.Name { + names[index] = name.Name + } + fmt.Fprintf(&builder, "event %s\n", strings.Join(names, " ")) + renderFullBlock(&builder, node.Body, " ") + } + } + return builder.String() +} + +func renderFullBlock(builder *strings.Builder, block ast.Block, indent string) { + for _, statement := range block.Statements { + switch node := statement.(type) { + case *ast.AssignmentStmt: + fmt.Fprintf(builder, "%sassign %s = %s\n", indent, fullVariable(node.Target), fullExpression(node.Value)) + case *ast.AddStmt: + fmt.Fprintf(builder, "%sadd %s\n", indent, fullExpression(node.Value)) + case *ast.ReturnStmt: + if node.Value == nil { + fmt.Fprintf(builder, "%sreturn\n", indent) + } else { + fmt.Fprintf(builder, "%sreturn %s\n", indent, fullExpression(node.Value)) + } + case *ast.StopStmt: + fmt.Fprintf(builder, "%sstop\n", indent) + case *ast.ExprStmt: + fmt.Fprintf(builder, "%sexpr %s\n", indent, fullExpression(node.Expression)) + case *ast.EffectStmt: + name := "" + if len(node.Tokens) > 0 { + name = node.Tokens[0].Lexeme + } + fmt.Fprintf(builder, "%seffect %s\n", indent, name) + case *ast.IfStmt: + fmt.Fprintf(builder, "%sif %s\n", indent, fullExpression(node.Condition)) + renderFullBlock(builder, node.Then, indent+" ") + for _, clause := range node.ElseIf { + fmt.Fprintf(builder, "%selse if %s\n", indent, fullExpression(clause.Condition)) + renderFullBlock(builder, clause.Body, indent+" ") + } + if node.Else != nil { + fmt.Fprintf(builder, "%selse\n", indent) + renderFullBlock(builder, *node.Else, indent+" ") + } + case *ast.LoopTimesStmt: + fmt.Fprintf(builder, "%sloop %s times\n", indent, fullExpression(node.Count)) + renderFullBlock(builder, node.Body, indent+" ") + case *ast.LoopRangeStmt: + fmt.Fprintf(builder, "%sloop numbers from %s to %s\n", indent, fullExpression(node.Start), fullExpression(node.End)) + renderFullBlock(builder, node.Body, indent+" ") + case *ast.LoopPlayersStmt: + fmt.Fprintf(builder, "%sloop players\n", indent) + renderFullBlock(builder, node.Body, indent+" ") + case *ast.LoopEntitiesStmt: + fmt.Fprintf(builder, "%sloop entities in radius %s around %s\n", indent, fullExpression(node.Radius), fullExpression(node.Around)) + renderFullBlock(builder, node.Body, indent+" ") + } + } +} + +func fullExpression(expression ast.Expression) string { + switch node := expression.(type) { + case *ast.NilLiteral: + return "nil" + case *ast.BoolLiteral: + return strconv.FormatBool(node.Value) + case *ast.IntLiteral: + return strconv.FormatInt(node.Value, 10) + case *ast.FloatLiteral: + return strconv.FormatFloat(node.Value, 'g', -1, 64) + case *ast.StringExpr: + return strconv.Quote(stringValue(node)) + case *ast.VariableExpr: + return fullVariable(node) + case *ast.CallExpr: + parts := make([]string, len(node.Callee.Parts)) + for index, part := range node.Callee.Parts { + parts[index] = part.Name + } + name := strings.Join(parts, ".") + if !node.ExplicitParens { + return name + } + arguments := make([]string, len(node.Arguments)) + for index, argument := range node.Arguments { + arguments[index] = fullExpression(argument) + } + return name + "(" + strings.Join(arguments, ", ") + ")" + case *ast.UnaryExpr: + return "(" + operatorShape(node.Operator) + " " + fullExpression(node.Operand) + ")" + case *ast.BinaryExpr: + return "(" + operatorShape(node.Operator) + " " + fullExpression(node.Left) + " " + fullExpression(node.Right) + ")" + case *ast.GroupExpr: + return "(group " + fullExpression(node.Expression) + ")" + case *ast.RangeExpr: + return "(range " + fullExpression(node.Start) + " " + fullExpression(node.End) + ")" + case *ast.PatternExpr: + parts := make([]string, len(node.Tokens)) + for index, tok := range node.Tokens { + parts[index] = tok.Lexeme + } + return "pattern(" + strings.Join(parts, " ") + ")" + default: + return fmt.Sprintf("%T", expression) + } +} + +func fullVariable(variable *ast.VariableExpr) string { + var builder strings.Builder + if variable.Qualifier != nil { + builder.WriteString(variable.Qualifier.Name) + builder.WriteByte('.') + } + builder.WriteByte('$') + if variable.Local { + builder.WriteByte('_') + } + builder.WriteString(variable.Name.Name) + for _, access := range variable.Accesses { + switch node := access.(type) { + case *ast.FieldAccess: + builder.WriteByte('.') + builder.WriteString(node.Field.Name) + case *ast.EmptyIndexAccess: + builder.WriteString("[]") + case *ast.IndexAccess: + builder.WriteByte('[') + builder.WriteString(fullExpression(node.Index)) + builder.WriteByte(']') + } + } + return builder.String() +} diff --git a/src/internal/parser/parser.go b/src/internal/parser/parser.go index 22c70a1..128eae4 100644 --- a/src/internal/parser/parser.go +++ b/src/internal/parser/parser.go @@ -16,10 +16,11 @@ type Result struct { } type parser struct { - file source.File - tokens []token.Token - current int - diagnostics []diagnostic.Diagnostic + file source.File + tokens []token.Token + current int + diagnostics []diagnostic.Diagnostic + expressionStops map[token.Type]bool } func Parse(file source.File, lexed lexer.Result) Result { @@ -235,6 +236,14 @@ func (parser *parser) peekNext() token.Token { return parser.tokens[parser.current+1] } +func (parser *parser) peekAt(distance int) token.Token { + index := parser.current + distance + if index >= len(parser.tokens) { + return parser.tokens[len(parser.tokens)-1] + } + return parser.tokens[index] +} + func (parser *parser) previous() token.Token { return parser.tokens[parser.current-1] } diff --git a/src/internal/parser/statement.go b/src/internal/parser/statement.go new file mode 100644 index 0000000..67a96a2 --- /dev/null +++ b/src/internal/parser/statement.go @@ -0,0 +1,317 @@ +package parser + +import ( + "github.com/puff-lang/puff/internal/ast" + "github.com/puff-lang/puff/internal/diagnostic" + "github.com/puff-lang/puff/internal/token" +) + +func (parser *parser) parseBlock(allowElse bool) ast.Block { + start := parser.peek().StartOffset + var statements []ast.Statement + parser.skipNewlines() + for !parser.atEnd() && !parser.check(token.End) { + if parser.check(token.Else) { + if allowElse { + break + } + parser.reportUnexpected(parser.peek(), "else can only appear inside an if block.") + parser.synchronizeLine() + parser.match(token.Newline) + continue + } + statement := parser.parseStatement() + if statement != nil { + statements = append(statements, statement) + } + parser.skipNewlines() + } + return ast.Block{ + NodeBase: parser.base(start, parser.peek().StartOffset), + Statements: statements, + } +} + +func (parser *parser) parseStatement() ast.Statement { + switch parser.peek().Type { + case token.Dollar: + return parser.parseVariableStatement(nil) + case token.Add: + return parser.parseAddStatement() + case token.If: + return parser.parseIfStatement() + case token.Loop: + return parser.parseLoopStatement() + case token.Return: + return parser.parseReturnStatement() + case token.Stop: + return parser.parseStopStatement() + } + + if parser.checkName() && parser.peekAt(1).Type == token.Dot && parser.peekAt(2).Type == token.Dollar { + qualifier := parser.parseIdentifier() + parser.advance() + return parser.parseVariableStatement(qualifier) + } + return parser.parseExpressionOrEffectStatement() +} + +func (parser *parser) parseVariableStatement(qualifier *ast.Identifier) ast.Statement { + start := parser.peek().StartOffset + expression := parser.parseVariable(qualifier) + variable, _ := expression.(*ast.VariableExpr) + if !parser.match(token.Equal) { + parser.requireLineEnd() + return &ast.ExprStmt{ + NodeBase: parser.base(start, variable.Span().EndOffset), + Expression: variable, + } + } + + value := parser.parseExpressionUntil(token.Newline) + end := parser.statementEnd(start, value) + parser.requireLineEnd() + return &ast.AssignmentStmt{ + NodeBase: parser.base(start, end), + Target: variable, + Value: value, + } +} + +func (parser *parser) parseAddStatement() ast.Statement { + start := parser.advance().StartOffset + value := parser.parseExpressionUntil(token.To) + if !parser.match(token.To) { + parser.reportExpected("to", "") + } + + var target ast.Assignable + if parser.check(token.Dollar) { + target, _ = parser.parseVariable(nil).(ast.Assignable) + } else if parser.checkName() { + target = parser.parseAccessExpression() + } else { + parser.reportExpected("assignable target", "") + parser.synchronizeLine() + } + end := start + if target != nil { + end = target.Span().EndOffset + } else if value != nil { + end = value.Span().EndOffset + } + parser.requireLineEnd() + return &ast.AddStmt{ + NodeBase: parser.base(start, end), + Value: value, + Target: target, + } +} + +func (parser *parser) parseAccessExpression() ast.Assignable { + start := parser.current + for !parser.check(token.Newline) && !parser.atEnd() { + parser.advance() + } + tokens := append([]token.Token(nil), parser.tokens[start:parser.current]...) + if len(tokens) == 0 { + return nil + } + return &ast.AccessExpr{ + NodeBase: parser.base(tokens[0].StartOffset, tokens[len(tokens)-1].EndOffset), + Tokens: tokens, + } +} + +func (parser *parser) parseIfStatement() ast.Statement { + start := parser.advance().StartOffset + condition := parser.parseExpressionUntil(token.Newline) + parser.requireLineEnd() + thenBlock := parser.parseBlock(true) + + var elseIf []ast.ElseIfClause + var elseBlock *ast.Block + for parser.match(token.Else) { + clauseStart := parser.previous().StartOffset + if parser.match(token.If) { + clauseCondition := parser.parseExpressionUntil(token.Newline) + parser.requireLineEnd() + body := parser.parseBlock(true) + elseIf = append(elseIf, ast.ElseIfClause{ + NodeBase: parser.base(clauseStart, body.Span().EndOffset), + Condition: clauseCondition, + Body: body, + }) + continue + } + + parser.requireLineEnd() + body := parser.parseBlock(false) + elseBlock = &body + break + } + + end := parser.consumeBlockEnd(start) + return &ast.IfStmt{ + NodeBase: parser.base(start, end), + Condition: condition, + Then: thenBlock, + ElseIf: elseIf, + Else: elseBlock, + } +} + +func (parser *parser) parseLoopStatement() ast.Statement { + start := parser.advance().StartOffset + switch { + case parser.match(token.Players): + parser.requireLineEnd() + body := parser.parseBlock(false) + end := parser.consumeBlockEnd(start) + return &ast.LoopPlayersStmt{NodeBase: parser.base(start, end), Body: body} + case parser.match(token.Numbers): + if !parser.match(token.From) { + parser.reportExpected("from", "") + } + rangeStart := parser.parseExpressionUntil(token.To, token.Newline) + hasTo := parser.match(token.To) + if !hasTo { + if rangeStart != nil { + parser.reportExpected("to", "") + } + parser.synchronizeLine() + } + var rangeEnd ast.Expression + if hasTo { + rangeEnd = parser.parseExpressionUntil(token.Newline) + } + parser.requireLineEnd() + body := parser.parseBlock(false) + end := parser.consumeBlockEnd(start) + return &ast.LoopRangeStmt{ + NodeBase: parser.base(start, end), + Start: rangeStart, + End: rangeEnd, + Body: body, + } + case parser.match(token.Entities): + if !parser.match(token.In) { + parser.reportExpected("in", "") + } + if !parser.match(token.Radius) { + parser.reportExpected("radius", "") + } + radius := parser.parseExpressionUntil(token.Around, token.Newline) + hasAround := parser.match(token.Around) + if !hasAround { + if radius != nil { + parser.reportExpected("around", "") + } + parser.synchronizeLine() + } + var around ast.Expression + if hasAround { + around = parser.parseExpressionUntil(token.Newline) + } + parser.requireLineEnd() + body := parser.parseBlock(false) + end := parser.consumeBlockEnd(start) + return &ast.LoopEntitiesStmt{ + NodeBase: parser.base(start, end), + Radius: radius, + Around: around, + Body: body, + } + default: + count := parser.parseExpressionUntil(token.Times, token.Newline) + if !parser.match(token.Times) { + if count != nil { + parser.reportExpected("times", "") + } + parser.synchronizeLine() + } + parser.requireLineEnd() + body := parser.parseBlock(false) + end := parser.consumeBlockEnd(start) + return &ast.LoopTimesStmt{ + NodeBase: parser.base(start, end), + Count: count, + Body: body, + } + } +} + +func (parser *parser) parseReturnStatement() ast.Statement { + startToken := parser.advance() + var value ast.Expression + if !parser.check(token.Newline) && !parser.atEnd() { + value = parser.parseExpressionUntil(token.Newline) + } + end := startToken.EndOffset + if value != nil { + end = value.Span().EndOffset + } + parser.requireLineEnd() + return &ast.ReturnStmt{ + NodeBase: parser.base(startToken.StartOffset, end), + Value: value, + } +} + +func (parser *parser) parseStopStatement() ast.Statement { + tok := parser.advance() + parser.requireLineEnd() + return &ast.StopStmt{NodeBase: parser.base(tok.StartOffset, tok.EndOffset)} +} + +func (parser *parser) parseExpressionOrEffectStatement() ast.Statement { + startIndex := parser.current + expression := parser.parseExpressionUntil(token.Newline) + if pattern, ok := expression.(*ast.PatternExpr); ok { + parser.synchronizeLine() + tokens := append([]token.Token(nil), parser.tokens[startIndex:parser.current]...) + end := pattern.Span().EndOffset + if len(tokens) > 0 { + end = tokens[len(tokens)-1].EndOffset + } + parser.requireLineEnd() + return &ast.EffectStmt{ + NodeBase: parser.base(parser.tokens[startIndex].StartOffset, end), + Tokens: tokens, + } + } + + end := parser.statementEnd(parser.tokens[startIndex].StartOffset, expression) + parser.requireLineEnd() + return &ast.ExprStmt{ + NodeBase: parser.base(parser.tokens[startIndex].StartOffset, end), + Expression: expression, + } +} + +func (parser *parser) consumeBlockEnd(openingStart int) int { + if parser.match(token.End) { + end := parser.previous().EndOffset + parser.requireLineEnd() + return end + } + eof := parser.peek() + parser.report( + diagnostic.CodeExpectedEnd, + `Expected "end" before end of file.`, + "Add end to close the block.", + eof.StartOffset, + eof.EndOffset, + ) + return eof.EndOffset +} + +func (parser *parser) statementEnd(start int, expression ast.Expression) int { + if expression != nil { + return expression.Span().EndOffset + } + if parser.current > 0 { + return parser.previous().EndOffset + } + return start +} diff --git a/src/internal/parser/statement_test.go b/src/internal/parser/statement_test.go new file mode 100644 index 0000000..60bc2da --- /dev/null +++ b/src/internal/parser/statement_test.go @@ -0,0 +1,136 @@ +package parser + +import ( + "testing" + + "github.com/puff-lang/puff/internal/ast" +) + +func TestParseStatementsAndBranches(t *testing.T) { + result := parseTestSource("statements.puff", ` +fun example +$_price = 50 +add 1 to $coins +if $coins >= 100 +return true +else if $coins >= 50 +return false +else +stop +end +shop.Run +send "Hello" to player +end +`) + + if len(result.Diagnostics) != 0 { + t.Fatalf("expected no diagnostics, got %#v", result.Diagnostics) + } + body := result.File.Declarations[0].(*ast.FunctionDecl).Body + if len(body.Statements) != 5 { + t.Fatalf("expected five statements, got %#v", body.Statements) + } + assignment := body.Statements[0].(*ast.AssignmentStmt) + if !assignment.Target.Local || assignment.Target.Name.Name != "price" { + t.Fatalf("unexpected assignment: %#v", assignment) + } + if body.Statements[1].(*ast.AddStmt).Target.(*ast.VariableExpr).Name.Name != "coins" { + t.Fatalf("unexpected add statement: %#v", body.Statements[1]) + } + conditional := body.Statements[2].(*ast.IfStmt) + if len(conditional.Then.Statements) != 1 || len(conditional.ElseIf) != 1 || len(conditional.Else.Statements) != 1 { + t.Fatalf("unexpected conditional: %#v", conditional) + } + if _, ok := body.Statements[3].(*ast.ExprStmt); !ok { + t.Fatalf("expected expression statement, got %T", body.Statements[3]) + } + if _, ok := body.Statements[4].(*ast.EffectStmt); !ok { + t.Fatalf("expected effect statement, got %T", body.Statements[4]) + } +} + +func TestParseAllLoopForms(t *testing.T) { + result := parseTestSource("loops.puff", ` +on load +loop 3 times +stop +end +loop numbers from 10 to 1 +stop +end +loop players +stop +end +loop entities in radius 10 around player +stop +end +end +`) + + if len(result.Diagnostics) != 0 { + t.Fatalf("expected no diagnostics, got %#v", result.Diagnostics) + } + body := result.File.Declarations[0].(*ast.EventDecl).Body + if len(body.Statements) != 4 { + t.Fatalf("expected four loops, got %#v", body.Statements) + } + if body.Statements[0].(*ast.LoopTimesStmt).Count.(*ast.IntLiteral).Value != 3 { + t.Fatalf("unexpected times loop: %#v", body.Statements[0]) + } + rangeLoop := body.Statements[1].(*ast.LoopRangeStmt) + if rangeLoop.Start.(*ast.IntLiteral).Value != 10 || rangeLoop.End.(*ast.IntLiteral).Value != 1 { + t.Fatalf("unexpected range loop: %#v", rangeLoop) + } + if _, ok := body.Statements[2].(*ast.LoopPlayersStmt); !ok { + t.Fatalf("expected players loop, got %T", body.Statements[2]) + } + entities := body.Statements[3].(*ast.LoopEntitiesStmt) + if entities.Radius.(*ast.IntLiteral).Value != 10 || expressionShape(entities.Around) != "player" { + t.Fatalf("unexpected entities loop: %#v", entities) + } +} + +func TestParseAddPatternTargetAndCondition(t *testing.T) { + result := parseTestSource("patterns.puff", ` +on join +add amount to coins of target +if coins of player >= 100 +stop +end +end +`) + + if len(result.Diagnostics) != 0 { + t.Fatalf("expected no diagnostics, got %#v", result.Diagnostics) + } + body := result.File.Declarations[0].(*ast.EventDecl).Body + add := body.Statements[0].(*ast.AddStmt) + if _, ok := add.Value.(*ast.CallExpr); !ok { + t.Fatalf("expected amount expression, got %T", add.Value) + } + if len(add.Target.(*ast.AccessExpr).Tokens) != 3 { + t.Fatalf("unexpected access pattern: %#v", add.Target) + } + condition := body.Statements[1].(*ast.IfStmt).Condition.(*ast.BinaryExpr) + if _, ok := condition.Left.(*ast.PatternExpr); !ok { + t.Fatalf("expected pattern condition, got %T", condition.Left) + } +} + +func TestParseRejectsMissingLoopOperands(t *testing.T) { + tests := []string{ + "on load\nloop times\nend\nend\n", + "on load\nloop numbers from\nend\nend\n", + "on load\nloop numbers from to 3\nend\nend\n", + "on load\nloop numbers from 1 to\nend\nend\n", + "on load\nloop entities in radius around player\nend\nend\n", + "on load\nloop entities in radius 1 around\nend\nend\n", + } + + for _, input := range tests { + result := parseTestSource("loop.puff", input) + if len(result.Diagnostics) != 1 || result.Diagnostics[0].Message != `Expected "expression".` { + t.Fatalf("source %q: expected missing expression diagnostic, got %#v", input, result.Diagnostics) + } + } +} diff --git a/src/internal/parser/testdata/full_program.golden b/src/internal/parser/testdata/full_program.golden new file mode 100644 index 0000000..539d0a8 --- /dev/null +++ b/src/internal/parser/testdata/full_program.golden @@ -0,0 +1,28 @@ +requirements 2 +declarations 7 +pub global $default_coins = 100 +global $shop.name = "Main Shop" +pub fun serverName() -> string + return "Lobby" +fun hasEnoughCoins(player: Player, price: int) -> bool + if (>= $player.coins price) + return true + else + return false +event load + effect send + effect send + expr shop.Run +event tick +event join + assign $_price = 50 + if hasEnoughCoins(player, $_price) + effect send + else if (>= $player.coins 10) + effect send + else + effect send + loop numbers from 1 to 3 + effect send + loop players + effect send diff --git a/src/internal/parser/testdata/full_program.puff b/src/internal/parser/testdata/full_program.puff new file mode 100644 index 0000000..dc3b1db --- /dev/null +++ b/src/internal/parser/testdata/full_program.puff @@ -0,0 +1,49 @@ +# namespace: example +# tags: load, tick + +require "abc/shop" +require "github.com/123/123" as lib123 + +pub $default_coins = 100 +$shop.name = "Main Shop" + +pub fun serverName -> string + return "Lobby" +end + +fun hasEnoughCoins(player: Player, price: int) -> bool + if $player.coins >= price + return true + else + return false + end +end + +on load + send "Loaded {{namespace}}: {serverName}" to console + send "Shop tax: {shop.$tax}" to console + shop.Run +end + +on tick +end + +on join + $_price = 50 + + if hasEnoughCoins(player, $_price) + send "You can buy this." to player + else if $player.coins >= 10 + send "You are close." to player + else + send "You need more coins." to player + end + + loop numbers from 1 to 3 + send "Index {loop.index}: {loop.value}" to player + end + + loop players + send "Welcome, {loop.player}" to console + end +end diff --git a/src/internal/parser/top_level.go b/src/internal/parser/top_level.go index d8ea2d9..d1b876f 100644 --- a/src/internal/parser/top_level.go +++ b/src/internal/parser/top_level.go @@ -53,7 +53,8 @@ func (parser *parser) parseFunction(public bool) *ast.FunctionDecl { } parser.requireLineEnd() - body, end := parser.scanBlock(start) + body := parser.parseBlock(false) + end := parser.consumeBlockEnd(start) return &ast.FunctionDecl{ NodeBase: parser.base(start, end), Public: public, @@ -146,7 +147,8 @@ func (parser *parser) parseEvent() *ast.EventDecl { } parser.requireLineEnd() - body, end := parser.scanBlock(start) + body := parser.parseBlock(false) + end := parser.consumeBlockEnd(start) return &ast.EventDecl{ NodeBase: parser.base(start, end), Name: name, @@ -159,20 +161,21 @@ func (parser *parser) parseGlobal(public bool) *ast.GlobalAssignment { if public { parser.advance() } - target := parser.parseGlobalVariable() + target, _ := parser.parseVariable(nil).(*ast.VariableExpr) + if target != nil && target.Local && target.Name.Name != "" { + parser.report( + diagnostic.CodeInvalidTopLevelStatement, + "Executable statements are not allowed at the top level.", + "Move this statement into an event or function.", + target.Span().StartOffset, + target.Span().EndOffset, + ) + } if !parser.match(token.Equal) { parser.reportExpected("=", "") } - - expressionStart := parser.current - end := parser.lineEndOffset(start) - for !parser.check(token.Newline) && !parser.atEnd() { - parser.advance() - } - var value ast.Expression - if expressionStart < parser.current { - value = parser.parseDeferredExpression(parser.tokens[expressionStart:parser.current]) - } + value := parser.parseExpressionUntil(token.Newline) + end := parser.statementEnd(start, value) parser.requireLineEnd() return &ast.GlobalAssignment{ @@ -183,173 +186,12 @@ func (parser *parser) parseGlobal(public bool) *ast.GlobalAssignment { } } -func (parser *parser) parseGlobalVariable() *ast.VariableExpr { - start := parser.peek().StartOffset - if !parser.match(token.Dollar) { - parser.reportExpected("$", "") - return nil - } - if !parser.checkName() { - parser.reportExpected("global variable name", "") - return nil - } - - name := parser.parseIdentifier() - var accesses []ast.VariableAccess - for { - switch { - case parser.match(token.Dot): - if !parser.checkName() { - parser.reportExpected("field name", "") - return &ast.VariableExpr{NodeBase: parser.base(start, parser.previous().EndOffset), Name: *name, Accesses: accesses} - } - field := parser.parseIdentifier() - accesses = append(accesses, &ast.FieldAccess{ - NodeBase: parser.base(parser.tokens[parser.current-2].StartOffset, field.Span().EndOffset), - Field: *field, - }) - case parser.match(token.LBracket): - accessStart := parser.previous().StartOffset - if parser.match(token.RBracket) { - accesses = append(accesses, &ast.EmptyIndexAccess{NodeBase: parser.base(accessStart, parser.previous().EndOffset)}) - continue - } - indexStart := parser.current - for !parser.check(token.RBracket) && !parser.check(token.Newline) && !parser.atEnd() { - parser.advance() - } - var index ast.Expression - if indexStart < parser.current { - index = parser.parseDeferredExpression(parser.tokens[indexStart:parser.current]) - } - if !parser.match(token.RBracket) { - parser.reportExpected("]", "") - } - accesses = append(accesses, &ast.IndexAccess{ - NodeBase: parser.base(accessStart, parser.previous().EndOffset), - Index: index, - }) - default: - end := name.Span().EndOffset - if len(accesses) > 0 { - end = accesses[len(accesses)-1].Span().EndOffset - } - return &ast.VariableExpr{ - NodeBase: parser.base(start, end), - Name: *name, - Accesses: accesses, - } - } - } -} - -func (parser *parser) parseDeferredExpression(tokens []token.Token) ast.Expression { - start := tokens[0].StartOffset - end := tokens[len(tokens)-1].EndOffset - base := parser.base(start, end) - - if len(tokens) == 1 { - switch tokens[0].Type { - case token.Nil: - return &ast.NilLiteral{NodeBase: base} - case token.True: - return &ast.BoolLiteral{NodeBase: base, Value: true} - case token.False: - return &ast.BoolLiteral{NodeBase: base, Value: false} - case token.Int: - value, _ := tokens[0].Value.(int) - return &ast.IntLiteral{NodeBase: base, Value: int64(value)} - case token.Float: - value, _ := tokens[0].Value.(float64) - return &ast.FloatLiteral{NodeBase: base, Value: value} - } - } - if tokens[0].Type == token.StringStart && tokens[len(tokens)-1].Type == token.StringEnd { - return parser.stringFromTokens(tokens) - } - if len(tokens) == 2 && tokens[0].Type == token.LBracket && tokens[1].Type == token.RBracket { - return &ast.ListExpr{NodeBase: base} - } - - return &ast.PatternExpr{NodeBase: base, Tokens: append([]token.Token(nil), tokens...)} -} - func (parser *parser) parseString() *ast.StringExpr { if !parser.check(token.StringStart) { parser.reportExpected("string", "") return nil } - - start := parser.current - parser.advance() - depth := 0 - for !parser.atEnd() { - switch parser.peek().Type { - case token.InterpStart: - depth++ - case token.InterpEnd: - depth-- - case token.StringEnd: - if depth == 0 { - parser.advance() - return parser.stringFromTokens(parser.tokens[start:parser.current]) - } - case token.Newline: - return parser.stringFromTokens(parser.tokens[start:parser.current]) - } - parser.advance() - } - return parser.stringFromTokens(parser.tokens[start:parser.current]) -} - -func (parser *parser) stringFromTokens(tokens []token.Token) *ast.StringExpr { - if len(tokens) == 0 { - return nil - } - - expression := &ast.StringExpr{ - NodeBase: parser.base(tokens[0].StartOffset, tokens[len(tokens)-1].EndOffset), - } - if len(tokens[0].Lexeme) > 0 { - expression.Quote = tokens[0].Lexeme[0] - } - - endIndex := len(tokens) - if tokens[len(tokens)-1].Type == token.StringEnd { - endIndex-- - } - for index := 1; index < endIndex; index++ { - tok := tokens[index] - switch tok.Type { - case token.StringText: - value, _ := tok.Value.(string) - expression.Parts = append(expression.Parts, &ast.StringText{ - NodeBase: parser.base(tok.StartOffset, tok.EndOffset), - Raw: tok.Lexeme, - Value: value, - }) - case token.InterpStart: - interpolationStart := index + 1 - for index+1 < len(tokens)-1 && tokens[index+1].Type != token.InterpEnd { - index++ - } - parts := tokens[interpolationStart : index+1] - var value ast.Expression - if len(parts) > 0 { - value = parser.parseDeferredExpression(parts) - } - end := tok.EndOffset - if index+1 < len(tokens) && tokens[index+1].Type == token.InterpEnd { - index++ - end = tokens[index].EndOffset - } - expression.Parts = append(expression.Parts, &ast.StringInterpolation{ - NodeBase: parser.base(tok.StartOffset, end), - Expression: value, - }) - } - } - + expression, _ := parser.parseStringExpression().(*ast.StringExpr) return expression } @@ -394,49 +236,6 @@ func (parser *parser) checkName() bool { } } -func (parser *parser) scanBlock(openingStart int) (ast.Block, int) { - bodyStart := parser.peek().StartOffset - depth := 0 - lineStart := true - for !parser.atEnd() { - tok := parser.peek() - if lineStart { - switch tok.Type { - case token.If, token.Loop: - depth++ - case token.Else: - if depth == 0 { - parser.reportUnexpected(tok, "else can only appear inside an if block.") - } - case token.End: - if depth == 0 { - endToken := parser.advance() - parser.requireLineEnd() - return ast.Block{NodeBase: parser.base(bodyStart, endToken.StartOffset)}, endToken.EndOffset - } - depth-- - } - lineStart = false - } - - if parser.match(token.Newline) { - lineStart = true - continue - } - parser.advance() - } - - eof := parser.peek() - parser.report( - diagnostic.CodeExpectedEnd, - `Expected "end" before end of file.`, - "Add end to close the block.", - eof.StartOffset, - eof.EndOffset, - ) - return ast.Block{NodeBase: parser.base(bodyStart, eof.StartOffset)}, eof.EndOffset -} - func (parser *parser) requireLineEnd() { if parser.match(token.Newline) || parser.atEnd() { return diff --git a/src/internal/parser/top_level_test.go b/src/internal/parser/top_level_test.go index cb00267..2997dad 100644 --- a/src/internal/parser/top_level_test.go +++ b/src/internal/parser/top_level_test.go @@ -12,7 +12,6 @@ import ( "github.com/puff-lang/puff/internal/diagnostic" "github.com/puff-lang/puff/internal/lexer" "github.com/puff-lang/puff/internal/source" - "github.com/puff-lang/puff/internal/token" ) func TestParseTopLevelGolden(t *testing.T) { @@ -30,7 +29,8 @@ func TestParseTopLevelGolden(t *testing.T) { if len(result.Diagnostics) != 0 { t.Fatalf("expected no diagnostics, got %#v", result.Diagnostics) } - if got := renderFile(result.File); got != string(want) { + wantText := strings.ReplaceAll(string(want), "\r\n", "\n") + if got := renderFile(result.File); got != wantText { t.Fatalf("unexpected AST\nwant:\n%s\ngot:\n%s", want, got) } if result.File.Metadata[0].Span().StartOffset != 0 || result.File.Metadata[0].Span().EndOffset != 20 { @@ -126,8 +126,8 @@ pub $tax = 0.1 } stats := result.File.Declarations[2].(*ast.GlobalAssignment) - index := stats.Target.Accesses[0].(*ast.IndexAccess).Index.(*ast.PatternExpr) - if len(index.Tokens) != 2 || index.Tokens[0].Type != token.Dollar || index.Tokens[1].Lexeme != "key" { + index := stats.Target.Accesses[0].(*ast.IndexAccess).Index.(*ast.VariableExpr) + if index.Name.Name != "key" || index.Local { t.Fatalf("unexpected index expression: %#v", index) } if _, ok := stats.Value.(*ast.NilLiteral); !ok {