diff --git a/README.md b/README.md index e066478..d1e5dc1 100644 --- a/README.md +++ b/README.md @@ -7,9 +7,9 @@ implementation. - **Language server**: completion, go to definition, find references, hover, diagnostics, rename, document symbols, and formatting — including **range formatting** (Format Selection). -- **Formatter**: a full rewrite of the old template-based formatter. It - preserves comments and blank lines, understands width, and is deterministic - and idempotent. +- **Formatter**: a full rewrite of the old template-based formatter. It is + lossless (comments, annotations, and blank lines survive everywhere), + width-aware, deterministic, and idempotent — properties enforced by fuzzing. Fork of https://github.com/joyme123/thrift-ls, parser + lexer + formatter rewritten, lsp overhauled @@ -28,6 +28,7 @@ a CLI formatter. thriftls [flags] run the language server (default) thriftls lsp [flags] run the language server thriftls format [flags] format a thrift file +thriftls dump [--ir] dump the parse tree and formatter IR ``` Run `thriftls --help` or `thriftls format --help` for the full flag list. @@ -114,17 +115,103 @@ Formatting flags: | `-w` | Overwrite the file with the formatted result | | `-d` | Print a diff instead of the formatted result | | `--printWidth` | Target line width (default 80) | -| `--indent` | Indentation: a literal like `" "` or `"\t"`, a number like `8`, or a legacy spec like `"2spaces"` | +| `--indent` | Indentation: a literal like `" "` or `"\t"` | | `--align` | `field`, `assign`, or `disable` | -| `--field-separator` | Struct/enum field separators: `add`, `remove`, `semicolon`, or `disable` (keep as written) | -| `--function-separator` | Service arg/throws separators: `add`, `remove`, `semicolon`, or `disable` (keep as written) | -| `--break-structs` | Always break struct/union/exception bodies onto multiple lines | -| `--break-enums` | Always break enum bodies onto multiple lines | +| `---separator` | Separators per construct (`struct`, `union`, `exception`, `enum`, `argument`, `throws`): `comma`, `semicolon`, `none`, or `preserve` (keep as written) | +| `--break-` | Always break the construct's bodies onto multiple lines (same constructs) | | `--config` | Path to a `thriftls.json` config file | | `-I` | Additional include path, like the thrift compiler's `-I` (repeatable) | Flags override the config file. +### Debugging: `dump` + +`thriftls dump` prints the parse tree — every token with its position, +blank-line count, and attached comment trivia, plus the node spans — which +is useful to understand how the lexer attached a comment or why the +formatter moved something: + +```bash +thriftls dump path/to/file.thrift +``` + +With `--ir`, it also builds the formatter's document IR, prints it (which +records the layout decisions on the groups), and dumps the IR tree showing +which groups broke and which stayed flat: + +```bash +thriftls dump --ir --printWidth 100 path/to/file.thrift +``` + +## Formatter behavior + +The formatter is **lossless**: comments, `@` annotations, and blank lines +survive formatting everywhere — including comments inside container types +(`map`), const values, and annotation parens. The +formatted output re-parses cleanly and formatting is idempotent and +deterministic; comment preservation, idempotency, and parseability are +enforced by a fuzzer over the full option space. + +### Conditional breaking, zig-style + +Like `zig fmt`, a **trailing delimiter** on the last item of a list decides +whether the list folds: + +```thrift +struct S { + 1: i32 a; + 2: string b; +} +``` + +stays multiline because the source ends the last field with `;`, while + +```thrift +struct S { + 1: i32 a + 2: string b +} +``` + +folds to `struct S { 1: i32 a 2: string b }` when it fits. The rule applies +to struct/union/exception bodies, enum bodies, function arguments, and +throws clauses. The `break.*` options force the multiline layout regardless +of the source. + +Note: a trailing delimiter only forces the multiline layout when the +separator mode actually emits it — `remove` drops separators, so it cannot +force a break (the output would not round-trip). + +### Width-aware folding + +Every group — struct bodies, function signatures, argument lists, throws +clauses, const lists, annotations — decides independently whether it fits +in the remaining width at its position. In particular, arguments and throws +fold independently: + +```thrift +service Processor { + string upload( + 1: string imageUrl, + 2: arguments.Size size, + 3: arguments.Identifier id, + ) throws (1: errors.ProcessingError err) +} +``` + +The arguments break (trailing commas), while the throws clause stays flat +because it fits on the closing paren's line. Comments or blank lines inside +a clause force it to break, without breaking the other clause. + +### Column alignment + +Struct/union/exception fields and enum values are column-aligned within +their group (`align: field`). Alignment groups split at blank lines and +comments, like whitespace. A group is aligned only when the padded columns +fit within `printWidth` — except layouts that were deliberately +column-aligned in the source, which are preserved even when they overflow. +Trailing comments may overflow their line without affecting alignment. + ## Configuration Configuration lives in a `thriftls.json` file, discovered by walking up from @@ -138,11 +225,17 @@ the file being formatted or the workspace root (like Biome). Set the "tabWidth": 4, "align": "field", "separators": { - "fields": "semicolon", - "functions": "add" + "structs": "semicolon", + "unions": "semicolon", + "exceptions": "semicolon", + "enums": "comma", + "arguments": "comma", + "throws": "comma" }, "break": { "structs": true, + "unions": true, + "exceptions": true, "enums": true }, "includePaths": ["/path/to/base"], @@ -159,13 +252,8 @@ position. Default: `80`. ### indent -Indentation can be written three ways: - -- a literal string: `" "` (two spaces) or `"\t"` (a tab) -- a number: `8` (eight spaces) -- a legacy spec: `"2spaces"`, `"1tab"` (kept as aliases) - -Default: `" "` (four spaces). +Indentation is a literal string of spaces or tabs, e.g. `" "` or +`"\t"`. Default: `" "` (four spaces). ### tabWidth @@ -179,37 +267,51 @@ Controls column alignment of struct/union/exception fields and enum values. - `assign`: Align the `=` sign for default values - `disable`: No alignment -### separators +See [Column alignment](#column-alignment) for how alignment interacts with +width, comments, and blank lines. -Controls trailing separators for the two field contexts independently. +### separators -`separators.fields` — struct/union/exception fields and enum values: +Controls trailing separators per construct, independently. The +`separators` object has one key per construct: `structs`, `unions`, +`exceptions`, `enums`, `arguments` (function arguments), and `throws` +(throws entries). Each accepts: -- `disable`: Keep as written (default) -- `add`: Always add trailing commas +- `comma`: Always add trailing commas - `semicolon`: Always add trailing semicolons -- `remove`: Remove trailing separators +- `none`: Remove trailing separators +- `preserve`: Keep as written (default) -`separators.functions` — service arguments and throws entries, with the -same values: +For example, semicolons in structs and commas in enums: -- `disable`: Keep as written (default) -- `add`: Always add trailing commas -- `semicolon`: Always add trailing semicolons -- `remove`: Remove trailing separators +```json +"separators": { + "structs": "semicolon", + "unions": "semicolon", + "exceptions": "semicolon", + "enums": "comma" +} +``` Broken (multiline) argument and throws blocks are column-aligned like struct fields, controlled by `align`. +The separator mode also interacts with [conditional +breaking](#conditional-breaking-zig-style): a mode that keeps or adds +trailing separators (`preserve`, `comma`, `semicolon`) lets a source +trailing delimiter force the multiline layout, while `none` always folds +when the group fits. Under `preserve`, a *mixed* separator pattern (some +fields separated, some not) also forces the multiline layout — a flat line +whose separators are inconsistently present looks broken. + ### break Forces layouts that would otherwise collapse to one line to stay -multiline. +multiline, regardless of the source's trailing delimiters. Like +`separators`, the `break` object has one key per construct: `structs`, +`unions`, `exceptions`, `enums`, `arguments`, `throws`. -- `break.structs`: Always break struct, union, and exception bodies -- `break.enums`: Always break enum bodies - -Both default to `false`. +All default to `false`. ### includePaths @@ -235,3 +337,15 @@ Controls logging verbosity (the server logs to `$TMPDIR/thriftls.log`): go test ./... # unit and fuzz regression tests bash tests/e2e/run-e2e.sh # end-to-end formatter tests ``` + +The formatter is fuzz-tested end to end: `FuzzFormat` checks that any clean +document formats without errors, keeps every comment, is idempotent and +deterministic across the whole option space. The lexer, parser, doc +printer, LSP offset mapper, and range formatting each have their own fuzz +targets; the corpus entries under `testdata/fuzz` are permanent regression +tests. + +`thriftls dump` (see above) is the debugging companion: it shows the parse +tree and the formatter's document IR with the layout decisions, so a +formatting issue can be pinned to the parser, the IR construction, or the +printer. diff --git a/diff.go b/diff.go index 730f2b2..55f58b3 100644 --- a/diff.go +++ b/diff.go @@ -47,6 +47,7 @@ func Diff(oldName string, old []byte, newName string, new []byte) []byte { if bytes.Equal(old, new) { return nil } + x := lines(old) y := lines(new) @@ -83,6 +84,7 @@ func Diff(oldName string, old []byte, newName string, new []byte) []byte { start.x-- start.y-- } + end := m for end.x < len(x) && end.y < len(y) && x[end.x] == y[end.y] { end.x++ @@ -95,6 +97,7 @@ func Diff(oldName string, old []byte, newName string, new []byte) []byte { ctext = append(ctext, "-"+s) count.x++ } + for _, s := range y[done.y:start.y] { ctext = append(ctext, "+"+s) count.y++ @@ -110,7 +113,9 @@ func Diff(oldName string, old []byte, newName string, new []byte) []byte { count.x++ count.y++ } + done = end + continue } @@ -122,6 +127,7 @@ func Diff(oldName string, old []byte, newName string, new []byte) []byte { count.x++ count.y++ } + done = pair{start.x + n, start.y + n} // Format and emit chunk. @@ -130,13 +136,17 @@ func Diff(oldName string, old []byte, newName string, new []byte) []byte { if count.x > 0 { chunk.x++ } + if count.y > 0 { chunk.y++ } + fmt.Fprintf(&out, "@@ -%d,%d +%d,%d @@\n", chunk.x, count.x, chunk.y, count.y) + for _, s := range ctext { out.WriteString(s) } + count.x = 0 count.y = 0 ctext = ctext[:0] @@ -154,6 +164,7 @@ func Diff(oldName string, old []byte, newName string, new []byte) []byte { count.x++ count.y++ } + done = end } @@ -172,6 +183,7 @@ func lines(x []byte) []string { // using the same text as BSD/GNU diff (including the leading backslash). l[len(l)-1] += "\n\\ No newline at end of file\n" } + return l } @@ -194,6 +206,7 @@ func tgs(x, y []string) []pair { m[s] = c - 1 } } + for _, s := range y { if c := m[s]; c > -8 { m[s] = c - 4 @@ -207,12 +220,14 @@ func tgs(x, y []string) []pair { // yi[i] = increasing indexes of unique strings in y. // inv[i] = index j such that x[xi[i]] = y[yi[j]]. var xi, yi, inv []int + for i, s := range y { if m[s] == -1+-4 { m[s] = len(yi) yi = append(yi, i) } } + for i, s := range x { if j, ok := m[s]; ok && j >= 0 { xi = append(xi, i) @@ -227,10 +242,12 @@ func tgs(x, y []string) []pair { J := inv n := len(xi) T := make([]int, n) + L := make([]int, n) for i := range T { T[i] = n + 1 } + for i := range n { k := sort.Search(n, func(k int) bool { return T[k] >= J[i] @@ -238,14 +255,17 @@ func tgs(x, y []string) []pair { T[k] = J[i] L[i] = k + 1 } + k := 0 for _, v := range L { if k < v { k = v } } + seq := make([]pair, 2+k) seq[1+k] = pair{len(x), len(y)} // sentinel at end + lastj := n for i := n - 1; i >= 0; i-- { if L[i] == k && J[i] < lastj { @@ -253,6 +273,8 @@ func tgs(x, y []string) []pair { k-- } } + seq[0] = pair{0, 0} // sentinel at start + return seq } diff --git a/doc/doc.go b/doc/doc.go index 5973bae..6fc9db6 100644 --- a/doc/doc.go +++ b/doc/doc.go @@ -34,13 +34,16 @@ 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 } diff --git a/doc/dump.go b/doc/dump.go new file mode 100644 index 0000000..0ec909a --- /dev/null +++ b/doc/dump.go @@ -0,0 +1,94 @@ +package doc + +import ( + "fmt" + "strings" +) + +// Dump renders the document IR as an indented tree, for debugging: every +// node with its type, text, group id and break state. Printing a doc +// mutates group break states in place, so dumping after Print shows the +// layout decisions. +func Dump(d Doc) string { + var b strings.Builder + dumpDoc(&b, d, "") + + return b.String() +} + +func dumpDoc(b *strings.Builder, d Doc, ind string) { + switch v := d.(type) { + case nil: + fmt.Fprintf(b, "%s\n", ind) + case Concat: + fmt.Fprintf(b, "%sConcat\n", ind) + + for _, c := range v { + dumpDoc(b, c, ind+" ") + } + case *group: + extra := "" + if v.id != 0 { + extra = fmt.Sprintf(" id=%d", v.id) + } + + if v.brk { + extra += " brk" + } + + if v.expanded != nil { + fmt.Fprintf(b, "%sConditionalGroup%s (%d states)\n", ind, extra, len(v.expanded)) + + for _, s := range v.expanded { + dumpDoc(b, s, ind+" ") + } + + return + } + + fmt.Fprintf(b, "%sGroup%s\n", ind, extra) + dumpDoc(b, v.doc, ind+" ") + case *ifBreak: + fmt.Fprintf(b, "%sIfBreak (group %d)\n", ind, v.groupID) + dumpDoc(b, v.breakDoc, ind+" [broken] ") + dumpDoc(b, v.flatDoc, ind+" [flat] ") + case *indent: + fmt.Fprintf(b, "%sIndent\n", ind) + dumpDoc(b, v.doc, ind+" ") + case *align: + fmt.Fprintf(b, "%sAlign %d\n", ind, v.n) + dumpDoc(b, v.doc, ind+" ") + case *lineSuffix: + fmt.Fprintf(b, "%sLineSuffix\n", ind) + dumpDoc(b, v.doc, ind+" ") + case lineSuffixBoundary: + fmt.Fprintf(b, "%sLineSuffixBoundary\n", ind) + case breakParent: + fmt.Fprintf(b, "%sBreakParent\n", ind) + case trim: + fmt.Fprintf(b, "%sTrim\n", ind) + case Text: + if v == "" { + fmt.Fprintf(b, "%sText \"\"\n", ind) + } else { + fmt.Fprintf(b, "%sText %q\n", ind, string(v)) + } + case LineDoc: + kind := "Line" + if v.Hard { + kind = "HardLine" + } + + if v.Soft { + kind += " (soft)" + } + + if v.Literal { + kind += " (literal)" + } + + fmt.Fprintf(b, "%s%s\n", ind, kind) + default: + fmt.Fprintf(b, "%s%T\n", ind, d) + } +} diff --git a/doc/print.go b/doc/print.go index 158c5bc..4d6b043 100644 --- a/doc/print.go +++ b/doc/print.go @@ -3,6 +3,7 @@ package doc import ( "fmt" "math" + "slices" "strings" ) @@ -47,8 +48,10 @@ func (i indentation) add(o Options) indentation { if o.Indent == "" { o.Indent = strings.Repeat(" ", o.TabWidth) } + i.value += o.Indent i.length += o.TabWidth + return i } @@ -56,11 +59,13 @@ func (i indentation) align(n int, o Options) indentation { if o.Indent != "\t" { i.value += strings.Repeat(" ", n) i.length += n + return i } // With tabs, a numeric alignment renders as one tab, matching Prettier. i.value += "\t" i.length += o.TabWidth + return i } @@ -71,18 +76,23 @@ func Print(d Doc, o Options) (string, error) { if o.PrintWidth <= 0 { return "", fmt.Errorf("doc: PrintWidth must be positive, got %d", o.PrintWidth) } + if o.TabWidth <= 0 { return "", fmt.Errorf("doc: TabWidth must be positive, got %d", o.TabWidth) } + if o.NewLine == "" { o.NewLine = "\n" } + if o.Indent == "" { o.Indent = strings.Repeat(" ", o.TabWidth) } propagateBreaks(d) + p := &printer{o: o, groupMode: map[int]mode{}} + return p.run(d) } @@ -104,13 +114,17 @@ func (p *printer) write(s string) { // columns were removed. func (p *printer) trim() int { n := 0 - for i := len(p.out) - 1; i >= 0; i-- { - if p.out[i] != ' ' && p.out[i] != '\t' { + + for _, v := range slices.Backward(p.out) { + if v != ' ' && v != '\t' { break } + n++ } + p.out = p.out[:len(p.out)-n] + return n } @@ -128,14 +142,15 @@ func (p *printer) run(d Doc) (string, error) { if v != "" { s := string(v) p.write(s) + if len(commands) > 0 { p.position += stringWidth(s) } } case Concat: - for i := len(v) - 1; i >= 0; i-- { - commands = append(commands, command{indentation: cmd.indentation, mode: cmd.mode, doc: v[i]}) + for _, v0 := range slices.Backward(v) { + commands = append(commands, command{indentation: cmd.indentation, mode: cmd.mode, doc: v0}) } case *indent: @@ -150,6 +165,7 @@ func (p *printer) run(d Doc) (string, error) { case *group: { gcmd := p.printGroup(cmd, v, commands) + commands = append(commands, gcmd) if v.id != 0 { p.groupMode[v.id] = gcmd.mode @@ -165,12 +181,14 @@ func (p *printer) run(d Doc) (string, error) { groupMode = modeFlat } } + var contents Doc if groupMode == modeBreak { contents = v.breakDoc } else { contents = v.flatDoc } + if contents != nil { commands = append(commands, command{indentation: cmd.indentation, mode: cmd.mode, doc: contents}) } @@ -191,12 +209,14 @@ func (p *printer) run(d Doc) (string, error) { p.write(" ") p.position++ } + break } // A hard line printed in flat mode invalidates earlier // measurements of enclosing groups; the next group must // remeasure. p.shouldRemeasure = true + fallthrough case modeBreak: @@ -204,12 +224,15 @@ func (p *printer) run(d Doc) (string, error) { // Print the pending end-of-line suffixes before the // newline, then the line itself. commands = append(commands, cmd) - for i := len(p.lineSuffix) - 1; i >= 0; i-- { - commands = append(commands, p.lineSuffix[i]) + for _, v := range slices.Backward(p.lineSuffix) { + commands = append(commands, v) } + p.lineSuffix = nil + break } + if v.Literal { p.write(newLine) p.position = 0 @@ -230,9 +253,10 @@ func (p *printer) run(d Doc) (string, error) { // Flush remaining suffixes at the end of the document, in case there // is no line break after them. if len(commands) == 0 && len(p.lineSuffix) > 0 { - for i := len(p.lineSuffix) - 1; i >= 0; i-- { - commands = append(commands, p.lineSuffix[i]) + for _, v := range slices.Backward(p.lineSuffix) { + commands = append(commands, v) } + p.lineSuffix = nil } } @@ -248,6 +272,7 @@ func (p *printer) printGroup(cmd command, g *group, rest []command) command { if g.brk { m = modeBreak } + return command{indentation: cmd.indentation, mode: m, doc: g.doc} } @@ -290,18 +315,22 @@ func (p *printer) fits(next command, rest []command, remainingWidth int, hasLine // output is only used for width counting because trim needs to look // backwards for spaces. var output strings.Builder + commands := []command{next} for remainingWidth >= 0 { if p.fitsDBG { fmt.Printf(" fitloop rem=%d cmds=%d\n", remainingWidth, len(commands)) } + if len(commands) == 0 { if restIndex == 0 { return true } + restIndex-- commands = append(commands, rest[restIndex]) + continue } @@ -313,18 +342,21 @@ func (p *printer) fits(next command, rest []command, remainingWidth int, hasLine case Text: if v != "" { s := string(v) + if hasPendingSpace { output.WriteString(" ") + remainingWidth-- hasPendingSpace = false } + output.WriteString(s) remainingWidth -= stringWidth(s) } case Concat: - for i := len(v) - 1; i >= 0; i-- { - commands = append(commands, command{indentation: cmd.indentation, mode: cmd.mode, doc: v[i]}) + for _, v0 := range slices.Backward(v) { + commands = append(commands, command{indentation: cmd.indentation, mode: cmd.mode, doc: v0}) } case *indent, *align: @@ -338,14 +370,17 @@ func (p *printer) fits(next command, rest []command, remainingWidth int, hasLine if v.brk && cmd.mode == modeFlat { return false } + groupMode := cmd.mode if v.brk { groupMode = modeBreak } + contents := v.doc if v.expanded != nil && groupMode == modeBreak { contents = v.expanded[len(v.expanded)-1] } + commands = append(commands, command{indentation: cmd.indentation, mode: groupMode, doc: contents}) case *ifBreak: @@ -357,12 +392,14 @@ func (p *printer) fits(next command, rest []command, remainingWidth int, hasLine groupMode = modeFlat } } + var contents Doc if groupMode == modeBreak { contents = v.breakDoc } else { contents = v.flatDoc } + if contents != nil { commands = append(commands, command{indentation: cmd.indentation, mode: cmd.mode, doc: contents}) } @@ -371,6 +408,7 @@ func (p *printer) fits(next command, rest []command, remainingWidth int, hasLine if cmd.mode == modeBreak || v.Hard { return true } + if !v.Soft { hasPendingSpace = true } @@ -384,6 +422,7 @@ func (p *printer) fits(next command, rest []command, remainingWidth int, hasLine } } } + return false } @@ -394,17 +433,21 @@ func contentsOf(d Doc) Doc { case *align: return v.doc } + return nil } func trimTrailingWidth(s string) int { n := 0 + for i := len(s) - 1; i >= 0; i-- { if s[i] != ' ' && s[i] != '\t' { break } + n++ } + return n } @@ -413,6 +456,7 @@ func trimTrailingWidth(s string) int { // place, matching Prettier's traversal. func propagateBreaks(d Doc) { visited := map[*group]bool{} + var stack []*group enter := func(d Doc) bool { @@ -424,8 +468,10 @@ func propagateBreaks(d Doc) { if visited[v] { return false } + visited[v] = true } + return true } exit := func(d Doc) { @@ -449,6 +495,7 @@ func traverseDoc(d Doc, enter func(Doc) bool, exit func(Doc), includeConditional if d == nil || !enter(d) { return } + switch v := d.(type) { case Concat: for _, part := range v { @@ -460,6 +507,7 @@ func traverseDoc(d Doc, enter func(Doc) bool, exit func(Doc), includeConditional traverseDoc(state, enter, exit, includeConditionalGroups) } } + traverseDoc(v.doc, enter, exit, includeConditionalGroups) case *indent: traverseDoc(v.doc, enter, exit, includeConditionalGroups) @@ -471,5 +519,6 @@ func traverseDoc(d Doc, enter func(Doc) bool, exit func(Doc), includeConditional case *lineSuffix: traverseDoc(v.doc, enter, exit, includeConditionalGroups) } + exit(d) } diff --git a/doc/print_fuzz_test.go b/doc/print_fuzz_test.go index d7c9e18..7f4a889 100644 --- a/doc/print_fuzz_test.go +++ b/doc/print_fuzz_test.go @@ -36,18 +36,22 @@ func FuzzPrint(f *testing.F) { if !ok { t.Skip("program exhausted") } + return d } opts := Options{PrintWidth: 1 + len(program)%80, Indent: " ", TabWidth: 2, NewLine: "\n"} + first, err := Print(build(), opts) if err != nil { t.Fatalf("Print: %v", err) } + second, err := Print(build(), opts) if err != nil { t.Fatalf("Print (second): %v", err) } + if first != second { t.Fatalf("Print is not deterministic:\n%q\n%q", first, second) } @@ -89,10 +93,12 @@ func buildDoc(program []byte, depth int) (Doc, bool) { switch op := program[0]; op { case opText: n := int(program[1%len(program)]) % 8 + text := make([]rune, 0, n) for i := range n { text = append(text, textChars[int(program[(2+i)%len(program)])%len(textChars)]) } + return Text(string(text)), true case opLine, opSoftLine, opHardLine: @@ -110,26 +116,32 @@ func buildDoc(program []byte, depth int) (Doc, bool) { if !ok { return Concat{}, false } + if op == opGroupBreak { return GroupBreak(inner), true } + return Group(inner), true case opConcat: n := int(program[1%len(program)]) % 5 parts := make([]Doc, 0, n) + offset := 2 for range n { if offset >= len(program) { return Concat(parts), true } + part, ok := buildDoc(program[offset:], depth+1) if !ok { return Concat(parts), true } + parts = append(parts, part) offset++ } + return Concat(parts), true case opIndent: @@ -137,6 +149,7 @@ func buildDoc(program []byte, depth int) (Doc, bool) { if !ok { return Concat{}, false } + return Indent(inner), true case opAlign: @@ -144,14 +157,17 @@ func buildDoc(program []byte, depth int) (Doc, bool) { if !ok { return Concat{}, false } + return Align(int(program[1%len(program)])%5, inner), true case opIfBreak: brk, ok1 := buildDoc(program[1:], depth+1) + flat, ok2 := buildDoc(program[2%len(program):], depth+1) if !ok1 || !ok2 { return Concat{}, false } + return IfBreak(brk, flat), true case opLineSuffix: @@ -159,6 +175,7 @@ func buildDoc(program []byte, depth int) (Doc, bool) { if !ok { return Concat{}, false } + return LineSuffix(inner), true case opConditional: @@ -166,10 +183,12 @@ func buildDoc(program []byte, depth int) (Doc, bool) { if !ok { return Concat{}, false } + second, ok := buildDoc(program[2%len(program):], depth+1) if !ok { return Concat{}, false } + return ConditionalGroup(0, first, second, GroupBreak(first)), true case opTrim: diff --git a/doc/print_test.go b/doc/print_test.go index 7597bc4..4c6b88e 100644 --- a/doc/print_test.go +++ b/doc/print_test.go @@ -14,10 +14,12 @@ func testOpts(width int) Options { func printT(t *testing.T, d Doc, o Options) string { t.Helper() + got, err := Print(d, o) if err != nil { t.Fatalf("Print: %v", err) } + return got } @@ -411,6 +413,7 @@ func TestPrintRemeasure(t *testing.T) { T("dd"), }) got := printT(t, doc, testOpts(8)) + want := "a\nbbbb\ncc dd" if got != want { t.Errorf("got %q, want %q", got, want) @@ -419,6 +422,7 @@ 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") @@ -519,6 +523,7 @@ func TestPrintIdempotentDocs(t *testing.T) { } for _, d := range docs { first := printT(t, d, testOpts(2)) + second := printT(t, d, testOpts(2)) if first != second { t.Errorf("not idempotent: %q vs %q", first, second) @@ -532,6 +537,7 @@ func TestDocMutability(t *testing.T) { build := func() Doc { return Group(Concat{T("a"), HardLine, T("b"), SoftLine, T("c")}) } + first := printT(t, build(), testOpts(80)) for i := range 10 { if got := printT(t, build(), testOpts(80)); got != first { diff --git a/doc/width.go b/doc/width.go index 0dde55c..65a354b 100644 --- a/doc/width.go +++ b/doc/width.go @@ -18,17 +18,21 @@ func stringWidth(s string) int { // Fast path: pure ASCII without control characters has width equal to // its byte length. ascii := true + for i := 0; i < len(s); i++ { if s[i] >= 0x80 || s[i] <= 0x1f || s[i] == 0x7f { ascii = false + break } } + if ascii { return len(s) } width := 0 + for _, r := range s { switch { case unicode.Is(unicode.Cc, r): @@ -41,6 +45,7 @@ func stringWidth(s string) int { width++ } } + return width } @@ -58,5 +63,6 @@ func isWide(r rune) bool { case width.EastAsianWide, width.EastAsianFullwidth: return true } + return false } diff --git a/formatter/body.go b/formatter/body.go index 1698303..3710d2d 100644 --- a/formatter/body.go +++ b/formatter/body.go @@ -5,137 +5,224 @@ import ( "github.com/karitham/thrift-ls/syntax" ) -// structLike formats struct, union, and exception declarations. +// structLike formats struct, union, and exception declarations. The header +// renders as a token run up to and including the open brace; the brace +// text itself is emitted by bracedBody. func (f *formatter) structLike(v *syntax.Struct) doc.Doc { + open := f.scanKind(v.TokStart(), v.TokEnd(), syntax.TokenLBrace) + close := f.scanKind(open+1, v.TokEnd(), syntax.TokenRBrace) + parts := []doc.Doc{ - doc.Text(v.Kind.String()), - doc.Text(" "), - doc.Text(v.Name.Text), - f.bracedBody(v.Fields), + f.emitTokens(v.TokStart(), open, emitOpts{skipText: map[int]bool{open: true}, breakSkip: true}), + f.bracedBody(v.Fields, open, close, close != v.TokEnd(), f.constructOf(v.Kind)), + } + if v.Annotations != nil { + parts = append(parts, f.breakBeforeAnnotations(close)) } - parts = append(parts, f.annotationsDoc(v.Annotations)) + + parts = append(parts, f.annotationsDoc(v.Annotations, v.Annotations != nil && v.Annotations.TokEnd() == v.TokEnd())) + parts = append(parts, f.afterAnnotations(v.Annotations, v.TokEnd())) + return doc.Concat(parts) } // enum formats an enum declaration. func (f *formatter) enum(v *syntax.Enum) doc.Doc { + open := f.scanKind(v.TokStart(), v.TokEnd(), syntax.TokenLBrace) + close := f.scanKind(open+1, v.TokEnd(), syntax.TokenRBrace) + parts := []doc.Doc{ - doc.Text("enum "), - doc.Text(v.Name.Text), - f.bracedEnumBody(v.Values), + f.emitTokens(v.TokStart(), open, emitOpts{skipText: map[int]bool{open: true}, breakSkip: true}), + f.bracedEnumBody(v.Values, open, close, close != v.TokEnd()), + } + if v.Annotations != nil { + parts = append(parts, f.breakBeforeAnnotations(close)) } - parts = append(parts, f.annotationsDoc(v.Annotations)) + + parts = append(parts, f.annotationsDoc(v.Annotations, v.Annotations != nil && v.Annotations.TokEnd() == v.TokEnd())) + parts = append(parts, f.afterAnnotations(v.Annotations, v.TokEnd())) + return doc.Concat(parts) } +// scanKind returns the first token of the given kind in [start, end], or +// end when not found. +func (f *formatter) scanKind(start, end int, kind syntax.TokenKind) int { + for i := start; i <= end; i++ { + if f.token(i).Kind == kind { + return i + } + } + + return end +} + // bracedBody formats "{ fields }": flat as "S { 1: i32 a }" when it fits, // otherwise one field per line. Field alignment and trailing separators -// switch on the body group breaking. -func (f *formatter) bracedBody(fields []*syntax.Field) doc.Doc { - if len(fields) == 0 { - return doc.Text(" {}") +// switch on the body group breaking. open and close are the brace token +// indices; closeTrailing reports whether the close brace is not the last +// token of the declaration (annotations follow), so its trailing trivia +// belongs here rather than to the declaration's trailing comments. +// constructOf returns the construct for the struct-like kind. +func (f *formatter) constructOf(kind syntax.StructKind) Construct { + switch kind { + case syntax.TokenUnion: + return ConstructUnion + case syntax.TokenException: + return ConstructException } + + return ConstructStruct +} + +func (f *formatter) bracedBody(fields []*syntax.Field, open, close int, closeTrailing bool, c Construct) doc.Doc { bodyID := f.id() - inner := append([]doc.Doc{doc.Line, f.fieldList(fields, bodyID)}, f.closingTrivia(fields[len(fields)-1])...) - open := append([]doc.Doc{doc.Text(" {")}, f.openTrivia(fields[0])...) + sepMode := f.opts.Separator.Get(c) + inner := append([]doc.Doc{doc.Line, f.fieldList(fields, bodyID, sepMode)}, f.closingTriviaAt(close)...) + closeBreak := doc.IfBreak(doc.SoftLine, doc.Text(" ")) + + if len(fields) == 0 { + inner = append([]doc.Doc{}, f.closingTriviaAt(close)...) + closeBreak = doc.IfBreak(doc.SoftLine, doc.Concat{}) + } + + openDoc := append([]doc.Doc{doc.Text(" {")}, f.openTriviaAt(open)...) + content := doc.Concat{ - doc.Concat(open), + doc.Concat(openDoc), doc.Indent(doc.Concat(inner)), - doc.IfBreak(doc.SoftLine, doc.Text(" ")), - doc.Text("}"), + closeBreak, + f.emitTokens(close, close, emitOpts{trailing: closeTrailing}), } - if f.opts.BreakStructs || hasTrailingDelim(fields[len(fields)-1].Sep) { + if len(fields) > 0 && (f.opts.Break.Get(c) || sepForcesBreak(sepsOfFields(fields), sepMode)) { // BreakParent inside the group forces it to the broken layout. content = doc.Concat{doc.BreakParent, content} } + return doc.GroupID(bodyID, content) } // bracedEnumBody is bracedBody for enum values. -func (f *formatter) bracedEnumBody(values []*syntax.EnumValue) doc.Doc { +func (f *formatter) bracedEnumBody(values []*syntax.EnumValue, open, close int, closeTrailing bool) doc.Doc { + bodyID := f.id() + inner := append([]doc.Doc{doc.Line, f.enumValueList(values, bodyID)}, f.closingTriviaAt(close)...) + closeBreak := doc.IfBreak(doc.SoftLine, doc.Text(" ")) + if len(values) == 0 { - return doc.Text(" {}") + inner = append([]doc.Doc{}, f.closingTriviaAt(close)...) + closeBreak = doc.IfBreak(doc.SoftLine, doc.Concat{}) } - bodyID := f.id() - inner := append([]doc.Doc{doc.Line, f.enumValueList(values, bodyID)}, f.closingTrivia(values[len(values)-1])...) - open := append([]doc.Doc{doc.Text(" {")}, f.openTrivia(values[0])...) + + openDoc := append([]doc.Doc{doc.Text(" {")}, f.openTriviaAt(open)...) + content := doc.Concat{ - doc.Concat(open), + doc.Concat(openDoc), doc.Indent(doc.Concat(inner)), - doc.IfBreak(doc.SoftLine, doc.Text(" ")), - doc.Text("}"), + closeBreak, + f.emitTokens(close, close, emitOpts{trailing: closeTrailing}), } - if f.opts.BreakEnums || hasTrailingDelim(values[len(values)-1].Sep) { + if len(values) > 0 && (f.opts.Break.Get(ConstructEnum) || sepForcesBreak(sepsOfValues(values), f.opts.Separator.Get(ConstructEnum))) { content = doc.Concat{doc.BreakParent, content} } + return doc.GroupID(bodyID, content) } -// closingTrivia returns the comments attached to the token that closes a -// body (the one right after the last item), as docs each ending with a hard -// line. The leading hard line forces the body to break so the comments stay -// inside; the caller's closing line provides the newline after the last one. -func (f *formatter) closingTrivia(last syntax.Node) []doc.Doc { - close := f.token(last.TokEnd() + 1) +// closingTriviaAt returns the comments attached before the closing token, +// as docs each ending with a hard line. The leading hard line forces the +// body to break so the comments stay inside; the caller's closing line +// provides the newline after the last one. +func (f *formatter) closingTriviaAt(close int) []doc.Doc { + tok := f.token(close) + var parts []doc.Doc - if len(close.Leading) > 0 { + if len(tok.Leading) > 0 { parts = append(parts, doc.HardLine) + prevBlank := 0 - for i, c := range close.Leading { + for i, c := range tok.Leading { parts = append(parts, f.blankLineDocs(c.BlankLinesBefore-prevBlank, doc.HardLine)...) prevBlank = c.BlankLinesBefore - parts = append(parts, doc.Text(c.Text)) - if i < len(close.Leading)-1 { + + parts = append(parts, doc.Text(trimComment(c.Text))) + if i < len(tok.Leading)-1 { parts = append(parts, doc.HardLine) } } } + return parts } -// openTrivia returns the trailing comments of the token that opens a body -// (the one right before the first item), as line-suffix docs, plus a break -// parent so the body goes multiline. Empty when there are none. -func (f *formatter) openTrivia(first syntax.Node) []doc.Doc { - open := f.token(first.TokStart() - 1) +// openTriviaAt returns the trailing comments of the opening token as +// line-suffix docs, plus a break parent so the body goes multiline. Empty +// when there are none. +func (f *formatter) openTriviaAt(open int) []doc.Doc { + tok := f.token(open) + var parts []doc.Doc - if len(open.Trailing) > 0 { - for _, c := range open.Trailing { - parts = append(parts, doc.LineSuffix(doc.Text(" "+c.Text))) + + if len(tok.Trailing) > 0 { + for _, c := range tok.Trailing { + parts = append(parts, doc.LineSuffix(doc.Text(" "+trimComment(c.Text)))) } + parts = append(parts, doc.BreakParent) } + return parts } // service formats a service declaration. The body is always multiline: // functions are too complex to flatten. func (f *formatter) service(v *syntax.Service) doc.Doc { + open := f.scanKind(v.TokStart(), v.TokEnd(), syntax.TokenLBrace) + close := f.scanKind(open+1, v.TokEnd(), syntax.TokenRBrace) + var parts []doc.Doc + for i, fn := range v.Functions { if i > 0 { parts = append(parts, doc.HardLineNoBreak) parts = append(parts, f.blankLines(fn, doc.HardLineNoBreak)...) + } else { + parts = append(parts, f.blankLines(fn, doc.HardLineNoBreak)...) } + parts = append(parts, f.function(fn)) } + inner := doc.Concat{doc.Concat(parts), doc.Concat(f.closingTriviaAt(close))} + if len(v.Functions) > 0 { + // The first function starts its own line; the closing trivia + // provides the break before it for empty bodies, so the blank + // count does not double. + inner = doc.Concat{doc.HardLineNoBreak, doc.Concat(parts), doc.Concat(f.closingTriviaAt(close))} + } else if len(f.token(close).Leading) == 0 && f.token(close).BlankLinesBefore > 0 { + // Empty body with blank lines before the close and no comments: + // closingTriviaAt emits nothing, so preserve the blanks here. + inner = doc.Concat{doc.Concat(f.blankLineDocs(f.token(close).BlankLinesBefore, doc.HardLineNoBreak)), inner} + } + + openDoc := append([]doc.Doc{doc.Text(" {")}, f.openTriviaAt(open)...) body := doc.Concat{ - doc.Text(" {"), - doc.Indent(doc.Concat{doc.HardLineNoBreak, doc.Concat(parts)}), + doc.Concat(openDoc), + doc.Indent(inner), doc.HardLineNoBreak, - doc.Text("}"), + f.emitTokens(close, close, emitOpts{trailing: close != v.TokEnd()}), } out := []doc.Doc{ - doc.Text("service "), - doc.Text(v.Name.Text), + f.emitTokens(v.TokStart(), open, emitOpts{skipText: map[int]bool{open: true}, breakSkip: true}), + body, } - if v.Extends != nil { - out = append(out, doc.Text(" extends "), doc.Text(v.Extends.Text)) + if v.Annotations != nil { + out = append(out, f.breakBeforeAnnotations(close)) } - out = append(out, body) - out = append(out, f.annotationsDoc(v.Annotations)) + + out = append(out, f.annotationsDoc(v.Annotations, v.Annotations != nil && v.Annotations.TokEnd() == v.TokEnd())) + out = append(out, f.afterAnnotations(v.Annotations, v.TokEnd())) + return doc.Concat(out) } @@ -147,144 +234,266 @@ func (f *formatter) service(v *syntax.Service) doc.Doc { // throws clause is a sibling group, not an ancestor. func (f *formatter) function(v *syntax.Function) doc.Doc { parts := append(f.leadingComments(v), f.functionBody(v)) - parts = append(parts, f.trailingComments(v)...) + parts = append(parts, f.trailingComments(v, true)...) + return doc.Concat(parts) } func (f *formatter) functionBody(v *syntax.Function) doc.Doc { + // The header (up to the args open paren) renders as a token run, so + // comments between the header tokens are preserved by construction. + // The open paren's text is emitted by the args group; its trailing + // trivia belongs to openTrivia. + open := f.scanKind(v.TokStart(), v.TokEnd(), syntax.TokenLParen) + header := f.emitTokens(v.TokStart(), open, emitOpts{skipText: map[int]bool{open: true}, breakSkip: true}) + // Comments or blank lines in the arguments force the multiline layout: // the flat argument group would drop them. - if f.fieldsForcedBroken(v.Args) || hasTrailingDelim(lastSep(v.Args)) { - return f.functionBrokenArgs(v) + argsMode := f.opts.Separator.Get(ConstructArguments) + if f.fieldsForcedBroken(v.Args) || sepForcesBreak(sepsOfFields(v.Args), argsMode) || f.opts.Break.Get(ConstructArguments) { + return f.functionBrokenArgs(v, header) } // The argument group folds to "(a, b)" when it fits and unfolds to one - // field per line otherwise. - args := doc.Group(doc.IfBreak( - f.brokenParens("(", ")", v.Args), - doc.Concat{doc.Text("("), f.flatFieldsJoin(v.Args), doc.Text(")")}, - )) + // field per line otherwise, like the throws clause. + args := f.parenGroup(v.Args, open, parenClose(v.Args, open), false, argsMode) if v.Throws == nil { return doc.Group(doc.Concat{ - doc.Text(f.functionHeader(v)), + header, args, - f.annotationsDoc(v.Annotations), + f.functionTail(v, open), }) } - // The throws clause unfolds when the outer group breaks. A trailing - // delimiter, comment, or blank line inside it forces the unfold with a - // break parent; being a sibling of the args group, it cannot break the - // arguments themselves. - throws := doc.Concat{doc.Text(" throws "), f.brokenParens("(", ")", v.Throws.Fields)} - if f.fieldsForcedBroken(v.Throws.Fields) || hasTrailingDelim(lastSep(v.Throws.Fields)) { - throws = doc.Concat{throws, doc.BreakParent} - } - return doc.Group(doc.Concat{ - doc.Text(f.functionHeader(v)), + header, args, - doc.IfBreak( - throws, - doc.Concat{doc.Text(" throws ("), f.flatFieldsJoin(v.Throws.Fields), doc.Text(")")}, - ), - f.annotationsDoc(v.Annotations), + f.throwsGroup(v), + f.functionTail(v, open), }) } -// lastSep returns the separator of the last list item, or 0 for an empty -// list. -func lastSep(fields []*syntax.Field) syntax.TokenKind { +// parenGroup renders "(fields)" as its own group, folding independently: +// flat when it fits, one field per line otherwise. open and close are the +// paren token indices, whose trivia is preserved. forced requires the +// broken layout (comments, blank lines, or a trailing delimiter inside). +func (f *formatter) parenGroup(fields []*syntax.Field, open, close int, forced bool, sepMode SeparatorMode) doc.Doc { + broken := f.brokenParens(fields, open, close, sepMode) + if forced { + broken = doc.Concat{broken, doc.BreakParent} + } + + return doc.Group(doc.IfBreak( + broken, + doc.Concat{doc.Text("("), f.flatFieldsJoin(fields, sepMode), doc.Text(")")}, + )) +} + +// throwsGroup renders the throws clause with the same folding as the +// arguments, so it stays flat when it fits even if the arguments broke. +func (f *formatter) throwsGroup(v *syntax.Function) doc.Doc { + forced := f.fieldsForcedBroken(v.Throws.Fields) || sepForcesBreak(sepsOfFields(v.Throws.Fields), f.opts.Separator.Get(ConstructThrows)) || f.opts.Break.Get(ConstructThrows) + + return doc.Concat{doc.Text(" throws "), f.parenGroup(v.Throws.Fields, v.Throws.TokStart(), v.Throws.TokEnd(), forced, f.opts.Separator.Get(ConstructThrows))} +} + +// parenClose returns the close paren index matching the open paren at +// open, given the field list it encloses. +func parenClose(fields []*syntax.Field, open int) int { if len(fields) == 0 { - return 0 + return open + 1 } - return fields[len(fields)-1].Sep + + return fields[len(fields)-1].TokEnd() + 1 } -// hasTrailingDelim reports whether the separator of the last list item is -// present, i.e. the source carries a trailing delimiter. Like zig fmt, a -// trailing delimiter forces the list to unfold when formatting; without -// one, the list may fold when it fits. -func hasTrailingDelim(sep syntax.TokenKind) bool { - return sep != 0 +// sepForcesBreak reports whether the source's separators force the broken +// layout: a trailing delimiter on the last item under a mode that emits +// it, or a mixed separator pattern under preserve — a flat line whose +// separators are inconsistently present looks broken. +func sepForcesBreak(seps []syntax.TokenKind, mode SeparatorMode) bool { + if len(seps) > 0 && mode != SeparatorNone && seps[len(seps)-1] != 0 { + return true + } + + if mode != SeparatorPreserve || len(seps) < 2 { + return false + } + + want := seps[0] != 0 + for _, sep := range seps[1 : len(seps)-1] { + if (sep != 0) != want { + return true + } + } + + return false } -// functionHeader renders "[oneway] ". -func (f *formatter) functionHeader(v *syntax.Function) string { - out := "" - if v.Oneway != nil { - out += "oneway " +func sepsOfFields(fields []*syntax.Field) []syntax.TokenKind { + seps := make([]syntax.TokenKind, len(fields)) + for i, f := range fields { + seps[i] = f.Sep } - if v.Void != nil { - out += "void " - } else { - out += typeText(v.Type) + " " + + return seps +} + +func sepsOfValues(values []*syntax.EnumValue) []syntax.TokenKind { + seps := make([]syntax.TokenKind, len(values)) + for i, v := range values { + seps[i] = v.Sep } - return out + v.Name.Text + + return seps } // flatFieldsJoin joins fields with their separators on one line: each // field's own separator when preserving, or a single forced separator per -// the FunctionSeparator mode. -func (f *formatter) flatFieldsJoin(fields []*syntax.Field) doc.Doc { +// the sepMode. +func (f *formatter) flatFieldsJoin(fields []*syntax.Field, sepMode SeparatorMode) doc.Doc { parts := make([]doc.Doc, 0, len(fields)) for i, field := range fields { if i > 0 { parts = append(parts, doc.Concat{ - doc.Text(sepText(fields[i-1].Sep, f.opts.FunctionSeparator)), + doc.Text(sepText(fields[i-1].Sep, sepMode)), doc.Line, }) } - parts = append(parts, f.fieldContent(field, nil, false)) + + parts = append(parts, f.fieldContent(field, nil, false, sepMode)) } + return doc.Concat(parts) } // functionBrokenArgs renders the signature with arguments and throws both // broken, one per line. -func (f *formatter) functionBrokenArgs(v *syntax.Function) doc.Doc { +func (f *formatter) functionBrokenArgs(v *syntax.Function, header doc.Doc) doc.Doc { + open := f.scanKind(v.TokStart(), v.TokEnd(), syntax.TokenLParen) + parts := []doc.Doc{ - doc.Text(f.functionHeader(v)), - f.brokenParens("(", ")", v.Args), + header, + f.parenGroup(v.Args, open, parenClose(v.Args, open), true, f.opts.Separator.Get(ConstructArguments)), } if v.Throws != nil { - parts = append(parts, doc.Text(" throws "), f.brokenParens("(", ")", v.Throws.Fields)) + parts = append(parts, f.throwsGroup(v)) + } + + parts = append(parts, f.functionTail(v, open)) + + return doc.Concat(parts) +} + +// functionTail renders the tokens of the function after its args and +// throws clauses: the trailing trivia of the close parens, the +// annotations, and any stray tokens lenient sources leave — everything the +// structural layout does not emit itself. open is the args open paren. +func (f *formatter) functionTail(v *syntax.Function, open int) doc.Doc { + var parts []doc.Doc + + argsClose := parenClose(v.Args, open) + if argsClose < v.TokEnd() { + parts = append(parts, f.tailAfter(argsClose)...) + } + + idx := argsClose + 1 + + if v.Throws != nil { + throwsClose := v.Throws.TokEnd() + if throwsClose < v.TokEnd() { + parts = append(parts, f.tailAfter(throwsClose)...) + } + + idx = throwsClose + 1 + } + + if v.Annotations != nil { + if idx < v.Annotations.TokStart() { + parts = append(parts, f.emitTokens(idx, v.Annotations.TokStart()-1, emitOpts{leading: true})) + } + + parts = append(parts, f.annotationsDoc(v.Annotations, v.Annotations.TokEnd() == v.TokEnd())) + idx = v.Annotations.TokEnd() + 1 + } + + if idx <= v.TokEnd() { + // A line comment on the previous token (e.g. the annotations' + // close paren) would swallow the stray tokens. + if f.lineAfter(idx-1) || len(f.token(idx).Leading) > 0 { + parts = append(parts, doc.HardLine) + } + + parts = append(parts, f.emitTokens(idx, v.TokEnd(), emitOpts{leading: true})) } - parts = append(parts, f.annotationsDoc(v.Annotations)) + return doc.Concat(parts) } +// tailAfter renders the trailing trivia of a close paren followed by more +// tokens: the comments inline, with a hard line after line comments so +// nothing is swallowed. +func (f *formatter) tailAfter(idx int) []doc.Doc { + parts := []doc.Doc{} + for _, c := range f.token(idx).Trailing { + parts = append(parts, doc.Text(" "+trimComment(c.Text))) + } + + if f.lineAfter(idx) || len(f.token(idx+1).Leading) > 0 { + parts = append(parts, doc.HardLine) + } + + return parts +} + // brokenFields renders fields one per line, each with its trailing -// separator per the FieldSeparator option. Comments and blank lines inside +// separator per the sepMode option. Comments and blank lines inside // the list are preserved. -func (f *formatter) brokenFields(fields []*syntax.Field) doc.Doc { +func (f *formatter) brokenFields(fields []*syntax.Field, sepMode SeparatorMode) doc.Doc { var parts []doc.Doc + for i, field := range fields { if i > 0 { parts = append(parts, doc.HardLineNoBreak) parts = append(parts, f.blankLines(field, doc.HardLineNoBreak)...) + } else { + parts = append(parts, f.blankLines(field, doc.HardLineNoBreak)...) } + content := doc.Concat{ - f.fieldContent(field, f.alignmentFor(fields, i), true), - trailingSep(field.Sep, f.opts.FunctionSeparator), + f.fieldContent(field, f.alignmentFor(fields, i, sepMode), true, sepMode), + trailingSep(field.Sep, sepMode), } fieldDoc := append(f.leadingComments(field), content) - fieldDoc = append(fieldDoc, f.trailingComments(field)...) + fieldDoc = append(fieldDoc, f.trailingComments(field, sepEmits(field.Sep, sepMode))...) parts = append(parts, doc.Concat(fieldDoc)) } + return doc.Concat(parts) } // brokenParens renders "open, fields, close" one field per line, or just -// "openclose" when there are no fields. -func (f *formatter) brokenParens(open, close string, fields []*syntax.Field) doc.Doc { +// "openclose" when there are no fields. open and close are the paren token +// indices; their trivia is preserved even with no fields. +func (f *formatter) brokenParens(fields []*syntax.Field, open, close int, sepMode SeparatorMode) doc.Doc { if len(fields) == 0 { - return doc.Text(open + close) + closeDoc := f.emitTokens(close, close, emitOpts{leading: true}) + if f.lineAfter(open) || len(f.token(close).Leading) > 0 { + closeDoc = doc.Concat{doc.HardLine, closeDoc} + } + + return doc.Concat{ + doc.Text("("), + doc.Concat(f.openTriviaAt(open)), + closeDoc, + } } - inner := append([]doc.Doc{doc.HardLineNoBreak, f.brokenFields(fields)}, f.closingTrivia(fields[len(fields)-1])...) - parts := append([]doc.Doc{doc.Text(open)}, f.openTrivia(fields[0])...) - parts = append(parts, doc.Indent(doc.Concat(inner)), doc.HardLineNoBreak, doc.Text(close)) + + inner := append([]doc.Doc{doc.HardLineNoBreak, f.brokenFields(fields, sepMode)}, f.closingTriviaAt(close)...) + parts := append([]doc.Doc{doc.Text("(")}, f.openTriviaAt(open)...) + parts = append(parts, doc.Indent(doc.Concat(inner)), doc.HardLineNoBreak, doc.Text(")")) + return doc.Concat(parts) } @@ -299,18 +508,22 @@ func (f *formatter) fieldsForcedBroken(fields []*syntax.Field) bool { if len(f.token(fields[0].TokStart()-1).Trailing) > 0 { return true } + close := f.token(fields[len(fields)-1].TokEnd() + 1) if len(close.Leading) > 0 { return true } + for _, field := range fields { if f.blankBefore(field) >= 1 { return true } + if len(f.token(field.TokStart()).Leading) > 0 || len(f.token(field.TokEnd()).Trailing) > 0 { return true } } + return false } @@ -322,6 +535,7 @@ func (f *formatter) blankLines(n syntax.Node, line doc.Doc) []doc.Doc { if len(f.token(n.TokStart()).Leading) > 0 { return nil } + return f.blankLineDocs(f.blankBefore(n), line) } @@ -329,8 +543,9 @@ func (f *formatter) blankLines(n syntax.Node, line doc.Doc) []doc.Doc { // preserved exactly. func (f *formatter) blankLineDocs(count int, line doc.Doc) []doc.Doc { parts := make([]doc.Doc, 0, count) - for i := 0; i < count; i++ { + for range count { parts = append(parts, line) } + return parts } diff --git a/formatter/field.go b/formatter/field.go index 622b003..79ddf53 100644 --- a/formatter/field.go +++ b/formatter/field.go @@ -10,18 +10,23 @@ import ( // fieldList formats struct-like body fields. Fields are joined with their // separator when the body stays on one line and with newlines otherwise; in // break mode each field's own trailing separator is emitted, driven by the -// FieldSeparator option. Blank lines between fields are preserved and force +// sepMode option. Blank lines between fields are preserved and force // the body to break. Column alignment applies per blank-line group, // matching the previous formatter. -func (f *formatter) fieldList(fields []*syntax.Field, bodyID int) doc.Doc { +func (f *formatter) fieldList(fields []*syntax.Field, bodyID int, sepMode SeparatorMode) doc.Doc { var parts []doc.Doc + for i, field := range fields { if i > 0 { - parts = append(parts, fieldSep(fields[i-1].Sep, f.opts.FieldSeparator)) + parts = append(parts, fieldSep(fields[i-1].Sep, sepMode)) + parts = append(parts, f.blankLines(field, doc.HardLine)...) + } else { parts = append(parts, f.blankLines(field, doc.HardLine)...) } - parts = append(parts, f.field(field, f.alignmentFor(fields, i), bodyID)) + + parts = append(parts, f.field(field, f.alignmentFor(fields, i, sepMode), bodyID, sepMode)) } + return doc.Concat(parts) } @@ -45,112 +50,197 @@ func sepText(prevSep syntax.TokenKind, mode SeparatorMode) string { case SeparatorNone: return "" } + switch prevSep { case syntax.TokenComma: return "," case syntax.TokenSemicolon: return ";" } + return "" } // enumValueList formats enum bodies with the same layout as fieldList. func (f *formatter) enumValueList(values []*syntax.EnumValue, bodyID int) doc.Doc { var parts []doc.Doc + for i, value := range values { if i > 0 { - parts = append(parts, fieldSep(values[i-1].Sep, f.opts.FieldSeparator)) + parts = append(parts, fieldSep(values[i-1].Sep, f.opts.Separator.Get(ConstructEnum))) + parts = append(parts, f.blankLines(value, doc.HardLine)...) + } else { parts = append(parts, f.blankLines(value, doc.HardLine)...) } - parts = append(parts, f.enumValue(value, f.alignmentForEnum(values, i), bodyID)) + + parts = append(parts, f.enumValue(value, f.alignmentForEnum(values, i, f.opts.Separator.Get(ConstructEnum)), bodyID)) } + return doc.Concat(parts) } // groupedWith reports whether the node joins the alignment group of the -// item before it: no blank line and no comment sits between them. A comment -// in the gap is a visual break, like whitespace. -func (f *formatter) groupedWith(n syntax.Node) bool { - return f.blankBefore(n) < 1 && len(f.token(n.TokStart()).Leading) == 0 +// item before it: no blank line and no comment sits between them. A +// comment in the gap is a visual break, like whitespace. A comment +// trailing the previous item's separator only breaks the group when the +// separator is suppressed: the comment then moves to its own line and +// re-attaches, so the group must not depend on it. +func (f *formatter) groupedWith(prev, cur syntax.Node, sepMode SeparatorMode) bool { + // A comment trailing the previous item's separator only breaks the + // group when the separator is suppressed and the comment moves to its + // own line (re-attaching): the group must not depend on it. Comments + // that stay on the separator's line do not break the group. + prevTok := f.token(prev.TokEnd()) + sepBreak := isListSep(prevTok.Kind) && !sepEmits(prevTok.Kind, sepMode) && + (leadingLineComment(prevTok) || len(prevTok.Trailing) > 0 && f.lineAfter(prev.TokEnd()-1)) + + return f.blankBefore(cur) < 1 && + len(f.token(cur.TokStart()).Leading) == 0 && + !sepBreak } // alignmentFor returns the column alignment for field i, or nil when // alignment is disabled. Alignment is scoped to the blank-line and comment -// group the field belongs to. -func (f *formatter) alignmentFor(fields []*syntax.Field, i int) *columnAlign { +// group the field belongs to; sepMode is the separator mode the enclosing +// list uses, for the width gate. +func (f *formatter) alignmentFor(fields []*syntax.Field, i int, sepMode SeparatorMode) *columnAlign { if f.opts.Align == AlignDisable { return nil } + start := i - for start > 0 && f.groupedWith(fields[start]) { + for start > 0 && f.groupedWith(fields[start-1], fields[start], sepMode) { start-- } + end := i - for end+1 < len(fields) && f.groupedWith(fields[end+1]) { + for end+1 < len(fields) && f.groupedWith(fields[end], fields[end+1], sepMode) { end++ } + group := fields[start : end+1] + a := computeFieldAlign(group) - if !f.alignmentFits(group, a) && !f.sourceAligned(fields, start, end) { + if !f.alignmentFits(group, a, sepMode) && !f.sourceAligned(fields, start, end) { return nil } + return a } // alignmentFits reports whether column alignment keeps the structural // content of every line within printWidth: the padded columns plus, per -// field, its name and trailing separator. Trailing comments are excluded — -// a comment may overflow its line without affecting alignment. Default -// values and annotations are ignored; a line that long overflows regardless -// of alignment. -func (f *formatter) alignmentFits(fields []*syntax.Field, a *columnAlign) bool { +// field, its name and the trailing separator the given mode emits. +// Trailing comments are excluded — a comment may overflow its line without +// affecting alignment. Default values and annotations are ignored; a line +// that long overflows regardless of alignment. +func (f *formatter) alignmentFits(fields []*syntax.Field, a *columnAlign, sepMode SeparatorMode) bool { limit := f.opts.PrintWidth - f.opts.TabWidth + columns := a.idWidth + 1 if a.hasReq { columns += a.reqWidth + 1 } + if f.opts.Align == AlignField { columns += a.typeWidth + 1 } + longest := 0 + for _, field := range fields { - w := len(field.Name.Text) - if field.Sep != 0 { - w++ // trailing separator - } + w := len(field.Name.Text) + sepLen(field.Sep, sepMode) longest = maxInt(longest, w) } + return columns+longest <= limit } +// sepLen is the length of the trailing separator the mode emits after a +// field with the given source separator. +func sepLen(sep syntax.TokenKind, mode SeparatorMode) int { + switch mode { + case SeparatorComma, SeparatorSemicolon: + return 1 + case SeparatorNone: + return 0 + } + + if sep != 0 { + return 1 + } + + return 0 +} + // sourceAligned reports whether the group's field names start at the same -// column in the source. Such a layout is deliberate — preserving it matters -// more than printWidth — so alignment survives the width gate for it. +// column in the source, with at least one field padded beyond a single +// space after its type — a deliberately aligned layout, preserved even +// when it overflows printWidth. Accidental alignment (natural columns +// coinciding) does not count. func (f *formatter) sourceAligned(fields []*syntax.Field, start, end int) bool { col := 0 + padded := false + for i := start; i <= end; i++ { - c := f.token(fields[i].Name.TokStart()).Col + field := fields[i] + + c := f.token(field.Name.TokStart()).Col if col == 0 { col = c } else if c != col { return false } + + if f.namePadded(field) { + padded = true + } } - return true + + return padded } -func (f *formatter) alignmentForEnum(values []*syntax.EnumValue, i int) *columnAlign { +// namePadded reports whether more than one space separates the field name +// from its type in the source; a reference field (&) occupies one extra +// column. +func (f *formatter) namePadded(field *syntax.Field) bool { + prev := field.Name.TokStart() - 1 + if prev < 0 { + return false + } + + extra := 0 + if f.token(prev).Kind == syntax.TokenAmp { + extra = 1 // the & reference + + prev-- + if prev < 0 { + return false + } + } + + name := f.token(field.Name.TokStart()) + tok := f.token(prev) + // A single canonical space puts the name two columns past the type's + // end (the space at end+1); a reference adds one for the '&'. + return name.Col-(tok.Col+len(tok.Text)-1) > 2+extra +} + +func (f *formatter) alignmentForEnum(values []*syntax.EnumValue, i int, sepMode SeparatorMode) *columnAlign { if f.opts.Align == AlignDisable { return nil } + start := i - for start > 0 && f.groupedWith(values[start]) { + for start > 0 && f.groupedWith(values[start-1], values[start], sepMode) { start-- } + end := i - for end+1 < len(values) && f.groupedWith(values[end+1]) { + for end+1 < len(values) && f.groupedWith(values[end], values[end+1], sepMode) { end++ } + return computeEnumAlign(values[start : end+1]) } @@ -168,140 +258,210 @@ func maxInt(a, b int) int { if a > b { return a } + return b } // computeFieldAlign computes column widths for a group of fields. func computeFieldAlign(fields []*syntax.Field) *columnAlign { a := &columnAlign{} + for _, field := range fields { if field.FieldID != nil { a.idWidth = maxInt(a.idWidth, len(field.FieldID.Text)+1) // "N:" } + if field.Req != 0 { a.hasReq = true a.reqWidth = maxInt(a.reqWidth, len(field.Req.String())) } + a.typeWidth = maxInt(a.typeWidth, len(typeText(field.Type))) if field.Value != nil { a.nameWidth = maxInt(a.nameWidth, len(field.Name.Text)) } } + return a } // computeEnumAlign computes the name width for aligning '=' signs. func computeEnumAlign(values []*syntax.EnumValue) *columnAlign { a := &columnAlign{enumAssign: true} + for _, value := range values { if value.Value != nil { a.nameWidth = maxInt(a.nameWidth, len(value.Name.Text)) } } + return a } // field assembles a struct-like body field: leading comments, content, and // trailing comments. When the body breaks (referenced by bodyID), the // content switches to its column-aligned form with the trailing separator. -func (f *formatter) field(v *syntax.Field, align *columnAlign, bodyID int) doc.Doc { - content := f.fieldContent(v, align, false) +func (f *formatter) field(v *syntax.Field, align *columnAlign, bodyID int, sepMode SeparatorMode) doc.Doc { + content := f.fieldContent(v, align, false, sepMode) if bodyID != 0 { content = doc.IfBreakFor( - doc.Concat{f.fieldContent(v, align, true), trailingSep(v.Sep, f.opts.FieldSeparator)}, + doc.Concat{f.fieldContent(v, align, true, sepMode), trailingSep(v.Sep, sepMode)}, content, bodyID, ) } + parts := append(f.leadingComments(v), content) - parts = append(parts, f.trailingComments(v)...) + parts = append(parts, f.trailingComments(v, sepEmits(v.Sep, sepMode))...) + return doc.Concat(parts) } -// fieldContent renders the field line: id, requiredness, type, reference, -// name, default value, and annotations. padded selects the column-aligned -// form used when the enclosing body breaks; it has no effect when align is -// nil. -func (f *formatter) fieldContent(v *syntax.Field, align *columnAlign, padded bool) doc.Doc { +// emitWithAnnotations renders a token run split at the node's annotations, +// so they keep their foldable group instead of being inlined. +func (f *formatter) emitWithAnnotations(start, end int, ann *syntax.Annotations, o emitOpts) doc.Doc { + if ann == nil { + return f.emitTokens(start, end, o) + } + // The segment before the annotations owns its last token's trailing + // trivia (the annotations' close owns the node's trailing instead). + first := o + first.trailing = true + + parts := []doc.Doc{f.emitTokens(start, ann.TokStart()-1, first)} + if f.lineAfter(ann.TokStart()-1) || len(f.token(ann.TokStart()).Leading) > 0 { + parts = append(parts, doc.HardLine) + } + + parts = append(parts, f.annotationsDoc(ann, ann.TokEnd() == end)) + if ann.TokEnd() < end { + parts = append(parts, f.emitTokens(ann.TokEnd()+1, end, o)) + } + + return doc.Concat(parts) +} + +// fieldContent renders the field as a token run. padded selects the +// column-aligned form used when the enclosing body breaks; it has no +// effect when align is nil. The separator token's text is suppressed (the +// caller emits it), but its trivia is preserved. +func (f *formatter) fieldContent(v *syntax.Field, align *columnAlign, padded bool, sepMode SeparatorMode) doc.Doc { padded = padded && align != nil - var parts []doc.Doc - if v.FieldID != nil { - id := v.FieldID.Text + ":" - if padded { - id = padRight(id, align.idWidth) - } - parts = append(parts, doc.Text(id), doc.Text(" ")) + o := emitOpts{breakTrailing: true} + if v.Sep != 0 { + o.skipText = map[int]bool{v.TokEnd(): true} + o.breakSkip = sepEmits(v.Sep, sepMode) + } + + if padded { + o.pads, o.prefix = f.fieldPads(v, align) } - if v.Req != 0 || padded && align.hasReq { - req := "" - if v.Req != 0 { - req = v.Req.String() + return f.emitWithAnnotations(v.TokStart(), v.TokEnd(), v.Annotations, o) +} + +// fieldPads returns the alignment padding after each column token, or nil +// when the field carries comments that make the padded widths unknowable. +// The prefix is the leading padding of a field without an id whose +// requiredness column is empty. +func (f *formatter) fieldPads(v *syntax.Field, a *columnAlign) (map[int]string, string) { + start, end := v.TokStart(), v.TokEnd() + if v.Sep != 0 { + end-- + } + + for i := start + 1; i <= end; i++ { + if len(f.token(i).Leading) > 0 { + return nil, "" } - if padded { - req = padRight(req, align.reqWidth) + } + + for i := start; i < end; i++ { + if len(f.token(i).Trailing) > 0 { + return nil, "" } - parts = append(parts, doc.Text(req), doc.Text(" ")) } - typ := doc.Text(typeText(v.Type)) - if padded && f.opts.Align == AlignField { - typ = doc.Text(padRight(typeText(v.Type), align.typeWidth)) + pads := map[int]string{} + if v.FieldID != nil { + pads[v.TokStart()+1] = padRight("", a.idWidth-len(v.FieldID.Text)-1) } - parts = append(parts, typ, f.annotationsDoc(v.Type.Annotations), doc.Text(" ")) - if v.Reference { - parts = append(parts, doc.Text("&")) + if v.Req != 0 { + pads[v.Type.TokStart()-1] = padRight("", a.reqWidth-len(v.Req.String())) + } else if a.hasReq { + // The empty requiredness column: extend the id pad by one column + // plus the missing req width, or lead the field with it when there + // is no id token to pad after. + if v.FieldID != nil { + pads[v.TokStart()+1] += padRight("", a.reqWidth+1) + } else { + return nil, padRight("", a.reqWidth+1) + } } - name := doc.Text(v.Name.Text) - if padded && v.Value != nil && align.nameWidth > 0 { - // Fields with values have their name padded so the "=" signs - // align, in both align modes. - name = doc.Text(padRight(v.Name.Text, align.nameWidth)) + if f.opts.Align == AlignField { + end := v.Type.TokEnd() + if v.Type.Annotations != nil { + end = v.Type.Annotations.TokStart() - 1 + } + + pads[end] = padRight("", a.typeWidth-len(typeText(v.Type))) } - parts = append(parts, name) - if v.Value != nil { - parts = append(parts, doc.Text(" = "), f.constValue(v.Value)) + if v.Value != nil && a.nameWidth > 0 { + pads[v.Name.TokStart()] = padRight("", a.nameWidth-len(v.Name.Text)) } - parts = append(parts, f.annotationsDoc(v.Annotations)) - return doc.Concat(parts) + return pads, "" } // enumValue assembles an enum value with comments, aligning '=' signs when // the body breaks. func (f *formatter) enumValue(v *syntax.EnumValue, align *columnAlign, bodyID int) doc.Doc { - content := f.enumValueContent(v, align, false) + content := f.enumValueContent(v, align, false, f.opts.Separator.Get(ConstructEnum)) if bodyID != 0 { content = doc.IfBreakFor( - doc.Concat{f.enumValueContent(v, align, true), trailingSep(v.Sep, f.opts.FieldSeparator)}, + doc.Concat{f.enumValueContent(v, align, true, f.opts.Separator.Get(ConstructEnum)), trailingSep(v.Sep, f.opts.Separator.Get(ConstructEnum))}, content, bodyID, ) } + parts := append(f.leadingComments(v), content) - parts = append(parts, f.trailingComments(v)...) + parts = append(parts, f.trailingComments(v, sepEmits(v.Sep, f.opts.Separator.Get(ConstructEnum)))...) + return doc.Concat(parts) } -func (f *formatter) enumValueContent(v *syntax.EnumValue, align *columnAlign, padded bool) doc.Doc { +func (f *formatter) enumValueContent(v *syntax.EnumValue, align *columnAlign, padded bool, sepMode SeparatorMode) doc.Doc { padded = padded && align != nil - name := doc.Text(v.Name.Text) - if padded && align.enumAssign { - name = doc.Text(padRight(v.Name.Text, align.nameWidth)) + o := emitOpts{breakTrailing: true} + if v.Sep != 0 { + o.skipText = map[int]bool{v.TokEnd(): true} + o.breakSkip = sepEmits(v.Sep, sepMode) } - parts := []doc.Doc{name} - if v.Value != nil { - parts = append(parts, doc.Text(" = "), doc.Text(v.Value.Text)) + if padded && align.enumAssign && align.nameWidth > 0 { + o.pads = map[int]string{v.Name.TokStart(): padRight("", align.nameWidth-len(v.Name.Text))} } - parts = append(parts, f.annotationsDoc(v.Annotations)) - return doc.Concat(parts) + return f.emitWithAnnotations(v.TokStart(), v.TokEnd(), v.Annotations, o) +} + +// sepEmits reports whether the mode emits a non-empty trailing separator +// after a field with the given source separator. +func sepEmits(sep syntax.TokenKind, mode SeparatorMode) bool { + switch mode { + case SeparatorComma, SeparatorSemicolon: + return true + case SeparatorNone: + return false + } + + return sep != 0 } // trailingSep returns the trailing separator for the given original @@ -322,6 +482,7 @@ func trailingSep(sep syntax.TokenKind, mode SeparatorMode) doc.Doc { case syntax.TokenSemicolon: return doc.Text(";") } + return doc.Concat{} } } @@ -331,5 +492,6 @@ func padRight(s string, w int) string { if n := w - len(s); n > 0 { return s + strings.Repeat(" ", n) } + return s } diff --git a/formatter/format.go b/formatter/format.go index ddfdb93..b7246c1 100644 --- a/formatter/format.go +++ b/formatter/format.go @@ -11,6 +11,7 @@ package formatter import ( "errors" "fmt" + "strings" "github.com/karitham/thrift-ls/doc" "github.com/karitham/thrift-ls/syntax" @@ -45,6 +46,88 @@ const ( SeparatorNone ) +// Construct identifies a construct with per-construct options. +type Construct uint8 + +const ( + ConstructStruct Construct = iota + ConstructUnion + ConstructException + ConstructEnum + ConstructArguments + ConstructThrows +) + +// PerConstruct holds one option value per construct. +type PerConstruct[T any] struct { + Structs T + Unions T + Exceptions T + Enums T + Arguments T + Throws T +} + +// Get returns the value for the construct. +func (p PerConstruct[T]) Get(c Construct) T { + switch c { + case ConstructUnion: + return p.Unions + case ConstructException: + return p.Exceptions + case ConstructEnum: + return p.Enums + case ConstructArguments: + return p.Arguments + case ConstructThrows: + return p.Throws + } + + return p.Structs +} + +// Set assigns the value for the construct. +func (p *PerConstruct[T]) Set(c Construct, v T) { + switch c { + case ConstructUnion: + p.Unions = v + case ConstructException: + p.Exceptions = v + case ConstructEnum: + p.Enums = v + case ConstructArguments: + p.Arguments = v + case ConstructThrows: + p.Throws = v + default: + p.Structs = v + } +} + +// AllConstructs lists every construct, in config order. +var AllConstructs = []Construct{ + ConstructStruct, ConstructUnion, ConstructException, + ConstructEnum, ConstructArguments, ConstructThrows, +} + +// String returns the config key of the construct. +func (c Construct) String() string { + switch c { + case ConstructUnion: + return "unions" + case ConstructException: + return "exceptions" + case ConstructEnum: + return "enums" + case ConstructArguments: + return "arguments" + case ConstructThrows: + return "throws" + } + + return "structs" +} + // Options controls formatting behavior. Zero values mean defaults. type Options struct { // PrintWidth is the target line width. Must be positive. @@ -56,19 +139,12 @@ type Options struct { TabWidth int // Align controls column alignment (default AlignField). Align AlignMode - // FieldSeparator controls trailing separators after - // struct/union/exception fields and enum values (default + // Separator controls trailing separators per construct (default // SeparatorPreserve). - FieldSeparator SeparatorMode - // FunctionSeparator controls trailing separators after service - // arguments and throws entries (default SeparatorPreserve). - FunctionSeparator SeparatorMode - // BreakStructs forces struct, union, and exception bodies to the - // multiline layout, even when they fit on one line. - BreakStructs bool - // BreakEnums forces enum bodies to the multiline layout, even when - // they fit on one line. - BreakEnums bool + Separator PerConstruct[SeparatorMode] + // Break forces the multiline layout per construct, even when the body + // fits on one line. + Break PerConstruct[bool] // NoTrailingNewline suppresses the final newline that is otherwise // appended to the formatted output. NoTrailingNewline bool @@ -77,11 +153,18 @@ type Options struct { // DefaultOptions returns the default formatting options. func DefaultOptions() Options { return Options{ - PrintWidth: 80, - Indent: " ", - TabWidth: 4, - Align: AlignField, - FieldSeparator: SeparatorPreserve, + PrintWidth: 80, + Indent: " ", + TabWidth: 4, + Align: AlignField, + Separator: PerConstruct[SeparatorMode]{ + Structs: SeparatorPreserve, + Unions: SeparatorPreserve, + Exceptions: SeparatorPreserve, + Enums: SeparatorPreserve, + Arguments: SeparatorPreserve, + Throws: SeparatorPreserve, + }, } } @@ -91,12 +174,15 @@ func (o Options) normalize() Options { if o.PrintWidth <= 0 { o.PrintWidth = d.PrintWidth } + if o.Indent == "" { o.Indent = d.Indent } + if o.TabWidth <= 0 { o.TabWidth = d.TabWidth } + return o } @@ -107,14 +193,27 @@ func Format(d *syntax.Document, o Options) (string, error) { if d == nil { return "", errors.New("formatter: nil document") } + o = o.normalize() + return PrintIR(BuildIR(d, o), o) +} + +// BuildIR builds the document IR for the given options. The IR can be +// inspected with doc.Dump before printing; the printer mutates groups in +// place, so dump after PrintIR to see the layout decisions. +func BuildIR(d *syntax.Document, o Options) doc.Doc { f := &formatter{ doc: d, toks: d.Tokens, opts: o, } - ir := f.document() + + return f.document() +} + +// PrintIR prints the document IR. +func PrintIR(ir doc.Doc, o Options) (string, error) { return doc.Print(ir, doc.Options{ PrintWidth: o.PrintWidth, Indent: o.Indent, @@ -129,6 +228,7 @@ func FormatNode(d *syntax.Document, n syntax.Node, o Options) (string, error) { if d == nil || n == nil { return "", errors.New("formatter: nil document or node") } + o = o.normalize() f := &formatter{ @@ -137,6 +237,7 @@ func FormatNode(d *syntax.Document, n syntax.Node, o Options) (string, error) { opts: o, } ir := f.node(n) + return doc.Print(ir, doc.Options{ PrintWidth: o.PrintWidth, Indent: o.Indent, @@ -155,13 +256,181 @@ type formatter struct { // id returns a fresh non-zero group id for IfBreak references. func (f *formatter) id() int { f.nextID++ + return f.nextID } +// token returns the i-th token. func (f *formatter) token(i int) syntax.Token { return f.toks[i] } +// emitOpts controls token emission. +type emitOpts struct { + leading bool // emit the first token's leading trivia + trailing bool // emit the last token's trailing trivia + breakTrailing bool // line-comment trailing forces groups to break + skipText map[int]bool // tokens whose text and gap are suppressed + breakSkip bool // hard line before a skipped token whose text + // the caller emits (separators) + pads map[int]string // spaces inserted after a token, before its gap + prefix string // spaces emitted before the first token +} + +// emitTokens renders the tokens in [start, end] with their trivia, joined +// with canonical spacing. The first token's leading and last token's +// trailing trivia belong to the caller's comment helpers unless the +// corresponding flag is set. skipText suppresses separator tokens that the +// structural layout emits itself; pads widen alignment columns. +func (f *formatter) emitTokens(start, end int, o emitOpts) doc.Doc { + var parts []doc.Doc + if o.prefix != "" { + parts = append(parts, doc.Text(o.prefix)) + } + + for i := start; i <= end; i++ { + tok := f.token(i) + + skipped := o.skipText[i] + if i > start { + if skipped { + // Leading trivia always forces a hard line; a trailing + // line comment only when the caller emits the skipped + // token's text after it (separators), which would + // otherwise be swallowed. + if len(tok.Leading) > 0 || o.breakSkip && f.lineAfter(i-1) { + parts = append(parts, doc.HardLine) + } + } else { + parts = append(parts, f.tokenGap(f.token(i-1), tok)) + } + } + + if i > start || o.leading { + for j, c := range tok.Leading { + parts = append(parts, doc.Text(trimComment(c.Text))) + // The last comment's line end comes from the caller's + // structure for suppressed tokens, unless the caller + // emits the token's text after it. + if j < len(tok.Leading)-1 || !skipped || o.breakSkip { + parts = append(parts, doc.HardLine) + } + } + } + + if !skipped { + text := tok.Text + if tok.Kind == syntax.TokenAsync { + text = "oneway" + } + + parts = append(parts, doc.Text(text)) + } + + if !skipped && o.pads != nil { + if pad, ok := o.pads[i]; ok && pad != "" { + parts = append(parts, doc.Text(pad)) + } + } + + if i < end || o.trailing { + for _, c := range tok.Trailing { + parts = append(parts, doc.Text(" "+trimComment(c.Text))) + if o.breakTrailing && (c.Kind == syntax.TriviaLineComment || c.Kind == syntax.TriviaAnnotation) { + // A line comment must end its line: force the + // enclosing groups to break so nothing follows it. + parts = append(parts, doc.BreakParent) + } + } + } + } + + return doc.Concat(parts) +} + +// tokenGap returns the doc between two adjacent tokens: a line break when +// the source separated them with a line comment (which would swallow the +// next token on the same line), a canonical space otherwise. +func (f *formatter) tokenGap(prev, cur syntax.Token) doc.Doc { + for _, c := range prev.Trailing { + if c.Kind == syntax.TriviaLineComment || c.Kind == syntax.TriviaAnnotation { + return doc.HardLine + } + } + + if len(cur.Leading) > 0 { + return doc.HardLine + } + + return doc.Text(rawTokenGap(prev, cur)) +} + +// lineAfter reports whether the token ends its line with a line comment or +// annotation, which forces the next doc onto a new line. +func (f *formatter) lineAfter(i int) bool { + for _, c := range f.token(i).Trailing { + if c.Kind == syntax.TriviaLineComment || c.Kind == syntax.TriviaAnnotation { + return true + } + } + + return false +} + +// foldBreak is the foldable gap after an opening token or separating +// comma: a line in the broken layout, a space (or nothing) flat. A +// trailing line comment forces a hard line. +func (f *formatter) foldBreak(i int, flat string) doc.Doc { + if f.lineAfter(i) { + return doc.HardLine + } + + return doc.IfBreak(doc.Line, doc.Text(flat)) +} + +// commaSep renders a separating comma with its trivia: a hard line when +// the previous item ends its line with a comment (which would swallow the +// comma), then the comma and the foldable gap after it. +func (f *formatter) commaSep(comma int) []doc.Doc { + parts := []doc.Doc{} + if f.lineAfter(comma-1) || len(f.token(comma).Leading) > 0 { + parts = append(parts, doc.HardLine) + } + + parts = append(parts, f.emitTokens(comma, comma, emitOpts{leading: true, trailing: true})) + parts = append(parts, f.foldBreak(comma, " ")) + + return parts +} + +// rawTokenGap returns the canonical text between two adjacent tokens: +// opening and closing punctuation attaches to its neighbor, separators get +// a space after them, everything else is space-separated so tokens cannot +// merge ("const list" must not become "constlist"). +func rawTokenGap(prev, cur syntax.Token) string { + switch cur.Kind { + case syntax.TokenRBrace, syntax.TokenRParen, syntax.TokenRBracket, + syntax.TokenGt, syntax.TokenComma, syntax.TokenSemicolon, syntax.TokenColon: + return "" + } + + switch prev.Kind { + case syntax.TokenLBrace, syntax.TokenLParen, syntax.TokenLBracket, + syntax.TokenLt, syntax.TokenAmp: + return "" + } + + switch { + case cur.Kind == syntax.TokenLt && (prev.Kind == syntax.TokenMap || prev.Kind == syntax.TokenList || prev.Kind == syntax.TokenSet): + return "" // map<, list<, set< + case prev.Kind == syntax.TokenComma, prev.Kind == syntax.TokenSemicolon, + prev.Kind == syntax.TokenColon, prev.Kind == syntax.TokenEqual: + return " " + } + + return " " +} + // blankBefore returns the number of blank lines before the node's first // token. func (f *formatter) blankBefore(n syntax.Node) int { @@ -178,33 +447,72 @@ func (f *formatter) leadingComments(n syntax.Node) []doc.Doc { if len(tok.Leading) == 0 { return nil } + var parts []doc.Doc + prevBlank := 0 for _, c := range tok.Leading { parts = append(parts, f.blankLineDocs(c.BlankLinesBefore-prevBlank, doc.HardLine)...) prevBlank = c.BlankLinesBefore - parts = append(parts, doc.Text(c.Text), doc.HardLine) + parts = append(parts, doc.Text(trimComment(c.Text)), doc.HardLine) } + parts = append(parts, f.blankLineDocs(tok.BlankLinesBefore-prevBlank, doc.HardLine)...) + return parts } // trailingComments returns the comments attached after the node's last token -// on the same line, as line-suffix docs. The break parent forces the -// enclosing group to break so the comment stays on the node's own line. -func (f *formatter) trailingComments(n syntax.Node) []doc.Doc { +// on the same line. Block comments render as line-suffix docs; line comments +// and annotations render inline with a break parent, since a line suffix +// after them would merge into the comment's text. When the content before +// the separator already ends with a line comment, these comments cannot +// share the line and get their own lines instead. +func (f *formatter) trailingComments(n syntax.Node, sepEmitted bool) []doc.Doc { var parts []doc.Doc - for _, c := range f.token(n.TokEnd()).Trailing { - parts = append(parts, doc.LineSuffix(doc.Text(" "+c.Text)), doc.BreakParent) + // Comments attached to a separator share the separator's own line, + // unless the separator is not emitted (the mode drops it): then a line + // comment before it would swallow them, and they need their own lines. + last := f.token(n.TokEnd()) + + ownLine := !sepEmitted && (last.Kind == syntax.TokenComma || last.Kind == syntax.TokenSemicolon) && + (f.lineAfter(n.TokEnd()-1) || leadingLineComment(last)) + for _, c := range last.Trailing { + line := c.Kind == syntax.TriviaLineComment || c.Kind == syntax.TriviaAnnotation + if ownLine { + parts = append(parts, doc.HardLine, doc.Text(trimComment(c.Text)), doc.BreakParent) + + continue + } + + if line { + parts = append(parts, doc.Text(" "+trimComment(c.Text)), doc.BreakParent) + } else { + parts = append(parts, doc.LineSuffix(doc.Text(" "+trimComment(c.Text))), doc.BreakParent) + } } + return parts } +// leadingLineComment reports whether the token's leading trivia contains a +// line comment or annotation, which ends the previous line. +func leadingLineComment(tok syntax.Token) bool { + for _, c := range tok.Leading { + if c.Kind == syntax.TriviaLineComment || c.Kind == syntax.TriviaAnnotation { + return true + } + } + + return false +} + // node assembles a top-level node: its leading comments, its formatted // body, and its trailing comments. func (f *formatter) node(n syntax.Node) doc.Doc { parts := append(f.leadingComments(n), f.nodeBody(n)) - parts = append(parts, f.trailingComments(n)...) + parts = append(parts, f.trailingComments(n, true)...) + return doc.Concat(parts) } @@ -236,25 +544,37 @@ func (f *formatter) nodeBody(n syntax.Node) doc.Doc { // lines, blank lines preserved, and trailing comments. func (f *formatter) document() doc.Doc { var parts []doc.Doc + for i, n := range f.doc.Nodes { if i > 0 { parts = append(parts, doc.HardLine) parts = append(parts, f.blankLines(n, doc.HardLine)...) + } else if lead := f.token(n.TokStart()).Leading; len(lead) > 0 && lead[0].BlankLinesBefore > 0 { + // At file start the leading comments carry their blanks without + // a separator line, so N blanks would round-trip as N-1. The + // extra line keeps the count canonical: N blanks are N+1 + // newlines. + parts = append(parts, doc.HardLine) } + parts = append(parts, f.node(n)) } - // Comments at the end of the file attach to the EOF token. + // Comments at the end of the file attach to the EOF token. Like the + // first node's leading comments, a comment run at file start (no nodes) + // needs a separator line before its blanks to round-trip the count. eof := f.toks[len(f.toks)-1] if len(eof.Leading) > 0 { - if len(f.doc.Nodes) > 0 { + if len(f.doc.Nodes) > 0 || eof.Leading[0].BlankLinesBefore > 0 { parts = append(parts, doc.HardLine) } + prevBlank := 0 for i, c := range eof.Leading { parts = append(parts, f.blankLineDocs(c.BlankLinesBefore-prevBlank, doc.HardLine)...) prevBlank = c.BlankLinesBefore - parts = append(parts, doc.Text(c.Text)) + + parts = append(parts, doc.Text(trimComment(c.Text))) if i < len(eof.Leading)-1 { parts = append(parts, doc.HardLine) } @@ -264,73 +584,190 @@ func (f *formatter) document() doc.Doc { if !f.opts.NoTrailingNewline { parts = append(parts, doc.HardLine) } + return doc.Concat(parts) } // --- headers --------------------------------------------------------------- func (f *formatter) include(v *syntax.Include) doc.Doc { - return doc.Concat{doc.Text("include "), doc.Text(v.Path.Text)} + return f.emitTokens(v.TokStart(), v.TokEnd(), emitOpts{}) } func (f *formatter) cppInclude(v *syntax.CPPInclude) doc.Doc { - return doc.Concat{doc.Text("cpp_include "), doc.Text(v.Path.Text)} + return f.emitTokens(v.TokStart(), v.TokEnd(), emitOpts{}) } func (f *formatter) namespace(v *syntax.Namespace) doc.Doc { - parts := []doc.Doc{ - doc.Text("namespace "), - doc.Text(v.Scope.Text), - doc.Text(" "), - doc.Text(v.Name.Text), + end := v.TokEnd() + if v.Annotations != nil { + end = v.Annotations.TokStart() - 1 + } + + o := emitOpts{} + if v.Annotations != nil { + o.trailing = true + } + + parts := []doc.Doc{f.emitTokens(v.TokStart(), end, o)} + if v.Annotations != nil { + parts = append(parts, f.breakBeforeAnnotations(end)) } - parts = append(parts, f.annotationsDoc(v.Annotations)) + + parts = append(parts, f.annotationsDoc(v.Annotations, v.Annotations != nil && v.Annotations.TokEnd() == v.TokEnd())) + parts = append(parts, f.afterAnnotations(v.Annotations, v.TokEnd())) + return doc.Concat(parts) } func (f *formatter) typedef(v *syntax.Typedef) doc.Doc { - parts := []doc.Doc{ - doc.Text("typedef "), - f.fieldType(v.Type), - doc.Text(" "), - doc.Text(v.Name.Text), + end := v.TokEnd() + if v.Annotations != nil { + end = v.Annotations.TokStart() - 1 + } + + o := emitOpts{} + if v.Annotations != nil { + o.trailing = true + } + + parts := []doc.Doc{f.emitTokens(v.TokStart(), end, o)} + if v.Annotations != nil { + parts = append(parts, f.breakBeforeAnnotations(end)) } - parts = append(parts, f.annotationsDoc(v.Annotations)) + + parts = append(parts, f.annotationsDoc(v.Annotations, v.Annotations != nil && v.Annotations.TokEnd() == v.TokEnd())) + parts = append(parts, f.afterAnnotations(v.Annotations, v.TokEnd())) + return doc.Concat(parts) } func (f *formatter) constant(v *syntax.Const) doc.Doc { - return doc.Concat{ - doc.Text("const "), - f.fieldType(v.Type), - doc.Text(" "), - doc.Text(v.Name.Text), - doc.Text(" = "), - f.constValue(v.Value), + value := v.Value + if value == nil { + return f.emitTokens(v.TokStart(), v.TokEnd(), emitOpts{}) } + + eq := value.TokStart() - 1 + gap := f.tokenGap(f.token(eq), f.token(value.TokStart())) + + parts := []doc.Doc{ + f.emitTokens(v.TokStart(), eq, emitOpts{trailing: true}), + gap, + f.constValue(value, value.TokEnd() == v.TokEnd()), + } + if value.TokEnd() < v.TokEnd() { + // Stray tokens after the value (lenient sources): preserve them + // and their trivia, with a line break after the value's close. + stray := f.emitTokens(value.TokEnd()+1, v.TokEnd(), emitOpts{leading: true}) + if f.lineAfter(value.TokEnd()) || len(f.token(value.TokEnd()+1).Leading) > 0 { + stray = doc.Concat{doc.HardLine, stray} + } + + parts = append(parts, stray) + } + + return doc.Concat(parts) } // --- annotations ----------------------------------------------------------- +// breakBeforeAnnotations returns a hard line when the token at idx ends +// its line with a comment, or the annotations' first token carries leading +// trivia, so neither gets swallowed by the other. +func (f *formatter) breakBeforeAnnotations(idx int) doc.Doc { + if f.lineAfter(idx) || len(f.token(idx+1).Leading) > 0 { + return doc.HardLine + } + + return doc.Concat{} +} + +// afterAnnotations renders any tokens between the annotations and the +// node's end — stray separators lenient sources may leave — preserving +// their leading trivia and forcing a line break after the annotations' +// close when it ends its line with a comment. +func (f *formatter) afterAnnotations(a *syntax.Annotations, end int) doc.Doc { + if a == nil || a.TokEnd() >= end { + return doc.Concat{} + } + + parts := []doc.Doc{} + if f.lineAfter(a.TokEnd()) || len(f.token(a.TokEnd()+1).Leading) > 0 { + parts = append(parts, doc.HardLine) + } + + parts = append(parts, f.emitTokens(a.TokEnd()+1, end, emitOpts{leading: true})) + + return doc.Concat(parts) +} + // annotationsDoc returns an annotation group, or an empty doc when absent. -func (f *formatter) annotationsDoc(a *syntax.Annotations) doc.Doc { +// The group folds when it does not fit; the items and their separating +// commas render as token runs, so trivia inside the parens is preserved. +func (f *formatter) annotationsDoc(a *syntax.Annotations, isLast bool) doc.Doc { if a == nil { return doc.Concat{} } - items := make([]doc.Doc, 0, len(a.Items)) - for _, item := range a.Items { - var parts []doc.Doc - parts = append(parts, doc.Text(item.Name.Text)) - if item.Value != nil { - parts = append(parts, doc.Text(" = "), doc.Text(item.Value.Text)) + + open, close := a.TokStart(), a.TokEnd() + if len(a.Items) == 0 { + closeDoc := f.emitTokens(close, close, emitOpts{leading: true, trailing: !isLast}) + if f.lineAfter(open) || len(f.token(close).Leading) > 0 { + closeDoc = doc.Concat{doc.HardLine, closeDoc} + } + + return doc.Concat{ + doc.Text(" "), + f.emitTokens(open, open, emitOpts{leading: true, trailing: true}), + closeDoc, + } + } + + all := emitOpts{leading: true, trailing: true} + middle := make([]doc.Doc, 0, len(a.Items)*2) + last := open + + for i, item := range a.Items { + if i > 0 { + prev := a.Items[i-1] + if prev.Sep != 0 { + middle = append(middle, f.commaSep(prev.TokEnd())...) + } else { + // Lenient sources may omit separators; keep the items + // apart so their tokens cannot merge. + middle = append(middle, f.foldBreak(prev.TokEnd(), " ")) + } } - items = append(items, doc.Concat(parts)) + + end := item.TokEnd() + if item.Sep != 0 { + end-- + } + + middle = append(middle, f.emitTokens(item.TokStart(), end, all)) + last = end } + // A trailing comma after the last item (which may carry comments) is + // not between two items, so it is emitted here. + if lastItem := a.Items[len(a.Items)-1]; lastItem.Sep != 0 { + middle = append(middle, f.commaSep(lastItem.TokEnd())...) + last = lastItem.TokEnd() + } + group := doc.Group(doc.Concat{ - doc.Text("("), - doc.Indent(doc.Concat{doc.SoftLine, doc.Join(doc.Concat{doc.Text(","), doc.Line}, items)}), - doc.SoftLine, - doc.Text(")"), + f.emitTokens(open, open, all), + doc.Indent(doc.Concat{f.foldBreak(open, ""), doc.Concat(middle)}), + f.foldBreak(last, ""), + f.emitTokens(close, close, emitOpts{leading: true, trailing: !isLast}), }) + return doc.Concat{doc.Text(" "), group} } + +// trimComment returns the comment text without trailing whitespace, which +// the printer would trim at line ends anyway; emitting it untrimmed would +// skew the width measurement of enclosing groups. +func trimComment(text string) string { + return strings.TrimRight(text, " \t") +} diff --git a/formatter/format_fuzz_test.go b/formatter/format_fuzz_test.go index b582f16..45fa957 100644 --- a/formatter/format_fuzz_test.go +++ b/formatter/format_fuzz_test.go @@ -1,17 +1,22 @@ package formatter import ( + "reflect" + "strings" "testing" "github.com/karitham/thrift-ls/syntax" ) -// FuzzFormat checks the full formatting pipeline over arbitrary source: +// FuzzFormat checks the full formatting pipeline over arbitrary source and +// a full option set derived from the input bytes: // // - formatting never panics or loops forever // - formatted output of a clean document parses without errors // - formatting is idempotent: format(format(x)) == format(x) // - formatted output is deterministic +// - every comment and annotation survives formatting (lossless trivia) +// - in preserve mode, every field and enum separator survives func FuzzFormat(f *testing.F) { for _, seed := range []string{ "", @@ -24,6 +29,8 @@ func FuzzFormat(f *testing.F) { "include \"a.thrift\"\nnamespace go x", "struct S {", "service S {\n void f(1: i32 a // c\n )\n}", + "struct S {\n 1: map m\n}", + "const list L = [1, // mid\n 2]", } { f.Add(seed) } @@ -36,7 +43,8 @@ func FuzzFormat(f *testing.F) { } } - opts := testOpts(1 + len(src)%100) + opts := fuzzOpts(src) + out1, err := Format(doc, opts) if err != nil { t.Fatalf("Format: %v", err) @@ -50,26 +58,155 @@ func FuzzFormat(f *testing.F) { } } + // Losslessness: every comment and annotation text survives, in + // order. Right-trimmed because the printer trims line ends. + in := commentTexts(src) + if got := commentTexts(out1); !reflect.DeepEqual(in, got) { + t.Fatalf("comments lost or reordered:\nin: %q\nout: %q\ninput: %q\noutput: %q", in, got, src, out1) + } + + // In preserve mode every field and enum separator survives. + allPreserve := true + + for _, c := range AllConstructs { + if opts.Separator.Get(c) != SeparatorPreserve { + allPreserve = false + } + } + + if allPreserve { + in, errs := syntax.Parse([]byte(src)) + if hasParseErrors(errs) { + t.Fatalf("reparse of input failed: %v", errs) + } + + if got, _ := syntax.Parse([]byte(out1)); !reflect.DeepEqual(fieldSeps(in), fieldSeps(got)) { + t.Fatalf("separators changed in preserve mode:\nin: %v\nout: %v\ninput: %q\noutput: %q", + fieldSeps(in), fieldSeps(got), src, out1) + } + } + // Idempotency and determinism. out2, err := Format(doc, opts) if err != nil { t.Fatalf("Format (second): %v", err) } + if out1 != out2 { t.Fatalf("Format is not deterministic:\nfirst: %q\nsecond: %q", out1, out2) } + doc3, errs3 := syntax.Parse([]byte(out1)) for _, err := range errs3 { if err.Severity == syntax.SeverityError { t.Fatalf("re-parse failed: %v", errs3) } } + out3, err := Format(doc3, opts) if err != nil { t.Fatalf("Format (third): %v", err) } + if out1 != out3 { t.Fatalf("not idempotent:\nfirst: %q\nsecond: %q", out1, out3) } }) } + +// fuzzOpts derives a full option set deterministically from the input +// bytes, so every formatting path — alignment modes, separator modes, break +// flags, tab indents, width — is exercised by fuzzing. +func fuzzOpts(src string) Options { + var h [4]int + for i, b := range []byte(src) { + h[i%4] += int(b) + } + + o := DefaultOptions() + o.PrintWidth = 1 + h[0]%120 + + o.TabWidth = 1 + h[1]%8 + switch h[2] % 3 { + case 0: + o.Align = AlignField + case 1: + o.Align = AlignAssign + case 2: + o.Align = AlignDisable + } + + for i, c := range AllConstructs { + o.Separator.Set(c, SeparatorMode((h[i%4]+i)%4)) + o.Break.Set(c, h[(i+1)%4]%2 == 0) + } + + o.NoTrailingNewline = h[3]%3 == 0 + switch h[0] % 4 { + case 0: + o.Indent = " " + case 1: + o.Indent = "\t" + case 2: + o.Indent = " " + case 3: + o.Indent = "" + } + + return o +} + +// commentTexts returns every comment and annotation text in the source, in +// source order, right-trimmed to match the printer's line-end trimming. +func commentTexts(src string) []string { + toks, _ := syntax.Lex([]byte(src)) + + var texts []string + + for _, tok := range toks { + for _, tr := range tok.Leading { + texts = append(texts, strings.TrimRight(tr.Text, " \t")) + } + + for _, tr := range tok.Trailing { + texts = append(texts, strings.TrimRight(tr.Text, " \t")) + } + } + + return texts +} + +// fieldSeps returns the separator kind of every field and enum value in the +// document, in source order. +func fieldSeps(doc *syntax.Document) []syntax.TokenKind { + var seps []syntax.TokenKind + + add := func(f *syntax.Field) { seps = append(seps, f.Sep) } + + for _, n := range doc.Nodes { + switch v := n.(type) { + case *syntax.Struct: + for _, f := range v.Fields { + add(f) + } + case *syntax.Enum: + for _, ev := range v.Values { + seps = append(seps, ev.Sep) + } + case *syntax.Service: + for _, fn := range v.Functions { + for _, f := range fn.Args { + add(f) + } + + if fn.Throws != nil { + for _, f := range fn.Throws.Fields { + add(f) + } + } + } + } + } + + return seps +} diff --git a/formatter/format_test.go b/formatter/format_test.go index 6638c8b..38cc1e6 100644 --- a/formatter/format_test.go +++ b/formatter/format_test.go @@ -14,6 +14,7 @@ func hasParseErrors(errs []syntax.Error) bool { return true } } + return false } @@ -21,20 +22,24 @@ func hasParseErrors(errs []syntax.Error) bool { // test when parsing fails. func fmtSrc(t *testing.T, src string, opts Options) string { t.Helper() + got, err := Format(parseDoc(t, src), opts) if err != nil { t.Fatalf("Format: %v", err) } + return got } // parseDoc parses src and fails the test on parse errors. func parseDoc(t *testing.T, src string) *syntax.Document { t.Helper() + doc, errs := syntax.Parse([]byte(src)) if hasParseErrors(errs) { t.Fatalf("parse errors: %v", errs) } + return doc } @@ -43,22 +48,26 @@ func testOpts(width int) Options { o.PrintWidth = width o.Indent = " " o.TabWidth = 2 + return o } -// commaOpts returns testOpts at width with the given FieldSeparator. +// commaOpts returns testOpts at width with the given struct separator. func commaOpts(width int, mode SeparatorMode) Options { o := testOpts(width) - o.FieldSeparator = mode + o.Separator.Set(ConstructStruct, mode) + return o } // runCase formats, checks idempotency, and re-parses the output. func runCase(t *testing.T, src string, opts Options, want string) { t.Helper() + got := fmtSrc(t, src, opts) if got != want { t.Errorf("format mismatch\n got: %q\nwant: %q", got, want) + return } // Idempotency: formatting the output again must not change it. @@ -83,12 +92,14 @@ type formatCase struct { // runFormatCases runs width-based table cases through runCase. func runFormatCases(t *testing.T, cases []formatCase) { t.Helper() + for _, tt := range cases { t.Run(tt.name, func(t *testing.T) { width := tt.width if width == 0 { width = 80 } + runCase(t, tt.src, testOpts(width), tt.want) }) } @@ -236,6 +247,7 @@ func TestFormatConsts(t *testing.T) { if width == 0 { width = 40 } + runCase(t, tt.src, testOpts(width), tt.want) }) } @@ -397,10 +409,10 @@ func TestFormatFunctions(t *testing.T) { want: "service S {\n i32 getUser(1: i64 id, 2: string name) throws (\n NotFound e\n )\n}\n", }, { - name: "everything breaks when signature is long", + name: "args break but throws folds when it fits", src: "service S {\n i32 getUser(1: i64 id, 2: string name) throws (NotFound e)\n}", width: 45, - want: "service S {\n i32 getUser(\n 1: i64 id,\n 2: string name\n ) throws (\n NotFound e\n )\n}\n", + want: "service S {\n i32 getUser(\n 1: i64 id,\n 2: string name\n ) throws (NotFound e)\n}\n", }, { name: "args break without throws", @@ -524,6 +536,7 @@ func TestFormatOptions(t *testing.T) { opts: func() Options { o := testOpts(40) o.Align = AlignDisable + return o }(), src: "struct S {\n 1: required i64 id\n 2: string name\n}", @@ -534,6 +547,7 @@ func TestFormatOptions(t *testing.T) { opts: func() Options { o := testOpts(40) o.Align = AlignAssign + return o }(), src: "struct S {\n 1: i32 a = 5\n 2: string longer = \"x\"\n}", @@ -610,6 +624,7 @@ func TestFormatIsIdempotent(t *testing.T) { for i, src := range sources { t.Run("case-"+strconv.Itoa(i), func(t *testing.T) { first := fmtSrc(t, src, testOpts(40)) + second := fmtSrc(t, first, testOpts(40)) if first != second { t.Errorf("not idempotent:\nfirst: %q\nsecond: %q", first, second) @@ -626,10 +641,12 @@ struct User { doc := parseDoc(t, src) want := `// leading struct User { 1: required i64 id } (tag = "x")` + got, err := FormatNode(doc, doc.Structs()[0], testOpts(80)) if err != nil { t.Fatalf("FormatNode: %v", err) } + if got != want { t.Errorf("got %q, want %q", got, want) } @@ -646,8 +663,10 @@ func TestFormatSeparators(t *testing.T) { name: "fields semicolon, functions comma", opts: func() Options { o := testOpts(30) - o.FieldSeparator = SeparatorSemicolon - o.FunctionSeparator = SeparatorComma + o.Separator.Set(ConstructStruct, SeparatorSemicolon) + o.Separator.Set(ConstructArguments, SeparatorComma) + o.Separator.Set(ConstructThrows, SeparatorComma) + return o }(), src: "struct S {\n 1: i32 a\n 2: string b\n}\n\nservice F {\n void go(1: i32 x) throws (\n 1: E err\n )\n}", @@ -675,7 +694,9 @@ func TestFormatSeparators(t *testing.T) { name: "function comma add forces commas on throws", opts: func() Options { o := testOpts(30) - o.FunctionSeparator = SeparatorComma + o.Separator.Set(ConstructArguments, SeparatorComma) + o.Separator.Set(ConstructThrows, SeparatorComma) + return o }(), src: "service F {\n void go(1: i32 x) throws (\n 1: E err\n 2: F fail\n )\n}", @@ -685,7 +706,9 @@ func TestFormatSeparators(t *testing.T) { name: "function comma remove drops argument separators", opts: func() Options { o := testOpts(30) - o.FunctionSeparator = SeparatorNone + o.Separator.Set(ConstructArguments, SeparatorNone) + o.Separator.Set(ConstructThrows, SeparatorNone) + return o }(), src: "service F {\n void go(1: i32 x, 2: string y)\n}", @@ -710,7 +733,8 @@ func TestFormatAlwaysBreak(t *testing.T) { name: "break structs forces multiline", opts: func() Options { o := testOpts(80) - o.BreakStructs = true + o.Break.Set(ConstructStruct, true) + return o }(), src: "struct S {\n 1: i32 a\n}", @@ -720,7 +744,8 @@ func TestFormatAlwaysBreak(t *testing.T) { name: "break enums forces multiline", opts: func() Options { o := testOpts(80) - o.BreakEnums = true + o.Break.Set(ConstructEnum, true) + return o }(), src: "enum E {\n A,\n B\n}", @@ -730,7 +755,8 @@ func TestFormatAlwaysBreak(t *testing.T) { name: "break structs keeps empty bodies flat", opts: func() Options { o := testOpts(80) - o.BreakStructs = true + o.Break.Set(ConstructStruct, true) + return o }(), src: "struct S {}", @@ -740,7 +766,8 @@ func TestFormatAlwaysBreak(t *testing.T) { name: "break structs does not affect enums", opts: func() Options { o := testOpts(80) - o.BreakStructs = true + o.Break.Set(ConstructStruct, true) + return o }(), src: "struct S {\n 1: i32 a\n}\nenum E {\n A,\n B\n}", @@ -1097,3 +1124,193 @@ func TestFormatTrailingDelim(t *testing.T) { }) } } + +func TestFormatFileStartBlanks(t *testing.T) { + tests := []struct { + name string + src string + want string + }{ + { + name: "blank lines before a leading comment round-trip", + src: "\n\n// header\nstruct S {\n 1: i32 a\n}", + want: "\n\n// header\nstruct S { 1: i32 a }\n", + }, + { + name: "comment-only file with leading blanks round-trips", + src: "\n\n#0\n", + want: "\n\n#0\n", + }, + { + name: "comment at file start stays at line one", + src: "#0\n", + want: "#0\n", + }, + { + name: "carriage returns normalize to newlines", + src: "\r\r\r#0\n", + want: "\n\n\n#0\n", + }, + } + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + runCase(t, tt.src, testOpts(80), tt.want) + }) + } +} + +func TestFormatInnerTrivia(t *testing.T) { + tests := []struct { + name string + src string + want string + }{ + { + name: "comment inside a container type is preserved", + src: "struct S {\n 1: map m\n}", + want: "struct S {\n 1: map m\n}\n", + }, + { + name: "comment inside a const list is preserved", + src: "const list L = [1, // mid\n 2]", + want: "const list L = [\n 1, // mid\n 2\n]\n", + }, + { + name: "comment in empty body is preserved", + src: "struct A{\n#\n}", + want: "struct A {\n #\n}\n", + }, + { + name: "block comment inside type stays inline", + src: "struct S {\n 1: map m\n}", + want: "struct S { 1: map m }\n", + }, + { + name: "comment in header span is preserved", + src: "service S {\n void // c\n go(1: i32 a)\n}", + want: "service S {\n void // c\n go(1: i32 a)\n}\n", + }, + { + name: "empty body comment with trailing comment both kept", + src: "enum A{\n#\n}#", + want: "enum A {\n #\n} #\n", + }, + } + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + runCase(t, tt.src, testOpts(80), tt.want) + }) + } +} + +func TestFormatAlignmentNoFlip(t *testing.T) { + tests := []struct { + name string + width int + src string + want string + }{ + { + name: "accidentally aligned names do not force alignment", + width: 10, + src: "service S {\n void go(1: A00 A2 x000n0 A)\n}", + want: "service S {\n" + + " void go(\n" + + " 1: A00 A2\n" + + " x000n0 A\n" + + " )\n" + + "}\n", + }, + { + name: "deliberately padded names keep alignment", + width: 40, + src: "struct S {\n" + + " 1: federation.MobileSuitFrameType frame_type;\n" + + " 2: federation.PropulsionType propulsion_type;\n" + + "}", + want: "struct S {\n" + + " 1: federation.MobileSuitFrameType frame_type;\n" + + " 2: federation.PropulsionType propulsion_type;\n" + + "}\n", + }, + } + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + runCase(t, tt.src, testOpts(tt.width), tt.want) + }) + } +} + +func TestFormatThrowsFoldsWithBrokenArgs(t *testing.T) { + src := "service Processor {\n" + + " string upload(\n" + + " 1: string imageUrl,\n" + + " 2: arguments.Size size,\n" + + " 3: arguments.Identifier id,\n" + + " ) throws (1: errors.ProcessingError err)\n" + + "}" + want := "service Processor {\n" + + " string upload(\n" + + " 1: string imageUrl,\n" + + " 2: arguments.Size size,\n" + + " 3: arguments.Identifier id,\n" + + " ) throws (1: errors.ProcessingError err)\n" + + "}\n" + runCase(t, src, testOpts(80), want) +} + +func TestFormatSeparatorsPerConstruct(t *testing.T) { + opts := testOpts(80) + opts.Separator.Set(ConstructStruct, SeparatorSemicolon) + opts.Separator.Set(ConstructUnion, SeparatorSemicolon) + opts.Separator.Set(ConstructException, SeparatorSemicolon) + opts.Separator.Set(ConstructEnum, SeparatorComma) + + src := "struct S {\n 1: i32 a\n 2: i32 b\n}\n\nunion U {\n 1: i32 a\n 2: i32 b\n}\n\nexception X {\n 1: i32 a\n 2: i32 b\n}\n\nenum E {\n A\n B\n}" + want := "struct S { 1: i32 a; 2: i32 b }\n\nunion U { 1: i32 a; 2: i32 b }\n\nexception X { 1: i32 a; 2: i32 b }\n\nenum E { A, B }\n" + runCase(t, src, opts, want) +} + +func TestFormatPreserveSeparators(t *testing.T) { + tests := []struct { + name string + src string + want string + }{ + { + name: "mixed separators force the broken layout", + src: "struct S {\n 1: i32 a\n 2: string b;\n 3: bool c\n}", + want: "struct S {\n 1: i32 a\n 2: string b;\n 3: bool c\n}\n", + }, + { + name: "uniform separators fold flat", + src: "struct S {\n 1: i32 a;\n 2: string b\n}", + want: "struct S { 1: i32 a; 2: string b }\n", + }, + { + name: "no separators fold flat", + src: "enum E {\n A\n B\n}", + want: "enum E { A B }\n", + }, + { + name: "mixed separators in enum force the broken layout", + src: "enum E {\n A\n B,\n C\n}", + want: "enum E {\n A\n B,\n C\n}\n", + }, + { + name: "mixed separators in args force the broken layout", + src: "service F {\n void go(1: i32 a 2: string b; 3: bool c)\n}", + want: "service F {\n void go(\n 1: i32 a\n 2: string b;\n 3: bool c\n )\n}\n", + }, + { + name: "trailing delimiter on the last field forces broken", + src: "struct S {\n 1: i32 a\n 2: string b;\n}", + want: "struct S {\n 1: i32 a\n 2: string b;\n}\n", + }, + } + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + runCase(t, tt.src, testOpts(80), tt.want) + }) + } +} diff --git a/formatter/testdata/fuzz/FuzzFormat/0a8ba8406d390d27 b/formatter/testdata/fuzz/FuzzFormat/0a8ba8406d390d27 new file mode 100644 index 0000000..7c8a428 --- /dev/null +++ b/formatter/testdata/fuzz/FuzzFormat/0a8ba8406d390d27 @@ -0,0 +1,2 @@ +go test fuzz v1 +string("service A{A A()()#\n,}") diff --git a/formatter/testdata/fuzz/FuzzFormat/0add0550527ae8f3 b/formatter/testdata/fuzz/FuzzFormat/0add0550527ae8f3 new file mode 100644 index 0000000..147f7ed --- /dev/null +++ b/formatter/testdata/fuzz/FuzzFormat/0add0550527ae8f3 @@ -0,0 +1,2 @@ +go test fuzz v1 +string("enum A{\n#\n}#") diff --git a/formatter/testdata/fuzz/FuzzFormat/105a2eeb3cac4258 b/formatter/testdata/fuzz/FuzzFormat/105a2eeb3cac4258 new file mode 100644 index 0000000..ea9937d --- /dev/null +++ b/formatter/testdata/fuzz/FuzzFormat/105a2eeb3cac4258 @@ -0,0 +1,2 @@ +go test fuzz v1 +string("service x{A0 A(A&A AA A)}") diff --git a/formatter/testdata/fuzz/FuzzFormat/1a541db09a978464 b/formatter/testdata/fuzz/FuzzFormat/1a541db09a978464 new file mode 100644 index 0000000..660c078 --- /dev/null +++ b/formatter/testdata/fuzz/FuzzFormat/1a541db09a978464 @@ -0,0 +1,2 @@ +go test fuzz v1 +string("service C{ z00 a(1: A00 A2 x000n0 A) t x020 (B X)}") diff --git a/formatter/testdata/fuzz/FuzzFormat/1d204b998359871d b/formatter/testdata/fuzz/FuzzFormat/1d204b998359871d new file mode 100644 index 0000000..1ff0d76 --- /dev/null +++ b/formatter/testdata/fuzz/FuzzFormat/1d204b998359871d @@ -0,0 +1,2 @@ +go test fuzz v1 +string("typedef A A(A0 A)") diff --git a/formatter/testdata/fuzz/FuzzFormat/2ce6c4d11355e008 b/formatter/testdata/fuzz/FuzzFormat/2ce6c4d11355e008 new file mode 100644 index 0000000..d3dd0dd --- /dev/null +++ b/formatter/testdata/fuzz/FuzzFormat/2ce6c4d11355e008 @@ -0,0 +1,2 @@ +go test fuzz v1 +string("enum A{}\r#\r()") diff --git a/formatter/testdata/fuzz/FuzzFormat/2cea8d070cda7c22 b/formatter/testdata/fuzz/FuzzFormat/2cea8d070cda7c22 new file mode 100644 index 0000000..2113c3e --- /dev/null +++ b/formatter/testdata/fuzz/FuzzFormat/2cea8d070cda7c22 @@ -0,0 +1,2 @@ +go test fuzz v1 +string("enum A{A}(#\r)") diff --git a/formatter/testdata/fuzz/FuzzFormat/2f6cf61071db8ac9 b/formatter/testdata/fuzz/FuzzFormat/2f6cf61071db8ac9 new file mode 100644 index 0000000..35da92a --- /dev/null +++ b/formatter/testdata/fuzz/FuzzFormat/2f6cf61071db8ac9 @@ -0,0 +1,2 @@ +go test fuzz v1 +string("struct A{#\n}") diff --git a/formatter/testdata/fuzz/FuzzFormat/30bd06f692a2d082 b/formatter/testdata/fuzz/FuzzFormat/30bd06f692a2d082 new file mode 100644 index 0000000..ba72746 --- /dev/null +++ b/formatter/testdata/fuzz/FuzzFormat/30bd06f692a2d082 @@ -0,0 +1,2 @@ +go test fuzz v1 +string("typedef A A(A\n#\n,)") diff --git a/formatter/testdata/fuzz/FuzzFormat/317431e66ff11302 b/formatter/testdata/fuzz/FuzzFormat/317431e66ff11302 new file mode 100644 index 0000000..d055eb0 --- /dev/null +++ b/formatter/testdata/fuzz/FuzzFormat/317431e66ff11302 @@ -0,0 +1,2 @@ +go test fuzz v1 +string("service A{A A(\r\rA A)}") diff --git a/formatter/testdata/fuzz/FuzzFormat/347fd8c11b506572 b/formatter/testdata/fuzz/FuzzFormat/347fd8c11b506572 new file mode 100644 index 0000000..0802eae --- /dev/null +++ b/formatter/testdata/fuzz/FuzzFormat/347fd8c11b506572 @@ -0,0 +1,2 @@ +go test fuzz v1 +string("const listA=[00A0]") diff --git a/formatter/testdata/fuzz/FuzzFormat/4a948a7788ab09e2 b/formatter/testdata/fuzz/FuzzFormat/4a948a7788ab09e2 new file mode 100644 index 0000000..f1c1105 --- /dev/null +++ b/formatter/testdata/fuzz/FuzzFormat/4a948a7788ab09e2 @@ -0,0 +1,2 @@ +go test fuzz v1 +string("enum A02{A ,#\nA0=0}") diff --git a/formatter/testdata/fuzz/FuzzFormat/4d4e0091885f4958 b/formatter/testdata/fuzz/FuzzFormat/4d4e0091885f4958 new file mode 100644 index 0000000..1bb2938 --- /dev/null +++ b/formatter/testdata/fuzz/FuzzFormat/4d4e0091885f4958 @@ -0,0 +1,2 @@ +go test fuzz v1 +string("typedef A A#") diff --git a/formatter/testdata/fuzz/FuzzFormat/51dd8024287b00fd b/formatter/testdata/fuzz/FuzzFormat/51dd8024287b00fd new file mode 100644 index 0000000..6edc4ca --- /dev/null +++ b/formatter/testdata/fuzz/FuzzFormat/51dd8024287b00fd @@ -0,0 +1,2 @@ +go test fuzz v1 +string("typedef A A(A)#\n,") diff --git a/formatter/testdata/fuzz/FuzzFormat/5ad4ef301cf0051f b/formatter/testdata/fuzz/FuzzFormat/5ad4ef301cf0051f new file mode 100644 index 0000000..8600bfd --- /dev/null +++ b/formatter/testdata/fuzz/FuzzFormat/5ad4ef301cf0051f @@ -0,0 +1,2 @@ +go test fuzz v1 +string("const A A0= [0, #0\n ]") diff --git a/formatter/testdata/fuzz/FuzzFormat/6013bc9c1f58f67a b/formatter/testdata/fuzz/FuzzFormat/6013bc9c1f58f67a new file mode 100644 index 0000000..217232f --- /dev/null +++ b/formatter/testdata/fuzz/FuzzFormat/6013bc9c1f58f67a @@ -0,0 +1,2 @@ +go test fuzz v1 +string("enum A{A00020\n#\n,#\n}") diff --git a/formatter/testdata/fuzz/FuzzFormat/63a08e1d1b8eac17 b/formatter/testdata/fuzz/FuzzFormat/63a08e1d1b8eac17 new file mode 100644 index 0000000..438f574 --- /dev/null +++ b/formatter/testdata/fuzz/FuzzFormat/63a08e1d1b8eac17 @@ -0,0 +1,2 @@ +go test fuzz v1 +string("enum A{A#\r()}") diff --git a/formatter/testdata/fuzz/FuzzFormat/6bf5f083c01e4f46 b/formatter/testdata/fuzz/FuzzFormat/6bf5f083c01e4f46 new file mode 100644 index 0000000..72b437a --- /dev/null +++ b/formatter/testdata/fuzz/FuzzFormat/6bf5f083c01e4f46 @@ -0,0 +1,2 @@ +go test fuzz v1 +string("enum A#\r{}") diff --git a/formatter/testdata/fuzz/FuzzFormat/739cd322f5a4ce84 b/formatter/testdata/fuzz/FuzzFormat/739cd322f5a4ce84 new file mode 100644 index 0000000..808ab84 --- /dev/null +++ b/formatter/testdata/fuzz/FuzzFormat/739cd322f5a4ce84 @@ -0,0 +1,2 @@ +go test fuzz v1 +string("service C0{ X000 A0(1:A020 A, ) }") diff --git a/formatter/testdata/fuzz/FuzzFormat/786e75e57ee88aac b/formatter/testdata/fuzz/FuzzFormat/786e75e57ee88aac new file mode 100644 index 0000000..93c23f7 --- /dev/null +++ b/formatter/testdata/fuzz/FuzzFormat/786e75e57ee88aac @@ -0,0 +1,2 @@ +go test fuzz v1 +string("enum A{A#\n,}") diff --git a/formatter/testdata/fuzz/FuzzFormat/7d1b13ec2b6df311 b/formatter/testdata/fuzz/FuzzFormat/7d1b13ec2b6df311 new file mode 100644 index 0000000..ea27ae0 --- /dev/null +++ b/formatter/testdata/fuzz/FuzzFormat/7d1b13ec2b6df311 @@ -0,0 +1,2 @@ +go test fuzz v1 +string("enum A1{A#\n#\n,A0=0}") diff --git a/formatter/testdata/fuzz/FuzzFormat/7d3edf5144a837fe b/formatter/testdata/fuzz/FuzzFormat/7d3edf5144a837fe new file mode 100644 index 0000000..b0c088f --- /dev/null +++ b/formatter/testdata/fuzz/FuzzFormat/7d3edf5144a837fe @@ -0,0 +1,2 @@ +go test fuzz v1 +string("enum A{}(\r#\r)") diff --git a/formatter/testdata/fuzz/FuzzFormat/831742f9bbf009d6 b/formatter/testdata/fuzz/FuzzFormat/831742f9bbf009d6 new file mode 100644 index 0000000..c288bdf --- /dev/null +++ b/formatter/testdata/fuzz/FuzzFormat/831742f9bbf009d6 @@ -0,0 +1,2 @@ +go test fuzz v1 +string("enum A{}#\r()") diff --git a/formatter/testdata/fuzz/FuzzFormat/84ab644fa6ea62cc b/formatter/testdata/fuzz/FuzzFormat/84ab644fa6ea62cc new file mode 100644 index 0000000..bfcf30b --- /dev/null +++ b/formatter/testdata/fuzz/FuzzFormat/84ab644fa6ea62cc @@ -0,0 +1,2 @@ +go test fuzz v1 +string("enum A{A#\n,B00#\nA00}") diff --git a/formatter/testdata/fuzz/FuzzFormat/86a87b8702497a70 b/formatter/testdata/fuzz/FuzzFormat/86a87b8702497a70 new file mode 100644 index 0000000..3aaddfb --- /dev/null +++ b/formatter/testdata/fuzz/FuzzFormat/86a87b8702497a70 @@ -0,0 +1,2 @@ +go test fuzz v1 +string("service A{A A(\n#\n)}") diff --git a/formatter/testdata/fuzz/FuzzFormat/9407cafa157562e0 b/formatter/testdata/fuzz/FuzzFormat/9407cafa157562e0 new file mode 100644 index 0000000..3b5656f --- /dev/null +++ b/formatter/testdata/fuzz/FuzzFormat/9407cafa157562e0 @@ -0,0 +1,2 @@ +go test fuzz v1 +string("typedef A A()\n#\n,") diff --git a/formatter/testdata/fuzz/FuzzFormat/a45c7b6e01a81072 b/formatter/testdata/fuzz/FuzzFormat/a45c7b6e01a81072 new file mode 100644 index 0000000..5c030f1 --- /dev/null +++ b/formatter/testdata/fuzz/FuzzFormat/a45c7b6e01a81072 @@ -0,0 +1,2 @@ +go test fuzz v1 +string("service A{A A()#\n,}") diff --git a/formatter/testdata/fuzz/FuzzFormat/a691eaf46b192a96 b/formatter/testdata/fuzz/FuzzFormat/a691eaf46b192a96 new file mode 100644 index 0000000..a0bd071 --- /dev/null +++ b/formatter/testdata/fuzz/FuzzFormat/a691eaf46b192a96 @@ -0,0 +1,2 @@ +go test fuzz v1 +string("const A A=[0;#\n]") diff --git a/formatter/testdata/fuzz/FuzzFormat/aa884d274bb46803 b/formatter/testdata/fuzz/FuzzFormat/aa884d274bb46803 new file mode 100644 index 0000000..33b6dec --- /dev/null +++ b/formatter/testdata/fuzz/FuzzFormat/aa884d274bb46803 @@ -0,0 +1,2 @@ +go test fuzz v1 +string("service A{A A(#\n)}") diff --git a/formatter/testdata/fuzz/FuzzFormat/aac0fbb3722aa15e b/formatter/testdata/fuzz/FuzzFormat/aac0fbb3722aa15e new file mode 100644 index 0000000..25fb8bd --- /dev/null +++ b/formatter/testdata/fuzz/FuzzFormat/aac0fbb3722aa15e @@ -0,0 +1,2 @@ +go test fuzz v1 +string("enum A{A00=00A0}#0 ") diff --git a/formatter/testdata/fuzz/FuzzFormat/b06fd83df05bcd9d b/formatter/testdata/fuzz/FuzzFormat/b06fd83df05bcd9d new file mode 100644 index 0000000..8dd8cfb --- /dev/null +++ b/formatter/testdata/fuzz/FuzzFormat/b06fd83df05bcd9d @@ -0,0 +1,2 @@ +go test fuzz v1 +string("enum A{A#\n,#\nA020}") diff --git a/formatter/testdata/fuzz/FuzzFormat/bfa3d3c4dbe2b7b5 b/formatter/testdata/fuzz/FuzzFormat/bfa3d3c4dbe2b7b5 new file mode 100644 index 0000000..ddd4595 --- /dev/null +++ b/formatter/testdata/fuzz/FuzzFormat/bfa3d3c4dbe2b7b5 @@ -0,0 +1,2 @@ +go test fuzz v1 +string("\r\r\r#0") diff --git a/formatter/testdata/fuzz/FuzzFormat/caa3d663c4908eac b/formatter/testdata/fuzz/FuzzFormat/caa3d663c4908eac new file mode 100644 index 0000000..76676ca --- /dev/null +++ b/formatter/testdata/fuzz/FuzzFormat/caa3d663c4908eac @@ -0,0 +1,2 @@ +go test fuzz v1 +string("#00\nstruct A{#0\n}") diff --git a/formatter/testdata/fuzz/FuzzFormat/cc2a214d09780604 b/formatter/testdata/fuzz/FuzzFormat/cc2a214d09780604 new file mode 100644 index 0000000..1073542 --- /dev/null +++ b/formatter/testdata/fuzz/FuzzFormat/cc2a214d09780604 @@ -0,0 +1,2 @@ +go test fuzz v1 +string("const listA=[0[]#0\n]") diff --git a/formatter/testdata/fuzz/FuzzFormat/cc80f0214b0ceefe b/formatter/testdata/fuzz/FuzzFormat/cc80f0214b0ceefe new file mode 100644 index 0000000..25f578f --- /dev/null +++ b/formatter/testdata/fuzz/FuzzFormat/cc80f0214b0ceefe @@ -0,0 +1,2 @@ +go test fuzz v1 +string("const A A0=[0] #0\n,") diff --git a/formatter/testdata/fuzz/FuzzFormat/ce8a17c96204f573 b/formatter/testdata/fuzz/FuzzFormat/ce8a17c96204f573 new file mode 100644 index 0000000..f59907b --- /dev/null +++ b/formatter/testdata/fuzz/FuzzFormat/ce8a17c96204f573 @@ -0,0 +1,2 @@ +go test fuzz v1 +string("typedef A A(A#\n,A)") diff --git a/formatter/testdata/fuzz/FuzzFormat/d07e039d1e905df1 b/formatter/testdata/fuzz/FuzzFormat/d07e039d1e905df1 new file mode 100644 index 0000000..4c4366e --- /dev/null +++ b/formatter/testdata/fuzz/FuzzFormat/d07e039d1e905df1 @@ -0,0 +1,2 @@ +go test fuzz v1 +string("struct A{\n#\n}") diff --git a/formatter/testdata/fuzz/FuzzFormat/d2f769beff5f1fa5 b/formatter/testdata/fuzz/FuzzFormat/d2f769beff5f1fa5 new file mode 100644 index 0000000..b968b64 --- /dev/null +++ b/formatter/testdata/fuzz/FuzzFormat/d2f769beff5f1fa5 @@ -0,0 +1,2 @@ +go test fuzz v1 +string("enum A1{A#\n,#\nA0=0}") diff --git a/formatter/testdata/fuzz/FuzzFormat/d4e3e67833954e95 b/formatter/testdata/fuzz/FuzzFormat/d4e3e67833954e95 new file mode 100644 index 0000000..b25c779 --- /dev/null +++ b/formatter/testdata/fuzz/FuzzFormat/d4e3e67833954e95 @@ -0,0 +1,2 @@ +go test fuzz v1 +string("service C{ Az00 a(1: A00 A2 x000n7 A)X2 x0200(B X)}") diff --git a/formatter/testdata/fuzz/FuzzFormat/db9fec3fe1c71e35 b/formatter/testdata/fuzz/FuzzFormat/db9fec3fe1c71e35 new file mode 100644 index 0000000..e9373cc --- /dev/null +++ b/formatter/testdata/fuzz/FuzzFormat/db9fec3fe1c71e35 @@ -0,0 +1,2 @@ +go test fuzz v1 +string("include#\n\"\"") diff --git a/formatter/testdata/fuzz/FuzzFormat/e508a31dad0cd567 b/formatter/testdata/fuzz/FuzzFormat/e508a31dad0cd567 new file mode 100644 index 0000000..4cf7a9b --- /dev/null +++ b/formatter/testdata/fuzz/FuzzFormat/e508a31dad0cd567 @@ -0,0 +1,2 @@ +go test fuzz v1 +string("service C0{\n #00000000000000000000000000000000000000000000\n}") diff --git a/formatter/testdata/fuzz/FuzzFormat/e9fbe1131b5c94fc b/formatter/testdata/fuzz/FuzzFormat/e9fbe1131b5c94fc new file mode 100644 index 0000000..cc8a79b --- /dev/null +++ b/formatter/testdata/fuzz/FuzzFormat/e9fbe1131b5c94fc @@ -0,0 +1,2 @@ +go test fuzz v1 +string("enum A02{A#\n,#\nA0=0}") diff --git a/formatter/testdata/fuzz/FuzzFormat/f6c1a7b9dfe81b75 b/formatter/testdata/fuzz/FuzzFormat/f6c1a7b9dfe81b75 new file mode 100644 index 0000000..430cfc5 --- /dev/null +++ b/formatter/testdata/fuzz/FuzzFormat/f6c1a7b9dfe81b75 @@ -0,0 +1,2 @@ +go test fuzz v1 +string("typedef A A()#\n,") diff --git a/formatter/testdata/fuzz/FuzzFormat/f8c0a10561d92f6d b/formatter/testdata/fuzz/FuzzFormat/f8c0a10561d92f6d new file mode 100644 index 0000000..1d02092 --- /dev/null +++ b/formatter/testdata/fuzz/FuzzFormat/f8c0a10561d92f6d @@ -0,0 +1,2 @@ +go test fuzz v1 +string("enum A{A#0100\n,}") diff --git a/formatter/testdata/fuzz/FuzzFormat/fde2a86adaa7051d b/formatter/testdata/fuzz/FuzzFormat/fde2a86adaa7051d new file mode 100644 index 0000000..4775b59 --- /dev/null +++ b/formatter/testdata/fuzz/FuzzFormat/fde2a86adaa7051d @@ -0,0 +1,2 @@ +go test fuzz v1 +string("typedef A A(A,#\n)") diff --git a/formatter/testdata/fuzz/FuzzFormat/fe244eab1c16e814 b/formatter/testdata/fuzz/FuzzFormat/fe244eab1c16e814 new file mode 100644 index 0000000..8bba03c --- /dev/null +++ b/formatter/testdata/fuzz/FuzzFormat/fe244eab1c16e814 @@ -0,0 +1,2 @@ +go test fuzz v1 +string("enum A0{A0\n#\n,}") diff --git a/formatter/testdata/fuzz/FuzzFormat/ffa1225bd5a451d6 b/formatter/testdata/fuzz/FuzzFormat/ffa1225bd5a451d6 new file mode 100644 index 0000000..caf8470 --- /dev/null +++ b/formatter/testdata/fuzz/FuzzFormat/ffa1225bd5a451d6 @@ -0,0 +1,2 @@ +go test fuzz v1 +string("enum A1{A#\n,}") diff --git a/formatter/type.go b/formatter/type.go index d9c345a..8e689f7 100644 --- a/formatter/type.go +++ b/formatter/type.go @@ -3,25 +3,18 @@ package formatter import ( "strings" - "github.com/karitham/thrift-ls/doc" "github.com/karitham/thrift-ls/syntax" ) -// fieldType formats a type reference: base keyword, identifier, or -// container. Containers and base types carry optional annotations. -func (f *formatter) fieldType(t *syntax.FieldType) doc.Doc { - parts := []doc.Doc{doc.Text(typeText(t))} - parts = append(parts, f.annotationsDoc(t.Annotations)) - return doc.Concat(parts) -} - -// typeText renders a type reference as plain text, matching the compiler's -// spelling: containers use "map", optional cpp_type comes after the -// container keyword ("list cpp_type \"x\" "), matching the grammar. +// typeText renders a type reference as plain text for width measurement, +// matching the compiler's spelling: containers use "map", optional +// cpp_type comes after the container keyword ("list cpp_type \"x\" "). +// Rendering itself goes through token runs, which preserve trivia. func typeText(t *syntax.FieldType) string { if t == nil { return "" } + switch t.Kind { case syntax.TypeBase: return t.Base.String() @@ -34,6 +27,7 @@ func typeText(t *syntax.FieldType) string { case syntax.TypeSet: return containerText("set", t.CPPType, typeText(t.ValueType)) } + return "" } @@ -43,5 +37,6 @@ func containerText(kind string, cppType *syntax.Token, inner ...string) string { // The grammar reads cpp_type before '<': "list cpp_type \"x\" ". return prefix + " cpp_type " + cppType.Text + " <" + strings.Join(inner, ", ") + ">" } + return prefix + "<" + strings.Join(inner, ", ") + ">" } diff --git a/formatter/value.go b/formatter/value.go index a4eba24..e3922d9 100644 --- a/formatter/value.go +++ b/formatter/value.go @@ -5,50 +5,141 @@ import ( "github.com/karitham/thrift-ls/syntax" ) -// constValue formats a constant value. Scalars are emitted verbatim from -// their source text; lists and maps are groups that stay on one line when -// they fit and break with one entry per line otherwise. -func (f *formatter) constValue(v *syntax.ConstValue) doc.Doc { +// isListSep reports whether the token separates list items: a comma or a +// semicolon (lenient sources may use either). +func isListSep(kind syntax.TokenKind) bool { + return kind == syntax.TokenComma || kind == syntax.TokenSemicolon +} + +// constValue formats a constant value. Scalars render as a token run; +// lists and maps are groups that stay on one line when they fit and break +// with one entry per line otherwise. Every segment is a token run, so +// comments inside the value are preserved. isLast reports whether the +// value ends the enclosing declaration, in which case its trailing trivia +// belongs to the declaration's trailing comments. +func (f *formatter) constValue(v *syntax.ConstValue, isLast bool) doc.Doc { if v == nil { return doc.Concat{} } + switch v.Kind { case syntax.ValueList: - items := make([]doc.Doc, 0, len(v.List)) - for _, item := range v.List { - items = append(items, f.constValue(item)) + return f.constList(v, isLast) + case syntax.ValueMap: + return f.constMap(v, isLast) + default: + o := emitOpts{leading: true} + if !isLast { + o.trailing = true } - return f.bracketList("[", "]", items) - case syntax.ValueMap: - entries := make([]doc.Doc, 0, len(v.Map)) - for _, entry := range v.Map { - entries = append(entries, doc.Concat{ - f.constValue(entry.Key), - doc.Text(": "), - f.constValue(entry.Value), - }) + return f.emitTokens(v.TokStart(), v.TokEnd(), o) + } +} + +// constList formats "[ items ]" as a foldable group. +func (f *formatter) constList(v *syntax.ConstValue, isLast bool) doc.Doc { + open, close := v.TokStart(), v.TokEnd() + + all := emitOpts{leading: true, trailing: true} + if len(v.List) == 0 { + closeDoc := f.emitTokens(close, close, emitOpts{leading: true, trailing: !isLast}) + if f.lineAfter(open) || len(f.token(close).Leading) > 0 { + closeDoc = doc.Concat{doc.HardLine, closeDoc} } - return f.bracketList("{", "}", entries) - default: - return doc.Text(v.Text) + return doc.Concat{ + f.emitTokens(open, open, all), + closeDoc, + } } + + middle := make([]doc.Doc, 0, len(v.List)*2) + last := open + + for i, item := range v.List { + if i > 0 { + prevEnd := v.List[i-1].TokEnd() + if isListSep(f.token(prevEnd + 1).Kind) { + middle = append(middle, f.commaSep(prevEnd+1)...) + } else { + // Lenient sources may omit separators; keep the items + // apart so their tokens cannot merge. + middle = append(middle, f.foldBreak(prevEnd, " ")) + } + } + + middle = append(middle, f.constValue(item, false)) + last = item.TokEnd() + } + // A trailing comma after the last item (which may carry comments) is + // not between two items, so it is emitted here. + if isListSep(f.token(last + 1).Kind) { + middle = append(middle, f.commaSep(last+1)...) + last++ + } + + closeOpts := emitOpts{leading: true} + if !isLast { + closeOpts.trailing = true + } + + return doc.Group(doc.Concat{ + f.emitTokens(open, open, all), + doc.Indent(doc.Concat{f.foldBreak(open, ""), doc.Concat(middle)}), + f.foldBreak(last, ""), + f.emitTokens(close, close, closeOpts), + }) } -// bracketList formats [open, items, close] as a group: flat when it fits, -// otherwise one item per line, indented. Empty lists stay on one line. -func (f *formatter) bracketList(open, close string, items []doc.Doc) doc.Doc { - if len(items) == 0 { - return doc.Text(open + close) +// constMap formats "{ key: value, ... }" as a foldable group. +func (f *formatter) constMap(v *syntax.ConstValue, isLast bool) doc.Doc { + open, close := v.TokStart(), v.TokEnd() + + all := emitOpts{leading: true, trailing: true} + if len(v.Map) == 0 { + closeDoc := f.emitTokens(close, close, emitOpts{leading: true, trailing: !isLast}) + if f.lineAfter(open) || len(f.token(close).Leading) > 0 { + closeDoc = doc.Concat{doc.HardLine, closeDoc} + } + + return doc.Concat{ + f.emitTokens(open, open, all), + closeDoc, + } + } + + middle := make([]doc.Doc, 0, len(v.Map)*2) + last := open + + for i, entry := range v.Map { + if i > 0 { + prevEnd := v.Map[i-1].Value.TokEnd() + if isListSep(f.token(prevEnd + 1).Kind) { + middle = append(middle, f.commaSep(prevEnd+1)...) + } else { + middle = append(middle, f.foldBreak(prevEnd, " ")) + } + } + + middle = append(middle, f.emitTokens(entry.Key.TokStart(), entry.Value.TokEnd(), all)) + last = entry.Value.TokEnd() } + + if isListSep(f.token(last + 1).Kind) { + middle = append(middle, f.commaSep(last+1)...) + last++ + } + + closeOpts := emitOpts{leading: true} + if !isLast { + closeOpts.trailing = true + } + return doc.Group(doc.Concat{ - doc.Text(open), - doc.Indent(doc.Concat{ - doc.SoftLine, - doc.Join(doc.Concat{doc.Text(","), doc.Line}, items), - }), - doc.SoftLine, - doc.Text(close), + f.emitTokens(open, open, all), + doc.Indent(doc.Concat{f.foldBreak(open, ""), doc.Concat(middle)}), + f.foldBreak(last, ""), + f.emitTokens(close, close, closeOpts), }) } diff --git a/log/log.go b/log/log.go index 0b150da..c3652f0 100644 --- a/log/log.go +++ b/log/log.go @@ -16,10 +16,12 @@ import ( // the old logrus levels so CLI flags keep their meaning. func Init(level int) { file := os.TempDir() + "/thriftls.log" + logFile, err := os.OpenFile(file, os.O_RDWR|os.O_CREATE|os.O_APPEND, 0o766) if err != nil { panic(err) } + slog.SetDefault(slog.New(slog.NewTextHandler(logFile, &slog.HandlerOptions{ Level: slogLevel(level), }))) diff --git a/lsp/cache/cache.go b/lsp/cache/cache.go index 1f2ab91..f63d606 100644 --- a/lsp/cache/cache.go +++ b/lsp/cache/cache.go @@ -33,6 +33,7 @@ func New(store *memoize.Store, includePaths []string) *Cache { IncludePaths: includePaths, memoizedFS: &memoizedFS{filesByID: map[FileID][]*DiskFile{}}, } + return c } diff --git a/lsp/cache/file.go b/lsp/cache/file.go index 5515d31..da9d2b8 100644 --- a/lsp/cache/file.go +++ b/lsp/cache/file.go @@ -112,7 +112,9 @@ type FilesMap struct { func (m *FilesMap) Get(key uri.URI) (FileHandle, bool) { m.mu.RLock() defer m.mu.RUnlock() + fh, ok := m.files[key] + return fh, ok } @@ -145,6 +147,7 @@ func (m *FilesMap) Clone() *FilesMap { for key := range m.files { newMap.files[key] = m.files[key] } + for key := range m.overlays { newMap.overlays[key] = m.overlays[key] } @@ -187,6 +190,7 @@ func FileChangeFromLSPDidChange(params *protocol.DidChangeTextDocumentParams) [] // semantics using the current full content. continue } + changes = append(changes, &FileChange{ URI: params.TextDocument.URI, Version: int(params.TextDocument.Version), @@ -194,5 +198,6 @@ func FileChangeFromLSPDidChange(params *protocol.DidChangeTextDocumentParams) [] From: FileChangeTypeDidChange, }) } + return changes } diff --git a/lsp/cache/file_posix.go b/lsp/cache/file_posix.go index 1203e6f..c61105d 100644 --- a/lsp/cache/file_posix.go +++ b/lsp/cache/file_posix.go @@ -13,7 +13,9 @@ func getFileID(filename string) (FileID, time.Time, error) { if err != nil { return FileID{}, time.Time{}, err } + stat := fi.Sys().(*syscall.Stat_t) + return FileID{ device: uint64(stat.Dev), // (int32 on darwin, uint64 on linux) inode: stat.Ino, diff --git a/lsp/cache/fs_memoized.go b/lsp/cache/fs_memoized.go index 50377e4..85c28db 100644 --- a/lsp/cache/fs_memoized.go +++ b/lsp/cache/fs_memoized.go @@ -65,6 +65,7 @@ func (fs *memoizedFS) ReadFile(ctx context.Context, uri uri.URI) (FileHandle, er recentlyModified := time.Since(mtime) < 2*time.Second fs.mu.Lock() + fhs, ok := fs.filesByID[id] if ok && fhs[0].modTime.Equal(mtime) { var fh *DiskFile @@ -72,6 +73,7 @@ func (fs *memoizedFS) ReadFile(ctx context.Context, uri uri.URI) (FileHandle, er for _, h := range fhs { if h.uri == uri { fh = h + break } } @@ -84,6 +86,7 @@ func (fs *memoizedFS) ReadFile(ctx context.Context, uri uri.URI) (FileHandle, er fs.filesByID[id] = fhs } fs.mu.Unlock() + return fh, nil } fs.mu.Unlock() @@ -101,6 +104,7 @@ func (fs *memoizedFS) ReadFile(ctx context.Context, uri uri.URI) (FileHandle, er delete(fs.filesByID, id) } fs.mu.Unlock() + return fh, nil } @@ -113,6 +117,7 @@ func readFile(ctx context.Context, uri uri.URI, mtime time.Time) (*DiskFile, err case <-ctx.Done(): return nil, ctx.Err() } + defer func() { <-ioLimit }() // It is possible that a race causes us to read a file with different file @@ -123,6 +128,7 @@ func readFile(ctx context.Context, uri uri.URI, mtime time.Time) (*DiskFile, err if err != nil { content = nil // just in case } + return &DiskFile{ modTime: mtime, uri: uri, diff --git a/lsp/cache/fs_overlay.go b/lsp/cache/fs_overlay.go index 3518321..1c84a4e 100644 --- a/lsp/cache/fs_overlay.go +++ b/lsp/cache/fs_overlay.go @@ -28,10 +28,12 @@ func NewOverlayFS(delegate FileSource) *overlayFS { func (fs *overlayFS) Overlays() []*Overlay { fs.mu.Lock() defer fs.mu.Unlock() + overlays := make([]*Overlay, 0, len(fs.overlays)) for _, overlay := range fs.overlays { overlays = append(overlays, overlay) } + return overlays } @@ -40,9 +42,11 @@ func (fs *overlayFS) ReadFile(ctx context.Context, uri uri.URI) (FileHandle, err fs.mu.Lock() overlay, ok := fs.overlays[uri] fs.mu.Unlock() + if ok { return overlay, nil } + return fs.delegate.ReadFile(ctx, uri) } @@ -50,16 +54,19 @@ func (fs *overlayFS) ReadFile(ctx context.Context, uri uri.URI) (FileHandle, err func (fs *overlayFS) Update(ctx context.Context, changes []*FileChange) error { for _, change := range changes { var base []byte + if change.From == FileChangeTypeDidChange { fh, err := fs.ReadFile(ctx, change.URI) if err != nil { return err } + base, err = fh.Content() if err != nil { return err } } + overlay := NewOverlay(change.URI, change.FullContent(base), int32(change.Version)) slog.Debug("new overlay content", "content", string(overlay.content), "uri", change.URI) @@ -68,6 +75,7 @@ func (fs *overlayFS) Update(ctx context.Context, changes []*FileChange) error { fs.overlays[change.URI] = overlay fs.mu.Unlock() } + return nil } diff --git a/lsp/cache/graph.go b/lsp/cache/graph.go index b64b4d1..9f6cea7 100644 --- a/lsp/cache/graph.go +++ b/lsp/cache/graph.go @@ -21,9 +21,11 @@ func (n *IncludeNode) Clone() *IncludeNode { if len(n.indegree) > 0 { newNode.indegree = make([]uri.URI, len(n.indegree)) } + if len(n.outdegree) > 0 { newNode.outdegree = make([]uri.URI, len(n.outdegree)) } + copy(newNode.indegree, n.indegree) copy(newNode.outdegree, n.outdegree) @@ -52,6 +54,7 @@ func NewIncludeGraph() *IncludeGraph { func (g *IncludeGraph) Get(file uri.URI) *IncludeNode { g.mu.RLock() defer g.mu.RUnlock() + return g.mapper[file] } @@ -60,6 +63,7 @@ func (g *IncludeGraph) Get(file uri.URI) *IncludeNode { func (g *IncludeGraph) Set(file uri.URI, includes []*syntax.Include, resolve func(cur uri.URI, includePath string) uri.URI) { g.mu.Lock() defer g.mu.Unlock() + includeURIs := make([]uri.URI, 0, len(includes)) for _, inc := range includes { if inc.Path == nil { @@ -69,6 +73,7 @@ func (g *IncludeGraph) Set(file uri.URI, includes []*syntax.Include, resolve fun includeURI := resolve(file, strings.Trim(inc.Path.Text, "\"'")) includeURIs = append(includeURIs, includeURI) } + sort.SliceStable(includeURIs, func(i, j int) bool { return includeURIs[i] < includeURIs[j] }) @@ -81,20 +86,25 @@ func (g *IncludeGraph) Set(file uri.URI, includes []*syntax.Include, resolve fun }) equal := true + for i := range includeURIs { if includeURIs[i] != node.outdegree[i] { equal = false + break } } + if equal { return } } + g.removeWithoutLock(file) } else { node = &IncludeNode{} } + for _, inc := range includeURIs { node.outdegree = append(node.outdegree, inc) @@ -103,6 +113,7 @@ func (g *IncludeGraph) Set(file uri.URI, includes []*syntax.Include, resolve fun outNode = &IncludeNode{} g.mapper[inc] = outNode } + outNode.indegree = append(outNode.indegree, file) } @@ -112,6 +123,7 @@ func (g *IncludeGraph) Set(file uri.URI, includes []*syntax.Include, resolve fun func (g *IncludeGraph) Remove(file uri.URI) { g.mu.Lock() defer g.mu.Unlock() + g.removeWithoutLock(file) } @@ -146,9 +158,11 @@ func (g *IncludeGraph) removeWithoutLock(file uri.URI) { if len(outNode.indegree) == 0 { outNode.indegree = nil } + break } } + if len(outNode.indegree) == 0 && len(outNode.outdegree) == 0 { delete(g.mapper, outFile) } diff --git a/lsp/cache/graph_test.go b/lsp/cache/graph_test.go index 1dafb6e..31ce238 100644 --- a/lsp/cache/graph_test.go +++ b/lsp/cache/graph_test.go @@ -24,6 +24,7 @@ func Test_IncludeGraph_Set(t *testing.T) { tmpDir, err := os.MkdirTemp("", "thrift-test") assert.NoError(t, err) + defer func() { _ = os.RemoveAll(tmpDir) }() baseDir := filepath.Join(tmpDir, "base") @@ -85,6 +86,7 @@ func Test_Graph(t *testing.T) { file1 := uri.MustParse("file:///tmp/model/user.thrift") file2 := uri.MustParse("file:///tmp/base.thrift") file3 := uri.MustParse("file:///tmp/addr.thrift") + graph.Set(file1, []*syntax.Include{ {Path: &syntax.Token{Text: "../base.thrift"}}, {Path: &syntax.Token{Text: "../addr.thrift"}}, @@ -119,6 +121,7 @@ func Test_Graph(t *testing.T) { assert.Equal(t, expectNode3, graph.Get("file:///tmp/addr.thrift"), "addr.thrift") graph.Remove(file1) + expectNode1 = nil expectNode2 = &IncludeNode{ indegree: []uri.URI{file3}, @@ -126,6 +129,7 @@ func Test_Graph(t *testing.T) { expectNode3 = &IncludeNode{ outdegree: []uri.URI{file2}, } + assert.Equal(t, expectNode1, graph.Get("file:///tmp/model/user.thrift"), "user.thrift") assert.Equal(t, expectNode2, graph.Get("file:///tmp/base.thrift"), "base.thrift") assert.Equal(t, expectNode3, graph.Get("file:///tmp/addr.thrift"), "addr.thrift") @@ -140,6 +144,7 @@ func Test_Graph(t *testing.T) { // the snapshot's resolver-based resolution. func resolveWithPaths(includePaths []string) func(uri.URI, string) uri.URI { r := resolver.NewWithFS(includePaths, resolver.FS()) + return func(cur uri.URI, includePath string) uri.URI { return uri.File(r.Resolve(cur.Path(), includePath)) } diff --git a/lsp/cache/parse.go b/lsp/cache/parse.go index 6e1c601..22e9224 100644 --- a/lsp/cache/parse.go +++ b/lsp/cache/parse.go @@ -34,6 +34,7 @@ func (c *ParseCaches) Set(filePath uri.URI, res *ParsedFile) { func (c *ParseCaches) Get(filePath uri.URI) *ParsedFile { c.mu.RLock() defer c.mu.RUnlock() + return c.caches[filePath] } @@ -53,6 +54,7 @@ func (c *ParseCaches) Clone() *ParseCaches { for i := range c.caches { clone[i] = c.caches[i] } + return &ParseCaches{caches: clone} } @@ -62,12 +64,15 @@ func (c *ParseCaches) Tokens() map[string]struct{} { } tokens := make(map[string]struct{}) + for _, parsed := range c.caches { if parsed.ast == nil { continue } + collectTokens(parsed.ast, tokens) } + c.tokens = tokens return tokens @@ -80,10 +85,12 @@ func (c *ParseCaches) TokensForFile(file uri.URI, getIncludes func(uri.URI) []ur visited := make(map[uri.URI]bool) var collect func(f uri.URI) + collect = func(f uri.URI) { if visited[f] { return } + visited[f] = true pf := c.Get(f) @@ -97,6 +104,7 @@ func (c *ParseCaches) TokensForFile(file uri.URI, getIncludes func(uri.URI) []ur } collect(file) + return tokens } @@ -104,6 +112,7 @@ func (c *ParseCaches) TokensForFile(file uri.URI, getIncludes func(uri.URI) []ur // names, field names, and identifier references. func collectTokens(ast *syntax.Document, tokens map[string]struct{}) { var walk func(n syntax.Node) + walk = func(n syntax.Node) { switch v := n.(type) { case *syntax.Identifier: @@ -112,31 +121,37 @@ func collectTokens(ast *syntax.Document, tokens map[string]struct{}) { if v.Kind == syntax.ValueIdent { tokens[v.Text] = struct{}{} } + for _, item := range v.List { walk(item) } + for _, entry := range v.Map { walk(entry.Key) walk(entry.Value) } case *syntax.Struct: walk(v.Name) + for _, f := range v.Fields { walk(f) } case *syntax.Service: walk(v.Name) + for _, fn := range v.Functions { walk(fn) } case *syntax.Enum: walk(v.Name) + for _, ev := range v.Values { walk(ev) } case *syntax.Field: walk(v.Type) walk(v.Name) + if v.Value != nil { walk(v.Value) } @@ -144,10 +159,13 @@ func collectTokens(ast *syntax.Document, tokens map[string]struct{}) { if v.Type != nil { walk(v.Type) } + walk(v.Name) + for _, a := range v.Args { walk(a) } + if v.Throws != nil { for _, f := range v.Throws.Fields { walk(f) @@ -157,9 +175,11 @@ func collectTokens(ast *syntax.Document, tokens map[string]struct{}) { if v.Ident != nil { walk(v.Ident) } + if v.KeyType != nil { walk(v.KeyType) } + if v.ValueType != nil { walk(v.ValueType) } @@ -207,6 +227,7 @@ func (p *ParsedFile) AggregatedError() error { if len(p.errs) == 0 { return nil } + return fmt.Errorf("aggregated error: %v", p.errs) } diff --git a/lsp/cache/parse_test.go b/lsp/cache/parse_test.go index 1cfff88..01fa109 100644 --- a/lsp/cache/parse_test.go +++ b/lsp/cache/parse_test.go @@ -10,6 +10,7 @@ func TestParse(t *testing.T) { type args struct { fh FileHandle } + tests := []struct { name string args args diff --git a/lsp/cache/resolver_test.go b/lsp/cache/resolver_test.go index 09d38dd..93ca719 100644 --- a/lsp/cache/resolver_test.go +++ b/lsp/cache/resolver_test.go @@ -15,6 +15,7 @@ import ( func TestResolver(t *testing.T) { tmpDir, err := os.MkdirTemp("", "resolver-test") assert.NoError(t, err) + defer func() { _ = os.RemoveAll(tmpDir) }() baseDir := filepath.Join(tmpDir, "base") @@ -107,7 +108,7 @@ func TestResolver(t *testing.T) { } result := resolver.GetIncludePath(doc, "shared") - assert.Equal(t, "", result) + assert.Empty(t, result) }, }, { diff --git a/lsp/cache/session.go b/lsp/cache/session.go index c3e8de0..cc9e414 100644 --- a/lsp/cache/session.go +++ b/lsp/cache/session.go @@ -44,10 +44,13 @@ func NewSession(cache *Cache) *Session { func (s *Session) Initialize(fn func()) { s.initializedMu.Lock() defer s.initializedMu.Unlock() + if s.initialized { return } + s.initialized = true + fn() } @@ -71,6 +74,7 @@ func (s *Session) ViewOf(fileURI uri.URI) (*View, error) { for i := range s.views { if s.views[i].ContainsFile(fileURI) { s.viewMap[fileURI] = s.views[i] + return s.views[i], nil } } @@ -78,6 +82,7 @@ func (s *Session) ViewOf(fileURI uri.URI) (*View, error) { for i := range s.views { if s.views[i].FileKnown(fileURI) { s.viewMap[fileURI] = s.views[i] + return s.views[i], nil } } diff --git a/lsp/cache/snapshot.go b/lsp/cache/snapshot.go index e157dc1..b0215f6 100644 --- a/lsp/cache/snapshot.go +++ b/lsp/cache/snapshot.go @@ -46,9 +46,11 @@ func (s *snapshotFS) Stat(name string) (fs.FileInfo, error) { if _, ok := s.ss.files.Get(uri.File(name)); ok { return snapshotFileInfo{}, nil } + if s.disk == nil { s.disk = resolver.FS() } + return fs.Stat(s.disk, name) } @@ -58,11 +60,14 @@ func (s *snapshotFS) Open(name string) (fs.File, error) { if err != nil { return nil, err } + return &snapshotFile{Reader: bytes.NewReader(content), info: snapshotFileInfo{}}, nil } + if s.disk == nil { s.disk = resolver.FS() } + return s.disk.Open(name) } @@ -93,6 +98,7 @@ func (r *Resolver) IncludePaths() []string { func (r *Resolver) ResolveInclude(cur uri.URI, includePath string) uri.URI { filePath := cur.Path() resolvedPath := r.central.Resolve(filePath, includePath) + return uri.File(resolvedPath) } @@ -109,12 +115,15 @@ func (r *Resolver) GetIncludePath(ast *syntax.Document, includeName string) stri if include.Path == nil { continue } + path := lsputils.IncludePathText(include) + name := getIncludeNameFromPath(path) if name == includeName { return path } } + return "" } @@ -125,6 +134,7 @@ func (r *Resolver) GetIncludeURI(cur uri.URI, ast *syntax.Document, includeName if path == "" { return "" } + return r.ResolveInclude(cur, path) } @@ -132,6 +142,7 @@ func (r *Resolver) GetIncludeURI(cur uri.URI, ast *syntax.Document, includeName func getIncludeNameFromPath(path string) string { items := strings.Split(path, "/") name := items[len(items)-1] + return strings.TrimSuffix(name, ".thrift") } @@ -176,6 +187,7 @@ func NewSnapshot(view *View, store *memoize.Store, includePaths []string) *Snaps func (s *Snapshot) Acquire() func() { s.refCount.Add(1) + return s.refCount.Done } @@ -201,10 +213,12 @@ func (s *Snapshot) ReadFile(ctx context.Context, uri uri.URI) (FileHandle, error } slog.Debug("snapshot read from fs") + fh, err := s.view.fs.ReadFile(ctx, uri) if err != nil { return nil, err } + s.files.Set(uri, fh) return fh, nil @@ -235,12 +249,14 @@ func (s *Snapshot) Parse(ctx context.Context, uri uri.URI) (*ParsedFile, error) pf, err := Parse(fh) if err != nil { slog.Debug("snapshot parse failed", "err", err) + return nil, err } if pf.AST() != nil { s.graph.Set(uri, pf.AST().Includes(), s.Resolver().ResolveInclude) } + s.parsedCache.Set(uri, pf) return pf, nil @@ -256,6 +272,7 @@ func (s *Snapshot) TokensForFile(file uri.URI) map[string]struct{} { if node == nil { return nil } + return node.OutDegree() }) } diff --git a/lsp/cache/view.go b/lsp/cache/view.go index 82ad1c2..97fbc56 100644 --- a/lsp/cache/view.go +++ b/lsp/cache/view.go @@ -71,7 +71,6 @@ func NewView(name string, folder uri.URI, fs FileSource, store *memoize.Store, i func (v *View) ContainsFile(uri uri.URI) bool { // folder: file:///workdir/ // file: file:///workdir/file.idl - folder := v.folder.Path() file := uri.Path() @@ -113,6 +112,7 @@ func (v *View) FileChange(ctx context.Context, changes []*FileChange, postFns .. // release previous snapshot v.snapshotRelease() v.snapshotMu.Lock() + v.snapshot = newSnapshot for _, change := range changes { v.snapshot.ForgetFile(change.URI) @@ -126,14 +126,17 @@ func (v *View) FileChange(ctx context.Context, changes []*FileChange, postFns .. // TODO(jpf): 异步 parse 和 completion 的顺序问题 // go func() { defer asyncRelease() + uris := make(map[uri.URI]struct{}) for _, change := range changes { uris[change.URI] = struct{}{} } + for uri := range uris { v.snapshotMu.Lock() _, err := v.snapshot.Parse(ctx, uri) v.snapshotMu.Unlock() + if err != nil { slog.Error("parse error", "err", err) } diff --git a/lsp/codejump.go b/lsp/codejump.go index 958670f..6936ac3 100644 --- a/lsp/codejump.go +++ b/lsp/codejump.go @@ -10,10 +10,12 @@ import ( func (s *Server) definition(ctx context.Context, params *protocol.DefinitionParams) (result []protocol.Location, err error) { file := params.TextDocument.URI + view, err := s.session.ViewOf(file) if err != nil { return nil, err } + ss, release := view.Snapshot() defer release() @@ -22,10 +24,12 @@ func (s *Server) definition(ctx context.Context, params *protocol.DefinitionPara func (s *Server) references(ctx context.Context, params *protocol.ReferenceParams) (result []protocol.Location, err error) { file := params.TextDocument.URI + view, err := s.session.ViewOf(file) if err != nil { return nil, err } + ss, release := view.Snapshot() defer release() @@ -34,10 +38,12 @@ func (s *Server) references(ctx context.Context, params *protocol.ReferenceParam func (s *Server) typeDefinition(ctx context.Context, params *protocol.TypeDefinitionParams) (result []protocol.Location, err error) { file := params.TextDocument.URI + view, err := s.session.ViewOf(file) if err != nil { return nil, err } + ss, release := view.Snapshot() defer release() diff --git a/lsp/codejump/cross_project_test.go b/lsp/codejump/cross_project_test.go index 90c83ee..e2fefea 100644 --- a/lsp/codejump/cross_project_test.go +++ b/lsp/codejump/cross_project_test.go @@ -47,16 +47,20 @@ func TestDefinitionCrossProjectInclude(t *testing.T) { if len(offset) > 0 { idx += offset[0] } + lineStart := idx for lineStart > 0 && appContent[lineStart-1] != '\n' { lineStart-- } + line := 0 + for i := 0; i < lineStart; i++ { if appContent[i] == '\n' { line++ } } + return protocol.Position{Line: uint32(line), Character: uint32(idx - lineStart)} } diff --git a/lsp/codejump/definition.go b/lsp/codejump/definition.go index ef77d11..0dc67da 100644 --- a/lsp/codejump/definition.go +++ b/lsp/codejump/definition.go @@ -16,6 +16,7 @@ import ( // a type reference, a constant value identifier, or a service reference. func Definition(ctx context.Context, ss *cache.Snapshot, file uri.URI, pos protocol.Position) (res []protocol.Location, err error) { res = make([]protocol.Location, 0) + pf, target, err := resolveTarget(ctx, ss, file, pos) if err != nil { return res, err @@ -29,6 +30,7 @@ func Definition(ctx context.Context, ss *cache.Snapshot, file uri.URI, pos proto case TargetService: return serviceDefinition(ctx, ss, file, pf, target) } + return res, err } @@ -50,19 +52,24 @@ func FindTypeDefinition(ctx context.Context, ss *cache.Snapshot, file uri.URI, a if dstException := GetExceptionNode(dstAst, identifier); dstException != nil { return astFile, dstException.Name, DefinitionException, nil } + if dstStruct := GetStructNode(dstAst, identifier); dstStruct != nil { return astFile, dstStruct.Name, DefinitionStruct, nil } + if dstEnum := GetEnumNode(dstAst, identifier); dstEnum != nil { return astFile, dstEnum.Name, DefinitionEnum, nil } + if dstUnion := GetUnionNode(dstAst, identifier); dstUnion != nil { return astFile, dstUnion.Name, DefinitionUnion, nil } + if dstTypedef := GetTypedefNode(dstAst, identifier); dstTypedef != nil { return astFile, dstTypedef.Name, DefinitionTypedef, nil } } + return file, nil, DefinitionNone, nil } @@ -72,6 +79,7 @@ func FindConstValueDefinition(ctx context.Context, ss *cache.Snapshot, file uri. if value == nil || value.Kind != syntax.ValueIdent { return "", nil, nil } + name := value.Text if name == "true" || name == "false" { return "", nil, nil @@ -87,10 +95,12 @@ func FindConstValueDefinition(ctx context.Context, ss *cache.Snapshot, file uri. if dstEnumValue := GetEnumValueIdentifierNode(dstAst, identifier); dstEnumValue != nil { return astFile, dstEnumValue, nil } + if constIdentifier := GetConstIdentifierNode(dstAst, identifier); constIdentifier != nil { return astFile, constIdentifier, nil } } + return file, nil, nil } @@ -112,6 +122,7 @@ func FindServiceDefinition(ctx context.Context, ss *cache.Snapshot, file uri.URI return astFile, dstService.Name, nil } } + return file, nil, nil } @@ -123,28 +134,35 @@ func parseDefinitionFile(ctx context.Context, ss *cache.Snapshot, file uri.URI) if err != nil { return nil, err } + if len(pf.Errors()) > 0 { slog.Error("parse error", "errs", pf.Errors()) } + if pf.AST() == nil { return nil, errNoAST } + return pf.AST(), nil } func typeNameDefinition(ctx context.Context, ss *cache.Snapshot, file uri.URI, pf *cache.ParsedFile, target *target) ([]protocol.Location, error) { ft := target.parent.(*syntax.FieldType) + astFile, id, _, err := FindTypeDefinition(ctx, ss, file, pf.AST(), ft) if err != nil { return nil, err } + if id == nil { return nil, nil } + loc, err := jumpInFile(ctx, ss, astFile, id) if err != nil { return nil, err } + return []protocol.Location{loc}, nil } @@ -153,13 +171,16 @@ func constValueDefinition(ctx context.Context, ss *cache.Snapshot, file uri.URI, if err != nil { return nil, err } + if id == nil { return nil, nil } + loc, err := jumpInFile(ctx, ss, astFile, id) if err != nil { return nil, err } + return []protocol.Location{loc}, nil } @@ -168,12 +189,15 @@ func serviceDefinition(ctx context.Context, ss *cache.Snapshot, file uri.URI, pf if err != nil { return nil, err } + if id == nil { return nil, nil } + loc, err := jumpInFile(ctx, ss, astFile, id) if err != nil { return nil, err } + return []protocol.Location{loc}, nil } diff --git a/lsp/codejump/definition_list_test.go b/lsp/codejump/definition_list_test.go index 8078da3..f39b071 100644 --- a/lsp/codejump/definition_list_test.go +++ b/lsp/codejump/definition_list_test.go @@ -38,6 +38,7 @@ const list my_list = [MyEnum.Value1, MyEnum.Value2]` file uri.URI pos protocol.Position } + tests := []struct { name string args args diff --git a/lsp/codejump/definition_test.go b/lsp/codejump/definition_test.go index 14334df..b0850bb 100644 --- a/lsp/codejump/definition_test.go +++ b/lsp/codejump/definition_test.go @@ -91,6 +91,7 @@ struct Person { file uri.URI pos protocol.Position } + tests := []struct { name string args args diff --git a/lsp/codejump/hits.go b/lsp/codejump/hits.go index 6017c4d..e6f8de7 100644 --- a/lsp/codejump/hits.go +++ b/lsp/codejump/hits.go @@ -16,5 +16,6 @@ func hits(hits []referenceHit) []protocol.Location { for _, h := range hits { out = append(out, h.loc) } + return out } diff --git a/lsp/codejump/hover.go b/lsp/codejump/hover.go index 88134fe..6075f8a 100644 --- a/lsp/codejump/hover.go +++ b/lsp/codejump/hover.go @@ -27,6 +27,7 @@ func Hover(ctx context.Context, ss *cache.Snapshot, file uri.URI, pos protocol.P case TargetService: return hoverService(ctx, ss, file, pf, target) } + return res, err } @@ -40,29 +41,35 @@ func hoverService(ctx context.Context, ss *cache.Snapshot, file uri.URI, pf *cac if err != nil || id == nil { return "", err } + dstAst, err := parseDefinitionFile(ctx, ss, astFile) if err != nil { return "", err } + svc := GetServiceNode(dstAst, id.Text) if svc == nil { return "", nil } + return formatNode(dstAst, svc) } func hoverDefinition(ctx context.Context, ss *cache.Snapshot, file uri.URI, pf *cache.ParsedFile, target *target) (string, error) { ft := target.parent.(*syntax.FieldType) + astFile, id, kind, err := FindTypeDefinition(ctx, ss, file, pf.AST(), ft) if err != nil || id == nil { return "", err } + dstAst, err := parseDefinitionFile(ctx, ss, astFile) if err != nil { return "", err } var node syntax.Node + switch kind { case DefinitionException: node = GetExceptionNode(dstAst, id.Text) @@ -75,9 +82,11 @@ func hoverDefinition(ctx context.Context, ss *cache.Snapshot, file uri.URI, pf * case DefinitionTypedef: node = GetTypedefNode(dstAst, id.Text) } + if node == nil { return "", nil } + return formatNode(dstAst, node) } @@ -86,6 +95,7 @@ func hoverConstValue(ctx context.Context, ss *cache.Snapshot, file uri.URI, pf * if err != nil || id == nil { return "", err } + dstAst, err := parseDefinitionFile(ctx, ss, astFile) if err != nil { return "", err @@ -94,8 +104,10 @@ func hoverConstValue(ctx context.Context, ss *cache.Snapshot, file uri.URI, pf * if dstEnum := GetEnumNodeByEnumValue(dstAst, id.Text); dstEnum != nil { return formatNode(dstAst, dstEnum) } + if dstConst := GetConstNode(dstAst, id.Text); dstConst != nil { return formatNode(dstAst, dstConst) } + return "", nil } diff --git a/lsp/codejump/hover_test.go b/lsp/codejump/hover_test.go index baf481f..e1b1030 100644 --- a/lsp/codejump/hover_test.go +++ b/lsp/codejump/hover_test.go @@ -41,10 +41,13 @@ const Color defaultColor = GREEN` if len(offset) > 0 { idx += offset[0] } + if idx < 0 { t.Fatalf("marker %q not found", marker) } + line, col := 0, 0 + for i := 0; i < idx; i++ { if file[i] == '\n' { line++ @@ -53,6 +56,7 @@ const Color defaultColor = GREEN` col++ } } + return protocol.Position{Line: uint32(line), Character: uint32(col)} } @@ -76,6 +80,7 @@ const Color defaultColor = GREEN` if err != nil { t.Fatalf("Hover: %v", err) } + if !strings.Contains(got, tt.want) { t.Errorf("hover text %q does not contain %q", got, tt.want) } @@ -87,10 +92,12 @@ func TestHoverUnresolvable(t *testing.T) { ss := cache.BuildSnapshotForTest([]*cache.FileChange{ {URI: "file:///tmp/test.thrift", Version: 0, Content: []byte("struct S {\n 1: Missing x\n}"), From: cache.FileChangeTypeDidOpen}, }) + got, err := Hover(t.Context(), ss, "file:///tmp/test.thrift", protocol.Position{Line: 1, Character: 8}) if err != nil { t.Fatalf("Hover: %v", err) } + if got != "" { t.Errorf("expected empty hover for undefined type, got %q", got) } diff --git a/lsp/codejump/reference.go b/lsp/codejump/reference.go index 867b460..8a67c69 100644 --- a/lsp/codejump/reference.go +++ b/lsp/codejump/reference.go @@ -26,6 +26,7 @@ var validReferenceDefinitionType = map[DefinitionKind]struct{}{ // the cursor: type definitions, constant values, enum values, and services. func Reference(ctx context.Context, ss *cache.Snapshot, file uri.URI, pos protocol.Position) (res []protocol.Location, err error) { res = make([]protocol.Location, 0) + pf, target, err := resolveTarget(ctx, ss, file, pos) if err != nil { return res, err @@ -37,26 +38,31 @@ func Reference(ctx context.Context, ss *cache.Snapshot, file uri.URI, pos protoc if err != nil { return nil, err } + return hits(refs), nil case TargetConstValue: refs, err := searchConstValueReferences(ctx, ss, file, pf, target) if err != nil { return nil, err } + return hits(refs), nil case TargetService: refs, err := searchServiceReferences(ctx, ss, file, target.identifier().Text) if err != nil { return nil, err } + return hits(refs), nil case TargetDefinition: refs, err := searchDefinitionReferences(ctx, ss, file, pf, target) if err != nil { return nil, err } + return hits(refs), nil } + return res, err } @@ -73,18 +79,22 @@ func searchDefinitionReferences(ctx context.Context, ss *cache.Snapshot, file ur switch parent.(type) { case *syntax.Const: typeName := fmt.Sprintf("%s.%s", lsputils.GetIncludeName(file), id.Text) + return searchConstValueIdentifierReferences(ctx, ss, file, typeName) case *syntax.EnumValue: enum, ok := grandparent(target.path).(*syntax.Enum) if !ok { return res, err } + typeName := fmt.Sprintf("%s.%s.%s", lsputils.GetIncludeName(file), enum.Name.Text, id.Text) + return searchConstValueIdentifierReferences(ctx, ss, file, typeName) case *syntax.Service: svcName := id.Text if strings.Contains(svcName, ".") { include, _ := lsputils.ParseIdent(file, pf.AST().Includes(), svcName) + resolver := ss.Resolver() if path := resolver.GetIncludePath(pf.AST(), include); path != "" { file = resolver.ResolveInclude(file, path) @@ -92,6 +102,7 @@ func searchDefinitionReferences(ctx context.Context, ss *cache.Snapshot, file ur } else { svcName = fmt.Sprintf("%s.%s", lsputils.GetIncludeName(file), svcName) } + return searchServiceReferences(ctx, ss, file, svcName) } @@ -99,10 +110,13 @@ func searchDefinitionReferences(ctx context.Context, ss *cache.Snapshot, file ur if !ok { return res, err } + if _, ok := validReferenceDefinitionType[kind]; !ok { return res, err } + typeName := fmt.Sprintf("%s.%s", lsputils.GetIncludeName(file), id.Text) + return searchIdentifierReferences(ctx, ss, file, typeName, kind) } @@ -110,6 +124,7 @@ func grandparent(path []syntax.Node) syntax.Node { if len(path) < 3 { return nil } + return path[len(path)-3] } @@ -134,12 +149,14 @@ func definitionKindOf(n syntax.Node) (DefinitionKind, bool) { case *syntax.Service: return DefinitionService, true } + return DefinitionNone, false } func searchTypeNameReferences(ctx context.Context, ss *cache.Snapshot, file uri.URI, pf *cache.ParsedFile, target *target) (res []referenceHit, err error) { res = make([]referenceHit, 0) ft := target.parent.(*syntax.FieldType) + typeName := typeReferenceName(ft) if typeName == "" || IsBasicType(typeName) { return res, err @@ -150,13 +167,16 @@ func searchTypeNameReferences(ctx context.Context, ss *cache.Snapshot, file uri. if err != nil { return res, err } + if identifierNode == nil { return res, err } + loc, err := jumpInFile(ctx, ss, definitionFile, identifierNode) if err != nil { return res, err } + res = append(res, referenceHit{loc: loc, text: identifierNode.Text}) // Search usages of the type name. @@ -164,7 +184,9 @@ func searchTypeNameReferences(ctx context.Context, ss *cache.Snapshot, file uri. if err != nil { return res, err } + res = append(res, locations...) + return res, err } @@ -175,6 +197,7 @@ func searchServiceReferences(ctx context.Context, ss *cache.Snapshot, file uri.U if err != nil { return nil, err } + res = append(res, locations...) for _, referenceFile := range referenceFiles(ss, file) { @@ -182,8 +205,10 @@ func searchServiceReferences(ctx context.Context, ss *cache.Snapshot, file uri.U if err != nil { return nil, err } + res = append(res, locations...) } + return res, err } @@ -194,9 +219,11 @@ func referenceFiles(ss *cache.Snapshot, file uri.URI) []uri.URI { if includeNode == nil { return nil } + if len(includeNode.InDegree()) == 0 && len(includeNode.OutDegree()) == 0 { ss.Graph().Debug() } + return includeNode.InDegree() } @@ -205,6 +232,7 @@ func searchServiceDefinitionReferences(ctx context.Context, ss *cache.Snapshot, if err != nil { return res, err } + if pf.AST() == nil { return res, err } @@ -213,8 +241,10 @@ func searchServiceDefinitionReferences(ctx context.Context, ss *cache.Snapshot, if svc.Extends == nil || svc.Extends.Text != svcName { continue } + res = append(res, referenceHit{loc: jump(file, pf.AST(), svc.Extends), text: svc.Extends.Text}) } + return res, err } @@ -226,6 +256,7 @@ func searchIdentifierReferences(ctx context.Context, ss *cache.Snapshot, file ur if err != nil { return nil, err } + res = append(res, locations...) for _, referenceFile := range referenceFiles(ss, file) { @@ -233,8 +264,10 @@ func searchIdentifierReferences(ctx context.Context, ss *cache.Snapshot, file ur if err != nil { return nil, err } + res = append(res, locations...) } + return res, err } @@ -246,6 +279,7 @@ func searchDefinitionIdentifierReferences(ctx context.Context, ss *cache.Snapsho if err != nil { return res, err } + if pf.AST() == nil { return res, err } @@ -254,19 +288,25 @@ func searchDefinitionIdentifierReferences(ctx context.Context, ss *cache.Snapsho if ft == nil || typeReferenceName(ft) != typeName { return } + res = append(res, referenceHit{loc: jump(file, pf.AST(), ft.Ident), text: ft.Ident.Text}) } + var searchFieldType func(ft *syntax.FieldType) + searchFieldType = func(ft *syntax.FieldType) { if ft == nil { return } + if ft.KeyType != nil { searchFieldType(ft.KeyType) } + if ft.ValueType != nil { searchFieldType(ft.ValueType) } + jumpFieldType(ft) } jumpField := func(field *syntax.Field) { @@ -282,11 +322,13 @@ func searchDefinitionIdentifierReferences(ctx context.Context, ss *cache.Snapsho for _, fn := range svc.Functions { searchFieldType(fn.Type) processStructLike(fn.Args) + if fn.Throws != nil { processStructLike(fn.Throws.Fields) } } } + if definitionType == DefinitionException { return res, err } @@ -294,18 +336,23 @@ func searchDefinitionIdentifierReferences(ctx context.Context, ss *cache.Snapsho for _, st := range pf.AST().Structs() { processStructLike(st.Fields) } + for _, st := range pf.AST().Unions() { processStructLike(st.Fields) } + for _, st := range pf.AST().Exceptions() { processStructLike(st.Fields) } + for _, typedef := range pf.AST().Typedefs() { searchFieldType(typedef.Type) } + for _, cst := range pf.AST().Consts() { searchFieldType(cst.Type) } + return res, err } @@ -317,20 +364,25 @@ func searchConstValueReferences(ctx context.Context, ss *cache.Snapshot, file ur if err != nil { return res, err } + if identifierNode == nil { return res, err } + loc, err := jumpInFile(ctx, ss, definitionFile, identifierNode) if err != nil { return res, err } + res = append(res, referenceHit{loc: loc, text: identifierNode.Text}) locations, err := searchConstValueIdentifierReferences(ctx, ss, definitionFile, value.Text) if err != nil { return res, err } + res = append(res, locations...) + return res, err } @@ -341,6 +393,7 @@ func searchConstValueIdentifierReferences(ctx context.Context, ss *cache.Snapsho if err != nil { return nil, err } + res = append(res, locations...) for _, referenceFile := range referenceFiles(ss, file) { @@ -348,8 +401,10 @@ func searchConstValueIdentifierReferences(ctx context.Context, ss *cache.Snapsho if err != nil { return nil, err } + res = append(res, locations...) } + return res, err } @@ -358,6 +413,7 @@ func searchConstValueIdentifierReference(ctx context.Context, ss *cache.Snapshot if err != nil { return res, err } + if pf.AST() == nil { return res, err } @@ -376,22 +432,28 @@ func searchConstValueIdentifierReference(ctx context.Context, ss *cache.Snapshot for _, st := range pf.AST().Structs() { processStructLike(st.Fields) } + for _, st := range pf.AST().Unions() { processStructLike(st.Fields) } + for _, st := range pf.AST().Exceptions() { processStructLike(st.Fields) } + for _, cst := range pf.AST().Consts() { jumpValue(cst.Value) } + for _, svc := range pf.AST().Services() { for _, fn := range svc.Functions { processStructLike(fn.Args) + if fn.Throws != nil { processStructLike(fn.Throws.Fields) } } } + return res, err } diff --git a/lsp/codejump/reference_test.go b/lsp/codejump/reference_test.go index 01a4a22..88b0f3f 100644 --- a/lsp/codejump/reference_test.go +++ b/lsp/codejump/reference_test.go @@ -73,6 +73,7 @@ const UserKind kind = "1" file uri.URI pos protocol.Position } + tests := []struct { name string args args diff --git a/lsp/codejump/rename.go b/lsp/codejump/rename.go index 7d89737..f6be1d1 100644 --- a/lsp/codejump/rename.go +++ b/lsp/codejump/rename.go @@ -24,8 +24,10 @@ func PrepareRename(ctx context.Context, ss *cache.Snapshot, file uri.URI, pos pr switch target.kind { case TargetDefinition, TargetConstValue, TargetService: rg := nodeRange(pf.AST(), target.node) + return &rg, nil } + return nil, fmt.Errorf("rename not supported at this position") } @@ -39,12 +41,14 @@ func Rename(ctx context.Context, ss *cache.Snapshot, file uri.URI, pos protocol. } var refs []referenceHit + switch target.kind { case TargetTypeName: ft := target.parent.(*syntax.FieldType) if typeReferenceName(ft) == "" || IsBasicType(typeReferenceName(ft)) { return nil, fmt.Errorf("rename not supported for basic types") } + refs, err = searchTypeNameReferences(ctx, ss, file, pf, target) if err != nil { return nil, err @@ -57,6 +61,7 @@ func Rename(ctx context.Context, ss *cache.Snapshot, file uri.URI, pos protocol. } else if id == nil { return nil, fmt.Errorf("definition not found") } + refs, err = searchConstValueReferences(ctx, ss, file, pf, target) if err != nil { return nil, err @@ -68,11 +73,13 @@ func Rename(ctx context.Context, ss *cache.Snapshot, file uri.URI, pos protocol. svcName = fmt.Sprintf("%s.%s", lsputils.GetIncludeName(file), svcName) } else { include, _ := lsputils.ParseIdent(file, pf.AST().Includes(), svcName) + resolver := ss.Resolver() if path := resolver.GetIncludePath(pf.AST(), include); path != "" { file = resolver.ResolveInclude(file, path) } } + refs, err = searchServiceReferences(ctx, ss, file, svcName) if err != nil { return nil, err @@ -105,11 +112,13 @@ func Rename(ctx context.Context, ss *cache.Snapshot, file uri.URI, pos protocol. // text becomes user.newtext. func convertHitsToWorkspaceEdit(refs []referenceHit, newName string) *protocol.WorkspaceEdit { changes := make(map[uri.URI][]protocol.TextEdit) + for i := range refs { text := newName if dot := strings.LastIndexByte(refs[i].text, '.'); dot >= 0 { text = refs[i].text[:dot+1] + newName } + changes[refs[i].loc.URI] = append(changes[refs[i].loc.URI], protocol.TextEdit{ Range: refs[i].loc.Range, NewText: text, diff --git a/lsp/codejump/rename_test.go b/lsp/codejump/rename_test.go index 7fb4b5c..76dae90 100644 --- a/lsp/codejump/rename_test.go +++ b/lsp/codejump/rename_test.go @@ -73,6 +73,7 @@ const UserKind kind = "1" file uri.URI pos protocol.Position } + tests := []struct { name string args args @@ -290,6 +291,7 @@ const UserKind kind = "1" pos protocol.Position newText string } + tests := []struct { name string args args diff --git a/lsp/codejump/target.go b/lsp/codejump/target.go index 21ac510..6d43ac3 100644 --- a/lsp/codejump/target.go +++ b/lsp/codejump/target.go @@ -46,6 +46,7 @@ func resolveTarget(ctx context.Context, ss *cache.Snapshot, file uri.URI, pos pr if err != nil { return nil, nil, err } + if pf.AST() == nil { return nil, nil, errNoAST } @@ -67,7 +68,9 @@ func resolveTarget(ctx context.Context, ss *cache.Snapshot, file uri.URI, pos pr if len(path) > 1 { t.parent = path[len(path)-2] } + t.kind = classify(t) + return pf, t, nil } @@ -84,12 +87,14 @@ func classify(t *target) TargetKind { case *syntax.Service: return TargetService } + return TargetDefinition case *syntax.ConstValue: if n.Kind == syntax.ValueIdent { return TargetConstValue } } + return TargetNone } @@ -99,6 +104,7 @@ func (t *target) identifier() *syntax.Identifier { if id, ok := t.node.(*syntax.Identifier); ok { return id } + return nil } @@ -119,15 +125,18 @@ func jumpInFile(ctx context.Context, ss *cache.Snapshot, file uri.URI, node synt if err != nil { return protocol.Location{}, err } + if pf.AST() == nil { return protocol.Location{}, errNoAST } + return jump(file, pf.AST(), node), nil } // nodeRange converts a node span to an LSP range. func nodeRange(doc *syntax.Document, node syntax.Node) protocol.Range { start, end := doc.Range(node) + return protocol.Range{ Start: protocol.Position{ Line: uint32(start.Line - 1), diff --git a/lsp/codejump/target_test.go b/lsp/codejump/target_test.go index 6645f65..e8ad4b2 100644 --- a/lsp/codejump/target_test.go +++ b/lsp/codejump/target_test.go @@ -41,11 +41,13 @@ service Svc extends Base { if len(offset) > 0 { idx += offset[0] } + if idx < 0 { t.Fatalf("marker %q not found", marker) } // Convert byte offset to line/character. line, col := 0, 0 + for i := 0; i < idx; i++ { if src[i] == '\n' { line++ @@ -54,6 +56,7 @@ service Svc extends Base { col++ } } + return protocol.Position{Line: uint32(line), Character: uint32(col)} } @@ -80,6 +83,7 @@ service Svc extends Base { if err != nil { t.Fatalf("resolveTarget: %v", err) } + if target.kind != tt.want { t.Errorf("kind = %v, want %v (node %T)", target.kind, tt.want, target.node) } @@ -102,6 +106,7 @@ func TestResolveTargetNoNode(t *testing.T) { if err != nil { t.Fatalf("resolveTarget: %v", err) } + if target.kind != TargetNone { t.Errorf("kind = %v, want TargetNone", target.kind) } @@ -113,5 +118,6 @@ func indexOf(s, sub string) int { return i } } + return -1 } diff --git a/lsp/codejump/type_definition.go b/lsp/codejump/type_definition.go index e54d5d2..f32de48 100644 --- a/lsp/codejump/type_definition.go +++ b/lsp/codejump/type_definition.go @@ -16,6 +16,7 @@ import ( // type. func TypeDefinition(ctx context.Context, ss *cache.Snapshot, file uri.URI, pos protocol.Position) (res []protocol.Location, err error) { res = make([]protocol.Location, 0) + pf, target, err := resolveTarget(ctx, ss, file, pos) if err != nil { return res, err @@ -31,17 +32,21 @@ func TypeDefinition(ctx context.Context, ss *cache.Snapshot, file uri.URI, pos p if err != nil { return nil, err } + if id == nil { return nil, nil } + loc, err := jumpInFile(ctx, ss, astFile, id) if err != nil { return nil, err } + return []protocol.Location{loc}, nil case TargetDefinition: return declarationTypeDefinition(ctx, ss, file, pf, target) } + return res, err } @@ -49,6 +54,7 @@ func TypeDefinition(ctx context.Context, ss *cache.Snapshot, file uri.URI, pos p // a field, typedef, function, or const under the cursor. func declarationTypeDefinition(ctx context.Context, ss *cache.Snapshot, file uri.URI, pf *cache.ParsedFile, target *target) ([]protocol.Location, error) { var ft *syntax.FieldType + switch parent := target.parent.(type) { case *syntax.Field: ft = parent.Type @@ -59,6 +65,7 @@ func declarationTypeDefinition(ctx context.Context, ss *cache.Snapshot, file uri case *syntax.Const: ft = parent.Type } + if ft == nil { return nil, nil } @@ -67,12 +74,15 @@ func declarationTypeDefinition(ctx context.Context, ss *cache.Snapshot, file uri if err != nil { return nil, err } + if id == nil { return nil, nil } + loc, err := jumpInFile(ctx, ss, astFile, id) if err != nil { return nil, err } + return []protocol.Location{loc}, nil } diff --git a/lsp/codejump/type_definition_test.go b/lsp/codejump/type_definition_test.go index d8b5498..f9c03da 100644 --- a/lsp/codejump/type_definition_test.go +++ b/lsp/codejump/type_definition_test.go @@ -72,6 +72,7 @@ const user.UserType usermale = "male" file uri.URI pos protocol.Position } + tests := []struct { name string args args diff --git a/lsp/codejump/utils.go b/lsp/codejump/utils.go index cd27cf4..c431b51 100644 --- a/lsp/codejump/utils.go +++ b/lsp/codejump/utils.go @@ -34,21 +34,25 @@ func definitionFiles(ctx context.Context, ss *cache.Snapshot, file uri.URI, ast include, _ := lsputils.ParseIdent(file, ast.Includes(), name) if include != "" { resolver := ss.Resolver() + path := resolver.GetIncludePath(ast, include) if path == "" { // Doesn't match any include path; treat as local. return []uri.URI{file} } + return []uri.URI{resolver.ResolveInclude(file, path)} } files := []uri.URI{file} resolver := ss.Resolver() + for _, inc := range ast.Includes() { if path := lsputils.IncludePathText(inc); path != "" { files = append(files, resolver.ResolveInclude(file, path)) } } + return files } @@ -57,11 +61,13 @@ func GetExceptionNode(ast *syntax.Document, name string) *syntax.Struct { if ast == nil { return nil } + for _, excep := range ast.Exceptions() { if excep.Name != nil && excep.Name.Text == name { return excep } } + return nil } @@ -70,11 +76,13 @@ func GetStructNode(ast *syntax.Document, name string) *syntax.Struct { if ast == nil { return nil } + for _, st := range ast.Structs() { if st.Name != nil && st.Name.Text == name { return st } } + return nil } @@ -83,11 +91,13 @@ func GetUnionNode(ast *syntax.Document, name string) *syntax.Struct { if ast == nil { return nil } + for _, st := range ast.Unions() { if st.Name != nil && st.Name.Text == name { return st } } + return nil } @@ -96,11 +106,13 @@ func GetEnumNode(ast *syntax.Document, name string) *syntax.Enum { if ast == nil { return nil } + for _, st := range ast.Enums() { if st.Name != nil && st.Name.Text == name { return st } } + return nil } @@ -110,10 +122,12 @@ func GetEnumNodeByEnumValue(ast *syntax.Document, enumValueName string) *syntax. if ast == nil { return nil } + enumName, _, found := strings.Cut(enumValueName, ".") if found { return GetEnumNode(ast, enumName) } + for _, enum := range ast.Enums() { for _, value := range enum.Values { if value.Name != nil && value.Name.Text == enumValueName { @@ -121,6 +135,7 @@ func GetEnumNodeByEnumValue(ast *syntax.Document, enumValueName string) *syntax. } } } + return nil } @@ -130,6 +145,7 @@ func GetEnumValueIdentifierNode(ast *syntax.Document, name string) *syntax.Ident if ast == nil { return nil } + enumName, identifier, found := strings.Cut(name, ".") if !found { // Bare name: search all enum values. @@ -140,18 +156,22 @@ func GetEnumValueIdentifierNode(ast *syntax.Document, name string) *syntax.Ident } } } + return nil } + for _, enum := range ast.Enums() { if enum.Name == nil || enum.Name.Text != enumName { continue } + for _, enumValue := range enum.Values { if enumValue.Name != nil && enumValue.Name.Text == identifier { return enumValue.Name } } } + return nil } @@ -160,11 +180,13 @@ func GetConstNode(ast *syntax.Document, name string) *syntax.Const { if ast == nil { return nil } + for _, cst := range ast.Consts() { if cst.Name != nil && cst.Name.Text == name { return cst } } + return nil } @@ -174,11 +196,13 @@ func GetConstIdentifierNode(ast *syntax.Document, name string) *syntax.Identifie if ast == nil { return nil } + for _, cst := range ast.Consts() { if cst.Name != nil && cst.Name.Text == name { return cst.Name } } + return nil } @@ -187,11 +211,13 @@ func GetTypedefNode(ast *syntax.Document, name string) *syntax.Typedef { if ast == nil { return nil } + for _, td := range ast.Typedefs() { if td.Name != nil && td.Name.Text == name { return td } } + return nil } @@ -200,11 +226,13 @@ func GetServiceNode(ast *syntax.Document, name string) *syntax.Service { if ast == nil { return nil } + for _, svc := range ast.Services() { if svc.Name != nil && svc.Name.Text == name { return svc } } + return nil } @@ -234,12 +262,14 @@ var containerType = map[string]struct{}{ // IsBasicType reports whether t is a built-in base type. func IsBasicType(t string) bool { _, ok := basicType[t] + return ok } // IsContainerType reports whether t is a container keyword. func IsContainerType(t string) bool { _, ok := containerType[t] + return ok } @@ -249,8 +279,10 @@ func typeReferenceName(ft *syntax.FieldType) string { if ft == nil { return "" } + if ft.Kind == syntax.TypeIdent && ft.Ident != nil { return ft.Ident.Text } + return "" } diff --git a/lsp/completion/completion_test.go b/lsp/completion/completion_test.go index 555c4f3..b7a9961 100644 --- a/lsp/completion/completion_test.go +++ b/lsp/completion/completion_test.go @@ -15,11 +15,13 @@ import ( // paths. func buildSnapshot(t *testing.T, includePaths []string, files ...*cache.FileChange) *cache.Snapshot { t.Helper() + store := &memoize.Store{} c := cache.New(store, nil) fs := cache.NewOverlayFS(c) _ = fs.Update(t.Context(), files) view := cache.NewView("test", uri.File("/tmp"), fs, store, includePaths) + return cache.NewSnapshot(view, store, includePaths) } @@ -80,13 +82,16 @@ const i32 LIMIT = 10` assert.NoError(t, err) cands := semanticCandidates(t.Context(), ss, uri.URI(tt.file), pf, pos) + got := make(map[string]bool) for _, c := range cands { got[c.showText] = true } + for _, w := range tt.want { assert.True(t, got[w], "missing candidate %q in %v", w, got) } + for _, nw := range tt.notWant { assert.False(t, got[nw], "unexpected candidate %q in %v", nw, got) } diff --git a/lsp/completion/semantic_completion.go b/lsp/completion/semantic_completion.go index 3aa8c27..a611a39 100644 --- a/lsp/completion/semantic_completion.go +++ b/lsp/completion/semantic_completion.go @@ -20,6 +20,7 @@ func semanticCandidates(ctx context.Context, ss *cache.Snapshot, file uri.URI, p if len(path) == 0 { return nil } + target := path[len(path)-1] switch n := target.(type) { @@ -34,6 +35,7 @@ func semanticCandidates(ctx context.Context, ss *cache.Snapshot, file uri.URI, p if len(path) < 2 { return nil } + switch parent := path[len(path)-2].(type) { case *syntax.FieldType: return typeCandidates(ctx, ss, file, parsedFile.AST()) @@ -53,6 +55,7 @@ func semanticCandidates(ctx context.Context, ss *cache.Snapshot, file uri.URI, p if n.Name != nil && pos.Offset < tokenOffset(parsedFile.AST(), n.Name) { return typeCandidates(ctx, ss, file, parsedFile.AST()) } + if n.Value == nil { return valueCandidates(ctx, ss, file, parsedFile.AST()) } @@ -69,6 +72,7 @@ func semanticCandidates(ctx context.Context, ss *cache.Snapshot, file uri.URI, p return typeCandidates(ctx, ss, file, parsedFile.AST()) } } + return nil } @@ -85,24 +89,30 @@ func typeCandidates(ctx context.Context, ss *cache.Snapshot, file uri.URI, doc * for _, st := range ast.Structs() { names[st.Name.Text] = struct{}{} } + for _, st := range ast.Unions() { names[st.Name.Text] = struct{}{} } + for _, st := range ast.Exceptions() { names[st.Name.Text] = struct{}{} } + for _, enum := range ast.Enums() { names[enum.Name.Text] = struct{}{} } + for _, td := range ast.Typedefs() { names[td.Name.Text] = struct{}{} } + for _, svc := range ast.Services() { names[svc.Name.Text] = struct{}{} } } collectTypeNames(doc) + for _, inc := range includedFiles(ss, file) { if pf, err := ss.Parse(ctx, inc); err == nil && pf.AST() != nil { collectTypeNames(pf.AST()) @@ -117,7 +127,9 @@ func typeCandidates(ctx context.Context, ss *cache.Snapshot, file uri.URI, doc * format: protocol.InsertTextFormatPlainText, }) } + sort.Slice(res, func(i, j int) bool { return res[i].showText < res[j].showText }) + return res } @@ -129,6 +141,7 @@ func valueCandidates(ctx context.Context, ss *cache.Snapshot, file uri.URI, doc for _, cst := range ast.Consts() { names[cst.Name.Text] = struct{}{} } + for _, enum := range ast.Enums() { names[enum.Name.Text] = struct{}{} for _, value := range enum.Values { @@ -139,6 +152,7 @@ func valueCandidates(ctx context.Context, ss *cache.Snapshot, file uri.URI, doc } collectValueNames(doc) + for _, inc := range includedFiles(ss, file) { if pf, err := ss.Parse(ctx, inc); err == nil && pf.AST() != nil { collectValueNames(pf.AST()) @@ -153,7 +167,9 @@ func valueCandidates(ctx context.Context, ss *cache.Snapshot, file uri.URI, doc format: protocol.InsertTextFormatPlainText, }) } + sort.Slice(res, func(i, j int) bool { return res[i].showText < res[j].showText }) + return res } @@ -161,22 +177,29 @@ func valueCandidates(ctx context.Context, ss *cache.Snapshot, file uri.URI, doc // include graph. func includedFiles(ss *cache.Snapshot, file uri.URI) []uri.URI { var out []uri.URI + visited := make(map[uri.URI]bool) + var visit func(f uri.URI) + visit = func(f uri.URI) { if visited[f] { return } + visited[f] = true + node := ss.Graph().Get(f) if node == nil { return } + for _, inc := range node.OutDegree() { out = append(out, inc) visit(inc) } } visit(file) + return out } diff --git a/lsp/completion/token_completion.go b/lsp/completion/token_completion.go index 46e61cc..23ab72f 100644 --- a/lsp/completion/token_completion.go +++ b/lsp/completion/token_completion.go @@ -70,6 +70,7 @@ func (c *TokenCompletion) Completion(ctx context.Context, ss *cache.Snapshot, cm if err != nil { return nil, rng, err } + if parsedFile.AST() == nil { return nil, rng, fmt.Errorf("parser ast failed") } @@ -90,11 +91,13 @@ func (c *TokenCompletion) Completion(ctx context.Context, ss *cache.Snapshot, cm // Include completion: the cursor is inside an include path literal. includePos := pos includePos.Col-- + includePath := parsedFile.AST().SearchNodePathByPosition(includePos) if items, includeRng, err := c.includeCompletion(ss, cmp.Fh.URI(), parsedFile.AST(), includePath); err == nil { candidates = append(candidates, items...) if len(items) > 0 { rng = includeRng + slog.Debug("include completion candidates", "candidates", candidates) } } @@ -104,12 +107,14 @@ func (c *TokenCompletion) Completion(ctx context.Context, ss *cache.Snapshot, cm if err != nil { return nil, rng, err } + var prefix []byte // get prefix by pos for i := pos.Offset - 1; i >= 0; i-- { if unicode.IsSpace(rune(content[i])) || content[i] == '.' || content[i] == '\'' || content[i] == '"' { prefix = content[i+1 : pos.Offset] rng.Start.Character = rng.Start.Character - uint32(len(prefix)) + break } } @@ -142,6 +147,7 @@ func (c *TokenCompletion) Completion(ctx context.Context, ss *cache.Snapshot, cm for i := range keywords { searchCandidate(i, keywords[i]) } + for i := range tokens { searchCandidate(i, protocol.InsertTextFormatPlainText) } @@ -151,13 +157,16 @@ func (c *TokenCompletion) Completion(ctx context.Context, ss *cache.Snapshot, cm sort.Slice(candidates, func(i, j int) bool { a, b := candidates[i].showText, candidates[j].showText aStarts := strings.HasPrefix(a, string(prefix)) + bStarts := strings.HasPrefix(b, string(prefix)) if aStarts != bStarts { return aStarts } + if len(a) != len(b) { return len(a) < len(b) } + return a < b }) @@ -182,6 +191,7 @@ func (c *TokenCompletion) includeCompletion(ss *cache.Snapshot, file uri.URI, do if len(path) == 0 { return res, rng, err } + include, ok := path[len(path)-1].(*syntax.Include) if !ok || include.Path == nil { return res, rng, err @@ -207,5 +217,6 @@ func (c *TokenCompletion) includeCompletion(ss *cache.Snapshot, file uri.URI, do res, err = ListDirAndFiles(currentDir, pathPrefix) slog.Debug("include completion", "res", res, "err", err) + return res, rng, err } diff --git a/lsp/completion/utils.go b/lsp/completion/utils.go index 94228bc..0e2cbb9 100644 --- a/lsp/completion/utils.go +++ b/lsp/completion/utils.go @@ -41,7 +41,9 @@ func ListDirAndFiles(dir, prefix string) (res []Candidate, err error) { if err != nil || baseDir == path { return nil } + slog.Debug("include completion name", "name", d.Name(), "prefix", filePrefix) + if strings.HasPrefix(d.Name(), filePrefix) { if d.IsDir() { res = append(res, Candidate{ @@ -61,6 +63,7 @@ func ListDirAndFiles(dir, prefix string) (res []Candidate, err error) { if d.IsDir() { return filepath.SkipDir } + return nil }) diff --git a/lsp/diagnostic.go b/lsp/diagnostic.go index 6f8f18d..26b4b07 100644 --- a/lsp/diagnostic.go +++ b/lsp/diagnostic.go @@ -21,6 +21,7 @@ func (s *Server) diagnostic(ctx context.Context, ss *cache.Snapshot, changeFile defer slog.Debug("diagnostic finished") diag := diagnostic.NewDiagnostic() + diagRes, err := diag.Diagnostic(ctx, ss, []uri.URI{changeFile.URI}) if err != nil { slog.Error("diagnostic failed", "err", err) @@ -29,15 +30,18 @@ func (s *Server) diagnostic(ctx context.Context, ss *cache.Snapshot, changeFile slog.Debug("publish diagnostic result", "count", len(diagRes)) var errs []error + for file, res := range diagRes { if res == nil { res = make([]protocol.Diagnostic, 0) } + params := &protocol.PublishDiagnosticsParams{ URI: file, Diagnostics: res, } slog.Debug("file diagnostics", "file", file, "diagnostics", res) + err = s.client.PublishDiagnostics(ctx, params) if err != nil { errs = append(errs, err) diff --git a/lsp/diagnostic/cycle_detect.go b/lsp/diagnostic/cycle_detect.go index f3b57b0..78bb900 100644 --- a/lsp/diagnostic/cycle_detect.go +++ b/lsp/diagnostic/cycle_detect.go @@ -20,6 +20,7 @@ func (c *CycleCheck) Diagnostic(ctx context.Context, ss *cache.Snapshot, changeF for _, file := range changeFiles { _ = getIncludes(ctx, ss, file, &includesMap) } + cyclePairs := cycleDetect(&includesMap) return cycleToDiagnosticItems(cyclePairs), nil @@ -45,6 +46,7 @@ func cyclePairToDiagnostic(pair CyclePair) protocol.Diagnostic { Source: protocol.NewOptional("thrift-ls"), Message: protocol.String(fmt.Sprintf("cycle dependency in %s", pair.include.file)), } + return res } @@ -82,19 +84,24 @@ func getIncludes(ctx context.Context, ss *cache.Snapshot, file uri.URI, includes pf, err := ss.Parse(ctx, file) if err != nil { slog.Error("parse failed", "file", file, "err", err) + return err } + if pf.AST() == nil { slog.Error("parse ast failed", "errs", pf.AggregatedError()) + return pf.AggregatedError() } includes := pf.AST().Includes() resolver := ss.Resolver() + for i := range includes { if includes[i].Path == nil { continue } + includeURI := resolver.ResolveIncludeWithText(file, lsputils.IncludePathText(includes[i])) (*includesMap)[file] = append((*includesMap)[file], Include{ file: includeURI, @@ -105,6 +112,7 @@ func getIncludes(ctx context.Context, ss *cache.Snapshot, file uri.URI, includes if _, ok := (*includesMap)[includeURI]; ok { continue } + _ = getIncludes(ctx, ss, includeURI, includesMap) } diff --git a/lsp/diagnostic/cycle_detect_test.go b/lsp/diagnostic/cycle_detect_test.go index 3bf5279..97e1edf 100644 --- a/lsp/diagnostic/cycle_detect_test.go +++ b/lsp/diagnostic/cycle_detect_test.go @@ -26,6 +26,7 @@ func Test_cycleDetect(t *testing.T) { type args struct { includesMap *map[uri.URI][]Include } + tests := []struct { name string args args @@ -64,6 +65,7 @@ func Test_cycleDetect(t *testing.T) { 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 }) @@ -72,6 +74,7 @@ func Test_cycleDetect(t *testing.T) { if got[i].file == got[j].file { return got[i].include.file < got[j].include.file } + return got[i].file < got[j].file }) @@ -115,6 +118,7 @@ include "./test/address.thrift"` t.Fatal(e) } } + return doc } userDoc := parseFor("file:///tmp/user.thrift", file1) @@ -142,6 +146,7 @@ include "./test/address.thrift"` file uri.URI includesMap *map[uri.URI][]Include } + tests := []struct { name string args args diff --git a/lsp/diagnostic/diagnostic.go b/lsp/diagnostic/diagnostic.go index a0b107d..f7334ec 100644 --- a/lsp/diagnostic/diagnostic.go +++ b/lsp/diagnostic/diagnostic.go @@ -36,20 +36,26 @@ func NewDiagnostic() Interface { func (d *Diagnostic) Diagnostic(ctx context.Context, ss *cache.Snapshot, changeFiles []uri.URI) (DiagnosticResult, error) { res := make(DiagnosticResult) + var errs []error + for _, impl := range registry { slog.Debug("diagnostic called", "impl", impl.Name()) + diagRes, err := impl.Diagnostic(ctx, ss, changeFiles) if err != nil { errs = append(errs, err) } + for key, items := range diagRes { res[key] = append(res[key], items...) } } + if len(errs) > 0 { return res, errors.Join(errs...) } + return res, nil } @@ -64,8 +70,10 @@ func tokenRange(doc *syntax.Document, tok *syntax.Token) protocol.Range { if tok == nil { return protocol.Range{} } + start := doc.TokenPosition(tokIndex(doc, tok)) end := doc.TokenEndPosition(tokIndex(doc, tok)) + return protocol.Range{ Start: protocol.Position{ Line: uint32(start.Line - 1), @@ -81,6 +89,7 @@ func tokenRange(doc *syntax.Document, tok *syntax.Token) protocol.Range { // nodeRange converts a node's span to an LSP range. func nodeRange(doc *syntax.Document, node syntax.Node) protocol.Range { start, end := doc.Range(node) + return protocol.Range{ Start: protocol.Position{ Line: uint32(start.Line - 1), @@ -106,5 +115,6 @@ func tokIndex(doc *syntax.Document, tok *syntax.Token) int { return i } } + return 0 } diff --git a/lsp/diagnostic/fieldid_check.go b/lsp/diagnostic/fieldid_check.go index bf11061..2179c55 100644 --- a/lsp/diagnostic/fieldid_check.go +++ b/lsp/diagnostic/fieldid_check.go @@ -19,11 +19,13 @@ type FieldIDCheck struct{} // throws field ids: they must be unique positive integers in [1, 32767]. func (c *FieldIDCheck) Diagnostic(ctx context.Context, ss *cache.Snapshot, changeFiles []uri.URI) (DiagnosticResult, error) { res := make(DiagnosticResult) + for _, file := range changeFiles { items, err := c.diagnostic(ctx, ss, file) if err != nil { return nil, err } + res[file] = items } @@ -39,6 +41,7 @@ func (c *FieldIDCheck) diagnostic(ctx context.Context, ss *cache.Snapshot, file if err != nil { return nil, err } + if pf.AST() == nil { return nil, errors.New("parse ast failed") } @@ -51,15 +54,18 @@ func (c *FieldIDCheck) diagnostic(ctx context.Context, ss *cache.Snapshot, file processStructLike := func(fields []*syntax.Field) { fieldIDSet := make(map[int][]*syntax.Field) + for i := range fields { field := fields[i] if field.FieldID == nil { continue } + value, err := strconv.ParseInt(field.FieldID.Text, 0, 32) if err != nil { continue } + fieldIDSet[int(value)] = append(fieldIDSet[int(value)], field) } @@ -78,6 +84,7 @@ func (c *FieldIDCheck) diagnostic(ctx context.Context, ss *cache.Snapshot, file if len(set) == 1 { continue } + for _, field := range set { ret = append(ret, protocol.Diagnostic{ Range: tokenRange(pf.AST(), field.FieldID), @@ -92,15 +99,19 @@ func (c *FieldIDCheck) diagnostic(ctx context.Context, ss *cache.Snapshot, file for _, st := range pf.AST().Structs() { processStructLike(st.Fields) } + for _, union := range pf.AST().Unions() { processStructLike(union.Fields) } + for _, excep := range pf.AST().Exceptions() { processStructLike(excep.Fields) } + for _, svc := range pf.AST().Services() { for _, fn := range svc.Functions { processStructLike(fn.Args) + if fn.Throws != nil { processStructLike(fn.Throws.Fields) } diff --git a/lsp/diagnostic/fieldid_check_test.go b/lsp/diagnostic/fieldid_check_test.go index bec69e3..7433d6c 100644 --- a/lsp/diagnostic/fieldid_check_test.go +++ b/lsp/diagnostic/fieldid_check_test.go @@ -48,11 +48,13 @@ service Demo { From: cache.FileChangeTypeDidOpen, }, }) + type args struct { ctx context.Context ss *cache.Snapshot changeFiles []uri.URI } + tests := []struct { name string c *FieldIDCheck diff --git a/lsp/diagnostic/parse.go b/lsp/diagnostic/parse.go index 60299e3..85f3eb2 100644 --- a/lsp/diagnostic/parse.go +++ b/lsp/diagnostic/parse.go @@ -18,10 +18,12 @@ func (p *Parse) Diagnostic(ctx context.Context, ss *cache.Snapshot, changeFiles var errs []error res := make(DiagnosticResult) + for _, uri := range changeFiles { parseRes, err := ss.Parse(ctx, uri) if err != nil { errs = append(errs, err) + continue } @@ -49,6 +51,7 @@ func syntaxErrorToDiagnostic(err syntax.Error) protocol.Diagnostic { if err.Severity == syntax.SeverityWarning { severity = protocol.DiagnosticSeverityWarning } + return protocol.Diagnostic{ Range: protocol.Range{ Start: protocol.Position{ diff --git a/lsp/diagnostic/semantic_analysis.go b/lsp/diagnostic/semantic_analysis.go index 4b161b6..292deed 100644 --- a/lsp/diagnostic/semantic_analysis.go +++ b/lsp/diagnostic/semantic_analysis.go @@ -18,11 +18,13 @@ type SemanticAnalysis struct{} func (s *SemanticAnalysis) Diagnostic(ctx context.Context, ss *cache.Snapshot, changeFiles []uri.URI) (DiagnosticResult, error) { res := make(DiagnosticResult) + for _, file := range changeFiles { items, err := s.diagnostic(ctx, ss, file) if err != nil { return nil, err } + res[file] = items } @@ -38,6 +40,7 @@ func (s *SemanticAnalysis) diagnostic(ctx context.Context, ss *cache.Snapshot, c if err != nil { return nil, err } + if pf.AST() == nil { return nil, errors.New("parse ast failed") } @@ -60,6 +63,7 @@ func (s *SemanticAnalysis) checkDefineConflict(ctx context.Context, pf *cache.Pa processStructLike := func(fields []*syntax.Field) { fieldMap := make(map[string]struct{}) + for i := range fields { field := fields[i] if _, exist := fieldMap[field.Name.Text]; exist { @@ -70,6 +74,7 @@ func (s *SemanticAnalysis) checkDefineConflict(ctx context.Context, pf *cache.Pa Message: protocol.String("field name conflict with other field"), }) } + fieldMap[field.Name.Text] = struct{}{} } } @@ -85,6 +90,7 @@ func (s *SemanticAnalysis) checkDefineConflict(ctx context.Context, pf *cache.Pa Message: protocol.String(fmt.Sprintf("%s name conflict with other %s", kind, previous)), }) } + definitionNameMap[name] = kind } @@ -92,20 +98,25 @@ func (s *SemanticAnalysis) checkDefineConflict(ctx context.Context, pf *cache.Pa processDefinition(st.Name.Text, st.Name, "struct") processStructLike(st.Fields) } + for _, union := range pf.AST().Unions() { processDefinition(union.Name.Text, union.Name, "union") processStructLike(union.Fields) } + for _, excep := range pf.AST().Exceptions() { processDefinition(excep.Name.Text, excep.Name, "exception") processStructLike(excep.Fields) } + for _, enum := range pf.AST().Enums() { processDefinition(enum.Name.Text, enum.Name, "enum") } + for _, cst := range pf.AST().Consts() { processDefinition(cst.Name.Text, cst.Name, "const") } + for _, td := range pf.AST().Typedefs() { processDefinition(td.Name.Text, td.Name, "typedef") } @@ -123,8 +134,10 @@ func (s *SemanticAnalysis) checkDefineConflict(ctx context.Context, pf *cache.Pa Message: protocol.String("function name conflict with other function"), }) } + fnMap[fn.Name.Text] = struct{}{} processStructLike(fn.Args) + if fn.Throws != nil { processStructLike(fn.Throws.Fields) } @@ -160,22 +173,27 @@ func (s *SemanticAnalysis) checkDefinitionExist(ctx context.Context, ss *cache.S for _, st := range pf.AST().Structs() { processStructLike(st.Fields) } + for _, union := range pf.AST().Unions() { processStructLike(union.Fields) } + for _, excep := range pf.AST().Exceptions() { processStructLike(excep.Fields) } + for _, cst := range pf.AST().Consts() { items := s.checkConstValueExist(ctx, ss, file, pf, cst.Value) ret = append(ret, items...) } + for _, svc := range pf.AST().Services() { for _, fn := range svc.Functions { items := s.checkTypeExist(ctx, ss, file, pf, fn.Type) ret = append(ret, items...) processStructLike(fn.Args) + if fn.Throws != nil { processStructLike(fn.Throws.Fields) } @@ -213,6 +231,7 @@ func (s *SemanticAnalysis) checkConstValueMatchType(pf *cache.ParsedFile, field if field.Value == nil { return nil } + expect := typeName(field.Type) value := field.Value valueKind := value.Kind @@ -228,8 +247,10 @@ func (s *SemanticAnalysis) checkConstValueMatchType(pf *cache.ParsedFile, field if expect != "bool" { return mismatchDiagnostic(pf.AST(), field, expect, "bool") } + return nil } + switch expect { case "i8", "i16", "i32", "i64": default: @@ -255,6 +276,7 @@ func sameKind(expect string, kind syntax.ConstValueKind) bool { case syntax.ValueDouble: return expect == "double" } + return false } @@ -273,6 +295,7 @@ func kindName(kind syntax.ConstValueKind) string { case syntax.ValueIdent: return "identifier" } + return "unknown" } @@ -291,6 +314,7 @@ func typeName(ft *syntax.FieldType) string { if ft == nil { return "" } + switch ft.Kind { case syntax.TypeIdent: return ft.Ident.Text @@ -303,6 +327,7 @@ func typeName(ft *syntax.FieldType) string { case syntax.TypeBase: return ft.Base.String() } + return "" } @@ -312,6 +337,7 @@ func (s *SemanticAnalysis) checkTypeExist(ctx context.Context, ss *cache.Snapsho if ft == nil { return res } + switch ft.Kind { case syntax.TypeMap, syntax.TypeList, syntax.TypeSet: return s.checkContainerTypeExist(ctx, ss, file, pf, ft) @@ -328,6 +354,7 @@ func (s *SemanticAnalysis) checkTypeExist(ctx context.Context, ss *cache.Snapsho }) } } + return res } @@ -337,8 +364,10 @@ func (s *SemanticAnalysis) checkContainerTypeExist(ctx context.Context, if ft.KeyType != nil { res = append(res, s.checkTypeExist(ctx, ss, file, pf, ft.KeyType)...) } + if ft.ValueType != nil { res = append(res, s.checkTypeExist(ctx, ss, file, pf, ft.ValueType)...) } + return res } diff --git a/lsp/diagnostic/semantic_analysis_test.go b/lsp/diagnostic/semantic_analysis_test.go index 3ca5575..2c7cc90 100644 --- a/lsp/diagnostic/semantic_analysis_test.go +++ b/lsp/diagnostic/semantic_analysis_test.go @@ -61,11 +61,13 @@ struct TestUUID { From: cache.FileChangeTypeDidOpen, }, }) + type args struct { ctx context.Context ss *cache.Snapshot changeFiles []uri.URI } + tests := []struct { name string args args diff --git a/lsp/format.go b/lsp/format.go index 3688144..a2817ec 100644 --- a/lsp/format.go +++ b/lsp/format.go @@ -18,6 +18,7 @@ func (s *Server) formatting(ctx context.Context, params *protocol.DocumentFormat document := params.TextDocument fileURI := document.URI + view, err := s.session.ViewOf(fileURI) if err != nil { return nil, err @@ -40,6 +41,7 @@ func (s *Server) formatting(ctx context.Context, params *protocol.DocumentFormat if err != nil { return nil, err } + if len(pf.Errors()) > 0 || pf.AST() == nil { return nil, pf.AggregatedError() } @@ -80,6 +82,7 @@ func (s *Server) formatting(ctx context.Context, params *protocol.DocumentFormat // produced and the request is a no-op. func (s *Server) rangeFormatting(ctx context.Context, params *protocol.DocumentRangeFormattingParams) (result []protocol.TextEdit, err error) { fileURI := params.TextDocument.URI + view, err := s.session.ViewOf(fileURI) if err != nil { return nil, err @@ -99,14 +102,17 @@ func (s *Server) rangeFormatting(ctx context.Context, params *protocol.DocumentR } mp := mapper.NewMapper(fileURI, content) + start, err := mp.LSPPosToParserPosition(lspPosition(params.Range.Start)) if err != nil { return nil, nil } + end, err := mp.LSPPosToParserPosition(lspPosition(params.Range.End)) if err != nil { return nil, nil } + rs, re := start.Offset, end.Offset if rs >= re { return nil, nil @@ -126,6 +132,7 @@ func (s *Server) rangeFormatting(ctx context.Context, params *protocol.DocumentR if err != nil { return nil, nil } + endPos, err := mp.OffsetToLSPPosition(re) if err != nil { return nil, nil @@ -170,6 +177,7 @@ func formatRangeText(content []byte, rs, re int, opts formatter.Options) (newTex rs = lineStart(content, rs) re = lineEnd(content, re) rs = skipBlankLinesForward(content, rs, re) + re = skipBlankLinesBackward(content, rs, re) if rs >= re { return "", rs, re, false @@ -182,6 +190,7 @@ func formatRangeText(content []byte, rs, re int, opts formatter.Options) (newTex } slice := content[rs:re] + doc, errs := syntax.Parse(slice) if len(errs) > 0 { return "", rs, re, false @@ -196,6 +205,7 @@ func formatRangeText(content []byte, rs, re int, opts formatter.Options) (newTex // starts at a line start and ends just before its last line's newline. formatted = strings.TrimLeft(formatted, "\n") formatted = strings.TrimRight(formatted, "\r\n") + return formatted, rs, re, true } @@ -204,6 +214,7 @@ func lineStart(content []byte, offset int) int { if i := bytes.LastIndexByte(content[:offset], '\n'); i != -1 { return i + 1 } + return 0 } @@ -213,6 +224,7 @@ func lineEnd(content []byte, offset int) int { if i := bytes.IndexByte(content[offset:], '\n'); i != -1 { return offset + i } + return len(content) } @@ -222,7 +234,9 @@ func blankLineBefore(content []byte, offset int) bool { if offset == 0 { return true } + start := lineStart(content, offset-1) + return len(bytes.TrimSpace(content[start:offset])) == 0 } @@ -232,7 +246,9 @@ func blankLineAfter(content []byte, offset int) bool { if offset == len(content) { return true } + end := lineEnd(content, offset+1) + return len(bytes.TrimSpace(content[offset+1:end])) == 0 } @@ -243,11 +259,14 @@ func skipBlankLinesForward(content []byte, offset, limit int) int { if end >= limit { break } + if len(bytes.TrimSpace(content[offset:end])) > 0 { break } + offset = end + 1 } + return offset } @@ -255,15 +274,13 @@ func skipBlankLinesForward(content []byte, offset, limit int) int { // past blank lines, stopping at start. func skipBlankLinesBackward(content []byte, start, offset int) int { for offset > start { - lineStart := lineStart(content, offset-1) if len(bytes.TrimSpace(content[lineStart:offset])) > 0 { break } - offset = lineStart - 1 - if offset < 0 { - offset = 0 - } + + offset = max(lineStart-1, 0) } + return offset } diff --git a/lsp/format_range_fuzz_test.go b/lsp/format_range_fuzz_test.go index c88070a..789804d 100644 --- a/lsp/format_range_fuzz_test.go +++ b/lsp/format_range_fuzz_test.go @@ -33,10 +33,12 @@ func FuzzFormatRangeText(f *testing.F) { opts := formatter.Options{} // Clamp the fuzzed offsets into the content. limit := len(content) + 1 + rs = rs % limit if rs < 0 { rs += limit } + re = re % limit if re < 0 { re += limit @@ -51,9 +53,11 @@ func FuzzFormatRangeText(f *testing.F) { if outRS < 0 || outRE > len(content) || outRS >= outRE { t.Fatalf("invalid accepted range [%d, %d) for content %q", outRS, outRE, content) } + if outRS != 0 && content[outRS-1] != '\n' { t.Fatalf("accepted range starts mid-line at %d in %q", outRS, content) } + if outRE != len(content) && content[outRE] != '\n' { t.Fatalf("accepted range ends mid-line at %d in %q", outRE, content) } @@ -68,6 +72,7 @@ func FuzzFormatRangeText(f *testing.F) { spliced = append(spliced, content[outRE:]...) origErrs := errorMessages(syntax.Parse(content)) + splicedErrs := errorMessages(syntax.Parse(spliced)) for msg, splicedCount := range splicedErrs { if splicedCount > origErrs[msg] { @@ -82,10 +87,12 @@ func FuzzFormatRangeText(f *testing.F) { func errorMessages(doc *syntax.Document, errs []syntax.Error) map[string]int { _ = doc counts := make(map[string]int) + for _, err := range errs { if err.Severity == syntax.SeverityError { counts[err.Message]++ } } + return counts } diff --git a/lsp/format_range_server_test.go b/lsp/format_range_server_test.go index 49fb0ea..d54515f 100644 --- a/lsp/format_range_server_test.go +++ b/lsp/format_range_server_test.go @@ -47,6 +47,7 @@ struct C { 3: i64 c } }) assert.NoError(t, err) assert.Len(t, edits, 1) + return edits[0].NewText } diff --git a/lsp/format_range_test.go b/lsp/format_range_test.go index eb62c6f..4c919d8 100644 --- a/lsp/format_range_test.go +++ b/lsp/format_range_test.go @@ -143,6 +143,7 @@ func TestFormatRangeTextSplice(t *testing.T) { if gotRS != strings.Index(content, "struct B") { t.Errorf("effective start = %d, want %d", gotRS, strings.Index(content, "struct B")) } + if got := content[gotRE]; got != '\n' { t.Errorf("effective end = %d, want a newline position (got %q)", gotRE, got) } @@ -153,10 +154,12 @@ func TestFormatRangeTextSplice(t *testing.T) { if len(errs) > 0 { t.Fatalf("parse errors: %v", errs) } + want, err := formatter.Format(doc, opts) if err != nil { t.Fatalf("whole-doc format: %v", err) } + if spliced != want { t.Errorf("splice mismatch\n got: %q\nwant: %q", spliced, want) } diff --git a/lsp/hover.go b/lsp/hover.go index 1fc7888..c3bc466 100644 --- a/lsp/hover.go +++ b/lsp/hover.go @@ -11,10 +11,12 @@ import ( func (s *Server) hover(ctx context.Context, params *protocol.HoverParams) (*protocol.Hover, error) { file := params.TextDocument.URI + view, err := s.session.ViewOf(file) if err != nil { return nil, err } + ss, release := view.Snapshot() defer release() @@ -31,6 +33,7 @@ func (s *Server) hover(ctx context.Context, params *protocol.HoverParams) (*prot if strings.HasPrefix(content, "\n") { markdown_prefix = "```thrift" } + markdown_suffix := "\n```" if strings.HasSuffix(content, "\n") { markdown_suffix = "```" diff --git a/lsp/impl.go b/lsp/impl.go index b8eade0..202cf50 100644 --- a/lsp/impl.go +++ b/lsp/impl.go @@ -31,10 +31,12 @@ func (s *Server) didOpen(ctx context.Context, params *protocol.DidOpenTextDocume s.session.Initialize(func() { file := change.URI + dirPos := strings.LastIndexByte(string(file), '/') if dirPos == -1 { return } + dir := file[0:dirPos] s.walkFoldersThriftFile(dir) }) @@ -60,6 +62,7 @@ func (s *Server) openFile(ctx context.Context, change *cache.FileChange) error { view.FileChange(ctx, []*cache.FileChange{change}, func() { ss, release := view.Snapshot() defer release() + err := s.diagnostic(ctx, ss, change) if err != nil { slog.Error("diagnostic error", "err", err) @@ -77,6 +80,7 @@ func (s *Server) didChange(ctx context.Context, params *protocol.DidChangeTextDo document := params.TextDocument fileURI := document.URI + view, err := s.session.ViewOf(fileURI) if err != nil { return err @@ -85,6 +89,7 @@ func (s *Server) didChange(ctx context.Context, params *protocol.DidChangeTextDo view.FileChange(ctx, changes, func() { ss, release := view.Snapshot() defer release() + for i := range changes { err := s.diagnostic(ctx, ss, changes[i]) if err != nil { @@ -122,6 +127,7 @@ func toLspCompletionList(items []*completion.CompletionItem, rng protocol.Range) list := &protocol.CompletionList{ IsIncomplete: true, } + for i := range items { item := protocol.CompletionItem{ Label: items[i].Label, @@ -140,11 +146,13 @@ func toLspCompletionList(items []*completion.CompletionItem, rng protocol.Range) } list.Items = append(list.Items, item) } + return list } func (s *Server) getFileContext(ctx context.Context, uri uri.URI) (ss *cache.Snapshot, release func(), fh cache.FileHandle, err error) { var view *cache.View + view, err = s.session.ViewOf(uri) if err != nil { return ss, release, fh, err @@ -155,6 +163,7 @@ func (s *Server) getFileContext(ctx context.Context, uri uri.URI) (ss *cache.Sna fh, err = ss.ReadFile(ctx, uri) if err != nil { release() + return ss, release, fh, err } diff --git a/lsp/impl_test.go b/lsp/impl_test.go index 5c38a63..75a5f44 100644 --- a/lsp/impl_test.go +++ b/lsp/impl_test.go @@ -16,6 +16,7 @@ func Test_DidOpen(t *testing.T) { ctx := t.Context() fileURI, err := uri.Parse("file:///tmp/file.thrift") assert.NoError(t, err) + fileContent := ` include "base.thrift" @@ -42,7 +43,7 @@ struct Test { fh, err := srv.session.ReadFile(ctx, fileURI) assert.NoError(t, err) - assert.Equal(t, int(fh.Version()), 0) + assert.Equal(t, 0, int(fh.Version())) gotContent, err := fh.Content() assert.NoError(t, err) assert.Equal(t, gotContent, []byte(fileContent)) @@ -52,6 +53,7 @@ func Test_DidChange(t *testing.T) { ctx := t.Context() fileURI, err := uri.Parse("file:///tmp/file.thrift") assert.NoError(t, err) + fileContentInit := ` include "base.thrift" @@ -102,7 +104,7 @@ struct Test { fh, err := srv.session.ReadFile(ctx, fileURI) assert.NoError(t, err) - assert.Equal(t, int(fh.Version()), 1) + assert.Equal(t, 1, int(fh.Version())) gotContent, err := fh.Content() assert.NoError(t, err) assert.Equal(t, gotContent, []byte(fileContent)) @@ -185,8 +187,8 @@ struct Test { assert.IsType(t, &protocol.CompletionList{}, completionResult) completionList := completionResult.(*protocol.CompletionList) - assert.True(t, len(completionList.Items) > 0) - assert.True(t, len(completionList.Items) <= 10) + assert.NotEmpty(t, completionList.Items) + assert.LessOrEqual(t, len(completionList.Items), 10) assert.Equal(t, tt.wantLabel, completionList.Items[0].Label) preselect, _ := completionList.Items[0].Preselect.Get() assert.Equal(t, tt.wantPreselect, preselect) @@ -286,6 +288,7 @@ struct Test { assert.IsType(t, &protocol.CompletionList{}, completionResult) completionList := completionResult.(*protocol.CompletionList) + labels := make([]string, len(completionList.Items)) for i, item := range completionList.Items { labels[i] = item.Label @@ -395,6 +398,7 @@ struct Other { assert.IsType(t, &protocol.CompletionList{}, completionResult) completionList := completionResult.(*protocol.CompletionList) + labels := make([]string, len(completionList.Items)) for i, item := range completionList.Items { labels[i] = item.Label diff --git a/lsp/include_paths_test.go b/lsp/include_paths_test.go index 128bab1..39e4c5d 100644 --- a/lsp/include_paths_test.go +++ b/lsp/include_paths_test.go @@ -30,6 +30,7 @@ func TestServerIncludePathsFlow(t *testing.T) { srv.session.CreateView(uri.File(dir)) view, err := srv.session.ViewOf(uri.File(filepath.Join(dir, "app.thrift"))) assert.NoError(t, err) + snapshot, release := view.Snapshot() defer release() diff --git a/lsp/initialize.go b/lsp/initialize.go index 300e070..0e5fb82 100644 --- a/lsp/initialize.go +++ b/lsp/initialize.go @@ -23,6 +23,7 @@ func (s *Server) initialize(ctx context.Context, params *protocol.InitializePara folders = append(folders, ws.URI) } } + if len(folders) == 0 { //nolint:staticcheck // intentional handling of legacy client params rootURI := params.RootURI @@ -33,12 +34,14 @@ func (s *Server) initialize(ctx context.Context, params *protocol.InitializePara rootURI = &r } } + if rootURI != nil { folders = append(folders, *rootURI) } } slog.Debug("initialized folders", "folders", folders) + if len(folders) > 0 { s.session.Initialize(func() { for i := range folders { @@ -55,6 +58,7 @@ func (s *Server) walkFoldersThriftFile(folder uri.URI) { // WalkDir walk files with lexical order _ = filepath.WalkDir(folder.Path(), func(path string, d fs.DirEntry, err error) error { slog.Debug("walking", "path", path) + if err != nil { return nil } @@ -69,6 +73,7 @@ func (s *Server) walkFoldersThriftFile(folder uri.URI) { fileURI := uri.File(path) slog.Debug("file path", "uri", fileURI) + if err := s.openFile(context.TODO(), &cache.FileChange{ URI: fileURI, Version: 0, diff --git a/lsp/lsputils/utils.go b/lsp/lsputils/utils.go index e48b4f7..d15d8c0 100644 --- a/lsp/lsputils/utils.go +++ b/lsp/lsputils/utils.go @@ -15,16 +15,19 @@ import ( // for example: file uri is file:///base.thrift, then `base` is include name func GetIncludeName(file uri.URI) string { fileName := file.Path() + index := strings.LastIndexByte(fileName, filepath.Separator) if index == -1 { return fileName } + fileName = string(fileName[index+1:]) index = strings.LastIndexByte(fileName, '.') if index == -1 { return fileName } + return string(fileName[0:index]) } @@ -34,6 +37,7 @@ func IncludePathText(inc *syntax.Include) string { if inc == nil || inc.Path == nil { return "" } + return strings.Trim(inc.Path.Text, "\"'") } @@ -45,11 +49,14 @@ func GetIncludePath(ast *syntax.Document, includeName string) string { if path == "" { continue } + items := strings.Split(path, "/") + path = items[len(items)-1] if !strings.HasSuffix(path, ".thrift") { continue } + name := strings.TrimSuffix(path, ".thrift") if name == includeName { return IncludePathText(include) diff --git a/lsp/lsputils/utils_test.go b/lsp/lsputils/utils_test.go index 0342e2b..9e30991 100644 --- a/lsp/lsputils/utils_test.go +++ b/lsp/lsputils/utils_test.go @@ -16,6 +16,7 @@ func Test_IncludeURI(t *testing.T) { cur uri.URI includePath string } + tests := []struct { name string args args @@ -67,6 +68,7 @@ include "../../user.extra.thrift" service Demo { user.Test Api(1:user.Test2 arg1, 2:user.Test3 arg2) throws (1:user.Error1 err) }` + ast, errs := syntax.Parse([]byte(file)) for _, e := range errs { if e.Severity == syntax.SeverityError { @@ -78,6 +80,7 @@ service Demo { ast *syntax.Document includeName string } + tests := []struct { name string args args @@ -111,6 +114,7 @@ func TestGetIncludeName(t *testing.T) { type args struct { file uri.URI } + tests := []struct { name string args args @@ -150,6 +154,7 @@ func TestIncludeNames(t *testing.T) { cur uri.URI includes []*syntax.Include } + tests := []struct { name string args args @@ -192,9 +197,9 @@ func TestIncludeURIWithPaths(t *testing.T) { // shared.thrift (exists) // service/ // order.thrift (exists) - tmpDir, err := os.MkdirTemp("", "thrift-test") assert.NoError(t, err) + defer func() { _ = os.RemoveAll(tmpDir) }() baseDir := filepath.Join(tmpDir, "base") @@ -267,6 +272,7 @@ func TestParseIdent(t *testing.T) { includes []*syntax.Include identifier string } + tests := []struct { name string args args diff --git a/lsp/mapper/mapper.go b/lsp/mapper/mapper.go index ea658fd..6d7fa14 100644 --- a/lsp/mapper/mapper.go +++ b/lsp/mapper/mapper.go @@ -34,11 +34,13 @@ func NewMapper(fileURI uri.URI, content []byte) *Mapper { func (m *Mapper) initLineStart() { m.lineInit.Do(func() { nlines := bytes.Count(m.content, []byte("\n")) + m.lineStart = make([]int, 1, nlines+1) // initially []int{0} for offset, b := range m.content { if b == '\n' { m.lineStart = append(m.lineStart, offset+1) } + if b >= utf8.RuneSelf { m.nonASCII = true } @@ -63,14 +65,12 @@ func (m *Mapper) GetLSPEndPosition() types.Position { // position (0-based line, UTF-16 code-unit column). func (m *Mapper) OffsetToLSPPosition(offset int) (types.Position, error) { m.initLineStart() + if offset < 0 || offset > len(m.content) { return types.Position{}, fmt.Errorf("invalid offset: %d, total content: %d", offset, len(m.content)) } - line := sort.Search(len(m.lineStart), func(i int) bool { return m.lineStart[i] > offset }) - 1 - if line < 0 { - line = 0 - } + line := max(sort.Search(len(m.lineStart), func(i int) bool { return m.lineStart[i] > offset })-1, 0) return types.Position{ Line: uint32(line), @@ -81,6 +81,7 @@ func (m *Mapper) OffsetToLSPPosition(offset int) (types.Position, error) { // convert from utf16-based to rune-based position func (m *Mapper) LSPPosToParserPosition(pos types.Position) (syntax.Position, error) { m.initLineStart() + line := int(pos.Line) + 1 if line > len(m.lineStart) { return syntax.InvalidPosition, fmt.Errorf("invalid position line, request line: %d, total line: %d", line, len(m.lineStart)) @@ -88,10 +89,12 @@ func (m *Mapper) LSPPosToParserPosition(pos types.Position) (syntax.Position, er if !m.nonASCII { col := int(pos.Character) + 1 + offset := m.lineStart[pos.Line] + int(pos.Character) if offset > len(m.content) { return syntax.InvalidPosition, fmt.Errorf("invalid position offset: %d, total content: %d, %s", offset, len(m.content), string(m.content)) } + var lineLength int if int(pos.Line+1) >= len(m.lineStart) { lineLength = len(m.content) - m.lineStart[pos.Line] @@ -111,37 +114,45 @@ func (m *Mapper) LSPPosToParserPosition(pos types.Position) (syntax.Position, er } lineStart := m.lineStart[pos.Line] + lineEnd := 0 if int(pos.Line) == len(m.lineStart)-1 { lineEnd = len(m.content) } else { lineEnd = m.lineStart[pos.Line+1] } + lineBytes := m.content[lineStart:lineEnd] utf16Col := 0 bytesCol := 0 + for len(lineBytes) > 0 { if utf16Col >= int(pos.Character) { break } + if lineBytes[0] < utf8.RuneSelf { utf16Col++ lineBytes = lineBytes[1:] bytesCol++ + continue } r, size := utf8.DecodeRune(lineBytes) + utf16Col++ if r >= 0x10000 { utf16Col++ } + lineBytes = lineBytes[size:] bytesCol += size } runeLen := utf8.RuneCount(m.content[lineStart : lineStart+bytesCol]) + offset := lineStart + bytesCol if offset > len(m.content) { return syntax.InvalidPosition, errors.New("invalid position character") @@ -164,10 +175,12 @@ func utf16Count(contents []byte) int { utf16Len := 0 for len(contents) > 0 { utf16Len++ + r, size := utf8.DecodeRune(contents) if r >= 0x10000 { utf16Len++ } + contents = contents[size:] } diff --git a/lsp/mapper/mapper_fuzz_test.go b/lsp/mapper/mapper_fuzz_test.go index deb2bab..912d682 100644 --- a/lsp/mapper/mapper_fuzz_test.go +++ b/lsp/mapper/mapper_fuzz_test.go @@ -27,24 +27,29 @@ func FuzzOffsetRoundTrip(f *testing.F) { f.Fuzz(func(t *testing.T, content []byte, offset int) { limit := len(content) + 1 + offset = offset % limit if offset < 0 { offset += limit } + if offset < len(content) && !utf8.RuneStart(content[offset]) { // Byte offset inside a multi-byte rune: not representable. return } m := NewMapper(uri.File("/tmp/fuzz.thrift"), content) + pos, err := m.OffsetToLSPPosition(offset) if err != nil { t.Fatalf("OffsetToLSPPosition(%d): %v", offset, err) } + back, err := m.LSPPosToParserPosition(pos) if err != nil { t.Fatalf("LSPPosToParserPosition(%+v): %v (content %q)", pos, err, content) } + if back.Offset != offset { t.Fatalf("round trip mismatch: %d -> %+v -> %d (content %q)", offset, pos, back.Offset, content) } diff --git a/lsp/mapper/mapper_test.go b/lsp/mapper/mapper_test.go index 2a22070..a23503f 100644 --- a/lsp/mapper/mapper_test.go +++ b/lsp/mapper/mapper_test.go @@ -15,6 +15,7 @@ func TestMapper_LSPPosToParserPosition(t *testing.T) { fileURI uri.URI content []byte } + type args struct { pos types.Position } @@ -185,6 +186,7 @@ func Test_utf16Count(t *testing.T) { type args struct { contents []byte } + tests := []struct { name string args args diff --git a/lsp/memoize/promise.go b/lsp/memoize/promise.go index 52378b3..6da5fdb 100644 --- a/lsp/memoize/promise.go +++ b/lsp/memoize/promise.go @@ -86,6 +86,7 @@ func NewPromise(debug string, function Function) *Promise { if function == nil { panic("nil function") } + return &Promise{ debug: debug, function: function, @@ -107,9 +108,11 @@ const ( func (p *Promise) Cached() any { p.mu.Lock() defer p.mu.Unlock() + if p.state == stateCompleted { return p.value } + return nil } @@ -128,6 +131,7 @@ func (p *Promise) Get(ctx context.Context, arg any) (any, error) { if ctx.Err() != nil { return nil, ctx.Err() } + p.mu.Lock() switch p.state { case stateIdle: @@ -136,6 +140,7 @@ func (p *Promise) Get(ctx context.Context, arg any) (any, error) { return p.wait(ctx) case stateCompleted: defer p.mu.Unlock() + return p.value, nil default: panic("unknown state") @@ -164,6 +169,7 @@ func (p *Promise) run(ctx context.Context, arg any) (any, error) { if childCtx.Err() != nil { return } + v := function(childCtx, arg) if childCtx.Err() != nil { return @@ -199,13 +205,16 @@ func (p *Promise) wait(ctx context.Context) (any, error) { case <-done: p.mu.Lock() defer p.mu.Unlock() + if p.state == stateCompleted { return p.value, nil } + return nil, nil case <-ctx.Done(): p.mu.Lock() defer p.mu.Unlock() + p.waiters-- if p.waiters == 0 && p.state == stateRunning { p.cancel() @@ -214,6 +223,7 @@ func (p *Promise) wait(ctx context.Context) (any, error) { p.done = nil p.cancel = nil } + return nil, ctx.Err() } } diff --git a/lsp/memoize/store.go b/lsp/memoize/store.go index f56860b..5bb3cb0 100644 --- a/lsp/memoize/store.go +++ b/lsp/memoize/store.go @@ -37,22 +37,28 @@ type Store struct { // store. func (store *Store) Promise(key any, function Function) (*Promise, func()) { store.promisesMu.Lock() + p, ok := store.promises[key] if !ok { p = NewPromise(reflect.TypeOf(key).String(), function) + if store.promises == nil { store.promises = map[any]*Promise{} } + store.promises[key] = p } + p.refcount++ store.promisesMu.Unlock() - var released int32 + var released atomic.Int32 + release := func() { - if !atomic.CompareAndSwapInt32(&released, 0, 1) { + if !released.CompareAndSwap(0, 1) { panic("release called more than once") } + store.promisesMu.Lock() p.refcount-- @@ -76,6 +82,7 @@ func (s *Store) Stats() map[reflect.Type]int { for k := range s.promises { result[reflect.TypeOf(k)]++ } + return result } diff --git a/lsp/rename.go b/lsp/rename.go index 01150de..9c3ed76 100644 --- a/lsp/rename.go +++ b/lsp/rename.go @@ -10,10 +10,12 @@ import ( func (s *Server) prepareRename(ctx context.Context, params *protocol.PrepareRenameParams) (*protocol.Range, error) { file := params.TextDocument.URI + view, err := s.session.ViewOf(file) if err != nil { return nil, err } + ss, release := view.Snapshot() defer release() @@ -22,10 +24,12 @@ func (s *Server) prepareRename(ctx context.Context, params *protocol.PrepareRena func (s *Server) rename(ctx context.Context, params *protocol.RenameParams) (*protocol.WorkspaceEdit, error) { file := params.TextDocument.URI + view, err := s.session.ViewOf(file) if err != nil { return nil, err } + ss, release := view.Snapshot() defer release() diff --git a/lsp/server.go b/lsp/server.go index 7a1ae11..ffb4c17 100644 --- a/lsp/server.go +++ b/lsp/server.go @@ -30,6 +30,7 @@ func NewServer(c *cache.Cache, client protocol.Client, formatOpts formatter.Opti func (s *Server) Initialize(ctx context.Context, params *protocol.InitializeParams) (result *protocol.InitializeResult, err error) { slog.Debug("Initialize called") defer slog.Debug("Initialize finished") + return s.initialize(ctx, params) } @@ -76,6 +77,7 @@ func (s *Server) ColorPresentation(ctx context.Context, params *protocol.ColorPr func (s *Server) Completion(ctx context.Context, params *protocol.CompletionParams) (result protocol.CompletionResult, err error) { slog.Debug("Completion called") defer slog.Debug("Completion finished") + return s.completion(ctx, params) } @@ -90,13 +92,16 @@ func (s *Server) Declaration(ctx context.Context, params *protocol.DeclarationPa func (s *Server) Definition(ctx context.Context, params *protocol.DefinitionParams) (result protocol.DefinitionResult, err error) { slog.Debug("Definition called") defer slog.Debug("Definition finished") + res, err := s.definition(ctx, params) + return protocol.LocationSlice(res), err } func (s *Server) DidChange(ctx context.Context, params *protocol.DidChangeTextDocumentParams) (err error) { slog.Debug("DidChange called") defer slog.Debug("DidChange finished") + return s.didChange(ctx, params) } @@ -119,6 +124,7 @@ func (s *Server) DidClose(ctx context.Context, params *protocol.DidCloseTextDocu func (s *Server) DidOpen(ctx context.Context, params *protocol.DidOpenTextDocumentParams) (err error) { slog.Debug("DidOpen called") defer slog.Debug("DidOpen finished") + return s.didOpen(ctx, params) } @@ -145,6 +151,7 @@ func (s *Server) DocumentLinkResolve(ctx context.Context, params *protocol.Docum func (s *Server) DocumentSymbol(ctx context.Context, params *protocol.DocumentSymbolParams) (result protocol.DocumentSymbolResult, err error) { slog.Debug("DocumentSymbol called") defer slog.Debug("DocumentSymbol finished") + return s.documentSymbol(ctx, params) } @@ -159,12 +166,14 @@ func (s *Server) FoldingRanges(ctx context.Context, params *protocol.FoldingRang func (s *Server) Formatting(ctx context.Context, params *protocol.DocumentFormattingParams) (result []protocol.TextEdit, err error) { slog.Debug("Formatting called") defer slog.Debug("Formatting finished") + return s.formatting(ctx, params) } func (s *Server) Hover(ctx context.Context, params *protocol.HoverParams) (result *protocol.Hover, err error) { slog.Debug("hover called") defer slog.Debug("hover finished") + return s.hover(ctx, params) } @@ -179,6 +188,7 @@ func (s *Server) OnTypeFormatting(ctx context.Context, params *protocol.Document func (s *Server) PrepareRename(ctx context.Context, params *protocol.PrepareRenameParams) (result protocol.PrepareRenameResult, err error) { slog.Debug("PrepareRename called") defer slog.Debug("PrepareRename finished") + return s.prepareRename(ctx, params) } @@ -189,12 +199,14 @@ func (s *Server) RangeFormatting(ctx context.Context, params *protocol.DocumentR func (s *Server) References(ctx context.Context, params *protocol.ReferenceParams) (result []protocol.Location, err error) { slog.Debug("References called") defer slog.Debug("References finished") + return s.references(ctx, params) } func (s *Server) Rename(ctx context.Context, params *protocol.RenameParams) (result *protocol.WorkspaceEdit, err error) { slog.Debug("Rename called") defer slog.Debug("Rename finished") + return s.rename(ctx, params) } @@ -209,7 +221,9 @@ func (s *Server) Symbols(ctx context.Context, params *protocol.WorkspaceSymbolPa func (s *Server) TypeDefinition(ctx context.Context, params *protocol.TypeDefinitionParams) (result protocol.DefinitionResult, err error) { slog.Debug("TypeDefinition called") defer slog.Debug("TypeDefinition finished") + res, err := s.typeDefinition(ctx, params) + return protocol.LocationSlice(res), err } diff --git a/lsp/stream.go b/lsp/stream.go index 19a8758..6e47253 100644 --- a/lsp/stream.go +++ b/lsp/stream.go @@ -48,5 +48,6 @@ func (s *StreamServer) ServeStream(ctx context.Context, conn jsonrpc2.Conn) erro ), )) <-conn.Done() + return conn.Err() } diff --git a/lsp/symbols.go b/lsp/symbols.go index d222ff7..259a098 100644 --- a/lsp/symbols.go +++ b/lsp/symbols.go @@ -10,10 +10,12 @@ import ( func (s *Server) documentSymbol(ctx context.Context, params *protocol.DocumentSymbolParams) (result protocol.DocumentSymbolSlice, err error) { file := params.TextDocument.URI + view, err := s.session.ViewOf(file) if err != nil { return nil, err } + ss, release := view.Snapshot() defer release() diff --git a/lsp/symbols/document.go b/lsp/symbols/document.go index 797c7fa..df3eee9 100644 --- a/lsp/symbols/document.go +++ b/lsp/symbols/document.go @@ -13,10 +13,12 @@ import ( // DocumentSymbols returns the document symbols of a file, in source order. func DocumentSymbols(ctx context.Context, ss *cache.Snapshot, file uri.URI) []*protocol.DocumentSymbol { res := make([]*protocol.DocumentSymbol, 0) + pf, err := ss.Parse(ctx, file) if err != nil { return res } + if pf.AST() == nil { return res } @@ -28,6 +30,7 @@ func DocumentSymbols(ctx context.Context, ss *cache.Snapshot, file uri.URI) []*p res = append(res, child) } } + return res } @@ -38,7 +41,9 @@ func nameRange(doc *syntax.Document, id *syntax.Identifier) protocol.Range { if id == nil { return protocol.Range{} } + start, end := doc.Range(id) + return protocol.Range{ Start: protocol.Position{ Line: uint32(start.Line - 1), @@ -72,6 +77,7 @@ func nodeSymbol(doc *syntax.Document, node syntax.Node) *protocol.DocumentSymbol case *syntax.Service: return serviceSymbol(doc, v) } + return nil } @@ -89,6 +95,7 @@ func structSymbol(doc *syntax.Document, st *syntax.Struct, detail string, kind p res.Children = append(res.Children, *child) } } + return res } @@ -109,6 +116,7 @@ func enumSymbol(doc *syntax.Document, enum *syntax.Enum) *protocol.DocumentSymbo } res.Children = append(res.Children, *child) } + return res } @@ -128,6 +136,7 @@ func serviceSymbol(doc *syntax.Document, svc *syntax.Service) *protocol.Document } res.Children = append(res.Children, *child) } + return res } diff --git a/main.go b/main.go index 614c900..775826b 100644 --- a/main.go +++ b/main.go @@ -11,6 +11,7 @@ import ( "github.com/urfave/cli/v3" + "github.com/karitham/thrift-ls/doc" "github.com/karitham/thrift-ls/formatter" tlog "github.com/karitham/thrift-ls/log" "github.com/karitham/thrift-ls/lsp" @@ -43,6 +44,27 @@ func main() { Flags: formatFlags(), Action: formatAction, }, + { + Name: "dump", + Usage: "dump the parse tree and document IR of a thrift file", + ArgsUsage: "", + Flags: []cli.Flag{ + &cli.BoolFlag{ + Name: "ir", + Usage: "also dump the formatted document IR with layout decisions", + }, + &cli.BoolFlag{ + Name: "ast", + Usage: "dump only the parse tree (tokens, trivia, node spans)", + }, + &cli.IntFlag{ + Name: "printWidth", + Usage: "line width for the IR dump", + Value: 80, + }, + }, + Action: dumpAction, + }, }, } @@ -70,9 +92,22 @@ func lspFlags() []cli.Flag { } } +// constructFlags maps the per-construct format flag names to constructs. +var constructFlags = []struct { + name string + construct formatter.Construct +}{ + {"struct", formatter.ConstructStruct}, + {"union", formatter.ConstructUnion}, + {"exception", formatter.ConstructException}, + {"enum", formatter.ConstructEnum}, + {"argument", formatter.ConstructArguments}, + {"throws", formatter.ConstructThrows}, +} + // formatFlags are the flags of the format subcommand. func formatFlags() []cli.Flag { - return []cli.Flag{ + flags := []cli.Flag{ &cli.BoolFlag{ Name: "w", Usage: "overwrite the file with the formatted result", @@ -87,28 +122,12 @@ func formatFlags() []cli.Flag { }, &cli.StringFlag{ Name: "indent", - Usage: `indentation: a literal like " " or "\t", a number like 8, or a legacy spec like "2spaces"`, + Usage: `indentation: a literal like " " or "\t"`, }, &cli.StringFlag{ Name: "align", Usage: `align fields: "field", "assign", or "disable"`, }, - &cli.StringFlag{ - Name: "field-separator", - Usage: `struct/enum field separators: "add", "remove", "semicolon", or "disable" to keep as written`, - }, - &cli.StringFlag{ - Name: "function-separator", - Usage: `service argument and throws separators: "add", "remove", "semicolon", or "disable" to keep as written`, - }, - &cli.BoolFlag{ - Name: "break-structs", - Usage: "always break struct, union, and exception bodies onto multiple lines", - }, - &cli.BoolFlag{ - Name: "break-enums", - Usage: "always break enum bodies onto multiple lines", - }, &cli.StringFlag{ Name: "config", Usage: "path to a thriftls.json config file", @@ -118,22 +137,39 @@ func formatFlags() []cli.Flag { Usage: "additional include path, like the thrift compiler (repeatable)", }, } + for _, cf := range constructFlags { + flags = append(flags, + &cli.StringFlag{ + Name: cf.name + "-separator", + Usage: fmt.Sprintf("%s separators: \"comma\", \"semicolon\", \"none\", or \"preserve\" to keep as written", cf.name), + }, + &cli.BoolFlag{ + Name: "break-" + cf.name, + Usage: fmt.Sprintf("always break %s bodies onto multiple lines", cf.name), + }, + ) + } + + return flags } // lspAction serves the language server on stdio. func lspAction(ctx context.Context, cmd *cli.Command) error { cfg := loadConfig(cmd.String("config"), ".") patch := options.Effective(cfg) + cli, err := lspPatch(cmd) if err != nil { return err } + patch = cli.Apply(patch) logLevelValue := 3 if patch.LogLevel != nil { logLevelValue = *patch.LogLevel } + tlog.Init(logLevelValue) fopts, err := patch.Formatter() @@ -149,79 +185,144 @@ func lspAction(ctx context.Context, cmd *cli.Command) error { ss := lsp.NewStreamServer(lspOpts) stream := jsonrpc2.NewStream(fakenet.NewConn("stdio", os.Stdin, os.Stdout)) conn := jsonrpc2.NewConn(stream) + err = ss.ServeStream(ctx, conn) if errors.Is(err, io.EOF) { return nil } + return err } // formatAction formats a single thrift file. func formatAction(ctx context.Context, cmd *cli.Command) error { file := cmd.Args().First() + cli, err := formatPatch(cmd) if err != nil { return err } + return formatFile(file, cmd.Bool("w"), cmd.Bool("d"), cmd.String("config"), cli) } +// dumpAction prints the parse tree, and optionally the formatted document +// IR with the printer's layout decisions. +func dumpAction(ctx context.Context, cmd *cli.Command) error { + file := cmd.Args().First() + if file == "" { + return errors.New("must specify a thrift file to dump, e.g. thriftls dump file.thrift") + } + + src, err := os.ReadFile(file) + if err != nil { + return err + } + + parsed, errs := syntax.Parse(src) + if parseErrors(errs) { + return fmt.Errorf("%s: file does not parse:\n%s", file, formatErrors(errs)) + } + + if !cmd.Bool("ir") { + fmt.Print(syntax.Dump(parsed)) + + return nil + } + + // The IR dump reflects the printer's layout decisions, so print first + // (the printer mutates the groups in place), then dump the tree. + fopts := formatter.DefaultOptions() + fopts.PrintWidth = cmd.Int("printWidth") + + ir := formatter.BuildIR(parsed, fopts) + if !cmd.Bool("ast") { + fmt.Print(syntax.Dump(parsed)) + fmt.Println("--- IR ---") + } + + if _, err := formatter.PrintIR(ir, fopts); err != nil { + return err + } + + fmt.Print(doc.Dump(ir)) + fmt.Println("--- output ---") + + formatted, err := formatter.Format(parsed, fopts) + if err != nil { + return err + } + + fmt.Print(formatted) + + return nil +} + // lspPatch builds an options patch from the explicitly set lsp flags. func lspPatch(cmd *cli.Command) (options.Patch, error) { p := options.Patch{} + if cmd.IsSet("logLevel") { v := cmd.Int("logLevel") p.LogLevel = &v } + if paths := cmd.StringSlice("I"); len(paths) > 0 { p.IncludePaths = &paths } + return p, nil } // formatPatch builds an options patch from the explicitly set format flags. func formatPatch(cmd *cli.Command) (options.Patch, error) { p := options.Patch{} + if cmd.IsSet("printWidth") { v := cmd.Int("printWidth") p.PrintWidth = &v } + if cmd.IsSet("indent") { ind, err := options.ParseIndentValue(cmd.String("indent")) if err != nil { return options.Patch{}, err } + p.Indent = &ind } + if cmd.IsSet("align") { v := cmd.String("align") p.Align = &v } - if cmd.IsSet("field-separator") { - v := cmd.String("field-separator") - p.Separators = &options.Separators{Fields: &v} - } - if cmd.IsSet("function-separator") { - v := cmd.String("function-separator") - if p.Separators == nil { - p.Separators = &options.Separators{} + + for _, cf := range constructFlags { + if cmd.IsSet(cf.name + "-separator") { + v := cmd.String(cf.name + "-separator") + + if p.Separators == nil { + p.Separators = &options.Separators{} + } + + p.Separators.Set(cf.construct, &v) } - p.Separators.Functions = &v - } - if cmd.IsSet("break-structs") { - v := cmd.Bool("break-structs") - p.Break = &options.Break{Structs: &v} - } - if cmd.IsSet("break-enums") { - v := cmd.Bool("break-enums") - if p.Break == nil { - p.Break = &options.Break{} + + if cmd.IsSet("break-" + cf.name) { + v := cmd.Bool("break-" + cf.name) + + if p.Break == nil { + p.Break = &options.Break{} + } + + p.Break.Set(cf.construct, &v) } - p.Break.Enums = &v } + if paths := cmd.StringSlice("I"); len(paths) > 0 { p.IncludePaths = &paths } + return p, nil } @@ -230,18 +331,22 @@ func formatPatch(cmd *cli.Command) (options.Patch, error) { func loadConfig(path, dir string) *options.Patch { if path == "" { var err error + path, err = options.FindConfig(dir) if err != nil { fatal(err) } } + if path == "" { return nil } + cfg, err := options.Load(path) if err != nil { fatal(err) } + return cfg } @@ -261,6 +366,7 @@ func formatFile(file string, write, diffOut bool, configPath string, cli options if err != nil { return err } + absFile, err := filepath.Abs(file) if err != nil { return err @@ -269,17 +375,18 @@ func formatFile(file string, write, diffOut bool, configPath string, cli options cfg := loadConfig(configPath, filepath.Dir(absFile)) patch := options.Effective(cfg) patch = cli.Apply(patch) + fopts, err := patch.Formatter() if err != nil { return err } - doc, errs := syntax.Parse(src) + parsed, errs := syntax.Parse(src) if parseErrors(errs) { return fmt.Errorf("%s: file does not parse:\n%s", file, formatErrors(errs)) } - out, err := formatter.Format(doc, fopts) + out, err := formatter.Format(parsed, fopts) if err != nil { return fmt.Errorf("%s: %w", file, err) } @@ -295,12 +402,15 @@ func formatFile(file string, write, diffOut bool, configPath string, cli options if info, err := os.Stat(file); err == nil { perms = info.Mode() } + return os.WriteFile(file, []byte(out), perms) case diffOut: fmt.Print(string(Diff("old", src, "new", []byte(out)))) + return nil default: fmt.Print(out) + return nil } } @@ -311,16 +421,19 @@ func parseErrors(errs []syntax.Error) bool { return true } } + return false } func formatErrors(errs []syntax.Error) string { var b strings.Builder + for _, e := range errs { if e.Severity == syntax.SeverityError { fmt.Fprintf(&b, " %s\n", e) } } + return b.String() } @@ -328,5 +441,6 @@ func derefStrings(p *[]string) []string { if p == nil { return nil } + return *p } diff --git a/options/options.go b/options/options.go index 03affd1..772a809 100644 --- a/options/options.go +++ b/options/options.go @@ -16,7 +16,6 @@ import ( "os" "path/filepath" "slices" - "strconv" "strings" "github.com/karitham/thrift-ls/formatter" @@ -25,22 +24,98 @@ import ( // ConfigFileName is the JSON config file name. const ConfigFileName = "thriftls.json" -// Separators configures trailing separators for the two field contexts. +// Separators configures trailing separators per construct. A nil value is +// unset. type Separators struct { - // Fields controls separators after struct/union/exception fields and - // enum values. - Fields *string `json:"fields"` - // Functions controls separators after service arguments and throws - // entries. - Functions *string `json:"functions"` + Structs *string `json:"structs"` + Unions *string `json:"unions"` + Exceptions *string `json:"exceptions"` + Enums *string `json:"enums"` + Arguments *string `json:"arguments"` + Throws *string `json:"throws"` } -// Break configures layouts that are forced multiline. +// Get returns the value for the construct. +func (s Separators) Get(c formatter.Construct) *string { + switch c { + case formatter.ConstructUnion: + return s.Unions + case formatter.ConstructException: + return s.Exceptions + case formatter.ConstructEnum: + return s.Enums + case formatter.ConstructArguments: + return s.Arguments + case formatter.ConstructThrows: + return s.Throws + } + + return s.Structs +} + +// Set assigns the value for the construct. +func (s *Separators) Set(c formatter.Construct, v *string) { + switch c { + case formatter.ConstructUnion: + s.Unions = v + case formatter.ConstructException: + s.Exceptions = v + case formatter.ConstructEnum: + s.Enums = v + case formatter.ConstructArguments: + s.Arguments = v + case formatter.ConstructThrows: + s.Throws = v + default: + s.Structs = v + } +} + +// Break configures layouts that are forced multiline per construct. A nil +// value is unset. type Break struct { - // Structs forces struct, union, and exception bodies multiline. - Structs *bool `json:"structs"` - // Enums forces enum bodies multiline. - Enums *bool `json:"enums"` + Structs *bool `json:"structs"` + Unions *bool `json:"unions"` + Exceptions *bool `json:"exceptions"` + Enums *bool `json:"enums"` + Arguments *bool `json:"arguments"` + Throws *bool `json:"throws"` +} + +// Get returns the value for the construct. +func (b Break) Get(c formatter.Construct) *bool { + switch c { + case formatter.ConstructUnion: + return b.Unions + case formatter.ConstructException: + return b.Exceptions + case formatter.ConstructEnum: + return b.Enums + case formatter.ConstructArguments: + return b.Arguments + case formatter.ConstructThrows: + return b.Throws + } + + return b.Structs +} + +// Set assigns the value for the construct. +func (b *Break) Set(c formatter.Construct, v *bool) { + switch c { + case formatter.ConstructUnion: + b.Unions = v + case formatter.ConstructException: + b.Exceptions = v + case formatter.ConstructEnum: + b.Enums = v + case formatter.ConstructArguments: + b.Arguments = v + case formatter.ConstructThrows: + b.Throws = v + default: + b.Structs = v + } } // Patch is a partial set of options; nil fields are unset. @@ -62,43 +137,51 @@ func (p Patch) Apply(base Patch) Patch { if p.PrintWidth != nil { out.PrintWidth = p.PrintWidth } + if p.Indent != nil { out.Indent = p.Indent } + if p.TabWidth != nil { out.TabWidth = p.TabWidth } + if p.Align != nil { out.Align = p.Align } + if p.Separators != nil { if out.Separators == nil { out.Separators = &Separators{} } - if p.Separators.Fields != nil { - out.Separators.Fields = p.Separators.Fields - } - if p.Separators.Functions != nil { - out.Separators.Functions = p.Separators.Functions + + for _, c := range formatter.AllConstructs { + if v := p.Separators.Get(c); v != nil { + out.Separators.Set(c, v) + } } } + if p.Break != nil { if out.Break == nil { out.Break = &Break{} } - if p.Break.Structs != nil { - out.Break.Structs = p.Break.Structs - } - if p.Break.Enums != nil { - out.Break.Enums = p.Break.Enums + + for _, c := range formatter.AllConstructs { + if v := p.Break.Get(c); v != nil { + out.Break.Set(c, v) + } } } + if p.IncludePaths != nil { out.IncludePaths = p.IncludePaths } + if p.LogLevel != nil { out.LogLevel = p.LogLevel } + return out } @@ -108,7 +191,15 @@ func Default() Patch { indent := Indent{Value: " ", Width: 4} tabWidth := 4 align := "field" - separators := Separators{Fields: new("disable"), Functions: new("disable")} + separators := Separators{ + Structs: new("preserve"), + Unions: new("preserve"), + Exceptions: new("preserve"), + Enums: new("preserve"), + Arguments: new("preserve"), + Throws: new("preserve"), + } + return Patch{ PrintWidth: &printWidth, Indent: &indent, @@ -123,35 +214,30 @@ func (p Patch) Validate() error { if p.PrintWidth != nil && *p.PrintWidth <= 0 { return errors.New("printWidth must be positive") } + if p.TabWidth != nil && *p.TabWidth <= 0 { return errors.New("tabWidth must be positive") } - if p.Align != nil && !oneOf(*p.Align, "field", "assign", "disable") { + + if p.Align != nil && !slices.Contains([]string{"field", "assign", "disable"}, *p.Align) { return fmt.Errorf("align must be one of \"field\", \"assign\", \"disable\", got %q", *p.Align) } + if p.Separators != nil { - for _, v := range []struct { - name string - value *string - }{ - {"separators.fields", p.Separators.Fields}, - {"separators.functions", p.Separators.Functions}, - } { - if v.value != nil && !oneOf(*v.value, "add", "remove", "semicolon", "disable", "preserve") { - return fmt.Errorf("%s must be one of \"add\", \"remove\", \"semicolon\", \"disable\" (keep as written), got %q", v.name, *v.value) + for _, c := range formatter.AllConstructs { + if v := p.Separators.Get(c); v != nil && !slices.Contains([]string{"comma", "semicolon", "none", "preserve"}, *v) { + return fmt.Errorf("separators.%s must be one of \"comma\", \"semicolon\", \"none\", \"preserve\" (keep as written), got %q", c, *v) } } } + if p.Indent != nil { if p.Indent.Width <= 0 || !isWhitespaceOnly(p.Indent.Value) { return errors.New("indent must be a string of spaces or tabs") } } - return nil -} -func oneOf(s string, options ...string) bool { - return slices.Contains(options, s) + return nil } // Formatter converts the patch to formatter options, validating first. @@ -159,17 +245,21 @@ func (p Patch) Formatter() (formatter.Options, error) { if err := p.Validate(); err != nil { return formatter.Options{}, err } + o := formatter.DefaultOptions() if p.PrintWidth != nil { o.PrintWidth = *p.PrintWidth } + if p.Indent != nil { o.Indent = p.Indent.Value o.TabWidth = p.Indent.Width } + if p.TabWidth != nil { o.TabWidth = *p.TabWidth } + if p.Align != nil { switch *p.Align { case "field": @@ -180,22 +270,23 @@ func (p Patch) Formatter() (formatter.Options, error) { o.Align = formatter.AlignDisable } } + if p.Separators != nil { - if p.Separators.Fields != nil { - o.FieldSeparator = separatorMode(*p.Separators.Fields) - } - if p.Separators.Functions != nil { - o.FunctionSeparator = separatorMode(*p.Separators.Functions) + for _, c := range formatter.AllConstructs { + if v := p.Separators.Get(c); v != nil { + o.Separator.Set(c, separatorMode(*v)) + } } } + if p.Break != nil { - if p.Break.Structs != nil { - o.BreakStructs = *p.Break.Structs - } - if p.Break.Enums != nil { - o.BreakEnums = *p.Break.Enums + for _, c := range formatter.AllConstructs { + if v := p.Break.Get(c); v != nil { + o.Break.Set(c, *v) + } } } + return o, nil } @@ -203,47 +294,40 @@ func (p Patch) Formatter() (formatter.Options, error) { // value is validated before this is called. func separatorMode(s string) formatter.SeparatorMode { switch s { - case "add": + case "comma": return formatter.SeparatorComma - case "remove": - return formatter.SeparatorNone case "semicolon": return formatter.SeparatorSemicolon - default: // "disable", "preserve" + case "none": + return formatter.SeparatorNone + default: // "preserve" return formatter.SeparatorPreserve } } // Indent is a resolved indentation: the string emitted for one level and // its display width. It is set from a config value that may be a literal -// string of spaces or tabs, a number of spaces, or a legacy spec like -// "2spaces" or "1tab". +// string of spaces or tabs, or a number of spaces. type Indent struct { Value string // the indentation string, spaces or tabs Width int // display width of one level } -// UnmarshalJSON accepts a number (spaces), a literal string of spaces or -// tabs, or a legacy spec string. +// UnmarshalJSON accepts a number (spaces) or a literal string of spaces +// or tabs. func (i *Indent) UnmarshalJSON(data []byte) error { - var n int - if err := json.Unmarshal(data, &n); err == nil { - ind, err := ParseIndentValue(strconv.Itoa(n)) - if err != nil { - return err - } - *i = ind - return nil - } var s string if err := json.Unmarshal(data, &s); err != nil { - return errors.New("indent must be a string of spaces or tabs, a number, or a legacy spec like \"2spaces\"") + return errors.New("indent must be a string of spaces or tabs") } + ind, err := ParseIndentValue(s) if err != nil { return err } + *i = ind + return nil } @@ -252,31 +336,29 @@ func (i *Indent) UnmarshalJSON(data []byte) error { // " " literal spaces, used as written // "\t" literal tabs, used as written // "8" a number of spaces -// "2spaces", "1tab", "tab" legacy specs, kept as aliases // // An empty spec yields the default of four spaces. func ParseIndentValue(s string) (Indent, error) { if s == "" { return Indent{Value: " ", Width: 4}, nil } - if n, err := strconv.Atoi(s); err == nil { - if n <= 0 { - return Indent{}, errors.New("indent must be a positive number of spaces") - } - return Indent{Value: strings.Repeat(" ", n), Width: n}, nil - } + if isWhitespaceOnly(s) { spaces := strings.Count(s, " ") + tabs := strings.Count(s, "\t") if spaces > 0 && tabs > 0 { return Indent{}, fmt.Errorf("indent %q mixes spaces and tabs", s) } + if tabs > 0 { return Indent{Value: s, Width: tabs * 4}, nil } + return Indent{Value: s, Width: spaces}, nil } - return ParseLegacyIndent(s) + + return Indent{}, errors.New("indent must be a string of spaces or tabs") } func isWhitespaceOnly(s string) bool { @@ -285,36 +367,8 @@ func isWhitespaceOnly(s string) bool { return false } } - return true -} -// ParseLegacyIndent parses the legacy indent specs ("4spaces", "1tab", -// "2tabs", "tab"), kept for compatibility. -func ParseLegacyIndent(s string) (Indent, error) { - lower := strings.ToLower(s) - num := 1 - unit := "" - for _, suffix := range []string{"spaces", "space", "tabs", "tab"} { - if strings.HasSuffix(lower, suffix) { - unit = suffix - prefix := strings.TrimSuffix(lower, suffix) - if prefix != "" { - n, err := strconv.Atoi(prefix) - if err != nil || n <= 0 { - return Indent{}, fmt.Errorf("invalid indent %q: use a literal like \" \", a number like 8, or a legacy spec like \"2spaces\"", s) - } - num = n - } - break - } - } - if unit == "" { - return Indent{}, fmt.Errorf("invalid indent %q: use a literal like \" \", a number like 8, or a legacy spec like \"2spaces\"", s) - } - if strings.HasPrefix(unit, "tab") { - return Indent{Value: strings.Repeat("\t", num), Width: num * 4}, nil - } - return Indent{Value: strings.Repeat(" ", num), Width: num}, nil + return true } // Load reads and parses a config file. Unknown keys are rejected so that @@ -324,12 +378,16 @@ func Load(path string) (*Patch, error) { if err != nil { return nil, err } + var p Patch + dec := json.NewDecoder(bytes.NewReader(data)) dec.DisallowUnknownFields() + if err := dec.Decode(&p); err != nil { return nil, fmt.Errorf("options: %s: %w", path, err) } + if err := p.Validate(); err != nil { return nil, fmt.Errorf("options: %s: %w", path, err) } @@ -345,8 +403,10 @@ func Load(path string) (*Patch, error) { abs = append(abs, filepath.Join(filepath.Dir(path), ip)) } } + p.IncludePaths = &abs } + return &p, nil } @@ -357,6 +417,7 @@ func FindConfig(dir string) (string, error) { if path := os.Getenv("THRIFTLS_CONFIG"); path != "" { return path, nil } + for d := dir; ; d = filepath.Dir(d) { path := filepath.Join(d, ConfigFileName) if _, err := os.Stat(path); err == nil { @@ -364,6 +425,7 @@ func FindConfig(dir string) (string, error) { } else if !os.IsNotExist(err) { return "", err } + if d == filepath.Dir(d) { return "", nil } @@ -377,5 +439,6 @@ func Effective(cfg *Patch) Patch { if cfg != nil { p = cfg.Apply(p) } + return p } diff --git a/options/options_test.go b/options/options_test.go index 262a2bf..8426198 100644 --- a/options/options_test.go +++ b/options/options_test.go @@ -21,17 +21,8 @@ func TestParseIndentValue(t *testing.T) { {"literal four spaces", " ", Indent{" ", 4}, false}, {"literal tab", "\t", Indent{"\t", 4}, false}, {"literal two tabs", "\t\t", Indent{"\t\t", 8}, false}, - {"number", "8", Indent{" ", 8}, false}, - {"number one", "1", Indent{" ", 1}, false}, - {"legacy spaces", "2spaces", Indent{" ", 2}, false}, - {"legacy space singular", "1space", Indent{" ", 1}, false}, - {"legacy tab", "1tab", Indent{"\t", 4}, false}, - {"legacy bare tab", "tab", Indent{"\t", 4}, false}, - {"legacy two tabs", "2tabs", Indent{"\t\t", 8}, false}, - {"zero number", "0", Indent{}, true}, {"mixed spaces and tabs", " \t", Indent{}, true}, {"garbage", "banana", Indent{}, true}, - {"negative", "-2", Indent{}, true}, } for _, tt := range tests { t.Run(tt.name, func(t *testing.T) { @@ -40,11 +31,14 @@ func TestParseIndentValue(t *testing.T) { if err == nil { t.Fatalf("expected error, got %+v", got) } + return } + if err != nil { t.Fatalf("unexpected error: %v", err) } + if got != tt.want { t.Errorf("got %+v, want %+v", got, tt.want) } @@ -60,8 +54,6 @@ func TestIndentUnmarshal(t *testing.T) { }{ {"string spaces", `" "`, Indent{" ", 2}}, {"string tab", `"\t"`, Indent{"\t", 4}}, - {"number", `8`, Indent{" ", 8}}, - {"legacy", `"2spaces"`, Indent{" ", 2}}, } for _, tt := range tests { t.Run(tt.name, func(t *testing.T) { @@ -69,6 +61,7 @@ func TestIndentUnmarshal(t *testing.T) { if err := json.Unmarshal([]byte(tt.json), &i); err != nil { t.Fatalf("unmarshal: %v", err) } + if i != tt.want { t.Errorf("got %+v, want %+v", i, tt.want) } @@ -87,6 +80,7 @@ func TestPatchApply(t *testing.T) { if got.PrintWidth == nil || *got.PrintWidth != 100 { t.Errorf("PrintWidth not overridden: %v", got.PrintWidth) } + if got.Align == nil || *got.Align != "field" { t.Errorf("Align should stay from base: %v", got.Align) } @@ -105,8 +99,8 @@ func TestPatchValidate(t *testing.T) { {"bad printWidth", Patch{PrintWidth: intPtr(0)}, true}, {"bad tabWidth", Patch{TabWidth: intPtr(-1)}, true}, {"bad align", Patch{Align: strPtr("sideways")}, true}, - {"bad comma", Patch{Separators: &Separators{Fields: strPtr("maybe")}}, true}, - {"preserve alias", Patch{Separators: &Separators{Fields: strPtr("preserve")}}, false}, + {"bad comma", Patch{Separators: &Separators{Structs: strPtr("maybe")}}, true}, + {"preserve alias", Patch{Separators: &Separators{Structs: strPtr("preserve")}}, false}, {"bad indent value", Patch{Indent: &Indent{Value: "x", Width: 1}}, true}, } for _, tt := range tests { @@ -115,6 +109,7 @@ func TestPatchValidate(t *testing.T) { if tt.wantErr && err == nil { t.Error("expected error") } + if !tt.wantErr && err != nil { t.Errorf("unexpected error: %v", err) } @@ -125,31 +120,37 @@ func TestPatchValidate(t *testing.T) { func TestPatchFormatter(t *testing.T) { indent := Indent{Value: " ", Width: 2} p := Patch{Indent: &indent, PrintWidth: new(100)} + o, err := p.Formatter() if err != nil { t.Fatalf("Formatter: %v", err) } + if o.PrintWidth != 100 || o.Indent != " " || o.TabWidth != 2 { t.Errorf("got %+v", o) } - if o.Align != formatter.AlignField || o.FieldSeparator != formatter.SeparatorPreserve { + + if o.Align != formatter.AlignField || o.Separator.Get(formatter.ConstructStruct) != formatter.SeparatorPreserve { t.Errorf("defaults wrong: %+v", o) } - comma := "add" + comma := "comma" align := "assign" - p = Patch{Separators: &Separators{Fields: &comma}, Align: &align} + p = Patch{Separators: &Separators{Structs: &comma}, Align: &align} + o, err = p.Formatter() if err != nil { t.Fatalf("Formatter: %v", err) } - if o.FieldSeparator != formatter.SeparatorComma || o.Align != formatter.AlignAssign { + + if o.Separator.Get(formatter.ConstructStruct) != formatter.SeparatorComma || o.Align != formatter.AlignAssign { t.Errorf("got %+v", o) } } func TestFindConfig(t *testing.T) { dir := t.TempDir() + sub := filepath.Join(dir, "a", "b") if err := os.MkdirAll(sub, 0o755); err != nil { t.Fatal(err) @@ -166,6 +167,7 @@ func TestFindConfig(t *testing.T) { if err := os.WriteFile(cfgPath, []byte("{}"), 0o644); err != nil { t.Fatal(err) } + got, err = FindConfig(sub) if err != nil || got != cfgPath { t.Fatalf("FindConfig = %q, %v; want %q", got, err, cfgPath) @@ -176,6 +178,7 @@ func TestFindConfig(t *testing.T) { if err := os.WriteFile(near, []byte("{}"), 0o644); err != nil { t.Fatal(err) } + got, err = FindConfig(sub) if err != nil || got != near { t.Fatalf("FindConfig = %q, %v; want %q", got, err, near) @@ -185,6 +188,7 @@ func TestFindConfig(t *testing.T) { func TestLoadAndEffective(t *testing.T) { dir := t.TempDir() cfgPath := filepath.Join(dir, "thriftls.json") + content := `{ "printWidth": 100, "indent": " ", @@ -198,9 +202,11 @@ func TestLoadAndEffective(t *testing.T) { if err != nil { t.Fatalf("Load: %v", err) } + if cfg.PrintWidth == nil || *cfg.PrintWidth != 100 { t.Errorf("printWidth = %v, want 100", cfg.PrintWidth) } + if cfg.Indent == nil || cfg.Indent.Value != " " { t.Errorf("indent = %+v, want two spaces", cfg.Indent) } @@ -210,9 +216,10 @@ func TestLoadAndEffective(t *testing.T) { t.Errorf("effective printWidth = %v, want 100", p.PrintWidth) } // Unset config fields keep their defaults. - if p.Separators == nil || p.Separators.Fields == nil || *p.Separators.Fields != "disable" { - t.Errorf("comma = %v, want default disable", p.Separators) + if p.Separators == nil || p.Separators.Structs == nil || *p.Separators.Structs != "preserve" { + t.Errorf("comma = %v, want default preserve", p.Separators) } + if p.Indent == nil || p.Indent.Value != " " { t.Errorf("indent = %+v, want two spaces", p.Indent) } @@ -222,6 +229,7 @@ func TestLoadAndEffective(t *testing.T) { if d.PrintWidth == nil || *d.PrintWidth != 80 { t.Errorf("default printWidth = %v, want 80", d.PrintWidth) } + if d.Indent == nil || d.Indent.Value != " " { t.Errorf("default indent = %+v, want four spaces", d.Indent) } @@ -232,6 +240,7 @@ func TestLoadRejectsUnknownOverrideKeys(t *testing.T) { // rather than silently ignoring per-file settings. dir := t.TempDir() cfgPath := filepath.Join(dir, "thriftls.json") + content := `{ "printWidth": 100, "overrides": [ @@ -241,6 +250,7 @@ func TestLoadRejectsUnknownOverrideKeys(t *testing.T) { if err := os.WriteFile(cfgPath, []byte(content), 0o644); err != nil { t.Fatal(err) } + if _, err := Load(cfgPath); err == nil { t.Fatal("Load accepted a config with an overrides key") } @@ -253,33 +263,39 @@ func TestPatchSeparatorModes(t *testing.T) { field formatter.SeparatorMode function formatter.SeparatorMode }{ - {"add", formatter.SeparatorComma, formatter.SeparatorComma}, - {"remove", formatter.SeparatorNone, formatter.SeparatorNone}, + {"comma", formatter.SeparatorComma, formatter.SeparatorComma}, + {"none", formatter.SeparatorNone, formatter.SeparatorNone}, {"semicolon", formatter.SeparatorSemicolon, formatter.SeparatorSemicolon}, - {"disable", formatter.SeparatorPreserve, formatter.SeparatorPreserve}, + {"preserve", formatter.SeparatorPreserve, formatter.SeparatorPreserve}, {"preserve", formatter.SeparatorPreserve, formatter.SeparatorPreserve}, } for _, tt := range tests { t.Run(tt.value, func(t *testing.T) { - p := Patch{Separators: &Separators{Fields: &tt.value, Functions: &tt.value}} + p := Patch{Separators: &Separators{Structs: &tt.value, Unions: &tt.value, Exceptions: &tt.value, Enums: &tt.value, Arguments: &tt.value, Throws: &tt.value}} + o, err := p.Formatter() if err != nil { t.Fatalf("Formatter: %v", err) } - if o.FieldSeparator != tt.field || o.FunctionSeparator != tt.function { - t.Errorf("value %q: field=%v function=%v", tt.value, o.FieldSeparator, o.FunctionSeparator) + + for _, c := range formatter.AllConstructs { + if o.Separator.Get(c) != tt.field { + t.Errorf("value %q: construct %s = %v, want %v", tt.value, c, o.Separator.Get(c), tt.field) + } } }) } // The two options map independently. - semicolon, add := "semicolon", "add" - p := Patch{Separators: &Separators{Fields: &semicolon, Functions: &add}} + semicolon, comma := "semicolon", "comma" + p := Patch{Separators: &Separators{Structs: &semicolon, Enums: &semicolon, Arguments: &comma, Throws: &comma}} + o, err := p.Formatter() if err != nil { t.Fatalf("Formatter: %v", err) } - if o.FieldSeparator != formatter.SeparatorSemicolon || o.FunctionSeparator != formatter.SeparatorComma { + + if o.Separator.Get(formatter.ConstructStruct) != formatter.SeparatorSemicolon || o.Separator.Get(formatter.ConstructArguments) != formatter.SeparatorComma { t.Errorf("independent mapping failed: %+v", o) } } @@ -289,11 +305,13 @@ func TestPatchBreak(t *testing.T) { trueVal, falseVal := true, false p := Patch{Break: &Break{Structs: &trueVal, Enums: &falseVal}} + o, err := p.Formatter() if err != nil { t.Fatalf("Formatter: %v", err) } - if !o.BreakStructs || o.BreakEnums { + + if !o.Break.Get(formatter.ConstructStruct) || o.Break.Get(formatter.ConstructEnum) { t.Errorf("break mapping wrong: %+v", o) } @@ -302,7 +320,10 @@ func TestPatchBreak(t *testing.T) { if err != nil { t.Fatalf("Formatter: %v", err) } - if o.BreakStructs || o.BreakEnums { - t.Errorf("breaks should default to false: %+v", o) + + for _, c := range formatter.AllConstructs { + if o.Break.Get(c) { + t.Errorf("breaks should default to false for %s: %+v", c, o) + } } } diff --git a/resolver/integration_test.go b/resolver/integration_test.go index 89b39df..e3d3b32 100644 --- a/resolver/integration_test.go +++ b/resolver/integration_test.go @@ -21,6 +21,7 @@ func TestResolver_Integration_NestedIncludes(t *testing.T) { // Create base.thrift (no includes) baseFile := filepath.Join(includeDir, "base.thrift") + baseContent := `namespace * base struct BaseID { @@ -32,6 +33,7 @@ struct BaseID { // Create middle.thrift (includes base) middleFile := filepath.Join(includeDir, "middle.thrift") + middleContent := `include "base.thrift" namespace * middle @@ -46,6 +48,7 @@ struct UserID { // Create main.thrift (includes middle, which includes base) mainFile := filepath.Join(tmpDir, "main.thrift") + mainContent := `include "middle.thrift" namespace * main @@ -65,6 +68,7 @@ struct User { if filename != middleFile { t.Errorf("expected %q, got %q", middleFile, filename) } + content, err := os.ReadFile(filename) if err != nil { t.Fatalf("failed to read middle.thrift: %v", err) @@ -77,9 +81,11 @@ struct User { t.Fatalf("failed to parse middle.thrift: %v", errs) } } + if len(middleDoc.Includes()) == 0 { t.Fatal("expected middle.thrift to have includes") } + includePath := middleDoc.Includes()[0].Path.Text if strings.Trim(includePath, "\"'") != "base.thrift" { t.Errorf("expected include path 'base.thrift', got %q", includePath) @@ -96,6 +102,7 @@ struct User { if err != nil { t.Fatalf("failed to read base.thrift: %v", err) } + if string(content2) != baseContent { t.Errorf("base.thrift content mismatch") } @@ -110,12 +117,14 @@ func TestResolver_Integration_ResolutionOrder(t *testing.T) { if err := os.MkdirAll(includeDir, 0o755); err != nil { t.Fatal(err) } + if err := os.MkdirAll(srcDir, 0o755); err != nil { t.Fatal(err) } // Create a version in include directory includeVersion := filepath.Join(includeDir, "shared.thrift") + includeContent := `namespace * shared struct SharedInInclude { @@ -127,6 +136,7 @@ struct SharedInInclude { // Create a different version in src directory srcVersion := filepath.Join(srcDir, "shared.thrift") + srcContent := `namespace * shared struct SharedInSrc { @@ -138,6 +148,7 @@ struct SharedInSrc { // Create main file in src mainFile := filepath.Join(srcDir, "main.thrift") + mainContent := `include "shared.thrift" namespace * test @@ -154,6 +165,7 @@ struct Data { // Test: IncludeCall should find include path first, not relative includeCall := r.IncludeCall(mainFile) + filename, content, err := includeCall("shared.thrift") if err != nil { t.Fatalf("failed to resolve shared.thrift: %v", err) @@ -163,6 +175,7 @@ struct Data { if filename != includeVersion { t.Errorf("expected include path version %q, got %q", includeVersion, filename) } + if string(content) != includeContent { t.Error("expected content from include directory") } @@ -174,6 +187,7 @@ func TestResolver_Integration_RelativeFallback(t *testing.T) { // Create file in root (no include directories) localFile := filepath.Join(tmpDir, "local.thrift") + localContent := `namespace * local struct LocalData { @@ -185,6 +199,7 @@ struct LocalData { // Create main file that references local mainFile := filepath.Join(tmpDir, "main.thrift") + mainContent := `include "local.thrift" namespace * main @@ -201,6 +216,7 @@ struct Container { // Test: Should fall back to relative resolution includeCall := r.IncludeCall(mainFile) + filename, content, err := includeCall("local.thrift") if err != nil { t.Fatalf("failed to resolve local.thrift: %v", err) @@ -209,6 +225,7 @@ struct Container { if filename != localFile { t.Errorf("expected %q, got %q", localFile, filename) } + if string(content) != localContent { t.Error("content mismatch") } @@ -223,12 +240,14 @@ func TestResolver_Integration_MultipleIncludePaths(t *testing.T) { if err := os.MkdirAll(includeDir1, 0o755); err != nil { t.Fatal(err) } + if err := os.MkdirAll(includeDir2, 0o755); err != nil { t.Fatal(err) } // File only in include path 2 file2 := filepath.Join(includeDir2, "unique.thrift") + content2 := `namespace * unique struct UniqueInDir2 { @@ -240,6 +259,7 @@ struct UniqueInDir2 { // File in both include paths (dir1 should win) fileBoth1 := filepath.Join(includeDir1, "both.thrift") + contentBoth1 := `namespace * both struct BothFromDir1 { @@ -250,6 +270,7 @@ struct BothFromDir1 { } fileBoth2 := filepath.Join(includeDir2, "both.thrift") + contentBoth2 := `namespace * both struct BothFromDir2 { @@ -260,6 +281,7 @@ struct BothFromDir2 { } mainFile := filepath.Join(tmpDir, "main.thrift") + mainContent := `include "unique.thrift" include "both.thrift"` if err := os.WriteFile(mainFile, []byte(mainContent), 0o644); err != nil { @@ -275,6 +297,7 @@ include "both.thrift"` if err != nil { t.Fatalf("failed to resolve unique.thrift: %v", err) } + if filename != file2 { t.Errorf("expected %q, got %q", file2, filename) } @@ -284,6 +307,7 @@ include "both.thrift"` if err != nil { t.Fatalf("failed to resolve both.thrift: %v", err) } + if filename2 != fileBoth1 { t.Errorf("expected first include path %q, got %q", fileBoth1, filename2) } @@ -300,6 +324,7 @@ func TestResolver_Integration_DeeplyNestedIncludes(t *testing.T) { // d.thrift (no includes) dFile := filepath.Join(includeDir, "d.thrift") + dContent := `namespace * d struct D { @@ -311,6 +336,7 @@ struct D { // c.thrift includes d cFile := filepath.Join(includeDir, "c.thrift") + cContent := `include "d.thrift" namespace * c @@ -324,6 +350,7 @@ struct C { // b.thrift includes c bFile := filepath.Join(includeDir, "b.thrift") + bContent := `include "c.thrift" namespace * b @@ -337,6 +364,7 @@ struct B { // a.thrift includes b aFile := filepath.Join(includeDir, "a.thrift") + aContent := `include "b.thrift" namespace * a @@ -357,16 +385,19 @@ struct A { if err != nil { return nil, err } + doc, errs := syntax.Parse(content) for _, e := range errs { if e.Severity == syntax.SeverityError { return nil, fmt.Errorf("parse %s: %v", file, errs) } } + var paths []string for _, inc := range doc.Includes() { paths = append(paths, strings.Trim(inc.Path.Text, "\"'")) } + return paths, nil } @@ -376,9 +407,11 @@ struct A { if err != nil { t.Fatal(err) } + if len(includes) != 1 { t.Fatalf("%s: expected one include, got %v", chain[i], includes) } + next := r.Resolve(chain[i], includes[0]) if next != chain[i+1] { t.Errorf("expected %q, got %q", chain[i+1], next) diff --git a/resolver/resolver.go b/resolver/resolver.go index 2f75ec8..b8d84b4 100644 --- a/resolver/resolver.go +++ b/resolver/resolver.go @@ -75,6 +75,7 @@ func (r *Resolver) Resolve(currentFile, includePath string) string { // exists reports whether the file exists on the resolver's filesystem. func (r *Resolver) exists(path string) bool { _, err := fs.Stat(r.fsys, path) + return err == nil } @@ -84,10 +85,12 @@ type IncludeCall func(include string) (filename string, content []byte, err erro // ResolveContent resolves an include path and reads the file content. func (r *Resolver) ResolveContent(currentFile, includePath string) (filename string, content []byte, err error) { filename = r.Resolve(currentFile, includePath) + content, err = fs.ReadFile(r.fsys, filename) if err != nil { return filename, nil, err } + return filename, content, nil } @@ -100,9 +103,11 @@ func (r *Resolver) IncludeCall(initialFile string) IncludeCall { candidatePath := filepath.Join(ip, include) if r.exists(candidatePath) { content, err = fs.ReadFile(r.fsys, candidatePath) + return candidatePath, content, err } } + return r.ResolveContent(initialFile, include) } } diff --git a/resolver/resolver_test.go b/resolver/resolver_test.go index 7042517..26de238 100644 --- a/resolver/resolver_test.go +++ b/resolver/resolver_test.go @@ -17,6 +17,7 @@ import ( // process working directory. func TestConfigRelativeIncludePaths(t *testing.T) { dir := t.TempDir() + cfgPath := filepath.Join(dir, "thriftls.json") if err := os.WriteFile(cfgPath, []byte(`{"includePaths": ["project/base"]}`), 0o644); err != nil { t.Fatal(err) @@ -26,9 +27,11 @@ func TestConfigRelativeIncludePaths(t *testing.T) { if err != nil { t.Fatalf("Load: %v", err) } + if cfg.IncludePaths == nil || len(*cfg.IncludePaths) != 1 { t.Fatalf("includePaths = %v", cfg.IncludePaths) } + want := filepath.Join(dir, "project", "base") if got := (*cfg.IncludePaths)[0]; got != want { t.Errorf("includePaths[0] = %q, want %q", got, want) @@ -40,6 +43,7 @@ func TestConfigRelativeIncludePaths(t *testing.T) { baseFile := filepath.Join(dir, "project", "base", "types.thrift") fsys := absMapFS{baseFile: []byte("struct T {}")} r := NewWithFS(*cfg.IncludePaths, fsys) + cur := filepath.Join(dir, "project", "app.thrift") if got := r.Resolve(cur, "types.thrift"); got != baseFile { t.Errorf("Resolve = %q, want %q", got, baseFile) @@ -54,6 +58,7 @@ func (m absMapFS) Stat(name string) (fs.FileInfo, error) { if _, ok := m[name]; !ok { return nil, &fs.PathError{Op: "stat", Path: name, Err: fs.ErrNotExist} } + return absMapFileInfo{name: name}, nil } @@ -62,6 +67,7 @@ func (m absMapFS) Open(name string) (fs.File, error) { if !ok { return nil, &fs.PathError{Op: "open", Path: name, Err: fs.ErrNotExist} } + return &absMapFile{Reader: bytes.NewReader(data), info: absMapFileInfo{name: name}}, nil } @@ -86,14 +92,17 @@ func (f *absMapFile) Close() error { return nil } func TestConfigAbsoluteIncludePaths(t *testing.T) { dir := t.TempDir() cfgPath := filepath.Join(dir, "thriftls.json") + abs := filepath.Join(dir, "elsewhere") if err := os.WriteFile(cfgPath, []byte(`{"includePaths": ["`+abs+`"]}`), 0o644); err != nil { t.Fatal(err) } + cfg, err := options.Load(cfgPath) if err != nil { t.Fatalf("Load: %v", err) } + if got := (*cfg.IncludePaths)[0]; got != abs { t.Errorf("includePaths[0] = %q, want %q", got, abs) } diff --git a/syntax/accessors.go b/syntax/accessors.go index ec46b03..acfdcce 100644 --- a/syntax/accessors.go +++ b/syntax/accessors.go @@ -7,109 +7,129 @@ package syntax // Includes returns the thrift include headers in source order. func (d *Document) Includes() []*Include { var out []*Include + for _, n := range d.Nodes { if v, ok := n.(*Include); ok { out = append(out, v) } } + return out } // CPPIncludes returns the cpp_include headers in source order. func (d *Document) CPPIncludes() []*CPPInclude { var out []*CPPInclude + for _, n := range d.Nodes { if v, ok := n.(*CPPInclude); ok { out = append(out, v) } } + return out } // Namespaces returns the namespace headers in source order. func (d *Document) Namespaces() []*Namespace { var out []*Namespace + for _, n := range d.Nodes { if v, ok := n.(*Namespace); ok { out = append(out, v) } } + return out } // Structs returns the struct declarations in source order. func (d *Document) Structs() []*Struct { var out []*Struct + for _, n := range d.Nodes { if v, ok := n.(*Struct); ok && v.Kind == StructDecl { out = append(out, v) } } + return out } // Unions returns the union declarations in source order. func (d *Document) Unions() []*Struct { var out []*Struct + for _, n := range d.Nodes { if v, ok := n.(*Struct); ok && v.Kind == UnionDecl { out = append(out, v) } } + return out } // Exceptions returns the exception declarations in source order. func (d *Document) Exceptions() []*Struct { var out []*Struct + for _, n := range d.Nodes { if v, ok := n.(*Struct); ok && v.Kind == ExceptionDecl { out = append(out, v) } } + return out } // Enums returns the enum declarations in source order. func (d *Document) Enums() []*Enum { var out []*Enum + for _, n := range d.Nodes { if v, ok := n.(*Enum); ok { out = append(out, v) } } + return out } // Services returns the service declarations in source order. func (d *Document) Services() []*Service { var out []*Service + for _, n := range d.Nodes { if v, ok := n.(*Service); ok { out = append(out, v) } } + return out } // Consts returns the const declarations in source order. func (d *Document) Consts() []*Const { var out []*Const + for _, n := range d.Nodes { if v, ok := n.(*Const); ok { out = append(out, v) } } + return out } // Typedefs returns the typedef declarations in source order. func (d *Document) Typedefs() []*Typedef { var out []*Typedef + for _, n := range d.Nodes { if v, ok := n.(*Typedef); ok { out = append(out, v) } } + return out } diff --git a/syntax/ast_visit.go b/syntax/ast_visit.go index 30af45e..e04525c 100644 --- a/syntax/ast_visit.go +++ b/syntax/ast_visit.go @@ -6,6 +6,7 @@ package syntax func (d *Document) SearchNodePathByPosition(pos Position) []Node { var path []Node d.searchNodePath(d, pos, &path) + return path } @@ -13,6 +14,7 @@ func (d *Document) searchNodePath(root Node, pos Position, path *[]Node) { if !d.Contains(root, pos) { return } + *path = append(*path, root) for _, child := range nodeChildren(root) { d.searchNodePath(child, pos, path) @@ -36,6 +38,7 @@ func nodeChildren(n Node) []Node { for _, value := range v.Values { out = append(out, value) } + return out case *EnumValue: return []Node{v.Name} @@ -44,65 +47,79 @@ func nodeChildren(n Node) []Node { for _, field := range v.Fields { out = append(out, field) } + return out case *Service: out := []Node{v.Name} if v.Extends != nil { out = append(out, v.Extends) } + for _, fn := range v.Functions { out = append(out, fn) } + return out case *Function: out := []Node{v.Name} if v.Type != nil { out = append(out, v.Type) } + for _, arg := range v.Args { out = append(out, arg) } + if v.Throws != nil { out = append(out, v.Throws) } + return out case *Throws: out := make([]Node, 0, len(v.Fields)) for _, field := range v.Fields { out = append(out, field) } + return out case *Field: out := []Node{v.Type, v.Name} if v.Value != nil { out = append(out, v.Value) } + return out case *FieldType: var out []Node if v.Ident != nil { out = append(out, v.Ident) } + if v.KeyType != nil { out = append(out, v.KeyType) } + if v.ValueType != nil { out = append(out, v.ValueType) } + return out case *ConstValue: var out []Node for _, item := range v.List { out = append(out, item) } + for _, entry := range v.Map { out = append(out, entry.Key, entry.Value) } + return out case *Namespace: return []Node{v.Name} case *Include, *CPPInclude, *Identifier: return nil } + return nil } diff --git a/syntax/dump.go b/syntax/dump.go new file mode 100644 index 0000000..759eec8 --- /dev/null +++ b/syntax/dump.go @@ -0,0 +1,55 @@ +package syntax + +import ( + "fmt" + "strings" +) + +// Dump renders a parsed document as a debug tree: every token with its +// kind, position, blank-line count, and attached trivia, followed by the +// node spans. Deterministic and stable for a given input, so dumps can be +// diffed across versions. +func Dump(d *Document) string { + var b strings.Builder + for i, tok := range d.Tokens { + fmt.Fprintf(&b, "tok %3d %-16s line=%-3d col=%-3d blb=%d %q\n", + i, tok.Kind, tok.Line, tok.Col, tok.BlankLinesBefore, tok.Text) + + for _, tr := range tok.Leading { + fmt.Fprintf(&b, " leading %-18s %q\n", tr.Kind, tr.Text) + } + + for _, tr := range tok.Trailing { + fmt.Fprintf(&b, " trailing %-18s %q\n", tr.Kind, tr.Text) + } + } + + for i, n := range d.Nodes { + fmt.Fprintf(&b, "node %3d %-20s [%d..%d]\n", i, nodeName(n), n.TokStart(), n.TokEnd()) + } + + return b.String() +} + +func nodeName(n Node) string { + switch v := n.(type) { + case *Include: + return "Include" + case *CPPInclude: + return "CPPInclude" + case *Namespace: + return "Namespace" + case *Const: + return "Const" + case *Typedef: + return "Typedef" + case *Enum: + return "Enum" + case *Struct: + return v.Kind.String() + case *Service: + return "Service" + default: + return fmt.Sprintf("%T", n) + } +} diff --git a/syntax/lexer.go b/syntax/lexer.go index 11eecfc..26bd4a2 100644 --- a/syntax/lexer.go +++ b/syntax/lexer.go @@ -122,12 +122,14 @@ var keywordNames = func() map[TokenKind]string { for text, kind := range keywordKinds { m[kind] = text } + return m }() // isKeyword reports whether the token kind is one of the reserved words. func isKeyword(k TokenKind) bool { _, ok := keywordNames[k] + return ok } @@ -145,9 +147,11 @@ func (k TokenKind) String() string { if name, ok := keywordNames[k]; ok { return name } + if name, ok := tokenKindNames[k]; ok { return name } + return fmt.Sprintf("TokenKind(%d)", uint8(k)) } @@ -172,6 +176,7 @@ func (k TriviaKind) String() string { case TriviaAnnotation: return "annotation" } + return fmt.Sprintf("TriviaKind(%d)", uint8(k)) } @@ -236,6 +241,7 @@ func (e Error) Error() string { // lexing continues past errors so the parser can still recover. func Lex(src []byte) ([]Token, []Error) { l := &lexer{src: string(src)} + return l.run() } @@ -253,6 +259,7 @@ type srcPos struct { func (l *lexer) run() ([]Token, []Error) { l.line, l.col = 1, 1 + var tokens []Token for { @@ -260,6 +267,7 @@ func (l *lexer) run() ([]Token, []Error) { if n := len(tokens); n > 0 { prevLine = tokens[n-1].Line } + leading, trailing, blankLines := l.scanTrivia(prevLine) tok := l.scanToken() @@ -301,6 +309,7 @@ func (l *lexer) scanTrivia(prevLine int) (leading, trailing []Trivia, blankLines return leading, trailing, blankLines } } + return leading, trailing, blankLines } @@ -309,6 +318,7 @@ func (l *lexer) appendComment(leading, trailing []Trivia, prevLine, blankLines i if t.Line == prevLine { return leading, append(trailing, t) } + return append(leading, t), trailing } @@ -316,14 +326,18 @@ func (l *lexer) appendComment(leading, trailing []Trivia, prevLine, blankLines i // lines it contains (a blank line is a line containing only whitespace). func (l *lexer) scanWhitespace() int { newlines := 0 + for l.off < len(l.src) { switch l.src[l.off] { case '\n': newlines++ + l.advanceByte() case '\r': newlines++ + l.advanceByte() + if l.off < len(l.src) && l.src[l.off] == '\n' { l.advanceByte() } @@ -333,12 +347,15 @@ func (l *lexer) scanWhitespace() int { if newlines > 0 { return newlines - 1 } + return 0 } } + if newlines > 0 { return newlines - 1 } + return 0 } @@ -347,6 +364,7 @@ func (l *lexer) scanLineComment() Trivia { for l.off < len(l.src) && l.src[l.off] != '\n' && l.src[l.off] != '\r' { l.advanceRune() } + return l.finishTrivia(TriviaLineComment, start) } @@ -358,6 +376,7 @@ func (l *lexer) scanLineAnnotation() Trivia { for l.off < len(l.src) && l.src[l.off] != '\n' && l.src[l.off] != '\r' { l.advanceRune() } + return l.finishTrivia(TriviaAnnotation, start) } @@ -373,22 +392,28 @@ func (l *lexer) scanBlockComment() Trivia { for { if l.off >= len(l.src) { l.errorfAt(start, "unterminated comment") + break } + if l.src[l.off] == '*' && l.peekByte(1) == '/' { l.advanceByte() l.advanceByte() + kind := TriviaBlockComment if doc { kind = TriviaDocComment } + return l.finishTrivia(kind, start) } // An empty doc comment /**/ ends right after the opening /**. if doc && l.off == start.offset+3 && l.src[l.off] == '/' { l.advanceByte() + return l.finishTrivia(TriviaDocComment, start) } + l.advanceRune() } @@ -420,6 +445,7 @@ func (l *lexer) scanToken() Token { if tok, ok := l.scanNumber(); ok { return tok } + l.errorf("unexpected character %q", c) l.advanceRune() case c == '\'' || c == '"': @@ -440,6 +466,7 @@ func (l *lexer) scanToken() Token { if kind, ok := symbolKinds[c]; ok { return l.symbolToken(kind) } + l.errorf("unexpected character %q", c) l.advanceRune() } @@ -458,6 +485,7 @@ var symbolKinds = map[byte]TokenKind{ func (l *lexer) symbolToken(kind TokenKind) Token { start := l.pos() l.advanceByte() + return Token{Kind: kind, Text: l.src[start.offset:l.off], Offset: start.offset, Line: start.line, Col: start.col} } @@ -467,28 +495,38 @@ func (l *lexer) symbolToken(kind TokenKind) Token { func (l *lexer) scanIdentifier() Token { start := l.pos() dotted := false + l.advanceByte() // first character is [a-zA-Z_] + for l.off < len(l.src) { c := l.src[l.off] if isIdentPart(c) { l.advanceByte() + continue } + if c == '.' && isIdentStart(l.peekByte(1)) { dotted = true + l.advanceByte() // dot l.advanceByte() // first identifier character after dot + continue } + break } + text := l.src[start.offset:l.off] kind := TokenIdentifier + if !dotted { if k, ok := keywordKinds[text]; ok { kind = k } } + return Token{Kind: kind, Text: text, Offset: start.offset, Line: start.line, Col: start.col} } @@ -506,6 +544,7 @@ func (l *lexer) scanNumber() (Token, bool) { length := 0 kind := TokenIntConstant + switch { case hexLen > 0 && hexLen >= intLen && hexLen >= dubLen: length = hexLen @@ -521,6 +560,7 @@ func (l *lexer) scanNumber() (Token, bool) { for i := 0; i < length; i++ { l.advanceByte() } + return Token{Kind: kind, Text: rest[:length], Offset: start.offset, Line: start.line, Col: start.col}, true } @@ -531,18 +571,23 @@ func matchHex(s string) int { if i < len(s) && (s[i] == '+' || s[i] == '-') { i++ } + if i+2 > len(s) || s[i] != '0' || (s[i+1] != 'x' && s[i+1] != 'X') { return 0 } + i += 2 digits := 0 + for i < len(s) && isHexDigit(s[i]) { i++ digits++ } + if digits == 0 { return 0 } + return i } @@ -553,13 +598,16 @@ func matchInt(s string) int { if i < len(s) && (s[i] == '+' || s[i] == '-') { i++ } + start := i for i < len(s) && isDigit(s[i]) { i++ } + if i == start { return 0 } + return i } @@ -571,40 +619,52 @@ func matchDouble(s string) int { if i < len(s) && (s[i] == '+' || s[i] == '-') { i++ } + sawDigit := false + for i < len(s) && isDigit(s[i]) { i++ sawDigit = true } + if i < len(s) && s[i] == '.' { digits := 0 for j := i + 1; j < len(s) && isDigit(s[j]); j++ { digits++ } + if digits == 0 { return 0 } + sawDigit = true i += 1 + digits } + if i < len(s) && (s[i] == 'e' || s[i] == 'E') { j := i + 1 if j < len(s) && (s[j] == '+' || s[j] == '-') { j++ } + digits := 0 + for j < len(s) && isDigit(s[j]) { j++ digits++ } + if digits == 0 { return 0 } + i = j } + if !sawDigit { return 0 } + return i } @@ -621,22 +681,28 @@ func (l *lexer) scanString() Token { for { if l.off >= len(l.src) { l.errorfAt(start, "unterminated string literal") + break } + c := l.src[l.off] switch c { case quote: l.advanceByte() + return Token{Kind: TokenStringLiteral, Text: l.src[start.offset:l.off], Offset: start.offset, Line: start.line, Col: start.col} case '\n', '\r': l.errorfAt(start, "newline in string literal") + return Token{Kind: TokenStringLiteral, Text: l.src[start.offset:l.off], Offset: start.offset, Line: start.line, Col: start.col} case '\\': if l.off+1 >= len(l.src) { l.advanceByte() // consume the backslash with the string l.errorfAt(start, "unterminated string literal") + return Token{Kind: TokenStringLiteral, Text: l.src[start.offset:l.off], Offset: start.offset, Line: start.line, Col: start.col} } + esc := l.peekByte(1) switch esc { case 'r', 'n', 't', '"', '\'', '\\': @@ -653,6 +719,7 @@ func (l *lexer) scanString() Token { l.advanceRune() } } + return Token{Kind: TokenStringLiteral, Text: l.src[start.offset:l.off], Offset: start.offset, Line: start.line, Col: start.col} } @@ -660,6 +727,7 @@ func (l *lexer) peekByte(ahead int) byte { if l.off+ahead >= len(l.src) { return 0 } + return l.src[l.off+ahead] } @@ -679,6 +747,7 @@ func (l *lexer) advanceByte() { default: l.col++ } + l.off++ } @@ -699,6 +768,7 @@ func (l *lexer) advanceRune() { default: l.col++ } + l.off += size } diff --git a/syntax/lexer_fuzz_test.go b/syntax/lexer_fuzz_test.go index 03d6d86..006df41 100644 --- a/syntax/lexer_fuzz_test.go +++ b/syntax/lexer_fuzz_test.go @@ -60,22 +60,28 @@ func FuzzLex(f *testing.F) { if err.Offset < 0 || err.Offset > len(src) { t.Fatalf("error offset %d out of range", err.Offset) } + checkPos(t, srcStr, err.Offset, err.Line, err.Col, "error") } // Token and trivia invariants. prevEnd := 0 + for i, tok := range toks { if tok.Kind == TokenInvalid { t.Fatalf("token %d has invalid kind", i) } + if tok.Offset < 0 || tok.Offset+len(tok.Text) > len(src) { t.Fatalf("token %d (%s) spans outside the source", i, tok.Kind) } + if got := src[tok.Offset : tok.Offset+len(tok.Text)]; string(got) != tok.Text { t.Fatalf("token %d text %q does not match source %q", i, tok.Text, got) } + checkPos(t, srcStr, tok.Offset, tok.Line, tok.Col, "token") + if tok.BlankLinesBefore < 0 { t.Fatalf("token %d has negative BlankLinesBefore", i) } @@ -85,6 +91,7 @@ func FuzzLex(f *testing.F) { if tr.Offset < prevEnd || tr.Offset+len(tr.Text) > tok.Offset || tr.Offset+len(tr.Text) > len(src) { t.Fatalf("leading trivia %q of token %d outside the gap [%d, %d)", tr.Text, i, prevEnd, tok.Offset) } + checkPos(t, srcStr, tr.Offset, tr.Line, tr.Col, "leading trivia") } @@ -92,7 +99,9 @@ func FuzzLex(f *testing.F) { if tr.Offset < tok.Offset || tr.Offset+len(tr.Text) > len(src) { t.Fatalf("trailing trivia %q of token %d outside the source", tr.Text, i) } + checkPos(t, srcStr, tr.Offset, tr.Line, tr.Col, "trailing trivia") + if tr.Line != tok.Line { t.Fatalf("trailing trivia %q of token %d starts on line %d, token is on line %d", tr.Text, i, tr.Line, tok.Line) @@ -114,6 +123,7 @@ func FuzzLex(f *testing.F) { // columns count runes. func checkPos(t *testing.T, src string, offset, line, col int, what string) { t.Helper() + gotLine, gotCol := lineColAt(src, offset) if gotLine != line || gotCol != col { t.Fatalf("%s at offset %d: reported %d:%d, want %d:%d", what, offset, line, col, gotLine, gotCol) @@ -122,6 +132,7 @@ func checkPos(t *testing.T, src string, offset, line, col int, what string) { func lineColAt(src string, offset int) (line, col int) { line, col = 1, 1 + for i := 0; i < offset; { switch src[i] { case '\n': @@ -131,8 +142,10 @@ func lineColAt(src string, offset int) (line, col int) { case '\r': if i+1 < offset && src[i+1] == '\n' { i++ // \r\n: the \n resets the line + continue } + line++ col = 1 i++ @@ -142,5 +155,6 @@ func lineColAt(src string, offset int) (line, col int) { i += size } } + return line, col } diff --git a/syntax/lexer_test.go b/syntax/lexer_test.go index f3263ea..a26da00 100644 --- a/syntax/lexer_test.go +++ b/syntax/lexer_test.go @@ -16,6 +16,7 @@ func tokenSpecs(toks []Token) []tokSpec { for _, tok := range toks { specs = append(specs, tokSpec{tok.Kind, tok.Text}) } + return specs } @@ -174,6 +175,7 @@ func TestLexTokens(t *testing.T) { if len(errs) > 0 { t.Fatalf("unexpected errors: %v", errs) } + if got := tokenSpecs(toks); !reflect.DeepEqual(got, tt.want) { t.Errorf("tokens mismatch\n got: %v\nwant: %v", got, tt.want) } @@ -224,9 +226,11 @@ func TestLexPositions(t *testing.T) { if len(errs) > 0 { t.Fatalf("unexpected errors: %v", errs) } + if len(toks) != len(tt.want) { t.Fatalf("got %d tokens (%v), want %d", len(toks), tokenSpecs(toks), len(tt.want)) } + for i, want := range tt.want { got := toks[i] if got.Line != want.line || got.Col != want.col || got.Offset != want.offset { @@ -379,17 +383,21 @@ func TestLexTrivia(t *testing.T) { if len(errs) > 0 { t.Fatalf("unexpected errors: %v", errs) } + for _, check := range tt.checks { if check.idx >= len(toks) { t.Fatalf("token index %d out of range (%d tokens)", check.idx, len(toks)) } + tok := toks[check.idx] if got := triviaTexts(tok.Leading); !reflect.DeepEqual(got, check.leading) { t.Errorf("token %d leading: got %v, want %v", check.idx, got, check.leading) } + if got := triviaTexts(tok.Trailing); !reflect.DeepEqual(got, check.trailing) { t.Errorf("token %d trailing: got %v, want %v", check.idx, got, check.trailing) } + if tok.BlankLinesBefore != check.blankLinesBefore { t.Errorf("token %d blankLinesBefore: got %d, want %d", check.idx, tok.BlankLinesBefore, check.blankLinesBefore) } @@ -402,10 +410,12 @@ func triviaTexts(trivia []Trivia) []string { if len(trivia) == 0 { return nil } + texts := make([]string, 0, len(trivia)) for _, t := range trivia { texts = append(texts, t.Text) } + return texts } @@ -427,14 +437,17 @@ func TestLexTriviaKinds(t *testing.T) { if len(errs) > 0 { t.Fatalf("unexpected errors: %v", errs) } + eof := toks[len(toks)-1] if eof.Kind != TokenEOF { t.Fatalf("last token is %v, want eof", eof.Kind) } + var got []TriviaKind for _, trivia := range eof.Leading { got = append(got, trivia.Kind) } + if !reflect.DeepEqual(got, tt.want) { t.Errorf("trivia kinds: got %v, want %v", got, tt.want) } @@ -523,15 +536,18 @@ func TestLexErrors(t *testing.T) { if len(errs) != len(tt.wantErrs) { t.Fatalf("got %d errors (%v), want %d", len(errs), errs, len(tt.wantErrs)) } + for i, want := range tt.wantErrs { if got := errs[i].Error(); !strings.Contains(got, want) { t.Errorf("error %d: got %q, want substring %q", i, got, want) } } + got := make([]TokenKind, 0, len(toks)) for _, tok := range toks { got = append(got, tok.Kind) } + if !reflect.DeepEqual(got, tt.wantKinds) { t.Errorf("token kinds: got %v, want %v", got, tt.wantKinds) } diff --git a/syntax/parser.go b/syntax/parser.go index e50d4f4..92d195f 100644 --- a/syntax/parser.go +++ b/syntax/parser.go @@ -2,6 +2,7 @@ package syntax import ( "fmt" + "slices" "sort" ) @@ -13,6 +14,7 @@ func Parse(src []byte) (*Document, []Error) { doc, parseErrs := ParseTokens(toks) errs := append(lexErrs, parseErrs...) sort.SliceStable(errs, func(i, j int) bool { return errs[i].Offset < errs[j].Offset }) + return doc, errs } @@ -20,6 +22,7 @@ func Parse(src []byte) (*Document, []Error) { // reparse a file from its cached tokens without re-lexing. func ParseTokens(toks []Token) (*Document, []Error) { p := &parser{toks: toks} + return p.parseDocument(), p.errs } @@ -38,6 +41,7 @@ func (p *parser) at(k TokenKind) bool { return p.cur().Kind == k } func (p *parser) advance() *Token { t := &p.toks[p.pos] p.pos++ + return t } @@ -46,6 +50,7 @@ func (p *parser) accept(k TokenKind) *Token { if p.at(k) { return p.advance() } + return nil } @@ -54,20 +59,22 @@ func (p *parser) accept(k TokenKind) *Token { func (p *parser) expect(k TokenKind, what string) bool { if p.at(k) { p.advance() + return true } + p.errorfCur("expected %s, got %q", what, p.cur().Text) + return false } // synchronizeTo skips tokens until one of the given kinds or EOF. func (p *parser) synchronizeTo(kinds ...TokenKind) { for !p.at(TokenEOF) { - for _, k := range kinds { - if p.cur().Kind == k { - return - } + if slices.Contains(kinds, p.cur().Kind) { + return } + p.advance() } } @@ -147,51 +154,67 @@ func (p *parser) parseDocument() *Document { func (p *parser) parseInclude() *Include { n := &Include{nodeBase: nodeBase{first: p.pos}} p.advance() // include + if !p.at(TokenStringLiteral) { p.errorfCur("expected include path string, got %q", p.cur().Text) + if !p.at(TokenEOF) { p.advance() // consume the offending token so the top level can continue } + return nil } + n.Path = p.advance() n.last = p.pos - 1 + return n } func (p *parser) parseCPPInclude() *CPPInclude { n := &CPPInclude{nodeBase: nodeBase{first: p.pos}} p.advance() // cpp_include + if !p.at(TokenStringLiteral) { p.errorfCur("expected cpp_include path string, got %q", p.cur().Text) + if !p.at(TokenEOF) { p.advance() // consume the offending token so the top level can continue } + return nil } + n.Path = p.advance() n.last = p.pos - 1 + return n } func (p *parser) parseNamespace() *Namespace { n := &Namespace{nodeBase: nodeBase{first: p.pos}} p.advance() // namespace + switch { case p.at(TokenIdentifier), p.at(TokenStar): n.Scope = p.advance() default: p.errorfCur("expected namespace scope, got %q", p.cur().Text) p.synchronizeTo(TokenEOF) + return nil } + n.Name = p.expectIdentifier("namespace name") if n.Name == nil { p.synchronizeTo(TokenEOF) + return nil } + n.Annotations = p.parseAnnotationsIfPresent() n.last = p.pos - 1 + return n } @@ -200,70 +223,94 @@ func (p *parser) parseNamespace() *Namespace { func (p *parser) parseConst() *Const { n := &Const{nodeBase: nodeBase{first: p.pos}} p.advance() // const + n.Type = p.parseFieldType() if n.Type == nil { p.synchronizeTo(TokenEOF) + return nil } + n.Name = p.expectIdentifier("constant name") if n.Name == nil { p.synchronizeTo(TokenEOF) + return nil } + if !p.expect(TokenEqual, "'='") { p.synchronizeTo(TokenEOF) + return nil } + n.Value = p.parseConstValue() if n.Value == nil { p.synchronizeTo(TokenEOF) + return nil } + if sep := p.acceptSeparator(); sep != 0 { n.Sep = sep } + n.last = p.pos - 1 + return n } func (p *parser) parseTypedef() *Typedef { n := &Typedef{nodeBase: nodeBase{first: p.pos}} p.advance() // typedef + n.Type = p.parseFieldType() if n.Type == nil { p.synchronizeTo(TokenEOF) + return nil } + n.Name = p.expectIdentifier("typedef name") if n.Name == nil { p.synchronizeTo(TokenEOF) + return nil } + n.Annotations = p.parseAnnotationsIfPresent() if sep := p.acceptSeparator(); sep != 0 { n.Sep = sep } + n.last = p.pos - 1 + return n } func (p *parser) parseEnum() *Enum { n := &Enum{nodeBase: nodeBase{first: p.pos}} p.advance() // enum + n.Name = p.expectIdentifier("enum name") if n.Name == nil { p.synchronizeTo(TokenEOF) + return nil } + if p.accept(TokenLBrace) != nil { for !p.at(TokenRBrace) && !p.at(TokenEOF) { v := p.parseEnumValue() if v == nil { p.synchronizeTo(TokenRBrace) + continue } + n.Values = append(n.Values, v) } + if !p.at(TokenRBrace) { p.errorfCur("expected '}' to close enum, got %q", p.cur().Text) } else { @@ -272,42 +319,54 @@ func (p *parser) parseEnum() *Enum { } else { p.errorfCur("expected '{' after enum name, got %q", p.cur().Text) } + n.Annotations = p.parseAnnotationsIfPresent() n.last = p.pos - 1 + return n } func (p *parser) parseEnumValue() *EnumValue { if !p.at(TokenIdentifier) { p.errorfCur("expected enum value name, got %q", p.cur().Text) + return nil } + v := &EnumValue{nodeBase: nodeBase{first: p.pos}} + v.Name = p.identifier() if p.at(TokenEqual) { p.advance() + if !p.at(TokenIntConstant) { p.errorfCur("expected integer enum value, got %q", p.cur().Text) } else { v.Value = p.advance() } } + v.Annotations = p.parseAnnotationsIfPresent() if sep := p.acceptSeparator(); sep != 0 { v.Sep = sep } + v.last = p.pos - 1 + return v } func (p *parser) parseStruct() *Struct { n := &Struct{nodeBase: nodeBase{first: p.pos}, Kind: StructKind(p.cur().Kind)} p.advance() // struct | union | exception + n.Name = p.expectIdentifier("struct name") if n.Name == nil { p.synchronizeTo(TokenEOF) + return nil } + if p.accept(TokenLBrace) != nil { n.Fields = p.parseFieldList(TokenRBrace) if !p.at(TokenRBrace) { @@ -318,38 +377,49 @@ func (p *parser) parseStruct() *Struct { } else { p.errorfCur("expected '{' after struct name, got %q", p.cur().Text) } + n.Annotations = p.parseAnnotationsIfPresent() n.last = p.pos - 1 + return n } func (p *parser) parseService() *Service { n := &Service{nodeBase: nodeBase{first: p.pos}} p.advance() // service + n.Name = p.expectIdentifier("service name") if n.Name == nil { p.synchronizeTo(TokenEOF) + return nil } + if p.at(TokenExtends) { p.advance() + n.Extends = p.expectIdentifier("base service name") if n.Extends == nil { p.synchronizeTo(TokenLBrace, TokenEOF) + if !p.at(TokenLBrace) { return nil } } } + if p.accept(TokenLBrace) != nil { for !p.at(TokenRBrace) && !p.at(TokenEOF) { f := p.parseFunction() if f == nil { p.synchronizeTo(TokenRBrace) + continue } + n.Functions = append(n.Functions, f) } + if !p.at(TokenRBrace) { p.errorfCur("expected '}' to close service, got %q", p.cur().Text) } else { @@ -358,8 +428,10 @@ func (p *parser) parseService() *Service { } else { p.errorfCur("expected '{' after service name, got %q", p.cur().Text) } + n.Annotations = p.parseAnnotationsIfPresent() n.last = p.pos - 1 + return n } @@ -380,6 +452,7 @@ func (p *parser) parseFunction() *Function { f.Type = p.parseFieldType() if f.Type == nil { p.synchronizeTo(TokenLParen, TokenEOF) + return nil } } @@ -387,13 +460,16 @@ func (p *parser) parseFunction() *Function { f.Name = p.expectIdentifier("function name") if f.Name == nil { p.synchronizeTo(TokenLParen, TokenEOF) + return nil } if !p.expect(TokenLParen, "'('") { p.synchronizeTo(TokenEOF) + return nil } + f.Args = p.parseFieldList(TokenRParen) if !p.at(TokenRParen) { p.errorfCur("expected ')' to close arguments, got %q", p.cur().Text) @@ -403,17 +479,22 @@ func (p *parser) parseFunction() *Function { if p.at(TokenThrows) { p.advance() + if !p.expect(TokenLParen, "'(' after throws") { p.synchronizeTo(TokenEOF) + return nil } + f.Throws = &Throws{nodeBase: nodeBase{first: p.pos - 1}} + f.Throws.Fields = p.parseFieldList(TokenRParen) if !p.at(TokenRParen) { p.errorfCur("expected ')' to close throws, got %q", p.cur().Text) } else { p.advance() } + f.Throws.last = p.pos - 1 } @@ -421,7 +502,9 @@ func (p *parser) parseFunction() *Function { if sep := p.acceptSeparator(); sep != 0 { f.Sep = sep } + f.last = p.pos - 1 + return f } @@ -431,17 +514,22 @@ func (p *parser) parseFunction() *Function { // struct bodies, function arguments, and throws clauses. func (p *parser) parseFieldList(term TokenKind) []*Field { var fields []*Field + for !p.at(term) && !p.at(TokenEOF) { f, ok := p.parseField() if !ok { p.synchronizeTo(TokenComma, TokenSemicolon, term) + if !p.at(term) && !p.at(TokenEOF) { p.advance() // consume the stray separator, if any } + continue } + fields = append(fields, f) } + return fields } @@ -452,6 +540,7 @@ func (p *parser) parseField() (*Field, bool) { f.FieldID = p.advance() if !p.expect(TokenColon, "':' after field id") { p.synchronizeTo(TokenComma, TokenSemicolon, TokenRBrace, TokenRParen) + return f, false } } @@ -464,6 +553,7 @@ func (p *parser) parseField() (*Field, bool) { f.Type = p.parseFieldType() if f.Type == nil { p.synchronizeTo(TokenComma, TokenSemicolon, TokenRBrace, TokenRParen) + return f, false } @@ -480,14 +570,17 @@ func (p *parser) parseField() (*Field, bool) { } else { p.errorfCur("expected field name, got %q", p.cur().Text) p.synchronizeTo(TokenComma, TokenSemicolon, TokenRBrace, TokenRParen) + return f, false } if p.at(TokenEqual) { p.advance() + f.Value = p.parseConstValue() if f.Value == nil { p.synchronizeTo(TokenComma, TokenSemicolon, TokenRBrace, TokenRParen) + return f, false } } @@ -496,7 +589,9 @@ func (p *parser) parseField() (*Field, bool) { if sep := p.acceptSeparator(); sep != 0 { f.Sep = sep } + f.last = p.pos - 1 + return f, true } @@ -515,10 +610,12 @@ func (p *parser) parseFieldType() *FieldType { case TokenSet: t.Kind = TypeSet } + p.advance() if p.at(TokenCPPType) { p.advance() + if !p.at(TokenStringLiteral) { p.errorfCur("expected string literal after cpp_type, got %q", p.cur().Text) } else { @@ -528,6 +625,7 @@ func (p *parser) parseFieldType() *FieldType { if !p.expect(TokenLt, "'<' to open container type") { p.synchronizeTo(TokenComma, TokenSemicolon, TokenRBrace, TokenRParen) + return nil } @@ -536,6 +634,7 @@ func (p *parser) parseFieldType() *FieldType { if t.KeyType == nil { p.synchronizeTo(TokenComma, TokenGt) } + if p.accept(TokenComma) != nil { t.ValueType = p.parseFieldType() if t.ValueType == nil { @@ -543,6 +642,7 @@ func (p *parser) parseFieldType() *FieldType { } } else { p.errorfCur("expected ',' and a value type for map, got %q", p.cur().Text) + if !p.at(TokenGt) { p.synchronizeTo(TokenGt) } @@ -556,6 +656,7 @@ func (p *parser) parseFieldType() *FieldType { if !p.expect(TokenGt, "'>' to close container type") { p.synchronizeTo(TokenComma, TokenSemicolon, TokenRBrace, TokenRParen) + return nil } @@ -566,8 +667,10 @@ func (p *parser) parseFieldType() *FieldType { default: if !isBaseType(p.cur().Kind) { p.errorfCur("expected type, got %q", p.cur().Text) + return nil } + t.Kind = TypeBase t.Base = p.cur().Kind p.deprecationWarnings(p.cur()) @@ -578,7 +681,9 @@ func (p *parser) parseFieldType() *FieldType { if t.Kind != TypeIdent { t.Annotations = p.parseAnnotationsIfPresent() } + t.last = p.pos - 1 + return t } @@ -597,6 +702,7 @@ func isBaseType(k TokenKind) bool { TokenDouble, TokenString, TokenBinary, TokenSlist, TokenUUID: return true } + return false } @@ -624,19 +730,26 @@ func (p *parser) parseConstValue() *ConstValue { case TokenLBracket: v.Kind = ValueList + p.advance() + for !p.at(TokenRBracket) && !p.at(TokenEOF) { item := p.parseConstValue() if item == nil { p.synchronizeTo(TokenComma, TokenSemicolon, TokenRBracket) + if !p.at(TokenRBracket) && !p.at(TokenEOF) { p.advance() // consume the stray separator } + continue } + v.List = append(v.List, item) + p.acceptSeparator() } + if !p.at(TokenRBracket) { p.errorfCur("expected ']' to close list constant, got %q", p.cur().Text) } else { @@ -645,34 +758,47 @@ func (p *parser) parseConstValue() *ConstValue { case TokenLBrace: v.Kind = ValueMap + p.advance() + for !p.at(TokenRBrace) && !p.at(TokenEOF) { key := p.parseConstValue() if key == nil { p.synchronizeTo(TokenComma, TokenSemicolon, TokenRBrace) + if !p.at(TokenRBrace) && !p.at(TokenEOF) { p.advance() } + continue } + if !p.expect(TokenColon, "':' between map key and value") { p.synchronizeTo(TokenComma, TokenSemicolon, TokenRBrace) + if !p.at(TokenRBrace) && !p.at(TokenEOF) { p.advance() } + continue } + value := p.parseConstValue() if value == nil { p.synchronizeTo(TokenComma, TokenSemicolon, TokenRBrace) + if !p.at(TokenRBrace) && !p.at(TokenEOF) { p.advance() } + continue } + v.Map = append(v.Map, ConstMapEntry{Key: key, Value: value}) + p.acceptSeparator() } + if !p.at(TokenRBrace) { p.errorfCur("expected '}' to close map constant, got %q", p.cur().Text) } else { @@ -681,10 +807,12 @@ func (p *parser) parseConstValue() *ConstValue { default: p.errorfCur("expected constant value, got %q", p.cur().Text) + return nil } v.last = p.pos - 1 + return v } @@ -694,6 +822,7 @@ func (p *parser) parseAnnotationsIfPresent() *Annotations { if !p.at(TokenLParen) { return nil } + return p.parseAnnotations() } @@ -709,15 +838,20 @@ func (p *parser) parseAnnotations() *Annotations { if !p.at(TokenIdentifier) { p.errorfCur("expected annotation name, got %q", p.cur().Text) p.synchronizeTo(TokenComma, TokenSemicolon, TokenRParen) + if !p.at(TokenRParen) && !p.at(TokenEOF) { p.advance() } + continue } + item := &Annotation{nodeBase: nodeBase{first: p.pos}} + item.Name = p.identifier() if p.at(TokenEqual) { p.advance() + if !p.at(TokenStringLiteral) { p.errorfCur("expected string literal annotation value, got %q", p.cur().Text) p.synchronizeTo(TokenComma, TokenSemicolon, TokenRParen) @@ -725,9 +859,11 @@ func (p *parser) parseAnnotations() *Annotations { item.Value = p.advance() } } + if sep := p.acceptSeparator(); sep != 0 { item.Sep = sep } + item.last = p.pos - 1 a.Items = append(a.Items, item) } @@ -737,7 +873,9 @@ func (p *parser) parseAnnotations() *Annotations { } else { p.advance() } + a.last = p.pos - 1 + return a } @@ -748,14 +886,17 @@ func (p *parser) acceptSeparator() TokenKind { case TokenComma, TokenSemicolon: k := p.cur().Kind p.advance() + return k } + return 0 } func (p *parser) identifier() *Identifier { i := p.pos t := p.advance() + return &Identifier{nodeBase: nodeBase{first: i, last: i}, Text: t.Text} } @@ -764,7 +905,9 @@ func (p *parser) identifier() *Identifier { func (p *parser) expectIdentifier(what string) *Identifier { if !p.at(TokenIdentifier) { p.errorfCur("expected %s, got %q", what, p.cur().Text) + return nil } + return p.identifier() } diff --git a/syntax/parser_fuzz_test.go b/syntax/parser_fuzz_test.go index b93a8c9..f7b2a28 100644 --- a/syntax/parser_fuzz_test.go +++ b/syntax/parser_fuzz_test.go @@ -39,6 +39,7 @@ func FuzzParse(f *testing.F) { if doc == nil { t.Fatal("Parse returned nil document") } + checkDoc(t, doc) for i := 1; i < len(errs); i++ { diff --git a/syntax/parser_test.go b/syntax/parser_test.go index 22110ef..c2298bc 100644 --- a/syntax/parser_test.go +++ b/syntax/parser_test.go @@ -12,19 +12,24 @@ func isNilNode(n Node) bool { if n == nil { return true } + v := reflect.ValueOf(n) + return v.Kind() == reflect.Pointer && v.IsNil() } func parseOK(t *testing.T, src string) *Document { t.Helper() + doc, errs := Parse([]byte(src)) for _, err := range errs { if err.Severity == SeverityError { t.Fatalf("unexpected parse errors: %v", errs) } } + checkDoc(t, doc) + return doc } @@ -32,17 +37,23 @@ func parseOK(t *testing.T, src string) *Document { // token range is well-formed and its scalar texts match the token stream. func checkDoc(t *testing.T, doc *Document) { t.Helper() + var walk func(n Node) + walk = func(n Node) { if isNilNode(n) { return } + start, end := n.TokStart(), n.TokEnd() if start < 0 || end >= len(doc.Tokens) || start > end { t.Errorf("node %T has invalid token range [%d, %d] of %d tokens", n, start, end, len(doc.Tokens)) + return } + tokText := func(i int) string { return doc.Tokens[i].Text } + switch v := n.(type) { case *Identifier: if got := tokText(v.TokStart()); got != v.Text { @@ -52,9 +63,11 @@ func checkDoc(t *testing.T, doc *Document) { if v.Text != "" && v.Text != tokText(v.TokStart()) { t.Errorf("const value %q does not match token %q", v.Text, tokText(v.TokStart())) } + for _, item := range v.List { walk(item) } + for _, entry := range v.Map { walk(entry.Key) walk(entry.Value) @@ -75,6 +88,7 @@ func checkDoc(t *testing.T, doc *Document) { case *Enum: walk(v.Name) walk(v.Annotations) + for _, val := range v.Values { walk(val) } @@ -84,6 +98,7 @@ func checkDoc(t *testing.T, doc *Document) { case *Struct: walk(v.Name) walk(v.Annotations) + for _, f := range v.Fields { walk(f) } @@ -91,6 +106,7 @@ func checkDoc(t *testing.T, doc *Document) { walk(v.Name) walk(v.Extends) walk(v.Annotations) + for _, f := range v.Functions { walk(f) } @@ -98,9 +114,11 @@ func checkDoc(t *testing.T, doc *Document) { walk(v.Type) walk(v.Name) walk(v.Annotations) + for _, a := range v.Args { walk(a) } + if v.Throws != nil { for _, f := range v.Throws.Fields { walk(f) @@ -130,13 +148,16 @@ func checkDoc(t *testing.T, doc *Document) { // top returns the first top-level node of the given type. func top[T Node](t *testing.T, doc *Document) T { t.Helper() + for _, n := range doc.Nodes { if v, ok := n.(T); ok { return v } } + var zero T t.Fatalf("no node of type %T in document", zero) + return zero } @@ -158,25 +179,32 @@ func TestParseStructs(t *testing.T) { if s.Kind != StructDecl { t.Errorf("kind = %v, want struct", s.Kind) } + if got := s.Name.Text; got != "User" { t.Errorf("name = %q", got) } + if len(s.Fields) != 3 { t.Fatalf("fields = %d, want 3", len(s.Fields)) } + f := s.Fields[0] if f.FieldID == nil || f.FieldID.Text != "1" { t.Errorf("field 0 id = %v", f.FieldID) } + if f.Req != TokenRequired { t.Errorf("field 0 req = %v", f.Req) } + if f.Type.Kind != TypeBase || f.Type.Base != TokenI64 { t.Errorf("field 0 type = %+v", f.Type) } + if f.Name.Text != "id" { t.Errorf("field 0 name = %q", f.Name.Text) } + if s.Fields[2].Type.Kind != TypeList || s.Fields[2].Type.ValueType.Base != TokenI32 { t.Errorf("field 2 type = %+v", s.Fields[2].Type) } @@ -194,9 +222,11 @@ func TestParseStructs(t *testing.T) { if s.Fields[0].FieldID != nil { t.Errorf("implicit id should be nil, got %v", s.Fields[0].FieldID) } + if s.Fields[1].Value.Kind != ValueString || s.Fields[1].Value.Text != `"x"` { t.Errorf("field 1 value = %+v", s.Fields[1].Value) } + if s.Fields[2].Value.Kind != ValueInt || s.Fields[2].Value.Text != "42" { t.Errorf("field 2 value = %+v", s.Fields[2].Value) } @@ -225,6 +255,7 @@ func TestParseStructs(t *testing.T) { if s.Fields[0].Name.Text != "namespace" { t.Errorf("field 0 name = %q", s.Fields[0].Name.Text) } + if s.Fields[1].Name.Text != "cpp_include" { t.Errorf("field 1 name = %q", s.Fields[1].Name.Text) } @@ -240,6 +271,7 @@ func TestParseStructs(t *testing.T) { }`, check: func(t *testing.T, doc *Document) { s := top[*Struct](t, doc) + want := []string{"uuid", "id", "map", "string"} for i, name := range want { if s.Fields[i].Name.Text != name { @@ -286,9 +318,11 @@ func TestParseStructs(t *testing.T) { if s.Annotations == nil || len(s.Annotations.Items) != 1 { t.Fatalf("struct annotations = %+v", s.Annotations) } + if s.Annotations.Items[0].Name.Text != "struct_anno" || s.Annotations.Items[0].Value.Text != `"x"` { t.Errorf("struct annotation = %+v", s.Annotations.Items[0]) } + if s.Fields[0].Annotations == nil { t.Fatal("field annotations missing") } @@ -301,10 +335,12 @@ func TestParseStructs(t *testing.T) { }`, check: func(t *testing.T, doc *Document) { s := top[*Struct](t, doc) + ft := s.Fields[0].Type if ft.Annotations == nil || ft.Annotations.Items[0].Name.Text != "tag" { t.Errorf("type annotations = %+v", ft.Annotations) } + if s.Fields[0].Annotations != nil { t.Errorf("field annotations should be nil, got %+v", s.Fields[0].Annotations) } @@ -320,6 +356,7 @@ func TestParseStructs(t *testing.T) { if s.Fields[0].Type.Annotations != nil { t.Errorf("type annotations should be nil, got %+v", s.Fields[0].Type.Annotations) } + if s.Fields[0].Annotations == nil || s.Fields[0].Annotations.Items[0].Name.Text != "tag" { t.Errorf("field annotations = %+v", s.Fields[0].Annotations) } @@ -336,6 +373,7 @@ func TestParseStructs(t *testing.T) { if s.Fields[0].Type.CPPType == nil || s.Fields[0].Type.CPPType.Text != `"std::map"` { t.Errorf("map cpp_type = %v", s.Fields[0].Type.CPPType) } + if s.Fields[1].Type.CPPType == nil || s.Fields[1].Type.CPPType.Text != `"std::vector"` { t.Errorf("list cpp_type = %v", s.Fields[1].Type.CPPType) } @@ -355,12 +393,15 @@ exception E { if u.Kind != UnionDecl { t.Errorf("kind = %v, want union", u.Kind) } + var e *Struct + for _, n := range doc.Nodes { if s, ok := n.(*Struct); ok && s.Kind == ExceptionDecl { e = s } } + if e == nil || e.Name.Text != "E" { t.Errorf("exception missing: %+v", e) } @@ -411,10 +452,12 @@ const bool j = false`, for _, n := range doc.Nodes { kinds = append(kinds, n.(*Const).Value.Kind) } + want := []ConstValueKind{ValueInt, ValueInt, ValueInt, ValueDouble, ValueDouble, ValueDouble, ValueString, ValueString, ValueInt, ValueInt} if len(kinds) != len(want) { t.Fatalf("consts = %d, want %d", len(kinds), len(want)) } + for i := range want { if kinds[i] != want[i] { t.Errorf("const %d kind = %v, want %v", i, kinds[i], want[i]) @@ -424,6 +467,7 @@ const bool j = false`, if doc.Nodes[1].(*Const).Value.Text != "0xa1" { t.Errorf("hex text = %q", doc.Nodes[1].(*Const).Value.Text) } + if doc.Nodes[4].(*Const).Value.Text != "1.3333e11" { t.Errorf("double text = %q", doc.Nodes[4].(*Const).Value.Text) } @@ -438,9 +482,11 @@ const double c = -1.5`, if doc.Nodes[0].(*Const).Value.Kind != ValueIdent || doc.Nodes[0].(*Const).Value.Text != "OTHER_CONST" { t.Errorf("ident value = %+v", doc.Nodes[0].(*Const).Value) } + if doc.Nodes[1].(*Const).Value.Text != "-1" { t.Errorf("negative int = %q", doc.Nodes[1].(*Const).Value.Text) } + if doc.Nodes[2].(*Const).Value.Text != "-1.5" { t.Errorf("negative double = %q", doc.Nodes[2].(*Const).Value.Text) } @@ -465,12 +511,15 @@ const list> c = [[1], [2, 3]]`, if a.Kind != ValueList || len(a.List) != 3 { t.Fatalf("list a = %+v", a) } + if a.List[1].Text != "2" { t.Errorf("list item = %q", a.List[1].Text) } + if b := doc.Nodes[1].(*Const).Value; len(b.List) != 0 { t.Errorf("empty list = %+v", b) } + c := doc.Nodes[2].(*Const).Value if len(c.List) != 2 || len(c.List[1].List) != 2 { t.Errorf("nested list = %+v", c) @@ -486,9 +535,11 @@ const map b = {}`, if a.Kind != ValueMap || len(a.Map) != 2 { t.Fatalf("map a = %+v", a) } + if a.Map[0].Key.Text != `"x"` || a.Map[0].Value.Text != "1" { t.Errorf("map entry = %+v", a.Map[0]) } + if b := doc.Nodes[1].(*Const).Value; len(b.Map) != 0 { t.Errorf("empty map = %+v", b) } @@ -512,6 +563,7 @@ const map b = {}`, if len(c.Value.List) != 2 { t.Errorf("list = %+v", c.Value) } + if c.Sep != TokenSemicolon { t.Errorf("const sep = %v", c.Sep) } @@ -556,15 +608,19 @@ func TestParseEnums(t *testing.T) { if len(e.Values) != 4 { t.Fatalf("values = %d", len(e.Values)) } + if e.Values[0].Value == nil || e.Values[0].Value.Text != "1" { t.Errorf("A value = %v", e.Values[0].Value) } + if e.Values[1].Value != nil { t.Errorf("B should auto-increment, got %v", e.Values[1].Value) } + if e.Values[2].Value.Text != "0x10" { t.Errorf("C value = %v", e.Values[2].Value) } + if e.Values[3].Sep != 0 { t.Errorf("D sep = %v", e.Values[3].Sep) } @@ -631,20 +687,25 @@ func TestParseServices(t *testing.T) { if len(s.Functions) != 3 { t.Fatalf("functions = %d", len(s.Functions)) } + f := s.Functions[0] if f.Type.Kind != TypeIdent || f.Type.Ident.Text != "User" { t.Errorf("getUser type = %+v", f.Type) } + if f.Throws == nil || len(f.Throws.Fields) != 1 { t.Fatalf("throws = %+v", f.Throws) } + if f.Throws.Fields[0].Type.Ident.Text != "NotFound" { t.Errorf("throws type = %+v", f.Throws.Fields[0].Type) } + ping := s.Functions[1] if ping.Oneway == nil || ping.Oneway.Kind != TokenOneway || ping.Void == nil { t.Errorf("ping = %+v", ping) } + if len(s.Functions[2].Args) != 2 { t.Errorf("update args = %d", len(s.Functions[2].Args)) } @@ -681,6 +742,7 @@ func TestParseServices(t *testing.T) { if s.Functions[0].Annotations == nil { t.Error("f annotations missing") } + if s.Functions[1].Annotations == nil || s.Functions[1].Annotations.Items[0].Name.Text != "g_anno" { t.Errorf("g annotations = %+v", s.Functions[1].Annotations) } @@ -722,18 +784,22 @@ namespace * global`, if len(doc.Nodes) != 4 { t.Fatalf("nodes = %d: %+v", len(doc.Nodes), doc.Nodes) } + inc := doc.Nodes[0].(*Include) if inc.Path.Text != `"shared.thrift"` { t.Errorf("include path = %q", inc.Path.Text) } + cpp := doc.Nodes[1].(*CPPInclude) if cpp.Path.Text != `"base.h"` { t.Errorf("cpp include = %q", cpp.Path.Text) } + ns := doc.Nodes[2].(*Namespace) if ns.Scope.Text != "java" || ns.Name.Text != "com.example" { t.Errorf("namespace = %+v", ns) } + if doc.Nodes[3].(*Namespace).Scope.Kind != TokenStar { t.Errorf("star scope = %+v", doc.Nodes[3].(*Namespace).Scope) } @@ -774,6 +840,7 @@ typedef map> Index`, if td.Type.Base != TokenI64 || td.Name.Text != "Timestamp" { t.Errorf("typedef = %+v", td) } + if doc.Nodes[1].(*Typedef).Type.Kind != TypeMap { t.Errorf("container typedef = %+v", doc.Nodes[1].(*Typedef)) } @@ -810,6 +877,7 @@ func TestParseAnnotations(t *testing.T) { src: "typedef i32 T (foo)", check: func(t *testing.T, doc *Document) { td := doc.Nodes[0].(*Typedef) + a := td.Annotations.Items[0] if a.Name.Text != "foo" || a.Value != nil { t.Errorf("bare annotation = %+v", a) @@ -821,13 +889,16 @@ func TestParseAnnotations(t *testing.T) { src: `typedef i32 T (a = "1", b = "x", c = 'y')`, check: func(t *testing.T, doc *Document) { td := doc.Nodes[0].(*Typedef) + items := td.Annotations.Items if len(items) != 3 { t.Fatalf("items = %d", len(items)) } + if items[0].Value.Text != `"1"` || items[1].Value.Text != `"x"` || items[2].Value.Text != `'y'` { t.Errorf("values = %v, %v, %v", items[0].Value, items[1].Value, items[2].Value) } + if items[0].Sep != TokenComma || items[1].Sep != TokenComma { t.Errorf("seps = %v, %v", items[0].Sep, items[1].Sep) } @@ -838,10 +909,12 @@ func TestParseAnnotations(t *testing.T) { src: "typedef i32 T (a = \"1\"; b = \"2\" c = \"3\")", check: func(t *testing.T, doc *Document) { td := doc.Nodes[0].(*Typedef) + items := td.Annotations.Items if len(items) != 3 { t.Fatalf("items = %d", len(items)) } + if items[0].Sep != TokenSemicolon || items[1].Sep != 0 || items[2].Sep != 0 { t.Errorf("seps = %v, %v, %v", items[0].Sep, items[1].Sep, items[2].Sep) } @@ -949,15 +1022,19 @@ func TestParseErrors(t *testing.T) { if doc == nil { t.Fatal("parse returned nil document") } + var got []string + for _, err := range errs { if err.Severity == SeverityError { got = append(got, err.Error()) } } + if len(got) != len(tt.wantErrs) { t.Fatalf("got %d errors (%v), want %d", len(got), got, len(tt.wantErrs)) } + for i, want := range tt.wantErrs { if !strings.Contains(got[i], want) { t.Errorf("error %d = %q, want substring %q", i, got[i], want) @@ -993,15 +1070,19 @@ func TestParseWarnings(t *testing.T) { for _, tt := range tests { t.Run(tt.name, func(t *testing.T) { _, errs := Parse([]byte(tt.src)) + var warns []string + for _, err := range errs { if err.Severity == SeverityWarning { warns = append(warns, err.Error()) } } + if len(warns) != len(tt.wantMsg) { t.Fatalf("warnings = %v, want %d", warns, len(tt.wantMsg)) } + for i, want := range tt.wantMsg { if !strings.Contains(warns[i], want) { t.Errorf("warning %d = %q, want substring %q", i, warns[i], want) diff --git a/syntax/position.go b/syntax/position.go index a7204bd..79d14d1 100644 --- a/syntax/position.go +++ b/syntax/position.go @@ -21,6 +21,7 @@ func (p Position) IsValid() bool { // TokenPosition returns the start position of token i. func (d *Document) TokenPosition(i int) Position { t := d.Tokens[i] + return Position{Line: t.Line, Col: t.Col, Offset: t.Offset} } @@ -28,6 +29,7 @@ func (d *Document) TokenPosition(i int) Position { // never span lines, so the end position is on the same line. func (d *Document) TokenEndPosition(i int) Position { t := d.Tokens[i] + return Position{ Line: t.Line, Col: t.Col + utf8.RuneCountInString(t.Text), @@ -42,17 +44,20 @@ func (d *Document) TokenIndex(t *Token) int { if t == nil { return 0 } + for i := range d.Tokens { if &d.Tokens[i] == t { return i } } + return 0 } // TokenRange returns the span of a token. func (d *Document) TokenRange(t *Token) (start, end Position) { i := d.TokenIndex(t) + return d.TokenPosition(i), d.TokenEndPosition(i) } @@ -65,5 +70,6 @@ func (d *Document) Range(n Node) (start, end Position) { // Contains reports whether pos lies within the node's span, inclusive. func (d *Document) Contains(n Node, pos Position) bool { start, end := d.Range(n) + return pos.Offset >= start.Offset && pos.Offset <= end.Offset } diff --git a/syntax/position_test.go b/syntax/position_test.go index b6226a9..340638f 100644 --- a/syntax/position_test.go +++ b/syntax/position_test.go @@ -22,30 +22,39 @@ service Svc { void f() } if len(doc.Includes()) != 1 || doc.Includes()[0].Path.Text != `"a.thrift"` { t.Errorf("Includes = %+v", doc.Includes()) } + if len(doc.CPPIncludes()) != 1 { t.Errorf("CPPIncludes = %d", len(doc.CPPIncludes())) } + if len(doc.Namespaces()) != 1 { t.Errorf("Namespaces = %d", len(doc.Namespaces())) } + if len(doc.Consts()) != 1 || doc.Consts()[0].Name.Text != "C" { t.Errorf("Consts = %+v", doc.Consts()) } + if len(doc.Typedefs()) != 1 { t.Errorf("Typedefs = %d", len(doc.Typedefs())) } + if len(doc.Enums()) != 1 { t.Errorf("Enums = %d", len(doc.Enums())) } + if len(doc.Structs()) != 1 || doc.Structs()[0].Name.Text != "S" { t.Errorf("Structs = %+v", doc.Structs()) } + if len(doc.Unions()) != 1 { t.Errorf("Unions = %d", len(doc.Unions())) } + if len(doc.Exceptions()) != 1 || doc.Exceptions()[0].Name.Text != "X" { t.Errorf("Exceptions = %+v", doc.Exceptions()) } + if len(doc.Services()) != 1 { t.Errorf("Services = %d", len(doc.Services())) } @@ -73,6 +82,7 @@ func TestDocumentRanges(t *testing.T) { if !doc.Contains(s, Position{Line: 2, Col: 7, Offset: 18}) { t.Error("position inside struct should be contained") } + if doc.Contains(s, Position{Line: 5, Col: 1, Offset: 25}) { t.Error("position outside struct should not be contained") } @@ -98,6 +108,7 @@ const i32 X = SOME_VALUE if idx < 0 { t.Fatalf("text %q not found in source", text) } + return doc.TokenPosition(tokIndexAt(doc, idx)) } @@ -120,6 +131,7 @@ const i32 X = SOME_VALUE if len(path) == 0 { t.Fatal("empty path") } + deepest := path[len(path)-1] if got := typeName(deepest); got != tt.want { t.Errorf("deepest node = %s, want %s (path: %v)", got, tt.want, pathNames(path)) @@ -132,6 +144,7 @@ const i32 X = SOME_VALUE if len(path) < 3 { t.Fatalf("path too short: %v", pathNames(path)) } + if _, ok := path[len(path)-2].(*Field); !ok { t.Errorf("parent of field name should be Field: %v", pathNames(path)) } @@ -141,6 +154,7 @@ const i32 X = SOME_VALUE if len(path) < 3 { t.Fatalf("path too short: %v", pathNames(path)) } + if _, ok := path[len(path)-2].(*Struct); !ok { t.Errorf("parent of struct name should be Struct: %v", pathNames(path)) } @@ -150,8 +164,10 @@ const i32 X = SOME_VALUE if idx < 0 { t.Fatal("const not found") } + pos := doc.TokenPosition(tokIndexAt(doc, idx)) pos.Offset -= 1 // the last newline of the blank line before const + path = doc.SearchNodePathByPosition(pos) if len(path) != 1 { // only the document itself t.Errorf("expected document-only path, got %v", pathNames(path)) @@ -161,11 +177,13 @@ const i32 X = SOME_VALUE func tokIndexAt(doc *Document, offset int) int { // Find the last token starting at or before offset. idx := 0 + for i, tok := range doc.Tokens { if tok.Offset <= offset { idx = i } } + return idx } @@ -204,6 +222,7 @@ func typeName(n Node) string { case *CPPInclude: return "*syntax.CPPInclude" } + return "unknown" } @@ -212,6 +231,7 @@ func pathNames(path []Node) []string { for _, n := range path { names = append(names, typeName(n)) } + return names } @@ -221,5 +241,6 @@ func indexOf(s, sub string) int { return i } } + return -1 }