package loader import ( "testing" "github.com/marzeq/qk/attributes" "github.com/marzeq/qk/parser" "github.com/marzeq/qk/sema" "github.com/marzeq/qk/tokeniser" ) func TestCompileTimeFunctionRunsThroughRegularIR(t *testing.T) { source := ` module test @link( when twice(1) == 2 { system "selected", } else { @compiler_error("wrong conditional link branch") } ) let twice(value: i32): i32 = value * 2 let $Answer: i32 = twice(21) when Answer == 42 { let selected: i32 = Answer } else { @compiler_error("regular IR evaluation returned the wrong value") } @compiler_assert(twice(2) == 4, "ordinary function call failed") ` tokens, err := tokeniser.NewTokeniser(source, "staging.qk").Tokenise() if err != nil { t.Fatal(err) } root, err := parser.NewParser(tokens).Parse() if err != nil { t.Fatal(err) } if err := SelectCompileTime(root, "test", StageConfig{}); err != nil { t.Fatal(err) } for _, node := range root.Body { if _, unresolved := node.(*parser.WhenNode); unresolved { t.Fatal("compile-time selection left a when node in the module") } } answer, ok := root.Body[2].(*parser.DeclarationNode) if !ok { t.Fatalf("expected answer declaration, got %T", root.Body[2]) } literal, ok := answer.Value.(*parser.IntegerLiteralNode) if !ok || literal.Value != "42" { t.Fatalf("expected evaluated literal 42, got %#v", answer.Value) } module := root.Body[0].(*parser.ModuleNode) link := module.Attributes.Get(attributes.AttributeTypeLink).(attributes.ModuleAttributeLink) if len(link.Links) != 1 || link.Links[0].Value != "selected" { t.Fatalf("unexpected selected links: %#v", link.Links) } } func TestCompileTimeBindingExpandsInsideNestedArrayType(t *testing.T) { source := ` module test let $pipe_buffer_size = 32 let Pipe = type struct { x: f64 } let GameState = type struct { pipes: [pipe_buffer_size]Pipe } ` tokens, err := tokeniser.NewTokeniser(source, "array_length.qk").Tokenise() if err != nil { t.Fatal(err) } root, err := parser.NewParser(tokens).Parse() if err != nil { t.Fatal(err) } if err := SelectCompileTime(root, "test", StageConfig{}); err != nil { t.Fatal(err) } state := root.Body[3].(*parser.TypeAliasNode).Type.(*parser.StructTypeNode) array := state.Fields[0].Type.(*parser.ArrayTypeNode) literal, ok := array.Length.(*parser.IntegerLiteralNode) if !ok || literal.Value != "32" { t.Fatalf("expected expanded array length 32, got %#v", array.Length) } info := &ModuleInfo{Path: "test", Name: "test", Root: root} if errors, _ := RunSemanticModule(info, sema.NewAnalyser(), false, false); len(errors) != 0 { t.Fatal(errors[0]) } } func TestCompileTimeBindingsAreVisibleAcrossModuleFiles(t *testing.T) { parse := func(path, source string) *parser.RootNode { tokens, err := tokeniser.NewTokeniser(source, path).Tokenise() if err != nil { t.Fatal(err) } root, err := parser.NewParser(tokens).Parse() if err != nil { t.Fatal(err) } return root } consumer := parse("main.qk", "module test\nlet Buffer = type [pipe_buffer_size]u8\nlet use(): i32 = pipe_buffer_size\nlet make(): [pipe_buffer_size]u8 = [0; pipe_buffer_size]\n") provider := parse("state.qk", "module test\nlet $pipe_buffer_size = 32\n") if err := SelectCompileTimeRoots([]*parser.RootNode{consumer, provider}, "test", StageConfig{}); err != nil { t.Fatal(err) } if len(provider.Body) != 2 { t.Fatalf("provider retained %d nodes, want module and binding; consumer has %d", len(provider.Body), len(consumer.Body)) } array := consumer.Body[1].(*parser.TypeAliasNode).Type.(*parser.ArrayTypeNode) literal, ok := array.Length.(*parser.IntegerLiteralNode) if !ok || literal.Value != "32" { t.Fatalf("expected cross-file array length 32, got %#v", array.Length) } partials := make([]*PartialModuleInfo, 0, 2) for _, root := range []*parser.RootNode{consumer, provider} { partial, err := CollectModuleInfo(root, false) if err != nil { t.Fatal(err) } partial.Path = "test" partials = append(partials, partial) } modules, err := BuildModules(partials) if err != nil { t.Fatal(err) } if len(modules["test"].Root.Body) != 5 { t.Fatalf("merged module has %d nodes", len(modules["test"].Root.Body)) } if declaration, ok := modules["test"].Root.Body[4].(*parser.DeclarationNode); !ok || declaration.Name != "pipe_buffer_size" { t.Fatalf("merged provider declaration is %#v", modules["test"].Root.Body[4]) } if errors, _ := RunSemanticModule(modules["test"], sema.NewAnalyser(), false, false); len(errors) != 0 { declaration := modules["test"].Root.Body[4].(*parser.DeclarationNode) function := modules["test"].Root.Body[2].(*parser.FunctionDefNode) identifier := function.Body.(*parser.IdentifierNode) t.Fatalf("%v\nprovider symbol: %#v\nuse symbol: %#v", errors[0], declaration.Symbol, identifier.Symbol) } } func TestLocalCompileTimeCallDoesNotStageUnrelatedImportedRuntimeCode(t *testing.T) { source := ` module test import std.io let sum(x, y: i32): i32 = x + y let main() { let $x = sum(1, 2) std.io.print("The sum is: {}\n", x) } ` tokens, err := tokeniser.NewTokeniser(source, "local_staging.qk").Tokenise() if err != nil { t.Fatal(err) } root, err := parser.NewParser(tokens).Parse() if err != nil { t.Fatal(err) } if err := SelectCompileTime(root, "test", StageConfig{}); err != nil { t.Fatal(err) } main := root.Body[3].(*parser.FunctionDefNode) binding := main.Body.(*parser.BlockNode).Body[0].(*parser.DeclarationNode) literal, ok := binding.Value.(*parser.IntegerLiteralNode) if !ok || literal.Value != "3" { t.Fatalf("expected evaluated local binding 3, got %#v", binding.Value) } bindingType, ok := binding.TypeNode.(*parser.NamedTypeNode) if !ok || bindingType.Name != "i32" { t.Fatalf("expected evaluated local binding to retain i32, got %#v", binding.TypeNode) } }