diff --git a/src/ast/AST.go b/src/ast/AST.go index f493a08..30fef4e 100644 --- a/src/ast/AST.go +++ b/src/ast/AST.go @@ -47,6 +47,7 @@ type ( Token token.Token } Switch struct { + Head *expression.Expression Cases []Case } ) \ No newline at end of file diff --git a/src/ast/parseSwitch.go b/src/ast/parseSwitch.go index 4254d5a..5ef8bb5 100644 --- a/src/ast/parseSwitch.go +++ b/src/ast/parseSwitch.go @@ -2,13 +2,17 @@ package ast import ( "git.urbach.dev/cli/q/src/errors" + "git.urbach.dev/cli/q/src/expression" "git.urbach.dev/cli/q/src/fs" "git.urbach.dev/cli/q/src/token" ) func parseSwitch(tokens token.List, file *fs.File) (Node, error) { - blockStart := tokens.IndexKind(token.BlockStart) - blockEnd := tokens.LastIndexKind(token.BlockEnd) + var ( + head *expression.Expression + blockStart = tokens.IndexKind(token.BlockStart) + blockEnd = tokens.LastIndexKind(token.BlockEnd) + ) if blockStart == -1 { return nil, errors.NewAt(MissingBlockStart, file, tokens[0].End()) @@ -18,6 +22,12 @@ func parseSwitch(tokens token.List, file *fs.File) (Node, error) { return nil, errors.NewAt(MissingBlockEnd, file, tokens[len(tokens)-1].End()) } + headTokens := tokens[1:blockStart] + + if len(headTokens) > 0 { + head = expression.Parse(headTokens) + } + body := tokens[blockStart+1 : blockEnd] if len(body) == 0 { @@ -25,5 +35,5 @@ func parseSwitch(tokens token.List, file *fs.File) (Node, error) { } cases, err := parseCases(body, file) - return &Switch{Cases: cases}, err + return &Switch{Head: head, Cases: cases}, err } \ No newline at end of file diff --git a/src/core/compileSwitch.go b/src/core/compileSwitch.go index 304550a..133bfb4 100644 --- a/src/core/compileSwitch.go +++ b/src/core/compileSwitch.go @@ -8,13 +8,26 @@ import ( // compileSwitch compiles a multi-branch instruction. func (f *Function) compileSwitch(s *ast.Switch) error { f.Count.Switch++ - exitLabel := f.CreateLabel("switch.exit", f.Count.Switch) - exitBlock := ssa.NewBlock(exitLabel) + + var ( + head ssa.Value + err error + exitLabel = f.CreateLabel("switch.exit", f.Count.Switch) + exitBlock = ssa.NewBlock(exitLabel) + ) + + if s.Head != nil { + head, err = f.evaluateRight(s.Head) + + if err != nil { + return err + } + } for i, branch := range s.Cases { if branch.Condition == nil { before := f.Block().Identifiers.Before - err := f.compileAST(branch.Body) + err = f.compileAST(branch.Body) if err != nil { return err @@ -38,10 +51,34 @@ func (f *Function) compileSwitch(s *ast.Switch) error { elseBlock = exitBlock } - err := f.compileCondition(branch.Condition, thenBlock, elseBlock) + if head != nil { + caseValue, err := f.evaluateRight(branch.Condition) - if err != nil { - return err + if err != nil { + return err + } + + condition, err := f.equal(head, caseValue, branch.Condition.Source()) + + if err != nil { + return err + } + + block := f.Block() + block.AddSuccessor(thenBlock) + block.AddSuccessor(elseBlock) + + block.Append(&ssa.Branch{ + Condition: condition, + Then: thenBlock, + Else: elseBlock, + }) + } else { + err = f.compileCondition(branch.Condition, thenBlock, elseBlock) + + if err != nil { + return err + } } f.AddBlock(thenBlock) diff --git a/src/core/equal.go b/src/core/equal.go new file mode 100644 index 0000000..9de275a --- /dev/null +++ b/src/core/equal.go @@ -0,0 +1,31 @@ +package core + +import ( + "git.urbach.dev/cli/q/src/errors" + "git.urbach.dev/cli/q/src/ssa" + "git.urbach.dev/cli/q/src/token" + "git.urbach.dev/cli/q/src/types" +) + +// equal returns the binary operation to compare the left with the right value. +func (f *Function) equal(left ssa.Value, right ssa.Value, source ssa.Source) (ssa.Value, error) { + leftStructType, leftIsStruct := types.Unwrap(left.Type()).(*types.Struct) + rightStructType, rightIsStruct := types.Unwrap(right.Type()).(*types.Struct) + + if leftIsStruct && rightIsStruct && leftStructType == types.String && rightStructType == types.String { + return f.evaluateStringOp("equal", left, right, source) + } + + if leftIsStruct || rightIsStruct { + return nil, errors.New(InvalidStructOperation, f.File, source) + } + + comparison := f.Append(&ssa.BinaryOp{ + Left: left, + Right: right, + Op: token.Equal, + Source: source, + }) + + return comparison, nil +} \ No newline at end of file diff --git a/tests/switch-expression.q b/tests/switch-expression.q new file mode 100644 index 0000000..516a6a4 --- /dev/null +++ b/tests/switch-expression.q @@ -0,0 +1,35 @@ +main() { + c := 0 + + switch 42 { + 41 { c -= 1 } + } + + switch 42 { + 41 { c -= 1 } + _ { c += 1 } + } + + switch 42 { + 41 { c -= 1 } + 42 { c += 1 } + _ { c -= 1 } + } + + switch "b" { + "a" { c -= 1 } + } + + switch "b" { + "a" { c -= 1 } + _ { c += 1 } + } + + switch "b" { + "a" { c -= 1 } + "b" { c += 1 } + _ { c -= 1 } + } + + assert c == 4 +} \ No newline at end of file diff --git a/tests/tests_test.go b/tests/tests_test.go index e155ad4..db56f9a 100644 --- a/tests/tests_test.go +++ b/tests/tests_test.go @@ -49,6 +49,7 @@ var tests = []run{ {"branch-both", nil, "", "", 0}, {"jump-near", nil, "", "", 0}, {"switch", nil, "", "", 0}, + {"switch-expression", nil, "", "", 0}, {"phi", nil, "", "", 0}, {"phi-simple", nil, "", "", 0}, {"phi-advanced", nil, "", "", 0},