From ca889e760903622852d6c86154ee671a9443d056 Mon Sep 17 00:00:00 2001 From: karitham Date: Fri, 21 Aug 2026 22:56:28 +0200 Subject: [PATCH] pool format arenas, dedupe token walks, race CI - formatter: global format mutex becomes a sync.Pool of doc arenas, so concurrent formats stop serializing; pinned by concurrency tests under -race - syntax: export IsKeyword, IsTypeKeyword, PrevReal, NextReal; source, formatter, and parser drop their private copies - lsp/source: CycleCheck reads Snapshot.Dependents off the COW include graph instead of rebuilding its own reachability; checker-level tests pin code, severity, and reported edges - doc: delete the unused heap-constructor API; print tests port to the arena and catch Arena.Join(nil-parts) building a nil Doc concat that crashed the printer - main: add missing --set-separator / --break-set flags - ci: run go test -race --- .github/workflows/ci.yml | 2 +- doc/doc.go | 67 +------ doc/print_fuzz_test.go | 67 ++++--- doc/print_test.go | 299 +++++++++++++++------------- formatter/concurrent_test.go | 105 ++++++++++ formatter/format.go | 46 ++--- lsp/source/context.go | 40 +--- lsp/source/cycle_detect.go | 100 +++------- lsp/source/cycle_detect_test.go | 341 +++++++++++++++----------------- lsp/source/semantic.go | 32 +-- main.go | 1 + syntax/lexer.go | 38 +++- syntax/parser.go | 10 +- 13 files changed, 581 insertions(+), 567 deletions(-) create mode 100644 formatter/concurrent_test.go diff --git a/.github/workflows/ci.yml b/.github/workflows/ci.yml index a4b55cc..73f8f28 100644 --- a/.github/workflows/ci.yml +++ b/.github/workflows/ci.yml @@ -50,7 +50,7 @@ jobs: with: go-version-file: go.mod - name: test - run: go test ./... + run: go test -race ./... vsix: runs-on: ubuntu-latest diff --git a/doc/doc.go b/doc/doc.go index 839bf39..eb64b19 100644 --- a/doc/doc.go +++ b/doc/doc.go @@ -4,6 +4,10 @@ // that decide independently whether they fit, and conditional pieces. The // single printer turns any document into a string given a print width. // +// Documents are built through an Arena: it is a bump allocator whose nodes +// die with the print. The package's exported Doc values (Line, SoftLine, +// HardLine, ...) are shared and immutable and may be embedded anywhere. +// // This is a faithful port of Prettier's document algebra // (src/document/builders and src/document/printer). // @@ -24,19 +28,14 @@ type Text string func (Text) isDoc() {} -// textNode is the heap or arena allocated form of Text: a pointer, so -// boxing it into Doc does not allocate. +// textNode is the arena-allocated form of Text: a pointer, so boxing it +// into Doc does not allocate. type textNode struct { s string } func (textNode) isDoc() {} -// NewText returns a doc for literal output. The value type Text also -// exists for direct construction; prefer NewText or Arena.Text so the -// node is a pointer. -func NewText(s string) Doc { return &textNode{s: s} } - // Concat is a sequence of documents printed in order. type Concat []Doc @@ -49,24 +48,6 @@ type concatNode struct { func (concatNode) isDoc() {} -// Join returns a Concat of parts joined by sep. -func Join(sep Doc, parts []Doc) Doc { - if len(parts) == 0 { - return Concat(nil) - } - - out := make(Concat, 0, len(parts)*2-1) - for i, part := range parts { - if i > 0 { - out = append(out, sep) - } - - out = append(out, part) - } - - return out -} - // Arena is a bump allocator for doc nodes. A document built through an // arena allocates its nodes from a few growing regions instead of one // heap allocation per node, at the cost of the arena retaining the @@ -121,7 +102,7 @@ func (a *Arena) Concat(parts ...Doc) Doc { // Join returns a doc for parts joined by sep, allocated from the arena. func (a *Arena) Join(sep Doc, parts []Doc) Doc { if len(parts) == 0 { - return a.Concat(nil) + return a.Concat() } out := a.Parts(len(parts)*2 - 1) @@ -276,21 +257,6 @@ type group struct { func (*group) isDoc() {} -// Group wraps d in a group that breaks only when it does not fit. -func Group(d Doc) Doc { return &group{doc: d} } - -// GroupBreak wraps d in a group that always breaks. -func GroupBreak(d Doc) Doc { return &group{doc: d, brk: true} } - -// GroupID wraps d in a group with an ID so IfBreak can query its mode. -func GroupID(id int, d Doc) Doc { return &group{doc: d, id: id} } - -// ConditionalGroup tries each state in order (least expanded first) and -// prints the first that fits; the last state breaks if none fit. -func ConditionalGroup(id int, states ...Doc) Doc { - return &group{doc: states[0], id: id, expanded: states} -} - // ifBreak prints BreakDoc when the group it belongs to (or the group with // GroupID) is broken, and FlatDoc when it is flat. type ifBreak struct { @@ -301,23 +267,12 @@ type ifBreak struct { func (*ifBreak) isDoc() {} -// IfBreak builds an IfBreak for the innermost enclosing group. -func IfBreak(broken, flat Doc) Doc { return &ifBreak{breakDoc: broken, flatDoc: flat} } - -// IfBreakFor builds an IfBreak that follows the group with the given ID. -func IfBreakFor(broken, flat Doc, groupID int) Doc { - return &ifBreak{breakDoc: broken, flatDoc: flat, groupID: groupID} -} - type indent struct { doc Doc } func (*indent) isDoc() {} -// Indent increases the indentation of its contents by one level. -func Indent(d Doc) Doc { return &indent{doc: d} } - type align struct { n int doc Doc @@ -325,20 +280,12 @@ type align struct { func (*align) isDoc() {} -// Align indents its contents by n columns relative to the current -// indentation. With tabs enabled, n is rounded up to one tab. -func Align(n int, d Doc) Doc { return &align{n: n, doc: d} } - type lineSuffix struct { doc Doc } func (*lineSuffix) isDoc() {} -// LineSuffix prints its contents at the end of the current line, after the -// next line break (used for end-of-line comments). -func LineSuffix(d Doc) Doc { return &lineSuffix{doc: d} } - type lineSuffixBoundary struct{} func (lineSuffixBoundary) isDoc() {} diff --git a/doc/print_fuzz_test.go b/doc/print_fuzz_test.go index 7f4a889..04fda81 100644 --- a/doc/print_fuzz_test.go +++ b/doc/print_fuzz_test.go @@ -86,8 +86,17 @@ var textChars = []rune{'a', 'b', ' ', 'x', '日', '本', '😀', '\u0301', '\t', // buildDoc builds a document from the bytecode program, returning the doc // and whether the program was fully consumed. func buildDoc(program []byte, depth int) (Doc, bool) { + var a Arena + + return buildDocArena(&a, program, depth) +} + +// buildDocArena is buildDoc against an arena: it exists so fuzz seeds and +// regression corpus entries (which predate the arena-only API) keep +// building docs by value where convenient. +func buildDocArena(a *Arena, program []byte, depth int) (Doc, bool) { if depth > maxDepth || len(program) == 0 { - return Concat{}, false + return a.Concat(), false } switch op := program[0]; op { @@ -99,7 +108,7 @@ func buildDoc(program []byte, depth int) (Doc, bool) { text = append(text, textChars[int(program[(2+i)%len(program)])%len(textChars)]) } - return Text(string(text)), true + return a.Text(string(text)), true case opLine, opSoftLine, opHardLine: switch op { @@ -112,16 +121,16 @@ func buildDoc(program []byte, depth int) (Doc, bool) { } case opGroup, opGroupBreak: - inner, ok := buildDoc(program[1:], depth+1) + inner, ok := buildDocArena(a, program[1:], depth+1) if !ok { - return Concat{}, false + return a.Concat(), false } if op == opGroupBreak { - return GroupBreak(inner), true + return a.GroupBreak(inner), true } - return Group(inner), true + return a.Group(inner), true case opConcat: n := int(program[1%len(program)]) % 5 @@ -130,66 +139,66 @@ func buildDoc(program []byte, depth int) (Doc, bool) { offset := 2 for range n { if offset >= len(program) { - return Concat(parts), true + return a.Concat(parts...), true } - part, ok := buildDoc(program[offset:], depth+1) + part, ok := buildDocArena(a, program[offset:], depth+1) if !ok { - return Concat(parts), true + return a.Concat(parts...), true } parts = append(parts, part) offset++ } - return Concat(parts), true + return a.Concat(parts...), true case opIndent: - inner, ok := buildDoc(program[1:], depth+1) + inner, ok := buildDocArena(a, program[1:], depth+1) if !ok { - return Concat{}, false + return a.Concat(), false } - return Indent(inner), true + return a.Indent(inner), true case opAlign: - inner, ok := buildDoc(program[1:], depth+1) + inner, ok := buildDocArena(a, program[1:], depth+1) if !ok { - return Concat{}, false + return a.Concat(), false } - return Align(int(program[1%len(program)])%5, inner), true + return a.Align(int(program[1%len(program)])%5, inner), true case opIfBreak: - brk, ok1 := buildDoc(program[1:], depth+1) + brk, ok1 := buildDocArena(a, program[1:], depth+1) - flat, ok2 := buildDoc(program[2%len(program):], depth+1) + flat, ok2 := buildDocArena(a, program[2%len(program):], depth+1) if !ok1 || !ok2 { - return Concat{}, false + return a.Concat(), false } - return IfBreak(brk, flat), true + return a.IfBreak(brk, flat), true case opLineSuffix: - inner, ok := buildDoc(program[1:], depth+1) + inner, ok := buildDocArena(a, program[1:], depth+1) if !ok { - return Concat{}, false + return a.Concat(), false } - return LineSuffix(inner), true + return a.LineSuffix(inner), true case opConditional: - first, ok := buildDoc(program[1:], depth+1) + first, ok := buildDocArena(a, program[1:], depth+1) if !ok { - return Concat{}, false + return a.Concat(), false } - second, ok := buildDoc(program[2%len(program):], depth+1) + second, ok := buildDocArena(a, program[2%len(program):], depth+1) if !ok { - return Concat{}, false + return a.Concat(), false } - return ConditionalGroup(0, first, second, GroupBreak(first)), true + return a.ConditionalGroup(0, first, second, a.GroupBreak(first)), true case opTrim: return TrimDoc, true @@ -201,6 +210,6 @@ func buildDoc(program []byte, depth int) (Doc, bool) { return LineSuffixBoundary, true default: - return Concat{}, false + return a.Concat(), false } } diff --git a/doc/print_test.go b/doc/print_test.go index 4c6b88e..c0334b8 100644 --- a/doc/print_test.go +++ b/doc/print_test.go @@ -4,18 +4,18 @@ import ( "testing" ) -// T builds a Text from a string. -func T(s string) Doc { return Text(s) } - // testOpts are the defaults used by most cases. func testOpts(width int) Options { return Options{PrintWidth: width, Indent: " ", TabWidth: 2, NewLine: "\n"} } -func printT(t *testing.T, d Doc, o Options) string { +// printT builds the doc on a fresh arena and prints it. +func printT(t *testing.T, build func(a *Arena) Doc, o Options) string { t.Helper() - got, err := Print(d, o) + var a Arena + + got, err := Print(build(&a), o) if err != nil { t.Fatalf("Print: %v", err) } @@ -26,21 +26,21 @@ func printT(t *testing.T, d Doc, o Options) string { func TestPrintText(t *testing.T) { tests := []struct { name string - doc func() Doc + doc func(a *Arena) Doc want string }{ - {"empty", func() Doc { return Concat{} }, ""}, - {"text", func() Doc { return T("hello world") }, "hello world"}, - {"concat", func() Doc { return Concat{T("a"), T("b"), T("c")} }, "abc"}, - {"join", func() Doc { return Join(T(","), []Doc{T("a"), T("b"), T("c")}) }, "a,b,c"}, - {"join empty", func() Doc { return Join(T(","), nil) }, ""}, - {"group fits", func() Doc { return Group(Concat{T("hello"), T(" world")}) }, "hello world"}, - {"empty group", func() Doc { return Group(Concat{}) }, ""}, + {"empty", func(a *Arena) Doc { return a.Concat() }, ""}, + {"text", func(a *Arena) Doc { return a.Text("hello world") }, "hello world"}, + {"concat", func(a *Arena) Doc { return a.Concat(a.Text("a"), a.Text("b"), a.Text("c")) }, "abc"}, + {"join", func(a *Arena) Doc { return a.Join(a.Text(","), []Doc{a.Text("a"), a.Text("b"), a.Text("c")}) }, "a,b,c"}, + {"join empty", func(a *Arena) Doc { return a.Join(a.Text(","), nil) }, ""}, + {"group fits", func(a *Arena) Doc { return a.Group(a.Concat(a.Text("hello"), a.Text(" world"))) }, "hello world"}, + {"empty group", func(a *Arena) Doc { return a.Group(a.Concat()) }, ""}, } for _, tt := range tests { t.Run(tt.name, func(t *testing.T) { - if got := printT(t, tt.doc(), testOpts(80)); got != tt.want { + if got := printT(t, tt.doc, testOpts(80)); got != tt.want { t.Errorf("got %q, want %q", got, tt.want) } }) @@ -50,118 +50,120 @@ func TestPrintText(t *testing.T) { func TestPrintLineBreaking(t *testing.T) { tests := []struct { name string - doc func() Doc + doc func(a *Arena) Doc width int want string }{ { name: "line stays flat when it fits", - doc: func() Doc { return Group(Concat{T("a"), Line, T("b")}) }, + doc: func(a *Arena) Doc { return a.Group(a.Concat(a.Text("a"), Line, a.Text("b"))) }, width: 10, want: "a b", }, { name: "line breaks when it does not fit", - doc: func() Doc { return Group(Concat{T("a"), Line, T("b")}) }, + doc: func(a *Arena) Doc { return a.Group(a.Concat(a.Text("a"), Line, a.Text("b"))) }, width: 2, want: "a\nb", }, { name: "line breaks exactly at boundary", - doc: func() Doc { return Group(Concat{T("ab"), Line, T("c")}) }, + doc: func(a *Arena) Doc { return a.Group(a.Concat(a.Text("ab"), Line, a.Text("c"))) }, width: 3, want: "ab\nc", // "ab c" is 4 columns and does not fit }, { name: "softline is empty when flat", - doc: func() Doc { return Group(Concat{T("a"), SoftLine, T("b")}) }, + doc: func(a *Arena) Doc { return a.Group(a.Concat(a.Text("a"), SoftLine, a.Text("b"))) }, width: 80, want: "ab", }, { name: "softline breaks with the group", - doc: func() Doc { return Group(Concat{T("a"), SoftLine, T("b")}) }, + doc: func(a *Arena) Doc { return a.Group(a.Concat(a.Text("a"), SoftLine, a.Text("b"))) }, width: 1, want: "a\nb", }, { name: "canonical indent pattern", - doc: func() Doc { return Group(Concat{T("a"), Indent(Concat{Line, T("b")})}) }, + doc: func(a *Arena) Doc { return a.Group(a.Concat(a.Text("a"), a.Indent(a.Concat(Line, a.Text("b"))))) }, width: 10, want: "a b", }, { name: "canonical indent pattern breaks", - doc: func() Doc { return Group(Concat{T("a"), Indent(Concat{Line, T("b")})}) }, + doc: func(a *Arena) Doc { return a.Group(a.Concat(a.Text("a"), a.Indent(a.Concat(Line, a.Text("b"))))) }, width: 2, want: "a\n b", }, { name: "inner group stays flat inside broken outer", - doc: func() Doc { - return Group(Concat{ - T("a"), - Group(Concat{T("b"), SoftLine, T("c")}), + doc: func(a *Arena) Doc { + return a.Group(a.Concat( + a.Text("a"), + a.Group(a.Concat(a.Text("b"), SoftLine, a.Text("c"))), Line, - T("d"), - }) + a.Text("d"), + )) }, width: 4, want: "abc\nd", // the inner group fits in the remaining width }, { name: "inner group breaks when it does not fit", - doc: func() Doc { - return Group(Concat{ - T("aaaa"), - Group(Concat{T("bbbb"), SoftLine, T("cc")}), + doc: func(a *Arena) Doc { + return a.Group(a.Concat( + a.Text("aaaa"), + a.Group(a.Concat(a.Text("bbbb"), SoftLine, a.Text("cc"))), Line, - T("d"), - }) + a.Text("d"), + )) }, width: 8, want: "aaaabbbb\ncc\nd", // only the inner group's own line breaks }, { name: "inner group breaks when remaining width is too small", - doc: func() Doc { - return Group(Concat{ - T("a"), - Group(Concat{T("b"), SoftLine, T("c")}), + doc: func(a *Arena) Doc { + return a.Group(a.Concat( + a.Text("a"), + a.Group(a.Concat(a.Text("b"), SoftLine, a.Text("c"))), Line, - T("d"), - }) + a.Text("d"), + )) }, width: 2, want: "ab\nc\nd", }, { name: "hardline always breaks", - doc: func() Doc { return Group(Concat{T("a"), HardLineNoBreak, T("b")}) }, + doc: func(a *Arena) Doc { return a.Group(a.Concat(a.Text("a"), HardLineNoBreak, a.Text("b"))) }, width: 80, want: "a\nb", }, { - name: "hardline breaks enclosing group", - doc: func() Doc { return Group(Concat{T("a"), HardLine, T("b"), SoftLine, T("c")}) }, + name: "hardline breaks enclosing group", + doc: func(a *Arena) Doc { + return a.Group(a.Concat(a.Text("a"), HardLine, a.Text("b"), SoftLine, a.Text("c"))) + }, width: 80, want: "a\nb\nc", }, { name: "break propagates through nested groups", - doc: func() Doc { - return Group(Concat{ - Group(Concat{T("x"), HardLine, T("y"), SoftLine, T("z")}), + doc: func(a *Arena) Doc { + return a.Group(a.Concat( + a.Group(a.Concat(a.Text("x"), HardLine, a.Text("y"), SoftLine, a.Text("z"))), SoftLine, - T("w"), - }) + a.Text("w"), + )) }, width: 80, want: "x\ny\nz\nw", }, { name: "forced break group", - doc: func() Doc { return GroupBreak(Concat{T("a"), SoftLine, T("b")}) }, + doc: func(a *Arena) Doc { return a.GroupBreak(a.Concat(a.Text("a"), SoftLine, a.Text("b"))) }, width: 80, want: "a\nb", }, @@ -169,7 +171,7 @@ func TestPrintLineBreaking(t *testing.T) { for _, tt := range tests { t.Run(tt.name, func(t *testing.T) { - if got := printT(t, tt.doc(), testOpts(tt.width)); got != tt.want { + if got := printT(t, tt.doc, testOpts(tt.width)); got != tt.want { t.Errorf("got %q, want %q", got, tt.want) } }) @@ -179,40 +181,44 @@ func TestPrintLineBreaking(t *testing.T) { func TestPrintIfBreak(t *testing.T) { tests := []struct { name string - doc func() Doc + doc func(a *Arena) Doc width int want string }{ { - name: "flat takes flat contents", - doc: func() Doc { return Group(Concat{T("a"), IfBreak(T(","), T("")), T("b")}) }, + name: "flat takes flat contents", + doc: func(a *Arena) Doc { + return a.Group(a.Concat(a.Text("a"), a.IfBreak(a.Text(","), a.Text("")), a.Text("b"))) + }, width: 80, want: "ab", }, { - name: "broken takes break contents", - doc: func() Doc { return GroupBreak(Concat{T("a"), IfBreak(T(","), T("")), T("b")}) }, + name: "broken takes break contents", + doc: func(a *Arena) Doc { + return a.GroupBreak(a.Concat(a.Text("a"), a.IfBreak(a.Text(","), a.Text("")), a.Text("b"))) + }, width: 80, want: "a,b", }, { name: "ifBreak follows referenced group", - doc: func() Doc { - return Concat{ - GroupID(1, Concat{T("aaaa"), SoftLine, T("bbbb")}), - IfBreakFor(T(","), T(""), 1), - } + doc: func(a *Arena) Doc { + return a.Concat( + a.GroupID(1, a.Concat(a.Text("aaaa"), SoftLine, a.Text("bbbb"))), + a.IfBreakFor(a.Text(","), a.Text(""), 1), + ) }, width: 80, want: "aaaabbbb", }, { name: "ifBreak follows broken referenced group", - doc: func() Doc { - return Concat{ - GroupID(1, Concat{T("aaaa"), SoftLine, T("bbbb")}), - IfBreakFor(T(","), T(""), 1), - } + doc: func(a *Arena) Doc { + return a.Concat( + a.GroupID(1, a.Concat(a.Text("aaaa"), SoftLine, a.Text("bbbb"))), + a.IfBreakFor(a.Text(","), a.Text(""), 1), + ) }, width: 4, want: "aaaa\nbbbb,", @@ -221,7 +227,7 @@ func TestPrintIfBreak(t *testing.T) { for _, tt := range tests { t.Run(tt.name, func(t *testing.T) { - if got := printT(t, tt.doc(), testOpts(tt.width)); got != tt.want { + if got := printT(t, tt.doc, testOpts(tt.width)); got != tt.want { t.Errorf("got %q, want %q", got, tt.want) } }) @@ -253,12 +259,14 @@ func TestPrintConditionalGroup(t *testing.T) { for _, tt := range tests { t.Run(tt.name, func(t *testing.T) { - doc := ConditionalGroup(0, - Concat{T("a"), Line, T("b"), Line, T("c")}, - Concat{T("a"), Line, T("b"), HardLineNoBreak, T("c")}, - Concat{T("a"), HardLineNoBreak, T("b"), HardLineNoBreak, T("c")}, - ) - if got := printT(t, doc, testOpts(tt.width)); got != tt.want { + got := printT(t, func(a *Arena) Doc { + return a.ConditionalGroup(0, + a.Concat(a.Text("a"), Line, a.Text("b"), Line, a.Text("c")), + a.Concat(a.Text("a"), Line, a.Text("b"), HardLineNoBreak, a.Text("c")), + a.Concat(a.Text("a"), HardLineNoBreak, a.Text("b"), HardLineNoBreak, a.Text("c")), + ) + }, testOpts(tt.width)) + if got != tt.want { t.Errorf("got %q, want %q", got, tt.want) } }) @@ -268,34 +276,34 @@ func TestPrintConditionalGroup(t *testing.T) { func TestPrintIndentAndAlign(t *testing.T) { tests := []struct { name string - doc func() Doc + doc func(a *Arena) Doc want string }{ { name: "indent applies at line breaks", - doc: func() Doc { - return GroupBreak(Concat{T("a"), Indent(Concat{Line, T("b"), Line, T("c")})}) + doc: func(a *Arena) Doc { + return a.GroupBreak(a.Concat(a.Text("a"), a.Indent(a.Concat(Line, a.Text("b"), Line, a.Text("c"))))) }, want: "a\n b\n c", }, { name: "align by columns", - doc: func() Doc { - return GroupBreak(Concat{T("a"), Align(4, Concat{Line, T("b")})}) + doc: func(a *Arena) Doc { + return a.GroupBreak(a.Concat(a.Text("a"), a.Align(4, a.Concat(Line, a.Text("b"))))) }, want: "a\n b", }, { name: "align nests inside indent", - doc: func() Doc { - return GroupBreak(Concat{T("a"), Indent(Concat{Line, Align(2, Concat{T("b"), Line, T("c")})})}) + doc: func(a *Arena) Doc { + return a.GroupBreak(a.Concat(a.Text("a"), a.Indent(a.Concat(Line, a.Align(2, a.Concat(a.Text("b"), Line, a.Text("c"))))))) }, want: "a\n b\n c", }, { name: "nested indent accumulates", - doc: func() Doc { - return GroupBreak(Concat{T("a"), Indent(Indent(Concat{Line, T("b")}))}) + doc: func(a *Arena) Doc { + return a.GroupBreak(a.Concat(a.Text("a"), a.Indent(a.Indent(a.Concat(Line, a.Text("b")))))) }, want: "a\n b", }, @@ -303,7 +311,7 @@ func TestPrintIndentAndAlign(t *testing.T) { for _, tt := range tests { t.Run(tt.name, func(t *testing.T) { - if got := printT(t, tt.doc(), testOpts(80)); got != tt.want { + if got := printT(t, tt.doc, testOpts(80)); got != tt.want { t.Errorf("got %q, want %q", got, tt.want) } }) @@ -313,13 +321,13 @@ func TestPrintIndentAndAlign(t *testing.T) { func TestPrintTabIndent(t *testing.T) { tests := []struct { name string - doc func() Doc + doc func(a *Arena) Doc want string }{ { name: "tab indentation and measurement", - doc: func() Doc { - return GroupBreak(Concat{T("a"), Indent(Concat{Line, T("b")})}) + doc: func(a *Arena) Doc { + return a.GroupBreak(a.Concat(a.Text("a"), a.Indent(a.Concat(Line, a.Text("b"))))) }, want: "a\n\tb", }, @@ -327,7 +335,7 @@ func TestPrintTabIndent(t *testing.T) { for _, tt := range tests { t.Run(tt.name, func(t *testing.T) { o := Options{PrintWidth: 80, Indent: "\t", TabWidth: 4, NewLine: "\n"} - if got := printT(t, tt.doc(), o); got != tt.want { + if got := printT(t, tt.doc, o); got != tt.want { t.Errorf("got %q, want %q", got, tt.want) } }) @@ -337,29 +345,29 @@ func TestPrintTabIndent(t *testing.T) { func TestPrintLineSuffix(t *testing.T) { tests := []struct { name string - doc func() Doc + doc func(a *Arena) Doc want string }{ { name: "suffix prints before the line break", - doc: func() Doc { - return GroupBreak(Concat{T("a"), LineSuffix(T(" // c")), Line, T("b")}) + doc: func(a *Arena) Doc { + return a.GroupBreak(a.Concat(a.Text("a"), a.LineSuffix(a.Text(" // c")), Line, a.Text("b"))) }, want: "a // c\nb", }, { name: "suffix flushes at document end without a break", - doc: func() Doc { - return Concat{T("a"), LineSuffix(T(" // c"))} + doc: func(a *Arena) Doc { + return a.Concat(a.Text("a"), a.LineSuffix(a.Text(" // c"))) }, want: "a // c", }, { name: "boundary flushes pending suffixes", - doc: func() Doc { - return GroupBreak(Concat{ - T("a"), LineSuffix(T(" // first")), LineSuffixBoundary, - }) + doc: func(a *Arena) Doc { + return a.GroupBreak(a.Concat( + a.Text("a"), a.LineSuffix(a.Text(" // first")), LineSuffixBoundary, + )) }, // The boundary schedules a hard line that flushes the suffix and // ends the line, matching Prettier's boundary semantics. @@ -367,8 +375,8 @@ func TestPrintLineSuffix(t *testing.T) { }, { name: "suffix counts against width", - doc: func() Doc { - return Group(Concat{T("aaaa"), LineSuffix(T(" // c")), Line, T("b")}) + doc: func(a *Arena) Doc { + return a.Group(a.Concat(a.Text("aaaa"), a.LineSuffix(a.Text(" // c")), Line, a.Text("b"))) }, want: "aaaa b // c", // suffix width is not measured, like Prettier }, @@ -376,7 +384,7 @@ func TestPrintLineSuffix(t *testing.T) { for _, tt := range tests { t.Run(tt.name, func(t *testing.T) { - if got := printT(t, tt.doc(), testOpts(80)); got != tt.want { + if got := printT(t, tt.doc, testOpts(80)); got != tt.want { t.Errorf("got %q, want %q", got, tt.want) } }) @@ -386,16 +394,16 @@ func TestPrintLineSuffix(t *testing.T) { func TestPrintTrim(t *testing.T) { tests := []struct { name string - doc func() Doc + doc func(a *Arena) Doc want string }{ - {"trim removes trailing spaces", func() Doc { return Concat{T("a "), TrimDoc, T("b")} }, "ab"}, - {"trim removes trailing tabs", func() Doc { return Concat{T("a\t\t"), TrimDoc, T("b")} }, "ab"}, - {"trim without trailing whitespace", func() Doc { return Concat{T("a"), TrimDoc, T("b")} }, "ab"}, + {"trim removes trailing spaces", func(a *Arena) Doc { return a.Concat(a.Text("a "), TrimDoc, a.Text("b")) }, "ab"}, + {"trim removes trailing tabs", func(a *Arena) Doc { return a.Concat(a.Text("a\t\t"), TrimDoc, a.Text("b")) }, "ab"}, + {"trim without trailing whitespace", func(a *Arena) Doc { return a.Concat(a.Text("a"), TrimDoc, a.Text("b")) }, "ab"}, } for _, tt := range tests { t.Run(tt.name, func(t *testing.T) { - if got := printT(t, tt.doc(), testOpts(80)); got != tt.want { + if got := printT(t, tt.doc, testOpts(80)); got != tt.want { t.Errorf("got %q, want %q", got, tt.want) } }) @@ -406,13 +414,16 @@ func TestPrintRemeasure(t *testing.T) { // A hard line printed inside a flat-measured group invalidates the // measurement; the next group must remeasure instead of trusting the // flat shortcut. - doc := Group(Concat{ - T("a"), HardLineNoBreak, - Group(Concat{T("bbbb"), Line, T("cc")}), - Line, - T("dd"), - }) - got := printT(t, doc, testOpts(8)) + build := func(a *Arena) Doc { + return a.Group(a.Concat( + a.Text("a"), HardLineNoBreak, + a.Group(a.Concat(a.Text("bbbb"), Line, a.Text("cc"))), + Line, + a.Text("dd"), + )) + } + + got := printT(t, build, testOpts(8)) want := "a\nbbbb\ncc dd" if got != want { @@ -423,9 +434,9 @@ func TestPrintRemeasure(t *testing.T) { func TestPrintNewLineOption(t *testing.T) { o := Options{PrintWidth: 2, Indent: " ", TabWidth: 2, NewLine: "\r\n"} - doc := GroupBreak(Concat{T("a"), Line, T("b")}) - if got := printT(t, doc, o); got != "a\r\nb" { - t.Errorf("got %q, want %q", got, "a\r\nb") + doc := printT(t, func(a *Arena) Doc { return a.GroupBreak(a.Concat(a.Text("a"), Line, a.Text("b"))) }, o) + if doc != "a\r\nb" { + t.Errorf("got %q, want %q", doc, "a\r\nb") } } @@ -440,7 +451,7 @@ func TestPrintValidation(t *testing.T) { } for _, tt := range tests { t.Run(tt.name, func(t *testing.T) { - if _, err := Print(T("x"), tt.opts); err == nil { + if _, err := Print(Text("x"), tt.opts); err == nil { t.Error("expected validation error") } }) @@ -450,31 +461,31 @@ func TestPrintValidation(t *testing.T) { func TestPrintUnicodeWidth(t *testing.T) { tests := []struct { name string - doc func() Doc + doc func(a *Arena) Doc width int want string }{ { name: "wide characters count as two columns", - doc: func() Doc { return Group(Concat{T("日本語"), Line, T("x")}) }, + doc: func(a *Arena) Doc { return a.Group(a.Concat(a.Text("日本語"), Line, a.Text("x"))) }, width: 8, want: "日本語 x", }, { name: "wide characters break the group", - doc: func() Doc { return Group(Concat{T("日本語"), Line, T("x")}) }, + doc: func(a *Arena) Doc { return a.Group(a.Concat(a.Text("日本語"), Line, a.Text("x"))) }, width: 7, want: "日本語\nx", }, { name: "combining marks are zero width", - doc: func() Doc { return Group(Concat{T("e\u0301"), Line, T("x")}) }, + doc: func(a *Arena) Doc { return a.Group(a.Concat(a.Text("e\u0301"), Line, a.Text("x"))) }, width: 3, want: "e\u0301 x", }, { name: "emoji count as two columns", - doc: func() Doc { return Group(Concat{T("😀"), Line, T("x")}) }, + doc: func(a *Arena) Doc { return a.Group(a.Concat(a.Text("😀"), Line, a.Text("x"))) }, width: 4, want: "😀 x", }, @@ -482,7 +493,7 @@ func TestPrintUnicodeWidth(t *testing.T) { for _, tt := range tests { t.Run(tt.name, func(t *testing.T) { - if got := printT(t, tt.doc(), testOpts(tt.width)); got != tt.want { + if got := printT(t, tt.doc, testOpts(tt.width)); got != tt.want { t.Errorf("got %q, want %q", got, tt.want) } }) @@ -513,18 +524,33 @@ func TestStringWidth(t *testing.T) { func TestPrintIdempotentDocs(t *testing.T) { // These docs must print identically at the same width twice in a row; // break propagation must not change the result on the second pass. - docs := []Doc{ - Group(Concat{T("a"), Line, T("b")}), - Group(Concat{T("a"), HardLine, T("b"), SoftLine, T("c")}), - ConditionalGroup(0, - Concat{T("a"), Line, T("b"), Line, T("c")}, - Concat{T("a"), Line, T("b"), HardLineNoBreak, T("c")}, - ), + builds := []func(a *Arena) Doc{ + func(a *Arena) Doc { return a.Group(a.Concat(a.Text("a"), Line, a.Text("b"))) }, + func(a *Arena) Doc { + return a.Group(a.Concat(a.Text("a"), HardLine, a.Text("b"), SoftLine, a.Text("c"))) + }, + func(a *Arena) Doc { + return a.ConditionalGroup(0, + a.Concat(a.Text("a"), Line, a.Text("b"), Line, a.Text("c")), + a.Concat(a.Text("a"), Line, a.Text("b"), HardLineNoBreak, a.Text("c")), + ) + }, } - for _, d := range docs { - first := printT(t, d, testOpts(2)) + for _, build := range builds { + var a Arena + + d := build(&a) + + first, err := Print(d, testOpts(2)) + if err != nil { + t.Fatalf("Print: %v", err) + } + + second, err := Print(d, testOpts(2)) + if err != nil { + t.Fatalf("Print: %v", err) + } - second := printT(t, d, testOpts(2)) if first != second { t.Errorf("not idempotent: %q vs %q", first, second) } @@ -533,14 +559,19 @@ func TestPrintIdempotentDocs(t *testing.T) { func TestDocMutability(t *testing.T) { // Print mutates the doc (break propagation); printing a fresh doc each - // time must yield stable output. + // time must yield stable output. The arena is reset between prints, so + // the fresh doc reuses the arena's regions. + var a Arena + build := func() Doc { - return Group(Concat{T("a"), HardLine, T("b"), SoftLine, T("c")}) + a.Reset() + + return a.Group(a.Concat(a.Text("a"), HardLine, a.Text("b"), SoftLine, a.Text("c"))) } - first := printT(t, build(), testOpts(80)) + first := printT(t, func(*Arena) Doc { return build() }, testOpts(80)) for i := range 10 { - if got := printT(t, build(), testOpts(80)); got != first { + if got := printT(t, func(*Arena) Doc { return build() }, testOpts(80)); got != first { t.Fatalf("iteration %d: got %q, want %q", i, got, first) } } diff --git a/formatter/concurrent_test.go b/formatter/concurrent_test.go new file mode 100644 index 0000000..fc22b58 --- /dev/null +++ b/formatter/concurrent_test.go @@ -0,0 +1,105 @@ +package formatter + +import ( + "fmt" + "sync" + "testing" + + "github.com/karitham/thrift-ls/syntax" +) + +// TestFormatConcurrent pins the arena-pool invariant: Format must be safe +// to call from many goroutines at once. The pooled arena is the only +// shared state, so a data race here means the pool contract is broken. +// Run with -race in CI to make the check real. +func TestFormatConcurrent(t *testing.T) { + srcs := make([]string, 0, 16) + for i := range 16 { + srcs = append(srcs, fmt.Sprintf(`struct Data%d { + 1: required i32 id, + 2: optional string name, +} + +service Svc%d { + Data%d get(1: i32 id), +} +`, i, i, i)) + } + + opts := testOpts(80) + + var wg sync.WaitGroup + + errs := make([]error, len(srcs)) + + for i, src := range srcs { + wg.Go(func() { + doc, docErrs := syntax.Parse([]byte(src)) + if hasParseErrors(docErrs) { + errs[i] = fmt.Errorf("parse: %v", docErrs) + + return + } + + got, err := Format(doc, opts) + if err != nil { + errs[i] = err + + return + } + + // Re-parse: a corrupted arena (shared regions between + // goroutines) shows up as garbage output that no longer + // matches the input shape. + if _, reparses := syntax.Parse([]byte(got)); hasParseErrors(reparses) { + errs[i] = fmt.Errorf("output does not reparse: %q (%v)", got, reparses) + } + }) + } + + wg.Wait() + + for i, err := range errs { + if err != nil { + t.Errorf("goroutine %d: %v", i, err) + } + } +} + +// TestFormatConcurrentSameInput hammers one document from many goroutines: +// every result must be identical. Catches arena state leaking across calls +// even when each call is individually well-formed. +func TestFormatConcurrentSameInput(t *testing.T) { + src := `struct Item { + 1: required i64 id, + 2: map tags, +} +` + doc, errs := syntax.Parse([]byte(src)) + if hasParseErrors(errs) { + t.Fatalf("parse errors: %v", errs) + } + + opts := testOpts(80) + + want := fmtSrc(t, src, opts) + + var wg sync.WaitGroup + + for range 32 { + wg.Go(func() { + got, err := Format(doc, opts) + if err != nil { + t.Errorf("Format: %v", err) + + return + } + + if got != want { + t.Errorf("concurrent format mismatch:\n got: %q\nwant: %q", got, want) + } + }) + } + + wg.Wait() +} diff --git a/formatter/format.go b/formatter/format.go index c86220e..15f7134 100644 --- a/formatter/format.go +++ b/formatter/format.go @@ -218,14 +218,12 @@ func (o Options) normalize() Options { // Format renders a parsed document. The document must have been parsed // without errors; callers check the parse errors before formatting. Format // is deterministic and pure: it reads nothing but the document and options. -// formatArena is the shared node arena for Format and FormatNode. The -// doc IR dies with the print, so a single pooled arena is safe: regions -// are reused across calls and the garbage the GC would scan shrinks to -// the region overflows. -var ( - formatArena doc.Arena - formatMu sync.Mutex -) +// arenaPool pools the node arenas for Format and FormatNode. The doc IR +// dies with the print, so an arena is free for reuse the moment the call +// returns: regions are reused across calls and the garbage the GC would +// scan shrinks to the region overflows. The pool hands out one arena per +// goroutine, so concurrent formats do not serialize. +var arenaPool = sync.Pool{New: func() any { return new(doc.Arena) }} // Format renders the whole document. The arena is pooled: the returned // string is the only thing that outlives the call. @@ -236,19 +234,19 @@ func Format(d *syntax.Document, o Options) (string, error) { o = o.normalize() - formatMu.Lock() - defer formatMu.Unlock() + a := arenaPool.Get().(*doc.Arena) + defer arenaPool.Put(a) - formatArena.Reset() + a.Reset() f := &formatter{ - Arena: &formatArena, + Arena: a, doc: d, toks: d.Tokens, opts: o, } - return formatArena.Print(f.document(), printOptions(o)) + return a.Print(f.document(), printOptions(o)) } // printOptions maps formatter options to printer options. @@ -290,19 +288,19 @@ func FormatNode(d *syntax.Document, n syntax.Node, o Options) (string, error) { o = o.normalize() - formatMu.Lock() - defer formatMu.Unlock() + a := arenaPool.Get().(*doc.Arena) + defer arenaPool.Put(a) - formatArena.Reset() + a.Reset() f := &formatter{ - Arena: &formatArena, + Arena: a, doc: d, toks: d.Tokens, opts: o, } - return formatArena.Print(f.node(n), printOptions(o)) + return a.Print(f.node(n), printOptions(o)) } type formatter struct { @@ -359,21 +357,13 @@ func padAt(pads []padEntry, idx int) string { // prevReal returns the index of the previous real (non-comment) token // strictly before idx, or -1. func (f *formatter) prevReal(idx int) int { - for idx >= 0 && isComment(f.token(idx).Kind) { - idx-- - } - - return idx + return syntax.PrevReal(f.toks, idx) } // nextReal returns the index of the next real (non-comment) token at or // after idx. func (f *formatter) nextReal(idx int) int { - for idx < len(f.toks) && isComment(f.token(idx).Kind) { - idx++ - } - - return idx + return syntax.NextReal(f.toks, idx) } // emitTokens renders the tokens in [start, end] with the comments diff --git a/lsp/source/context.go b/lsp/source/context.go index b2ab6c4..eec23fb 100644 --- a/lsp/source/context.go +++ b/lsp/source/context.go @@ -66,9 +66,9 @@ func ResolveContext(doc *syntax.Document, pos syntax.Position) Context { // The token before the cursor: the token containing the cursor when the // cursor sits at its end, the previous token when mid-token. Comments // are skipped — the grammar slot is determined by the real tokens. - prevIdx := prevReal(doc.Tokens, atIdx) + prevIdx := syntax.PrevReal(doc.Tokens, atIdx) if at != nil && pos.Offset < at.Offset+len(at.Text) { - prevIdx = prevReal(doc.Tokens, atIdx-1) + prevIdx = syntax.PrevReal(doc.Tokens, atIdx-1) } c.Prefix, c.EditStart = prefixRange(doc, pos, atIdx, at) @@ -103,7 +103,7 @@ func ResolveContext(doc *syntax.Document, pos syntax.Position) Context { // before the struct member rule, so "{ |1:" is CtxFieldID, not a // member name position. if at != nil && at.Kind == syntax.TokenIntConstant { - if n := nextReal(doc.Tokens, atIdx+1); n < len(doc.Tokens) && doc.Tokens[n].Kind == syntax.TokenColon { + if n := syntax.NextReal(doc.Tokens, atIdx+1); n < len(doc.Tokens) && doc.Tokens[n].Kind == syntax.TokenColon { c.Kind = CtxFieldID return c @@ -201,7 +201,7 @@ func ResolveContext(doc *syntax.Document, pos syntax.Position) Context { // "songs.A" — a dotted identifier in a type slot: the type // provider scopes to the include and the edit replaces the // whole qualified prefix. - if strings.Contains(at.Text, ".") && typeSlotAfterIdent(doc, prevReal(doc.Tokens, atIdx-1)) { + if strings.Contains(at.Text, ".") && typeSlotAfterIdent(doc, syntax.PrevReal(doc.Tokens, atIdx-1)) { c.Kind = CtxType return c @@ -282,7 +282,7 @@ func identifierKind(path []syntax.Node, n *syntax.Identifier) ContextKind { // afterParenKind classifies the cursor right after '(' (prev is the opener). func afterParenKind(doc *syntax.Document, opener int) ContextKind { - prevIdx := prevReal(doc.Tokens, opener-1) + prevIdx := syntax.PrevReal(doc.Tokens, opener-1) if prevIdx < 0 { return CtxAnnotationKey } @@ -310,7 +310,7 @@ func afterParenKind(doc *syntax.Document, opener int) ContextKind { // insideParenKind classifies the cursor after ','/';' inside the group // opened at opener. func insideParenKind(doc *syntax.Document, opener int) ContextKind { - prevIdx := prevReal(doc.Tokens, opener-1) + prevIdx := syntax.PrevReal(doc.Tokens, opener-1) if prevIdx < 0 { return CtxAnnotationKey } @@ -334,7 +334,7 @@ func insideParenKind(doc *syntax.Document, opener int) ContextKind { // isThrowsGroup reports whether the group opened at opener is a throws // clause (which contains fields, not annotations). func isThrowsGroup(doc *syntax.Document, opener int) bool { - prevIdx := prevReal(doc.Tokens, opener-1) + prevIdx := syntax.PrevReal(doc.Tokens, opener-1) return prevIdx >= 0 && doc.Tokens[prevIdx].Kind == syntax.TokenThrows } @@ -344,7 +344,7 @@ func isThrowsGroup(doc *syntax.Document, opener int) bool { // after a field modifier or id colon, a const or typedef keyword, a // map/list/set opener, or a service function return. func typeSlotAfterIdent(doc *syntax.Document, idx int) bool { - prev := prevReal(doc.Tokens, idx-1) + prev := syntax.PrevReal(doc.Tokens, idx-1) if prev < 0 { return false } @@ -590,12 +590,12 @@ func containerKeywordBefore(doc *syntax.Document, brace int) (syntax.TokenKind, pastParens: } - j = prevReal(doc.Tokens, j) + j = syntax.PrevReal(doc.Tokens, j) if j < 1 || doc.Tokens[j].Kind != syntax.TokenIdentifier { return 0, false } - k := prevReal(doc.Tokens, j-1) + k := syntax.PrevReal(doc.Tokens, j-1) if k < 0 { return 0, false } @@ -618,26 +618,6 @@ func deepestNode(path []syntax.Node) syntax.Node { return path[len(path)-1] } -// prevReal returns the index of the previous non-comment token strictly -// before idx, or -1. Comments are stream tokens but never participate in -// the grammar, so every adjacency lookup skips them. -func prevReal(toks []syntax.Token, idx int) int { - for idx >= 0 && syntax.IsComment(toks[idx].Kind) { - idx-- - } - - return idx -} - -// nextReal returns the index of the next non-comment token at or after idx. -func nextReal(toks []syntax.Token, idx int) int { - for idx < len(toks) && syntax.IsComment(toks[idx].Kind) { - idx++ - } - - return idx -} - // tokenOffset returns the byte offset of the first token of n. func tokenOffset(doc *syntax.Document, n syntax.Node) int { return doc.TokenPosition(n.TokStart()).Offset diff --git a/lsp/source/cycle_detect.go b/lsp/source/cycle_detect.go index 3d8b6d1..abcfff3 100644 --- a/lsp/source/cycle_detect.go +++ b/lsp/source/cycle_detect.go @@ -4,6 +4,7 @@ import ( "context" "fmt" "log/slog" + "slices" "go.lsp.dev/protocol" "go.lsp.dev/uri" @@ -12,42 +13,50 @@ import ( "github.com/karitham/thrift-ls/syntax" ) +// CycleCheck reports include edges that close a cycle: the include +// X -> Y is reported when Y transitively includes X back. Cycles of any +// length are caught, including self-includes. type CycleCheck struct{} func (c *CycleCheck) Diagnostic(ctx context.Context, ss *cache.Snapshot, changeFiles []uri.URI) (DiagnosticResult, error) { - includesMap := make(map[uri.URI][]Include) + closure := make(map[uri.URI][]Include) for _, file := range changeFiles { - _ = getIncludes(ctx, ss, file, &includesMap) + _ = getIncludes(ctx, ss, file, &closure) } - cyclePairs := cycleDetect(includesMap) + diagnostics := make(DiagnosticResult) + for file, includes := range closure { + // Reachability comes from the snapshot's include graph: parsing + // the closure above registered exactly these edges via Register, + // so there is no second graph to keep in sync. Dependents is + // cycle-safe, so the walk terminates on the cycles it finds. + deps := ss.Dependents(file) + for _, inc := range includes { + if !slices.Contains(deps, inc.file) { + continue + } - return cycleToDiagnosticItems(cyclePairs), nil + diagnostics[file] = append(diagnostics[file], cycleDiagnostic(inc)) + } + } + + return diagnostics, nil } func (c *CycleCheck) Name() string { return "CycleCheck" } -func cycleToDiagnosticItems(pairs []CyclePair) DiagnosticResult { - diagnostics := make(DiagnosticResult) - for i := range pairs { - diagnostics[pairs[i].file] = append(diagnostics[pairs[i].file], cyclePairToDiagnostic(pairs[i])) - } - - return diagnostics -} - -func cyclePairToDiagnostic(pair CyclePair) protocol.Diagnostic { - res := protocol.Diagnostic{ - Range: nodeRange(pair.include.pf, pair.include.include), +// cycleDiagnostic builds the warning for one include edge that closes a +// cycle back to its including file. +func cycleDiagnostic(inc Include) protocol.Diagnostic { + return protocol.Diagnostic{ + Range: nodeRange(inc.pf, inc.include), Severity: protocol.DiagnosticSeverityWarning, Code: protocol.String(CodeIncludeCycle), Source: protocol.NewOptional("thrift-ls"), - Message: protocol.String(fmt.Sprintf("cycle dependency in %s", pair.include.file)), + Message: protocol.String(fmt.Sprintf("cycle dependency in %s", inc.file)), } - - return res } type Include struct { @@ -56,55 +65,10 @@ type Include struct { pf *cache.ParsedFile } -type CyclePair struct { - file uri.URI - include Include -} - -// cycleDetect returns every include edge that closes a cycle: the pair -// (file, include file->Y) is reported when Y transitively includes file. -// Cycles of any length are caught, including self-includes. -func cycleDetect(includesMap map[uri.URI][]Include) []CyclePair { - // reaches reports whether from can reach target via include edges, - // cycle-safe via the seen set. - var reaches func(from, target uri.URI, seen map[uri.URI]bool) bool - - reaches = func(from, target uri.URI, seen map[uri.URI]bool) bool { - if from == target { - return true - } - - if seen[from] { - return false - } - - seen[from] = true - - for _, inc := range includesMap[from] { - if reaches(inc.file, target, seen) { - return true - } - } - - return false - } - - cyclePairs := make([]CyclePair, 0) - - for file, includes := range includesMap { - for _, inc := range includes { - if reaches(inc.file, file, make(map[uri.URI]bool)) { - cyclePairs = append(cyclePairs, CyclePair{ - file: file, - include: inc, - }) - } - } - } - - return cyclePairs -} - +// getIncludes collects the include closure of file into includesMap: the +// include edges of every file reachable from file, parsed through the +// snapshot so the ParsedFiles (and the graph edges Register records) are +// shared with the rest of the analysis. func getIncludes(ctx context.Context, ss *cache.Snapshot, file uri.URI, includesMap *map[uri.URI][]Include) error { pf, err := ss.Parse(ctx, file) if err != nil { diff --git a/lsp/source/cycle_detect_test.go b/lsp/source/cycle_detect_test.go index da24809..51021b8 100644 --- a/lsp/source/cycle_detect_test.go +++ b/lsp/source/cycle_detect_test.go @@ -1,199 +1,225 @@ package source import ( - "context" - "sort" + "maps" + "slices" + "strings" "testing" "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" + "go.lsp.dev/protocol" "go.lsp.dev/uri" "github.com/karitham/thrift-ls/lsp/cache" "github.com/karitham/thrift-ls/options" ) -func Test_cycleDetect(t *testing.T) { - includesMap := map[uri.URI][]Include{ - "/user.thrift": { - Include{file: "/goods.thrift"}, - Include{file: "/address.thrift"}, - }, - "/goods.thrift": {Include{file: "/user.thrift"}}, - "/address.thrift": {Include{file: "/user.thrift"}}, - } +func buildSnapshotForTest(t *testing.T, files []*cache.FileChange) *cache.Snapshot { + t.Helper() + + c := cache.New() + fs := cache.NewOverlayFS(c) + _ = fs.Update(t.Context(), files) + + view := cache.NewView("file:///tmp", fs, nil, options.Patch{}) + ss := cache.NewSnapshot(view, nil) - type args struct { - includesMap map[uri.URI][]Include + return ss +} + +// cyclePair identifies one reported cycle include: the file containing the +// include statement and the resolved URI it points at. Diagnostics carry +// this as the message "cycle dependency in ". +type cyclePair struct { + from uri.URI + to uri.URI +} + +func sortedPairs(t *testing.T, res DiagnosticResult) []cyclePair { + t.Helper() + + pairs := make([]cyclePair, 0) + for file, diags := range res { + for _, d := range diags { + msg, ok := d.Message.(protocol.String) + require.True(t, ok, "diagnostic message must be protocol.String") + + pairs = append(pairs, cyclePair{ + from: file, + to: uri.URI(strings.TrimPrefix(string(msg), "cycle dependency in ")), + }) + } } - tests := []struct { - name string - args args - want []CyclePair - }{ - { - name: "cycle", - args: args{ - includesMap: includesMap, - }, - want: []CyclePair{ - { - file: "/user.thrift", - include: Include{ - file: "/goods.thrift", - }, - }, - { - file: "/goods.thrift", - include: Include{file: "/user.thrift"}, - }, - { - file: "/user.thrift", - include: Include{file: "/address.thrift"}, - }, - { - file: "/address.thrift", - include: Include{file: "/user.thrift"}, - }, - }, - }, + slices.SortFunc(pairs, func(a, b cyclePair) int { + if c := strings.Compare(string(a.from), string(b.from)); c != 0 { + return c + } + + return strings.Compare(string(a.to), string(b.to)) + }) + + if len(pairs) == 0 { + return nil } - for _, tt := range tests { - t.Run(tt.name, func(t *testing.T) { - sort.SliceStable(tt.want, func(i, j int) bool { - if tt.want[i].file == tt.want[j].file { - return tt.want[i].include.file < tt.want[j].include.file - } - return tt.want[i].file < tt.want[j].file - }) + return pairs +} - got := cycleDetect(tt.args.includesMap) - sort.SliceStable(got, func(i, j int) bool { - if got[i].file == got[j].file { - return got[i].include.file < got[j].include.file - } +// runCycleCheck builds a snapshot from an in-memory file tree rooted at +// file:///tmp and runs CycleCheck starting from root. It asserts every +// diagnostic carries the cycle code and warning severity, and returns the +// reported (from, to) pairs sorted. +func runCycleCheck(t *testing.T, files map[string]string, root string) []cyclePair { + t.Helper() - return got[i].file < got[j].file - }) + names := slices.Sorted(maps.Keys(files)) - assert.Equal(t, tt.want, got) + changes := make([]*cache.FileChange, 0, len(files)) + for _, name := range names { + changes = append(changes, &cache.FileChange{ + URI: uri.URI("file:///tmp/" + name), + Version: 0, + Content: []byte(files[name]), + From: cache.FileChangeTypeDidOpen, }) } + + ss := buildSnapshotForTest(t, changes) + + res, err := (&CycleCheck{}).Diagnostic(t.Context(), ss, []uri.URI{uri.URI("file:///tmp/" + root)}) + require.NoError(t, err) + + for file, diags := range res { + for _, d := range diags { + assert.Equal(t, protocol.DiagnosticSeverityWarning, d.Severity, file) + assert.Equal(t, protocol.String(CodeIncludeCycle), d.Code, file) + } + } + + return sortedPairs(t, res) } -// Test_cycleDetectN pins cycle detection on arbitrary-length cycles: an -// include edge X -> Y closes a cycle when Y transitively includes X, no -// matter the cycle length. The existing 2-cycle case is covered in -// Test_cycleDetect; these cases exercise longer cycles, self-includes, and -// acyclic graphs. -func Test_cycleDetectN(t *testing.T) { - // K-On themed include graph: the band's songs include each other's - // tabs, and the clubroom includes everything. +// TestCycleCheck pins cycle detection end to end: an include edge X -> Y +// closes a cycle when Y transitively includes X back, no matter the cycle +// length. Every reported edge must point back along the cycle; acyclic +// graphs and unresolvable includes must produce nothing. +func TestCycleCheck(t *testing.T) { const ( - tea = "/songs/tea_time.thrift" - git = "/songs/gitah.thrift" - bass = "/songs/mio.thrift" - drum = "/songs/ritsu.thrift" - club = "/clubroom.thrift" + a = "a.thrift" + b = "b.thrift" + c = "c.thrift" + d = "d.thrift" + club = "clubroom.thrift" + tea = "tea_time.thrift" + git = "gitah.thrift" + bass = "mio.thrift" + drum = "ritsu.thrift" ) tests := []struct { name string - graph map[uri.URI][]Include - want []CyclePair + files map[string]string + root string + want []cyclePair }{ { name: "acyclic chain", - graph: map[uri.URI][]Include{ - tea: {Include{file: git}}, - git: {Include{file: bass}}, - bass: {}, + files: map[string]string{ + a: `include "b.thrift"`, + b: `include "c.thrift"`, + c: ``, }, + root: a, want: nil, }, { - name: "2-cycle", - graph: map[uri.URI][]Include{ - tea: {Include{file: git}}, - git: {Include{file: tea}}, + name: "two file cycle", + files: map[string]string{ + a: `include "b.thrift"`, + b: `include "a.thrift"`, }, - want: []CyclePair{ - {file: tea, include: Include{file: git}}, - {file: git, include: Include{file: tea}}, + root: a, + want: []cyclePair{ + {"file:///tmp/" + a, "file:///tmp/" + b}, + {"file:///tmp/" + b, "file:///tmp/" + a}, }, }, { - name: "3-cycle", - graph: map[uri.URI][]Include{ - tea: {Include{file: git}}, - git: {Include{file: bass}}, - bass: {Include{file: tea}}, + name: "three file cycle", + files: map[string]string{ + a: `include "b.thrift"`, + b: `include "c.thrift"`, + c: `include "a.thrift"`, }, - want: []CyclePair{ - {file: tea, include: Include{file: git}}, - {file: git, include: Include{file: bass}}, - {file: bass, include: Include{file: tea}}, + root: a, + want: []cyclePair{ + {"file:///tmp/" + a, "file:///tmp/" + b}, + {"file:///tmp/" + b, "file:///tmp/" + c}, + {"file:///tmp/" + c, "file:///tmp/" + a}, }, }, { - name: "4-cycle", - graph: map[uri.URI][]Include{ - tea: {Include{file: git}}, - git: {Include{file: bass}}, - bass: {Include{file: drum}}, - drum: {Include{file: tea}}, + name: "four file cycle", + files: map[string]string{ + a: `include "b.thrift"`, + b: `include "c.thrift"`, + c: `include "d.thrift"`, + d: `include "a.thrift"`, }, - want: []CyclePair{ - {file: tea, include: Include{file: git}}, - {file: git, include: Include{file: bass}}, - {file: bass, include: Include{file: drum}}, - {file: drum, include: Include{file: tea}}, + root: a, + want: []cyclePair{ + {"file:///tmp/" + a, "file:///tmp/" + b}, + {"file:///tmp/" + b, "file:///tmp/" + c}, + {"file:///tmp/" + c, "file:///tmp/" + d}, + {"file:///tmp/" + d, "file:///tmp/" + a}, }, }, { - name: "self-include", - graph: map[uri.URI][]Include{ - tea: {Include{file: tea}}, + name: "self include", + files: map[string]string{ + a: `include "a.thrift"`, }, - want: []CyclePair{ - {file: tea, include: Include{file: tea}}, + root: a, + want: []cyclePair{ + {"file:///tmp/" + a, "file:///tmp/" + a}, }, }, { name: "diamond into a cycle", - graph: map[uri.URI][]Include{ - club: {Include{file: tea}, Include{file: git}}, - tea: {Include{file: bass}}, - git: {Include{file: drum}}, - bass: {Include{file: drum}, Include{file: club}}, - drum: {Include{file: tea}}, + files: map[string]string{ + club: "include \"" + tea + "\"\ninclude \"" + git + "\"", + tea: "include \"" + bass + "\"", + git: "include \"" + drum + "\"", + bass: "include \"" + drum + "\"\ninclude \"" + club + "\"", + drum: "include \"" + tea + "\"", }, - want: []CyclePair{ - {file: club, include: Include{file: tea}}, - {file: club, include: Include{file: git}}, - {file: tea, include: Include{file: bass}}, - {file: git, include: Include{file: drum}}, - {file: bass, include: Include{file: drum}}, - {file: bass, include: Include{file: club}}, - {file: drum, include: Include{file: tea}}, + root: club, + want: []cyclePair{ + {"file:///tmp/" + club, "file:///tmp/" + git}, + {"file:///tmp/" + club, "file:///tmp/" + tea}, + {"file:///tmp/" + git, "file:///tmp/" + drum}, + {"file:///tmp/" + bass, "file:///tmp/" + club}, + {"file:///tmp/" + bass, "file:///tmp/" + drum}, + {"file:///tmp/" + drum, "file:///tmp/" + tea}, + {"file:///tmp/" + tea, "file:///tmp/" + bass}, }, }, + { + name: "unresolvable include is not a cycle", + files: map[string]string{ + a: `include "ghost.thrift"`, + }, + root: a, + want: nil, + }, } for _, tt := range tests { t.Run(tt.name, func(t *testing.T) { - got := cycleDetect(tt.graph) - sort.SliceStable(got, func(i, j int) bool { - if got[i].file == got[j].file { - return got[i].include.file < got[j].include.file - } - - return got[i].file < got[j].file - }) - - assert.ElementsMatch(t, tt.want, got) + got := runCycleCheck(t, tt.files, tt.root) + assert.Equal(t, tt.want, got) }) } } @@ -252,47 +278,8 @@ include "./test/address.thrift"` includeMap := make(map[uri.URI][]Include) - type args struct { - ctx context.Context - ss *cache.Snapshot - file uri.URI - includesMap *map[uri.URI][]Include - } + err := getIncludes(t.Context(), ss, "file:///tmp/user.thrift", &includeMap) + require.NoError(t, err) - tests := []struct { - name string - args args - want *map[uri.URI][]Include - assertion assert.ErrorAssertionFunc - }{ - { - name: "normal", - args: args{ - ctx: t.Context(), - ss: ss, - file: "file:///tmp/user.thrift", - includesMap: &includeMap, - }, - want: &expectIncludeMap, - assertion: assert.NoError, - }, - } - for _, tt := range tests { - t.Run(tt.name, func(t *testing.T) { - tt.assertion(t, getIncludes(tt.args.ctx, tt.args.ss, tt.args.file, tt.args.includesMap)) - - assert.Equal(t, tt.want, tt.args.includesMap) - }) - } -} - -func buildSnapshotForTest(t *testing.T, files []*cache.FileChange) *cache.Snapshot { - c := cache.New() - fs := cache.NewOverlayFS(c) - _ = fs.Update(t.Context(), files) - - view := cache.NewView("file:///tmp", fs, nil, options.Patch{}) - ss := cache.NewSnapshot(view, nil) - - return ss + assert.Equal(t, expectIncludeMap, includeMap) } diff --git a/lsp/source/semantic.go b/lsp/source/semantic.go index aef57fe..0660315 100644 --- a/lsp/source/semantic.go +++ b/lsp/source/semantic.go @@ -125,11 +125,11 @@ func classifyToken(i int, tok syntax.Token, names map[int]int, types map[int]boo return tokNumber, true } - if isTypeKeyword(tok.Kind) { + if syntax.IsTypeKeyword(tok.Kind) { return tokType, true } - if isKeyword(tok.Kind) { + if syntax.IsKeyword(tok.Kind) { return tokKeyword, true } @@ -233,32 +233,4 @@ func typeReferences(doc *syntax.Document) map[int]bool { return types } -// isTypeKeyword reports whether the kind is a base or container type -// keyword, which always appears in a type position. -func isTypeKeyword(k syntax.TokenKind) bool { - switch k { - case syntax.TokenMap, syntax.TokenList, syntax.TokenSet, syntax.TokenVoid, - syntax.TokenBool, syntax.TokenByte, syntax.TokenI8, syntax.TokenI16, - syntax.TokenI32, syntax.TokenI64, syntax.TokenDouble, - syntax.TokenString, syntax.TokenBinary, syntax.TokenSlist, syntax.TokenUUID: - return true - } - - return false -} - // isKeyword reports whether the kind is a reserved word. -func isKeyword(k syntax.TokenKind) bool { - switch k { - case syntax.TokenInclude, syntax.TokenCPPInclude, syntax.TokenCPPType, - syntax.TokenNamespace, syntax.TokenStruct, syntax.TokenUnion, - syntax.TokenException, syntax.TokenService, syntax.TokenEnum, - syntax.TokenConst, syntax.TokenTypedef, syntax.TokenOneway, - syntax.TokenAsync, syntax.TokenThrows, syntax.TokenExtends, - syntax.TokenRequired, syntax.TokenOptional, syntax.TokenTrue, - syntax.TokenFalse: - return true - } - - return false -} diff --git a/main.go b/main.go index ed09ccc..abc26b1 100644 --- a/main.go +++ b/main.go @@ -123,6 +123,7 @@ var constructFlags = []struct { {"throws", formatter.ConstructThrows}, {"list", formatter.ConstructList}, {"map", formatter.ConstructMap}, + {"set", formatter.ConstructSet}, } // formatFlags are the flags of the format subcommand. diff --git a/syntax/lexer.go b/syntax/lexer.go index 2c41935..e846c37 100644 --- a/syntax/lexer.go +++ b/syntax/lexer.go @@ -133,13 +133,47 @@ var keywordNames = func() map[TokenKind]string { return m }() -// isKeyword reports whether the token kind is one of the reserved words. -func isKeyword(k TokenKind) bool { +// IsKeyword reports whether the token kind is one of the reserved words. +func IsKeyword(k TokenKind) bool { _, ok := keywordNames[k] return ok } +// IsTypeKeyword reports whether the token kind is a base or container +// type keyword, which always appears in a type position. +func IsTypeKeyword(k TokenKind) bool { + switch k { + case TokenMap, TokenList, TokenSet, TokenVoid, + TokenBool, TokenByte, TokenI8, TokenI16, TokenI32, TokenI64, + TokenDouble, TokenString, TokenBinary, TokenSlist, TokenUUID: + return true + } + + return false +} + +// PrevReal returns the index of the previous real (non-comment) token +// strictly before i, or -1. Comments are stream tokens but never +// participate in the grammar, so every adjacency lookup skips them. +func PrevReal(toks []Token, i int) int { + for i >= 0 && IsComment(toks[i].Kind) { + i-- + } + + return i +} + +// NextReal returns the index of the next real (non-comment) token at or +// after i. +func NextReal(toks []Token, i int) int { + for i < len(toks) && IsComment(toks[i].Kind) { + i++ + } + + return i +} + var tokenKindNames = map[TokenKind]string{ TokenInvalid: "invalid", TokenEOF: "eof", TokenIdentifier: "identifier", TokenIntConstant: "int", diff --git a/syntax/parser.go b/syntax/parser.go index dcb197f..0342195 100644 --- a/syntax/parser.go +++ b/syntax/parser.go @@ -35,13 +35,7 @@ type parser struct { // --- token helpers --------------------------------------------------------- // nextReal returns the index of the next non-comment token at or after i. -func (p *parser) nextReal(i int) int { - for i < len(p.toks) && IsComment(p.toks[i].Kind) { - i++ - } - - return i -} +func (p *parser) nextReal(i int) int { return NextReal(p.toks, i) } func (p *parser) cur() Token { return p.toks[p.nextReal(p.pos)] } @@ -551,7 +545,7 @@ func (p *parser) parseField() (*Field, bool) { // as uuid are both valid types and common field names, and the thrift // compiler accepts them there. Rejecting keywords as field names breaks // valid IDLs. - if p.at(TokenIdentifier) || isKeyword(p.cur().Kind) { + if p.at(TokenIdentifier) || IsKeyword(p.cur().Kind) { f.Name = p.identifier() } else { p.errorfCur("expected field name, got %q", p.cur().Text) -- 2.51.2