diff --git a/check_test.go b/check_test.go index 2e61ece..00a36aa 100644 --- a/check_test.go +++ b/check_test.go @@ -85,7 +85,7 @@ func Test_CheckMadeInAbyss(t *testing.T) { // its id (FieldIDCheck) and its name (DuplicateCheck). var repeatLine uint32 for _, d := range lints { - if strings.Contains(string(d.Message.(protocol.String)), "duplicate field repeat") { + if strings.Contains(message(d), "duplicate field repeat") { repeatLine = d.Range.Start.Line } } @@ -131,10 +131,15 @@ func corpusAbs(t *testing.T, name string) string { return p } +// message is the diagnostic message as plain text. +func message(d protocol.Diagnostic) string { + return string(d.Message.(protocol.String)) +} + // hasMessage reports whether any diagnostic carries msg. func hasMessage(diags []protocol.Diagnostic, msg string) bool { for _, d := range diags { - if strings.Contains(string(d.Message.(protocol.String)), msg) { + if strings.Contains(message(d), msg) { return true } } diff --git a/cli_test.go b/cli_test.go index b25993c..3a2e8b6 100644 --- a/cli_test.go +++ b/cli_test.go @@ -99,7 +99,7 @@ func Test_FormatCli_GoldenEnums(t *testing.T) { func Test_CheckCLI_MadeInAbyss(t *testing.T) { stdout, stderr, err := runCLI(t, "check", "tests/made-in-abyss") require.Error(t, err) - assert.Equal(t, "", stderr) + assert.Empty(t, stderr) assert.Contains(t, stdout, "lints.thrift:") assert.Contains(t, stdout, "cycle_a.thrift:") diff --git a/diff.go b/diff.go index 55f58b3..44ba4bf 100644 --- a/diff.go +++ b/diff.go @@ -77,7 +77,7 @@ func Diff(oldName string, old []byte, newName string, new []byte) []byte { // Expand matching lines as far possible, // establishing that x[start.x:end.x] == y[start.y:end.y]. - // Note that on the first (or last) iteration we may (or definitey do) + // Note that on the first (or last) iteration we may (or definitely do) // have an empty match: start.x==end.x and start.y==end.y. start := m for start.x > done.x && start.y > done.y && x[start.x-1] == y[start.y-1] { @@ -196,7 +196,7 @@ func lines(x []byte) []string { // Subsequence Problem,” Princeton TR #170 (January 1975), // available at https://research.swtch.com/tgs170.pdf. func tgs(x, y []string) []pair { - // Count the number of times each string appears in a and b. + // Count the number of times each string appears in x and y. // We only care about 0, 1, many, counted as 0, -1, -2 // for the x side and 0, -4, -8 for the y side. // Using negative numbers now lets us distinguish positive line numbers later. diff --git a/doc/dump.go b/doc/dump.go index 458f27d..1beab77 100644 --- a/doc/dump.go +++ b/doc/dump.go @@ -21,17 +21,9 @@ func dumpDoc(b *strings.Builder, d Doc, ind string) { case nil: fmt.Fprintf(b, "%s\n", ind) case Concat: - fmt.Fprintf(b, "%sConcat\n", ind) - - for _, c := range v { - dumpDoc(b, c, ind+" ") - } + dumpConcat(b, v, ind) case *concatNode: - fmt.Fprintf(b, "%sConcat\n", ind) - - for _, c := range v.parts { - dumpDoc(b, c, ind+" ") - } + dumpConcat(b, v.parts, ind) case *group: extra := "" if v.id != 0 { @@ -74,17 +66,9 @@ func dumpDoc(b *strings.Builder, d Doc, ind string) { 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)) - } + dumpText(b, ind, string(v)) case *textNode: - if v.s == "" { - fmt.Fprintf(b, "%sText \"\"\n", ind) - } else { - fmt.Fprintf(b, "%sText %q\n", ind, v.s) - } + dumpText(b, ind, v.s) case LineDoc: kind := "Line" if v.Hard { @@ -104,3 +88,21 @@ func dumpDoc(b *strings.Builder, d Doc, ind string) { fmt.Fprintf(b, "%s%T\n", ind, d) } } + +func dumpConcat(b *strings.Builder, parts []Doc, ind string) { + fmt.Fprintf(b, "%sConcat\n", ind) + + for _, c := range parts { + dumpDoc(b, c, ind+" ") + } +} + +func dumpText(b *strings.Builder, ind, s string) { + if s == "" { + fmt.Fprintf(b, "%sText \"\"\n", ind) + + return + } + + fmt.Fprintf(b, "%sText %q\n", ind, s) +} diff --git a/doc/print.go b/doc/print.go index f9bfa22..07375f3 100644 --- a/doc/print.go +++ b/doc/print.go @@ -153,7 +153,6 @@ func (p *printer) reset(o Options) { } type printer struct { - fitsDBG bool o Options position int out []byte @@ -290,33 +289,15 @@ func (p *printer) run(d Doc) (string, error) { p.position -= p.trim() case *group: - { - gcmd := p.printGroup(cmd, v, commands) + gcmd := p.printGroup(cmd, v, commands) - commands = append(commands, gcmd) - if v.id != 0 { - p.groupMode[v.id] = gcmd.mode - } + commands = append(commands, gcmd) + if v.id != 0 { + p.groupMode[v.id] = gcmd.mode } case *ifBreak: - groupMode := cmd.mode - if v.groupID != 0 { - if m, ok := p.groupMode[v.groupID]; ok { - groupMode = m - } else { - groupMode = modeFlat - } - } - - var contents Doc - if groupMode == modeBreak { - contents = v.breakDoc - } else { - contents = v.flatDoc - } - - if contents != nil { + if contents := p.ifBreakContents(v, cmd.mode); contents != nil { commands = append(commands, command{indentation: cmd.indentation, mode: cmd.mode, doc: contents}) } @@ -421,6 +402,25 @@ func (p *printer) run(d Doc) (string, error) { return string(p.out), nil } +// ifBreakContents returns the doc an IfBreak prints in a group whose mode is +// cmdMode: the mode of the named group when IfBreakFor names one, else the +// enclosing mode. +func (p *printer) ifBreakContents(v *ifBreak, cmdMode mode) Doc { + if v.groupID != 0 { + if m, ok := p.groupMode[v.groupID]; ok { + cmdMode = m + } else { + cmdMode = modeFlat + } + } + + if cmdMode == modeBreak { + return v.breakDoc + } + + return v.flatDoc +} + // printGroup decides whether g fits in the remaining width and returns the // command to print it. With expanded states it tries each state in order. func (p *printer) printGroup(cmd command, g *group, rest []command) command { @@ -481,10 +481,6 @@ func (p *printer) fits(next command, rest []command, remainingWidth int, hasLine defer func() { p.fitCommands = commands[:0] }() 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 @@ -564,23 +560,7 @@ func (p *printer) fits(next command, rest []command, remainingWidth int, hasLine commands = append(commands, command{indentation: cmd.indentation, mode: groupMode, doc: contents}) case *ifBreak: - groupMode := cmd.mode - if v.groupID != 0 { - if m, ok := p.groupMode[v.groupID]; ok { - groupMode = m - } else { - groupMode = modeFlat - } - } - - var contents Doc - if groupMode == modeBreak { - contents = v.breakDoc - } else { - contents = v.flatDoc - } - - if contents != nil { + if contents := p.ifBreakContents(v, cmd.mode); contents != nil { commands = append(commands, command{indentation: cmd.indentation, mode: cmd.mode, doc: contents}) } @@ -662,7 +642,7 @@ func propagateBreaks(d Doc) { } } } - traverseDoc(d, enter, exit, true) + traverseDoc(d, enter, exit) } func breakParentGroup(stack []*group) { @@ -671,7 +651,7 @@ func breakParentGroup(stack []*group) { } } -func traverseDoc(d Doc, enter func(Doc) bool, exit func(Doc), includeConditionalGroups bool) { +func traverseDoc(d Doc, enter func(Doc) bool, exit func(Doc)) { if d == nil || !enter(d) { return } @@ -679,29 +659,27 @@ func traverseDoc(d Doc, enter func(Doc) bool, exit func(Doc), includeConditional switch v := d.(type) { case Concat: for _, part := range v { - traverseDoc(part, enter, exit, includeConditionalGroups) + traverseDoc(part, enter, exit) } case *concatNode: for _, part := range v.parts { - traverseDoc(part, enter, exit, includeConditionalGroups) + traverseDoc(part, enter, exit) } case *group: - if includeConditionalGroups { - for _, state := range v.expanded { - traverseDoc(state, enter, exit, includeConditionalGroups) - } + for _, state := range v.expanded { + traverseDoc(state, enter, exit) } - traverseDoc(v.doc, enter, exit, includeConditionalGroups) + traverseDoc(v.doc, enter, exit) case *indent: - traverseDoc(v.doc, enter, exit, includeConditionalGroups) + traverseDoc(v.doc, enter, exit) case *align: - traverseDoc(v.doc, enter, exit, includeConditionalGroups) + traverseDoc(v.doc, enter, exit) case *ifBreak: - traverseDoc(v.breakDoc, enter, exit, includeConditionalGroups) - traverseDoc(v.flatDoc, enter, exit, includeConditionalGroups) + traverseDoc(v.breakDoc, enter, exit) + traverseDoc(v.flatDoc, enter, exit) case *lineSuffix: - traverseDoc(v.doc, enter, exit, includeConditionalGroups) + traverseDoc(v.doc, enter, exit) } exit(d) diff --git a/flake.nix b/flake.nix index c4d7862..eaa3549 100644 --- a/flake.nix +++ b/flake.nix @@ -12,7 +12,7 @@ thriftLs = pkgs: let - version = "0.1.3"; + version = "0.1.4"; in pkgs.buildGoModule { pname = "thrift-ls"; diff --git a/formatter/body.go b/formatter/body.go index 01df76a..a3f21ec 100644 --- a/formatter/body.go +++ b/formatter/body.go @@ -22,6 +22,17 @@ func (f *formatter) structLike(v *syntax.Struct) doc.Doc { return f.Concat(parts...) } +// bracedBody 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 bracedGroup. +func (f *formatter) bracedBody(fields []*syntax.Field, open, close int, closeTrailing bool, c Construct) doc.Doc { + bodyID := f.id() + sepMode := f.opts.Separator.Get(c) + forced := f.opts.Break.Get(c) || sepForcesBreak(sepsOfFields(fields), sepMode) + + return f.bracedGroup(f.fieldList(fields, bodyID, sepMode), bodyID, len(fields), open, close, closeTrailing, forced) +} + // enum formats an enum declaration. func (f *formatter) enum(v *syntax.Enum) doc.Doc { open := f.scanKind(v.TokStart(), v.TokEnd(), syntax.TokenLBrace) @@ -49,12 +60,6 @@ func (f *formatter) scanKind(start, end int, kind syntax.TokenKind) int { 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. 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 { @@ -67,15 +72,20 @@ func (f *formatter) constructOf(kind syntax.StructKind) Construct { return ConstructStruct } -func (f *formatter) bracedBody(fields []*syntax.Field, open, close int, closeTrailing bool, c Construct) doc.Doc { - bodyID := f.id() - sepMode := f.opts.Separator.Get(c) - inner := append([]doc.Doc{doc.Line, f.fieldList(fields, bodyID, sepMode)}, f.ownLineComments(close)...) - closeBreak := f.IfBreak(doc.SoftLine, f.Text(" ")) +// bracedGroup assembles "{ body }" from the prebuilt body list: flat as +// "S { 1: i32 a }" when it fits, otherwise one item per line. bodyID is +// the group id the list's IfBreakFor references; n is the item count; +// forced requires the broken layout. 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. +func (f *formatter) bracedGroup(body doc.Doc, bodyID, n, open, close int, closeTrailing, forced bool) doc.Doc { + inner := f.ownLineComments(close) + closeBreak := f.IfBreak(doc.SoftLine, f.Concat()) - if len(fields) == 0 { - inner = append([]doc.Doc{}, f.ownLineComments(close)...) - closeBreak = f.IfBreak(doc.SoftLine, f.Concat()) + if n > 0 { + inner = append([]doc.Doc{doc.Line, body}, inner...) + closeBreak = f.IfBreak(doc.SoftLine, f.Text(" ")) } openComments := f.sameLineComments(open) @@ -91,12 +101,9 @@ func (f *formatter) bracedBody(fields []*syntax.Field, open, close int, closeTra closeBreak, f.emitTokens(close, close, emitOpts{trailing: closeTrailing}), ) - if len(fields) > 0 && (f.opts.Break.Get(c) || sepForcesBreak(sepsOfFields(fields), sepMode)) { + if n > 0 && forced { // BreakParent inside the group forces it to the broken layout. - p := f.Parts(2) - p = append(p, doc.BreakParent) - p = append(p, content) - content = f.Concat(p...) + content = f.Concat(doc.BreakParent, content) } return f.GroupID(bodyID, content) @@ -105,35 +112,10 @@ func (f *formatter) bracedBody(fields []*syntax.Field, open, close int, closeTra // bracedEnumBody is bracedBody for enum values. 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.ownLineComments(close)...) - closeBreak := f.IfBreak(doc.SoftLine, f.Text(" ")) - - if len(values) == 0 { - inner = append([]doc.Doc{}, f.ownLineComments(close)...) - closeBreak = f.IfBreak(doc.SoftLine, f.Concat()) - } - - openComments := f.sameLineComments(open) + sepMode := f.opts.Separator.Get(ConstructEnum) + forced := f.opts.Break.Get(ConstructEnum) || sepForcesBreak(sepsOfValues(values), sepMode) - openDoc := append([]doc.Doc{f.Text(" {")}, openComments...) - if len(openComments) > 0 { - openDoc = append(openDoc, doc.BreakParent) - } - - content := f.Concat( - f.Concat(openDoc...), - f.Indent(f.Concat(inner...)), - closeBreak, - f.emitTokens(close, close, emitOpts{trailing: closeTrailing}), - ) - if len(values) > 0 && (f.opts.Break.Get(ConstructEnum) || sepForcesBreak(sepsOfValues(values), f.opts.Separator.Get(ConstructEnum))) { - p := f.Parts(2) - p = append(p, doc.BreakParent) - p = append(p, content) - content = f.Concat(p...) - } - - return f.GroupID(bodyID, content) + return f.bracedGroup(f.enumValueList(values, bodyID), bodyID, len(values), open, close, closeTrailing, forced) } // service formats a service declaration. The body is always multiline: @@ -149,24 +131,22 @@ func (f *formatter) service(v *syntax.Service) doc.Doc { // The separator line collapses when the previous function // ended with a line comment (which owns its line end). parts = append(parts, doc.Line) - parts = append(parts, f.blankLines(fn, doc.HardLine)...) - } else { - parts = append(parts, f.blankLines(fn, doc.HardLine)...) } + parts = append(parts, f.blankLines(fn, doc.HardLine)...) parts = append(parts, f.function(fn)) } p := f.Parts(2) p = append(p, f.Concat(parts...)) p = append(p, f.Concat(f.ownLineComments(close)...)) - inner := f.Concat(p...) + 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 = f.Concat(doc.Line, f.Concat(parts...), f.Concat(f.ownLineComments(close)...)) + inner = f.Concat(doc.Line, inner) } else if !f.hasOwnLineComments(close) && f.token(close).BlankLinesBefore > 0 { // Empty body with blank lines before the close and no comments: // the blanks round-trip through the close's own line. @@ -234,20 +214,14 @@ func (f *formatter) functionBody(v *syntax.Function) doc.Doc { // field per line otherwise, like the throws clause. args := f.parenGroup(v.Args, open, f.parenClose(v.Args, open), false, argsMode) - if v.Throws == nil { - return f.Group(f.Concat( - header, - args, - f.functionTail(v, open), - )) + parts := f.Parts(4) + parts = append(parts, header, args) + if v.Throws != nil { + parts = append(parts, f.throwsGroup(v)) } + parts = append(parts, f.functionTail(v, open)) - return f.Group(f.Concat( - header, - args, - f.throwsGroup(v), - f.functionTail(v, open), - )) + return f.Group(f.Concat(parts...)) } // parenGroup renders "(fields)" as its own group, folding independently: @@ -422,11 +396,9 @@ func (f *formatter) brokenFields(fields []*syntax.Field, sepMode SeparatorMode) // The separator line collapses when the previous field ended // with a line comment (which owns its line end). parts = append(parts, doc.Line) - parts = append(parts, f.blankLines(field, doc.HardLine)...) - } else { - parts = append(parts, f.blankLines(field, doc.HardLine)...) } + parts = append(parts, f.blankLines(field, doc.HardLine)...) parts = append(parts, f.fieldDoc(field, f.alignmentFor(fields, i, sepMode), 0, sepMode)) } diff --git a/formatter/field.go b/formatter/field.go index daf5186..bb07229 100644 --- a/formatter/field.go +++ b/formatter/field.go @@ -19,12 +19,10 @@ func (f *formatter) fieldList(fields []*syntax.Field, bodyID int, sepMode Separa for i, field := range fields { if i > 0 { parts = append(parts, f.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, sepMode), bodyID, sepMode)) + parts = append(parts, f.blankLines(field, doc.HardLine)...) + parts = append(parts, f.fieldDoc(field, f.alignmentFor(fields, i, sepMode), bodyID, sepMode)) } return f.Concat(parts...) @@ -72,11 +70,9 @@ func (f *formatter) enumValueList(values []*syntax.EnumValue, bodyID int) doc.Do for i, value := range values { if i > 0 { parts = append(parts, f.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.blankLines(value, doc.HardLine)...) parts = append(parts, f.enumValue(value, f.alignmentForEnum(values, i, f.opts.Separator.Get(ConstructEnum)), bodyID)) } @@ -323,35 +319,34 @@ func (f *formatter) nodeTrailingInline(end int, sep syntax.TokenKind, sepMode Se // forms on the body group's break state; with bodyID zero the broken form // renders directly (paren bodies are always broken). func (f *formatter) fieldDoc(v *syntax.Field, align *columnAlign, bodyID int, sepMode SeparatorMode) doc.Doc { + broken := f.Concat( + f.fieldContent(v, align, true, sepMode), + f.trailingSep(v.Sep, sepMode), + ) + content := f.fieldContent(v, align, false, sepMode) if bodyID != 0 { - broken := f.Parts(2) - broken = append(broken, f.fieldContent(v, align, true, sepMode)) - broken = append(broken, f.trailingSep(v.Sep, sepMode)) - - content = f.IfBreakFor(f.Concat(broken...), content, bodyID) + content = f.IfBreakFor(broken, content, bodyID) } else { - broken := f.Parts(2) - broken = append(broken, f.fieldContent(v, align, true, sepMode)) - broken = append(broken, f.trailingSep(v.Sep, sepMode)) - - content = f.Concat(broken...) + content = broken } parts := append(f.ownLineComments(v.TokStart()), content) - if f.nodeTrailingInline(v.TokEnd(), v.Sep, sepMode) { - parts = append(parts, f.sameLineComments(v.TokEnd())...) - } else { - parts = append(parts, f.suppressedSepComments(v.TokEnd())...) - } + parts = append(parts, f.itemTrailing(v.TokEnd(), v.Sep, sepMode)...) return f.Concat(parts...) } -// field assembles a struct-like body field, switching on the body group's -// break state. -func (f *formatter) field(v *syntax.Field, align *columnAlign, bodyID int, sepMode SeparatorMode) doc.Doc { - return f.fieldDoc(v, align, bodyID, sepMode) +// itemTrailing renders the same-line comments after the item's last token +// (its separator, when present): inline, or each on its own line when the +// separator text was dropped and the comments do not share the previous +// content's line. +func (f *formatter) itemTrailing(end int, sep syntax.TokenKind, sepMode SeparatorMode) []doc.Doc { + if f.nodeTrailingInline(end, sep, sepMode) { + return f.sameLineComments(end) + } + + return f.suppressedSepComments(end) } // emitWithAnnotations renders a token run split at the node's annotations, @@ -453,22 +448,23 @@ func (f *formatter) fieldPads(v *syntax.Field, a *columnAlign) ([]padEntry, stri // 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, f.opts.Separator.Get(ConstructEnum)) - if bodyID != 0 { - broken := f.Parts(2) - broken = append(broken, f.enumValueContent(v, align, true, f.opts.Separator.Get(ConstructEnum))) - broken = append(broken, f.trailingSep(v.Sep, f.opts.Separator.Get(ConstructEnum))) + sepMode := f.opts.Separator.Get(ConstructEnum) - content = f.IfBreakFor(f.Concat(broken...), content, bodyID) - } + broken := f.Concat( + f.enumValueContent(v, align, true, sepMode), + f.trailingSep(v.Sep, sepMode), + ) - parts := append(f.ownLineComments(v.TokStart()), content) - if f.nodeTrailingInline(v.TokEnd(), v.Sep, f.opts.Separator.Get(ConstructEnum)) { - parts = append(parts, f.sameLineComments(v.TokEnd())...) + content := f.enumValueContent(v, align, false, sepMode) + if bodyID != 0 { + content = f.IfBreakFor(broken, content, bodyID) } else { - parts = append(parts, f.suppressedSepComments(v.TokEnd())...) + content = broken } + parts := append(f.ownLineComments(v.TokStart()), content) + parts = append(parts, f.itemTrailing(v.TokEnd(), v.Sep, sepMode)...) + return f.Concat(parts...) } diff --git a/formatter/format.go b/formatter/format.go index a5cfed8..c86220e 100644 --- a/formatter/format.go +++ b/formatter/format.go @@ -342,11 +342,6 @@ type padEntry struct { text string } -// containsInt reports whether xs contains v. -func containsInt(xs []int, v int) bool { - return slices.Contains(xs, v) -} - // padAt returns the combined pads for the token index, or "". Multiple // entries at the same index (id pad + requiredness column) concatenate. func padAt(pads []padEntry, idx int) string { @@ -409,7 +404,7 @@ func (f *formatter) emitTokens(start, end int, o emitOpts) doc.Doc { continue } - skipped := containsInt(o.skipText, i) + skipped := slices.Contains(o.skipText, i) if !first { // Comments between the previous real token and this one @@ -446,23 +441,26 @@ func (f *formatter) emitTokens(start, end int, o emitOpts) doc.Doc { first = false } - if o.trailing { - // Same-line comments after the last real token. When the token's - // text is suppressed (the separator mode drops it), its same-line - // comments render inline only when they also share the previous - // content's line; otherwise they start their own line, so the - // output round-trips — the next emission would skip them as - // same-line with the suppressed token. - if containsInt(o.skipText, prev) && o.text == "" { - prevTok := f.prevReal(prev - 1) - if prevTok >= 0 && f.token(prevTok).Line == f.token(prev).Line { - parts = append(parts, f.sameLineComments(prev)...) - } else { - parts = append(parts, f.suppressedSepComments(prev)...) - } - } else { - parts = append(parts, f.sameLineComments(prev)...) - } + if !o.trailing { + return f.Concat(parts...) + } + + if !slices.Contains(o.skipText, prev) || o.text != "" { + parts = append(parts, f.sameLineComments(prev)...) + + return f.Concat(parts...) + } + + // Same-line comments after the last real token. When the token's text + // is suppressed (the separator mode drops it), its same-line comments + // render inline only when they also share the previous content's line; + // otherwise they start their own line, so the output round-trips — the + // next emission would skip them as same-line with the suppressed token. + prevTok := f.prevReal(prev - 1) + if prevTok >= 0 && f.token(prevTok).Line == f.token(prev).Line { + parts = append(parts, f.sameLineComments(prev)...) + } else { + parts = append(parts, f.suppressedSepComments(prev)...) } return f.Concat(parts...) @@ -601,39 +599,31 @@ func (f *formatter) cppInclude(v *syntax.CPPInclude) doc.Doc { } func (f *formatter) namespace(v *syntax.Namespace) doc.Doc { - end := v.TokEnd() - if v.Annotations != nil { - end = v.Annotations.TokStart() - 1 - } - - o := emitOpts{} - if v.Annotations != nil { - o.trailing = true - } - - parts := f.Parts(3) - parts = append(parts, f.emitTokens(v.TokStart(), end, o)) - 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 f.Concat(parts...) + return f.headerWithAnnotations(v.TokStart(), v.TokEnd(), v.Annotations) } func (f *formatter) typedef(v *syntax.Typedef) doc.Doc { - end := v.TokEnd() - if v.Annotations != nil { - end = v.Annotations.TokStart() - 1 + return f.headerWithAnnotations(v.TokStart(), v.TokEnd(), v.Annotations) +} + +// headerWithAnnotations emits the node's header tokens up to its +// annotations (which keep their own foldable group), and any stray tokens +// after them. +func (f *formatter) headerWithAnnotations(start, end int, ann *syntax.Annotations) doc.Doc { + headerEnd := end + if ann != nil { + headerEnd = ann.TokStart() - 1 } o := emitOpts{} - if v.Annotations != nil { + if ann != nil { o.trailing = true } parts := f.Parts(3) - parts = append(parts, f.emitTokens(v.TokStart(), end, o)) - parts = append(parts, f.annotationsDoc(v.Annotations, v.Annotations != nil && v.Annotations.TokEnd() == v.TokEnd())) - parts = append(parts, f.afterAnnotations(v.Annotations, v.TokEnd())) + parts = append(parts, f.emitTokens(start, headerEnd, o)) + parts = append(parts, f.annotationsDoc(ann, ann != nil && ann.TokEnd() == end)) + parts = append(parts, f.afterAnnotations(ann, end)) return f.Concat(parts...) } diff --git a/formatter/value.go b/formatter/value.go index 1511b4b..7f82cec 100644 --- a/formatter/value.go +++ b/formatter/value.go @@ -193,20 +193,18 @@ func (f *formatter) itemSep(sep int, mode SeparatorMode) []doc.Doc { text = "" } - if text == f.token(sep).Text { - p := f.Parts(2) - p = append(p, f.emitTokens(sep, sep, emitOpts{leading: true, trailing: true})) - p = append(p, f.foldBreak(sep, " ")) - - return p - } - // Forced separator differing from the source: the forced text replaces // the suppressed text inside the run, so the source token's comments // stay ordered around it — own-line comments before it, same-line // comments after. + o := emitOpts{leading: true, trailing: true} + if text != f.token(sep).Text { + o.skipText = []int{sep} + o.text = text + } + p := f.Parts(2) - p = append(p, f.emitTokens(sep, sep, emitOpts{leading: true, trailing: true, skipText: []int{sep}, text: text})) + p = append(p, f.emitTokens(sep, sep, o)) p = append(p, f.foldBreak(sep, " ")) return p diff --git a/lsp/cache/context.go b/lsp/cache/context.go index 03a9806..8b2d39f 100644 --- a/lsp/cache/context.go +++ b/lsp/cache/context.go @@ -20,7 +20,7 @@ func NewIncludeDeps() *IncludeDeps { return &IncludeDeps{graph: NewIncludeGraph()} } -// Includes returns the files file includes directly, in include order. +// Includes returns the files file includes directly, sorted ascending by URI. func (c *IncludeDeps) Includes(file uri.URI) []uri.URI { node := c.graph.Get(file) if node == nil { diff --git a/lsp/cache/graph.go b/lsp/cache/graph.go index ad27ca0..a8764e9 100644 --- a/lsp/cache/graph.go +++ b/lsp/cache/graph.go @@ -197,7 +197,6 @@ func (g *IncludeGraph) removeWithoutLock(file uri.URI) { continue } - // update outNode indegree for i := range outNode.indegree { if outNode.indegree[i] == file { outNode.indegree = append(outNode.indegree[0:i], outNode.indegree[i+1:]...) diff --git a/lsp/cache/parse_test.go b/lsp/cache/parse_test.go index d51ee04..13d38b4 100644 --- a/lsp/cache/parse_test.go +++ b/lsp/cache/parse_test.go @@ -8,21 +8,16 @@ import ( ) func TestParse(t *testing.T) { - type args struct { - fh FileHandle - } - tests := []struct { name string - args args + fh FileHandle assertion assert.ErrorAssertionFunc }{ { name: "normal", - args: args{ - fh: &Overlay{ - uri: "file:///tmp/types.thrift", - content: []byte(` + fh: &Overlay{ + uri: "file:///tmp/types.thrift", + content: []byte(` #include "base.thrift" struct Xtruct3 { @@ -31,18 +26,16 @@ struct Xtruct3 9: i32 i32_thing, 11: i64 i64_thing } - `), - version: 0, - }, + `), + version: 0, }, assertion: assert.NoError, }, { name: "invalid ast", - args: args{ - fh: &Overlay{ - uri: "file:///tmp/types.thrift", - content: []byte(` + fh: &Overlay{ + uri: "file:///tmp/types.thrift", + content: []byte(` #include "base.thrift" struct Xtruct3 { @@ -52,16 +45,15 @@ struct Xtruct3 11: i64 i64_thing, 12: } - `), - version: 0, - }, + `), + version: 0, }, assertion: assert.NoError, }, } for _, tt := range tests { t.Run(tt.name, func(t *testing.T) { - got, err := Parse(tt.args.fh) + got, err := Parse(tt.fh) tt.assertion(t, err) t.Logf("got: %v\n", got) }) diff --git a/lsp/cache/view.go b/lsp/cache/view.go index 0a4bc2a..8bebb7b 100644 --- a/lsp/cache/view.go +++ b/lsp/cache/view.go @@ -82,6 +82,12 @@ func (v *View) Config() options.Patch { return v.config } +// WalkFiles enumerates the view's file source under root: the disk in +// production, the in-memory tree in tests. +func (v *View) WalkFiles(ctx context.Context, root uri.URI, fn func(uri.URI) error) error { + return v.fs.WalkFiles(ctx, root, fn) +} + func (v *View) MarkFileKnown(fileURI uri.URI) { v.knownFilesMu.Lock() defer v.knownFilesMu.Unlock() diff --git a/lsp/mapper/apply.go b/lsp/mapper/apply.go index abf401d..1dd6c86 100644 --- a/lsp/mapper/apply.go +++ b/lsp/mapper/apply.go @@ -33,7 +33,6 @@ func (m *Mapper) ApplyEdits(edits []protocol.TextEdit) ([]byte, error) { all = append(all, pending{start, end, e.NewText}) } - // Later edits apply first so earlier offsets stay valid. sort.Slice(all, func(i, j int) bool { return all[i].start > all[j].start }) buf := m.content @@ -53,6 +52,9 @@ func (m *Mapper) ApplyEdits(edits []protocol.TextEdit) ([]byte, error) { // content. func (m *Mapper) offsetAt(pos protocol.Position) (int, error) { p, err := m.LSPPosToParserPosition(protocol.Position{Line: pos.Line, Character: pos.Character}) + if err != nil { + return 0, err + } - return p.Offset, err + return p.Offset, nil } diff --git a/lsp/mapper/mapper.go b/lsp/mapper/mapper.go index a09f8b1..a5c6e9e 100644 --- a/lsp/mapper/mapper.go +++ b/lsp/mapper/mapper.go @@ -2,7 +2,6 @@ package mapper import ( "bytes" - "errors" "fmt" "sort" "sync" @@ -21,7 +20,7 @@ type Mapper struct { nonASCII bool } -// NewMapper ... +// NewMapper returns a Mapper for the given document content. func NewMapper(content []byte) *Mapper { return &Mapper{ content: content, @@ -76,7 +75,8 @@ func (m *Mapper) OffsetToLSPPosition(offset int) (protocol.Position, error) { }, nil } -// convert from utf16-based to rune-based position +// LSPPosToParserPosition converts an LSP position (0-based line, UTF-16 +// code-unit column) to a parser position (1-based line, rune-based column). func (m *Mapper) LSPPosToParserPosition(pos protocol.Position) (syntax.Position, error) { m.initLineStart() @@ -85,23 +85,22 @@ func (m *Mapper) LSPPosToParserPosition(pos protocol.Position) (syntax.Position, return syntax.InvalidPosition, fmt.Errorf("invalid position line, request line: %d, total line: %d", line, len(m.lineStart)) } + lineStart := m.lineStart[pos.Line] + lineEnd := len(m.content) + if line < len(m.lineStart) { + lineEnd = m.lineStart[line] + } + if !m.nonASCII { col := int(pos.Character) + 1 - offset := m.lineStart[pos.Line] + int(pos.Character) + offset := lineStart + 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] - } else { - lineLength = m.lineStart[pos.Line+1] - m.lineStart[pos.Line] - } - - if col > lineLength+1 { // if line length is 0, col is 1 means col is at end of line - return syntax.InvalidPosition, fmt.Errorf("invalid position column: %d, line length: %d, %s", col, lineLength, string(m.content)) + if col > lineEnd-lineStart+1 { // if line length is 0, col is 1 means col is at end of line + return syntax.InvalidPosition, fmt.Errorf("invalid position column: %d, line length: %d, %s", col, lineEnd-lineStart, string(m.content)) } return syntax.Position{ @@ -111,25 +110,12 @@ func (m *Mapper) LSPPosToParserPosition(pos protocol.Position) (syntax.Position, }, nil } - 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 - } - + for len(lineBytes) > 0 && utf16Col < int(pos.Character) { if lineBytes[0] < utf8.RuneSelf { utf16Col++ lineBytes = lineBytes[1:] @@ -151,11 +137,6 @@ func (m *Mapper) LSPPosToParserPosition(pos protocol.Position) (syntax.Position, runeLen := utf8.RuneCount(m.content[lineStart : lineStart+bytesCol]) - offset := lineStart + bytesCol - if offset > len(m.content) { - return syntax.InvalidPosition, errors.New("invalid position character") - } - return syntax.Position{ Line: line, Col: runeLen + 1, diff --git a/lsp/server.go b/lsp/server.go index f6e6bcf..5074d0c 100644 --- a/lsp/server.go +++ b/lsp/server.go @@ -101,6 +101,7 @@ func (s *Server) addFolderView(folder uri.URI) *cache.View { // viewConfig resolves the config for a view rooted at folder: the pinned // --config file, or the nearest thrift-ls.json walking up, plus CLI flags. +// A folder with no usable config formats with defaults. func (s *Server) viewConfig(folder uri.URI) options.Patch { if s.configPath != "" { return s.cli.Apply(s.explicit) @@ -110,23 +111,29 @@ func (s *Server) viewConfig(folder uri.URI) options.Patch { if err != nil { slog.Error("config discovery failed", "dir", folder.FsPath(), "err", err) - return s.cli.Apply(options.Default()) + return s.defaultConfig() } if cfgPath == "" { - return s.cli.Apply(options.Default()) + return s.defaultConfig() } cfg, err := options.Load(cfgPath) if err != nil { slog.Error("config file rejected", "path", cfgPath, "err", err) - return s.cli.Apply(options.Default()) + return s.defaultConfig() } return s.cli.Apply(options.Effective(cfg)) } +// defaultConfig is the fallback for a folder without a usable config +// file: the defaults with the CLI overlay. +func (s *Server) defaultConfig() options.Patch { + return s.cli.Apply(options.Default()) +} + // applyLogLevel applies the first view config's log level; the logger is // process-wide, so later views keep it. func (s *Server) applyLogLevel(cfg options.Patch) { diff --git a/lsp/source/cross_reference_test.go b/lsp/source/cross_reference_test.go index 266b56d..898a586 100644 --- a/lsp/source/cross_reference_test.go +++ b/lsp/source/cross_reference_test.go @@ -48,7 +48,7 @@ func TestReferenceQualifiedCrossFile(t *testing.T) { mainFile := `include "federation.gundam.thrift" struct StrikeRouge { - 1: required federation.Gundam pack + 1: required federation.gundam.Gundam pack }` gundamFile := `struct Gundam { diff --git a/lsp/source/cycle_detect.go b/lsp/source/cycle_detect.go index 1486182..ee1bfe3 100644 --- a/lsp/source/cycle_detect.go +++ b/lsp/source/cycle_detect.go @@ -20,7 +20,7 @@ func (c *CycleCheck) Diagnostic(ctx context.Context, ss *cache.Snapshot, changeF _ = getIncludes(ctx, ss, file, &includesMap) } - cyclePairs := cycleDetect(&includesMap) + cyclePairs := cycleDetect(includesMap) return cycleToDiagnosticItems(cyclePairs), nil } @@ -64,7 +64,7 @@ type CyclePair struct { // cycleDetect returns every include edge that closes a cycle: the pair // (file, include file->Y) is reported when Y transitively includes file. // Cycles of any length are caught, including self-includes. -func cycleDetect(includesMap *map[uri.URI][]Include) []CyclePair { +func cycleDetect(includesMap map[uri.URI][]Include) []CyclePair { // reaches reports whether from can reach target via include edges, // cycle-safe via the seen set. var reaches func(from, target uri.URI, seen map[uri.URI]bool) bool @@ -80,7 +80,7 @@ func cycleDetect(includesMap *map[uri.URI][]Include) []CyclePair { seen[from] = true - for _, inc := range (*includesMap)[from] { + for _, inc := range includesMap[from] { if reaches(inc.file, target, seen) { return true } @@ -91,7 +91,7 @@ func cycleDetect(includesMap *map[uri.URI][]Include) []CyclePair { cyclePairs := make([]CyclePair, 0) - for file, includes := range *includesMap { + for file, includes := range includesMap { for _, inc := range includes { if reaches(inc.file, file, make(map[uri.URI]bool)) { cyclePairs = append(cyclePairs, CyclePair{ diff --git a/lsp/source/cycle_detect_test.go b/lsp/source/cycle_detect_test.go index 022829e..da24809 100644 --- a/lsp/source/cycle_detect_test.go +++ b/lsp/source/cycle_detect_test.go @@ -24,7 +24,7 @@ func Test_cycleDetect(t *testing.T) { } type args struct { - includesMap *map[uri.URI][]Include + includesMap map[uri.URI][]Include } tests := []struct { @@ -35,7 +35,7 @@ func Test_cycleDetect(t *testing.T) { { name: "cycle", args: args{ - includesMap: &includesMap, + includesMap: includesMap, }, want: []CyclePair{ { @@ -184,7 +184,7 @@ func Test_cycleDetectN(t *testing.T) { for _, tt := range tests { t.Run(tt.name, func(t *testing.T) { - got := cycleDetect(&tt.graph) + got := cycleDetect(tt.graph) sort.SliceStable(got, func(i, j int) bool { if got[i].file == got[j].file { return got[i].include.file < got[j].include.file diff --git a/lsp/source/field_qualifier_action.go b/lsp/source/field_qualifier_action.go index 297d4b7..226b359 100644 --- a/lsp/source/field_qualifier_action.go +++ b/lsp/source/field_qualifier_action.go @@ -18,6 +18,13 @@ import ( // it does not already carry: an unqualified field gets both, a required // field gets "Make field optional", and vice versa. Union fields never // offer "Make field required": unions have no required members. +// pickedFieldAction pairs an action with the declaration offset of its +// field, so the actions can be ordered into document order. +type pickedFieldAction struct { + offset int + code protocol.CodeAction +} + func MakeFieldQualifierAction(ctx context.Context, ss *cache.Snapshot, fh cache.FileHandle, rng protocol.Range) ([]protocol.CodeAction, error) { pf, err := ss.Parse(ctx, fh.URI()) if err != nil { @@ -28,12 +35,7 @@ func MakeFieldQualifierAction(ctx context.Context, ss *cache.Snapshot, fh cache. return nil, nil } - // Each picked action is kept beside the declaration offset of its - // field, so the final list can be sorted into document order. - var picked []struct { - offset int - code protocol.CodeAction - } + var picked []pickedFieldAction pf.AST().WalkFieldLists(func(fields []*syntax.Field, kind syntax.FieldListKind) { for _, field := range fields { @@ -45,10 +47,7 @@ func MakeFieldQualifierAction(ctx context.Context, ss *cache.Snapshot, fh cache. unionField := kind == syntax.UnionFields for _, qualifier := range fieldQualifiers(field, unionField) { - picked = append(picked, struct { - offset int - code protocol.CodeAction - }{ + picked = append(picked, pickedFieldAction{ offset: pf.AST().TokenPosition(field.TokStart()).Offset, code: fieldQualifierAction(pf, fh, field, qualifier), }) diff --git a/lsp/source/index.go b/lsp/source/index.go index 55d1c7e..ebd11d0 100644 --- a/lsp/source/index.go +++ b/lsp/source/index.go @@ -3,10 +3,8 @@ package source import ( "context" "fmt" - "io/fs" "log/slog" - "path/filepath" - "sort" + "slices" "strings" "go.lsp.dev/protocol" @@ -156,11 +154,17 @@ type Hit struct { File uri.URI Range protocol.Range Text string // as written: "User", "shared.User", "shared.thrift.User" + + // Kind is the grammar slot the reference sits in, so callers can tell + // type hits from value hits (e.g. for highlight kinds). + Kind cache.RefKind } // References returns every occurrence of name in file and in every file // that transitively includes it, restricted to the given reference kinds. -// The definition site is not included (no self-referencing hit). +// The definition site is not included (no self-referencing hit). Hits are +// matched by bare name only; use ReferencesTo for resolution-matched +// results. func (x *Index) References(ctx context.Context, file uri.URI, name string, kinds ...cache.RefKind) ([]Hit, error) { files := x.searchFiles(file) @@ -185,51 +189,72 @@ func (x *Index) References(ctx context.Context, file uri.URI, name string, kinds return out, nil } -// QualifiedValues returns value-position references whose qualifier is -// enumName: "Song.FUWA_FUWA_TIME" or "songs.Song.FUWA_FUWA_TIME", each -// hit covering only the enum segment so a rename rewrites the qualifier -// while keeping the member name. -func (x *Index) QualifiedValues(ctx context.Context, file uri.URI, enumName string) ([]Hit, error) { - files := x.searchFiles(file) - +// ReferencesTo returns every reference to def: in def.File and every file +// that transitively includes it. Hits are name- and resolution-matched, so +// same-named definitions elsewhere are not reported. +// +// For an enum def, value references qualified with the enum name +// ("Color.RED", "shared.Color.RED") are matched too, provided the +// qualifier resolves to this very enum; the hit covers only the enum +// segment, so a rename rewrites the qualifier while keeping the member +// name. +func (x *Index) ReferencesTo(ctx context.Context, def *Resolved, kinds ...cache.RefKind) ([]Hit, error) { var out []Hit - seen := map[uri.URI]bool{} - - for _, f := range files { - if seen[f] { - continue - } - - seen[f] = true + for _, f := range x.searchFiles(def.File) { pf, err := x.ss.Parse(ctx, f) if err != nil || pf.AST() == nil { continue } for _, r := range pf.Index().References() { - if r.Kind != cache.RefConstValue { + if r.Kind == cache.RefConstValue && def.Kind == DefinitionEnum && referenceKind(cache.RefConstValue, kinds) { + // Value position qualified with the enum: only the enum + // segment is rewritten. + seg, off, ok := enumSegment(r.Name, def.Name.Text) + if !ok { + continue + } + + // The qualifier ("Color", "shared.Color") must resolve to + // this very enum, not to a same-named one elsewhere. + qualifier := r.Name[:off+len(seg)] + enum, err := x.ResolveType(ctx, pf, typeReference(qualifier)) + if err != nil { + return nil, err + } + + if !sameDefinition(enum, def) { + continue + } + + out = append(out, enumSegmentHit(pf, r.Node, off, seg)) + continue } - seg, off, ok := enumSegment(r.Name, enumName) - if !ok { + if !referenceKind(r.Kind, kinds) { continue } - start, _ := pf.AST().Range(r.Node) + if bareName(r.Name) != bareName(def.Name.Text) { + continue + } - segStart := toLSPPosition(pf, syntax.Position{ - Line: start.Line, Col: start.Col, Offset: start.Offset + off, - }) - segEnd := toLSPPosition(pf, syntax.Position{ - Line: start.Line, Col: start.Col, Offset: start.Offset + off + len(seg), - }) + resolved, err := x.resolveReference(ctx, pf, r) + if err != nil { + return nil, err + } + + if !sameDefinition(resolved, def) { + continue + } out = append(out, Hit{ - File: f, - Range: protocol.Range{Start: segStart, End: segEnd}, - Text: seg, + File: pf.URI(), + Range: nodeRange(pf, r.Node), + Text: r.Name, + Kind: r.Kind, }) } } @@ -269,21 +294,19 @@ func (x *Index) FindInWorkspace(ctx context.Context, name string) (*Resolved, er } } - // Fallback to the old directory walk when KnownFiles is empty. - root := view.Folder().Path() + // Fallback to the directory walk when KnownFiles is empty: the walk + // goes through the view's file source (disk, or the in-memory tree in + // tests). + root := view.Folder() if root == "" { return nil, nil } - var files []string - - err := filepath.WalkDir(root, func(p string, d fs.DirEntry, err error) error { - if err != nil { - return nil - } + var files []uri.URI - if !d.IsDir() && strings.HasSuffix(d.Name(), ".thrift") { - files = append(files, p) + err := view.WalkFiles(ctx, root, func(u uri.URI) error { + if strings.HasSuffix(u.Path(), ".thrift") { + files = append(files, u) } return nil @@ -292,13 +315,14 @@ func (x *Index) FindInWorkspace(ctx context.Context, name string) (*Resolved, er return nil, nil } - sort.Strings(files) - for _, p := range files { - if include != "" && includeNameOf(uri.File(p)) != include { + slices.Sort(files) + + for _, f := range files { + if include != "" && includeNameOf(f) != include { continue } - pf, err := x.ss.Parse(ctx, uri.File(p)) + pf, err := x.ss.Parse(ctx, f) if err != nil || pf.AST() == nil { continue } @@ -334,12 +358,87 @@ func refKindsFor(k DefinitionKind) []cache.RefKind { // --- helpers --- -// searchFiles returns the file itself followed by its direct includers, -// deduplicated. This matches the existing reference-search file ordering. +// resolveReference resolves a reference to its definition, dispatching on +// the grammar slot the reference sits in. Unresolvable references (parse +// errors, unknown names) yield nil, not an error. +func (x *Index) resolveReference(ctx context.Context, pf *cache.ParsedFile, r cache.Reference) (*Resolved, error) { + switch r.Kind { + case cache.RefFieldType, cache.RefSignatureType: + ident, ok := r.Node.(*syntax.Identifier) + if !ok { + return nil, nil + } + + return x.ResolveType(ctx, pf, typeReference(ident.Text)) + case cache.RefConstValue: + v, ok := r.Node.(*syntax.ConstValue) + if !ok { + return nil, nil + } + + return x.ResolveValue(ctx, pf, v) + case cache.RefServiceExtends: + ident, ok := r.Node.(*syntax.Identifier) + if !ok { + return nil, nil + } + + return x.ResolveService(ctx, pf, ident) + } + + return nil, nil +} + +// typeReference builds a FieldType for a reference text, for resolution +// by name: the index resolves the text, the original node only carries it. +func typeReference(name string) *syntax.FieldType { + return &syntax.FieldType{Kind: syntax.TypeIdent, Ident: &syntax.Identifier{Text: name}} +} + +// sameDefinition reports whether resolved is the same definition as def: +// same file and same name. +func sameDefinition(resolved, def *Resolved) bool { + if resolved == nil || def == nil { + return false + } + + return resolved.File == def.File && resolved.Name.Text == def.Name.Text +} + +// referenceKind reports whether k is in kinds; an empty kinds matches all. +func referenceKind(k cache.RefKind, kinds []cache.RefKind) bool { + return len(kinds) == 0 || slices.Contains(kinds, k) +} + +// enumSegmentHit builds a hit covering one segment (off, seg) of the +// reference node's text, so a rename rewrites just that segment. +func enumSegmentHit(pf *cache.ParsedFile, node syntax.Node, off int, seg string) Hit { + start, _ := pf.AST().Range(node) + + segStart := toLSPPosition(pf, syntax.Position{ + Line: start.Line, Col: start.Col, Offset: start.Offset + off, + }) + segEnd := toLSPPosition(pf, syntax.Position{ + Line: start.Line, Col: start.Col, Offset: start.Offset + off + len(seg), + }) + + return Hit{ + File: pf.URI(), + Range: protocol.Range{Start: segStart, End: segEnd}, + Text: seg, + Kind: cache.RefConstValue, + } +} + +// searchFiles returns the file itself followed by its transitive +// dependents (every file that includes it, directly or through other +// includes), deduplicated. Resolution is transitive, so reference search +// must be too: a definition reached through a chain of includes is +// referenced from every file in the chain. func (x *Index) searchFiles(file uri.URI) []uri.URI { files := []uri.URI{file} - for _, dep := range x.ReferencingFiles(file) { + for _, dep := range x.ss.Dependents(file) { if dep != file { files = append(files, dep) } @@ -351,17 +450,10 @@ func (x *Index) searchFiles(file uri.URI) []uri.URI { // matches returns hits for references in pf whose kind is in kinds and // whose bare name equals bareName(name). func (x *Index) matches(pf *cache.ParsedFile, name string, kinds []cache.RefKind) []Hit { - kindSet := make(map[cache.RefKind]bool, len(kinds)) - for _, k := range kinds { - kindSet[k] = true - } - - haveKindSet := len(kinds) > 0 - var out []Hit for _, r := range pf.Index().References() { - if haveKindSet && !kindSet[r.Kind] { + if !referenceKind(r.Kind, kinds) { continue } @@ -373,6 +465,7 @@ func (x *Index) matches(pf *cache.ParsedFile, name string, kinds []cache.RefKind File: pf.URI(), Range: nodeRange(pf, r.Node), Text: r.Name, + Kind: r.Kind, }) } diff --git a/lsp/source/index_test.go b/lsp/source/index_test.go index 68c9036..4bf7c6f 100644 --- a/lsp/source/index_test.go +++ b/lsp/source/index_test.go @@ -111,12 +111,16 @@ func TestIndex_References_ConstValue(t *testing.T) { require.Len(t, hits, 1) } -func TestIndex_QualifiedValues(t *testing.T) { +func TestIndex_ReferencesToEnumValues(t *testing.T) { ctx := t.Context() ss := snap(t, "/t.thrift", "enum Color { RED = 0, BLUE = 1 }\nstruct Foo { 1: i32 id = Color.RED, }\nconst i32 C = Color.BLUE") - _ = parseOne(t, ss, fu("/t.thrift")) + pf := parseOne(t, ss, fu("/t.thrift")) + + def, err := NewIndex(ss).ResolveType(ctx, pf, ft("Color")) + require.NoError(t, err) + require.NotNil(t, def) - hits, err := NewIndex(ss).QualifiedValues(ctx, fu("/t.thrift"), "Color") + hits, err := NewIndex(ss).ReferencesTo(ctx, def, cache.RefFieldType, cache.RefSignatureType, cache.RefConstValue) require.NoError(t, err) require.Len(t, hits, 2) for _, h := range hits { diff --git a/lsp/source/provider.go b/lsp/source/provider.go index 52b6614..0c68580 100644 --- a/lsp/source/provider.go +++ b/lsp/source/provider.go @@ -126,14 +126,7 @@ func (annotationKeyProvider) Candidates(ctx context.Context, ss *cache.Snapshot, } } - var res []Candidate - for key := range keys { - res = append(res, Candidate{showText: key, insertText: key, format: protocol.InsertTextFormatPlainText}) - } - - sortCandidates(res) - - return res + return setCandidates(keys) } // annotationKeys collects the names of every annotation in the document: @@ -206,7 +199,13 @@ func (serviceExtendsProvider) Candidates(ctx context.Context, ss *cache.Snapshot } } - var res []Candidate + return setCandidates(names) +} + +// setCandidates converts a name set into alphabetically sorted candidates. +func setCandidates(names map[string]struct{}) []Candidate { + res := make([]Candidate, 0, len(names)) + for name := range names { res = append(res, Candidate{showText: name, insertText: name, format: protocol.InsertTextFormatPlainText}) } diff --git a/lsp/source/reference.go b/lsp/source/reference.go index 32ca54e..dc7bbd9 100644 --- a/lsp/source/reference.go +++ b/lsp/source/reference.go @@ -139,14 +139,13 @@ func searchTypeNameRefs(ctx context.Context, ix *Index, ss *cache.Snapshot, pf * hits := []indexHit{{loc: loc, text: def.Name.Text, kind: cache.RefFieldType}} - kinds := refKindsFor(def.Kind) - refs, err := ix.References(ctx, def.File, typeName, kinds...) + refs, err := ix.ReferencesTo(ctx, def, refKindsFor(def.Kind)...) if err != nil { return nil, err } for _, h := range refs { - hits = append(hits, indexHit{loc: protocol.Location{URI: h.File, Range: h.Range}, text: h.Text, kind: cache.RefFieldType}) + hits = append(hits, indexHit{loc: protocol.Location{URI: h.File, Range: h.Range}, text: h.Text, kind: h.Kind}) } return hits, nil @@ -172,13 +171,13 @@ func searchConstValueRefs(ctx context.Context, ix *Index, ss *cache.Snapshot, pf hits := []indexHit{{loc: loc, text: def.Name.Text, kind: cache.RefConstValue}} - refs, err := ix.References(ctx, def.File, value.Text, cache.RefConstValue) + refs, err := ix.ReferencesTo(ctx, def, cache.RefConstValue) if err != nil { return nil, err } for _, h := range refs { - hits = append(hits, indexHit{loc: protocol.Location{URI: h.File, Range: h.Range}, text: h.Text, kind: cache.RefConstValue}) + hits = append(hits, indexHit{loc: protocol.Location{URI: h.File, Range: h.Range}, text: h.Text, kind: h.Kind}) } return hits, nil @@ -196,14 +195,14 @@ func searchServiceRefs(ctx context.Context, ix *Index, ss *cache.Snapshot, file return nil, nil } - refs, err := ix.References(ctx, def.File, svcName, cache.RefServiceExtends) + refs, err := ix.ReferencesTo(ctx, def, cache.RefServiceExtends) if err != nil { return nil, err } hits := make([]indexHit, 0, len(refs)) for _, h := range refs { - hits = append(hits, indexHit{loc: protocol.Location{URI: h.File, Range: h.Range}, text: h.Text, kind: cache.RefServiceExtends}) + hits = append(hits, indexHit{loc: protocol.Location{URI: h.File, Range: h.Range}, text: h.Text, kind: h.Kind}) } return hits, nil @@ -218,20 +217,17 @@ func searchDefRefs(ctx context.Context, ix *Index, ss *cache.Snapshot, file uri. } parent := target.parent + + var def *Resolved + var kinds []cache.RefKind + switch parent.(type) { case *syntax.Const: - typeName := fmt.Sprintf("%s.%s", includeNameOf(file), id.Text) - - return valueRefHits(ctx, ix, file, typeName) + def = defFromNode(pf, parent) + kinds = []cache.RefKind{cache.RefConstValue} case *syntax.EnumValue: - enum, ok := grandparent(target.path).(*syntax.Enum) - if !ok { - return nil, nil - } - - typeName := fmt.Sprintf("%s.%s.%s", includeNameOf(file), enum.Name.Text, id.Text) - - return valueRefHits(ctx, ix, file, typeName) + def = defFromNode(pf, id) + kinds = []cache.RefKind{cache.RefConstValue} case *syntax.Service: svcName := id.Text if strings.Contains(svcName, ".") { @@ -246,68 +242,39 @@ func searchDefRefs(ctx context.Context, ix *Index, ss *cache.Snapshot, file uri. } return searchServiceRefs(ctx, ix, ss, file, svcName) - } - - kind, ok := definitionKindOf(parent) - if !ok { - return nil, nil - } - - if _, ok := validReferenceDefinitionType[kind]; !ok { - return nil, nil - } - - typeName := fmt.Sprintf("%s.%s", includeNameOf(file), id.Text) - - typeRefs, err := ix.References(ctx, file, typeName, refKindsFor(kind)...) - if err != nil { - return nil, err - } - - hits := make([]indexHit, 0, len(typeRefs)) - for _, h := range typeRefs { - hits = append(hits, indexHit{loc: protocol.Location{URI: h.File, Range: h.Range}, text: h.Text, kind: cache.RefFieldType}) - } + default: + kind, ok := definitionKindOf(parent) + if !ok { + return nil, nil + } - // Enum renames also touch value positions qualified with the enum name. - if kind == DefinitionEnum { - valRefs, err := ix.QualifiedValues(ctx, file, id.Text) - if err != nil { - return nil, err + if _, ok := validReferenceDefinitionType[kind]; !ok { + return nil, nil } - for _, h := range valRefs { - hits = append(hits, indexHit{loc: protocol.Location{URI: h.File, Range: h.Range}, text: h.Text, kind: cache.RefConstValue}) + def = defFromNode(pf, parent) + kinds = refKindsFor(kind) + + // Renaming an enum definition also touches value positions + // qualified with the enum name ("Color.RED"). + if kind == DefinitionEnum { + kinds = append(kinds, cache.RefConstValue) } } - return hits, nil -} - -// valueRefHits wraps value-kind reference lookups for consts and enum -// values. -func valueRefHits(ctx context.Context, ix *Index, file uri.URI, name string) ([]indexHit, error) { - refs, err := ix.References(ctx, file, name, cache.RefConstValue) + refs, err := ix.ReferencesTo(ctx, def, kinds...) if err != nil { return nil, err } hits := make([]indexHit, 0, len(refs)) for _, h := range refs { - hits = append(hits, indexHit{loc: protocol.Location{URI: h.File, Range: h.Range}, text: h.Text, kind: cache.RefConstValue}) + hits = append(hits, indexHit{loc: protocol.Location{URI: h.File, Range: h.Range}, text: h.Text, kind: h.Kind}) } return hits, nil } -func grandparent(path []syntax.Node) syntax.Node { - if len(path) < 3 { - return nil - } - - return path[len(path)-3] -} - // definitionKindOf maps a definition node to its kind. func definitionKindOf(n syntax.Node) (DefinitionKind, bool) { switch v := n.(type) { diff --git a/lsp/source/rename_correctness_test.go b/lsp/source/rename_correctness_test.go new file mode 100644 index 0000000..a2979f3 --- /dev/null +++ b/lsp/source/rename_correctness_test.go @@ -0,0 +1,150 @@ +package source + +import ( + "testing" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + "go.lsp.dev/protocol" + + "github.com/karitham/thrift-ls/lsp/cache" +) + +// TestRenameTransitiveInclude pins that a definition reached through a +// chain of includes (app → mid → base) is renamed everywhere it is +// referenced — including in files that include it only transitively. +func TestRenameTransitiveInclude(t *testing.T) { + ss := cache.BuildSnapshotForTest([]*cache.FileChange{ + { + URI: "file:///tmp/base.thrift", + Version: 0, + Content: []byte("struct User { 1: i32 id }\n"), + From: cache.FileChangeTypeDidOpen, + }, + { + URI: "file:///tmp/mid.thrift", + Version: 0, + Content: []byte("include \"base.thrift\"\n"), + From: cache.FileChangeTypeDidOpen, + }, + { + URI: "file:///tmp/app.thrift", + Version: 0, + Content: []byte("include \"mid.thrift\"\nstruct S { 1: User u }\n"), + From: cache.FileChangeTypeDidOpen, + }, + }) + + edit, err := Rename(t.Context(), ss, "file:///tmp/base.thrift", protocol.Position{Line: 0, Character: 7}, "Account") + require.NoError(t, err) + + assert.Equal(t, []protocol.TextEdit{{ + Range: protocol.Range{Start: protocol.Position{Line: 0, Character: 7}, End: protocol.Position{Line: 0, Character: 11}}, + NewText: "Account", + }}, edit.Changes["file:///tmp/base.thrift"], "the definition itself") + + assert.Equal(t, []protocol.TextEdit{{ + Range: protocol.Range{Start: protocol.Position{Line: 1, Character: 14}, End: protocol.Position{Line: 1, Character: 18}}, + NewText: "Account", + }}, edit.Changes["file:///tmp/app.thrift"], "the transitive reference in app.thrift must be renamed") +} + +// TestRenameResolutionMatched pins that renaming a definition leaves +// references to a same-named definition from another file untouched: +// matches are resolved to their actual definition, not matched by name. +func TestRenameResolutionMatched(t *testing.T) { + ss := cache.BuildSnapshotForTest([]*cache.FileChange{ + { + URI: "file:///tmp/base.thrift", + Version: 0, + Content: []byte("struct User { 1: i32 id }\n"), + From: cache.FileChangeTypeDidOpen, + }, + { + URI: "file:///tmp/app.thrift", + Version: 0, + Content: []byte("include \"base.thrift\"\nstruct User { 1: string name }\nstruct S { 1: base.User x, 2: User y }\n"), + From: cache.FileChangeTypeDidOpen, + }, + }) + + // Rename app.thrift's own User (line 1, the definition). + edit, err := Rename(t.Context(), ss, "file:///tmp/app.thrift", protocol.Position{Line: 1, Character: 7}, "Member") + require.NoError(t, err) + + // Only the unqualified reference and the definition change; the + // qualified reference to base.thrift's User does not. + assert.Equal(t, []protocol.TextEdit{ + {Range: protocol.Range{Start: protocol.Position{Line: 2, Character: 30}, End: protocol.Position{Line: 2, Character: 34}}, NewText: "Member"}, + {Range: protocol.Range{Start: protocol.Position{Line: 1, Character: 7}, End: protocol.Position{Line: 1, Character: 11}}, NewText: "Member"}, + }, edit.Changes["file:///tmp/app.thrift"]) + + assert.Empty(t, edit.Changes["file:///tmp/base.thrift"], "the same-named definition in base.thrift is untouched") +} + +// TestRenameEnumValueResolutionMatched pins that renaming an enum value +// only touches references that resolve to it: "colors.Palette.RED" must +// survive a rename of the local enum's RED. +func TestRenameEnumValueResolutionMatched(t *testing.T) { + ss := cache.BuildSnapshotForTest([]*cache.FileChange{ + { + URI: "file:///tmp/colors.thrift", + Version: 0, + Content: []byte("enum Palette { RED }\n"), + From: cache.FileChangeTypeDidOpen, + }, + { + URI: "file:///tmp/main.thrift", + Version: 0, + Content: []byte("include \"colors.thrift\"\nenum Local { RED }\nstruct S {\n 1: i32 a = Local.RED,\n 2: i32 b = colors.Palette.RED,\n}\n"), + From: cache.FileChangeTypeDidOpen, + }, + }) + + // Cursor on the local RED definition (line 1, char 13). + edit, err := Rename(t.Context(), ss, "file:///tmp/main.thrift", protocol.Position{Line: 1, Character: 13}, "CRIMSON") + require.NoError(t, err) + + var got []string + for _, te := range edit.Changes["file:///tmp/main.thrift"] { + got = append(got, te.NewText) + } + + // The Local.RED qualifier keeps the enum name, and the definition + // changes; colors.Palette.RED does not. + assert.Equal(t, []string{"Local.CRIMSON", "CRIMSON"}, got) +} + +// TestRenameEnumResolutionMatched pins that renaming an enum only touches +// value references qualified with that enum: same-named enums in other +// files are left alone. +func TestRenameEnumResolutionMatched(t *testing.T) { + ss := cache.BuildSnapshotForTest([]*cache.FileChange{ + { + URI: "file:///tmp/colors.thrift", + Version: 0, + Content: []byte("enum Color { RED }\n"), + From: cache.FileChangeTypeDidOpen, + }, + { + URI: "file:///tmp/main.thrift", + Version: 0, + Content: []byte("include \"colors.thrift\"\nenum Color { BLUE }\nstruct S {\n 1: i32 a = Color.BLUE,\n 2: i32 b = colors.Color.RED,\n}\n"), + From: cache.FileChangeTypeDidOpen, + }, + }) + + // Cursor on the local Color definition (line 1, char 5). + edit, err := Rename(t.Context(), ss, "file:///tmp/main.thrift", protocol.Position{Line: 1, Character: 5}, "Hue") + require.NoError(t, err) + + var got []string + for _, te := range edit.Changes["file:///tmp/main.thrift"] { + got = append(got, te.NewText) + } + + // The Color.BLUE qualifier and the definition change; colors.Color.RED + // is untouched. + assert.Equal(t, []string{"Hue", "Hue"}, got) + assert.Empty(t, edit.Changes["file:///tmp/colors.thrift"]) +} diff --git a/lsp/source/semantic.go b/lsp/source/semantic.go index 9a1f235..aef57fe 100644 --- a/lsp/source/semantic.go +++ b/lsp/source/semantic.go @@ -102,8 +102,8 @@ func Tokens(ctx context.Context, ss *cache.Snapshot, file uri.URI) ([]uint32, er return data, nil } -// classify maps a token to its semantic type. Definition names win over -// type keywords so a field named "string" stays a property; type +// classifyToken maps a token to its semantic type. Definition names win +// over type keywords so a field named "string" stays a property; type // references win over keywords so "string" in a type position is a type. func classifyToken(i int, tok syntax.Token, names map[int]int, types map[int]bool) (int, bool) { if syntax.IsComment(tok.Kind) { diff --git a/lsp/source/semantic_analysis.go b/lsp/source/semantic_analysis.go index 853c16b..0a87b42 100644 --- a/lsp/source/semantic_analysis.go +++ b/lsp/source/semantic_analysis.go @@ -58,9 +58,8 @@ func (s *SemanticAnalysis) diagnostic(ctx context.Context, ss *cache.Snapshot, c func (s *SemanticAnalysis) checkDefinitionExist(ctx context.Context, ss *cache.Snapshot, pf *cache.ParsedFile) []protocol.Diagnostic { ret := make([]protocol.Diagnostic, 0) - processStructLike := func(fields []*syntax.Field) { - for i := range fields { - field := fields[i] + processFields := func(fields []*syntax.Field) { + for _, field := range fields { items := s.checkTypeExist(ctx, ss, pf, field.Type) ret = append(ret, items...) @@ -77,7 +76,7 @@ func (s *SemanticAnalysis) checkDefinitionExist(ctx context.Context, ss *cache.S } pf.AST().WalkFieldLists(func(fields []*syntax.Field, _ syntax.FieldListKind) { - processStructLike(fields) + processFields(fields) }) for _, cst := range pf.AST().Consts() { diff --git a/lsp/source/semantic_based_completion.go b/lsp/source/semantic_based_completion.go index deb5917..9833ab9 100644 --- a/lsp/source/semantic_based_completion.go +++ b/lsp/source/semantic_based_completion.go @@ -22,7 +22,5 @@ func BuildCompletionItem(candidate Candidate) *CompletionItem { InsertText: candidate.insertText, InsertTextFormat: candidate.format, Kind: protocol.CompletionItemKindText, - Deprecated: false, - Documentation: "", } } diff --git a/lsp/source/semantic_completion.go b/lsp/source/semantic_completion.go index 3750b00..48263da 100644 --- a/lsp/source/semantic_completion.go +++ b/lsp/source/semantic_completion.go @@ -2,7 +2,6 @@ package source import ( "context" - "sort" "strings" "go.lsp.dev/protocol" @@ -43,7 +42,7 @@ func typeCandidates(ctx context.Context, ss *cache.Snapshot, file uri.URI, c Con }) } - sort.Slice(res, func(i, j int) bool { return res[i].showText < res[j].showText }) + sortCandidates(res) return res } @@ -51,14 +50,7 @@ func typeCandidates(ctx context.Context, ss *cache.Snapshot, file uri.URI, c Con names := make(map[string]struct{}) collectTypeNames(c.Doc, names) - var res []Candidate - for name := range names { - res = append(res, Candidate{ - showText: name, - insertText: name, - format: protocol.InsertTextFormatPlainText, - }) - } + res := setCandidates(names) // Types from included files are suggested with their include // qualifier: a bare reference to an imported type does not resolve. @@ -90,7 +82,7 @@ func typeCandidates(ctx context.Context, ss *cache.Snapshot, file uri.URI, c Con }) } - sort.Slice(res, func(i, j int) bool { return res[i].showText < res[j].showText }) + sortCandidates(res) return res } @@ -168,18 +160,7 @@ func valueCandidates(ctx context.Context, ss *cache.Snapshot, file uri.URI, doc } } - var res []Candidate - for name := range names { - res = append(res, Candidate{ - showText: name, - insertText: name, - format: protocol.InsertTextFormatPlainText, - }) - } - - sort.Slice(res, func(i, j int) bool { return res[i].showText < res[j].showText }) - - return res + return setCandidates(names) } // includedFiles returns the files transitively included by file, per the diff --git a/lsp/source/token_completion.go b/lsp/source/token_completion.go index 15744cf..0df0b7b 100644 --- a/lsp/source/token_completion.go +++ b/lsp/source/token_completion.go @@ -2,7 +2,7 @@ package source import ( "context" - "fmt" + "errors" "sort" "strings" @@ -65,7 +65,7 @@ func (c *TokenCompletion) Completion(ctx context.Context, ss *cache.Snapshot, cm } if parsedFile.AST() == nil { - return nil, protocol.Range{}, false, fmt.Errorf("parser ast failed") + return nil, protocol.Range{}, false, errors.New("parser ast failed") } pos, err := parsedFile.Mapper().LSPPosToParserPosition(cmp.Pos) @@ -170,7 +170,7 @@ func (c *TokenCompletion) Completion(ctx context.Context, ss *cache.Snapshot, cm truncated = true } - cursor := protocol.Position{Line: cmp.Pos.Line, Character: cmp.Pos.Character} + cursor := cmp.Pos rng := protocol.Range{End: cursor} if start, err := parsedFile.Mapper().OffsetToLSPPosition(editStart); err == nil { diff --git a/main.go b/main.go index 7a7c0e4..ed09ccc 100644 --- a/main.go +++ b/main.go @@ -178,12 +178,12 @@ func lspAction(ctx context.Context, cmd *cli.Command) error { cfg := loadConfig(cmd.String("config"), ".") patch := options.Effective(cfg) - cli, err := lspPatch(cmd) + cliPatch, err := lspPatch(cmd) if err != nil { return err } - patch = cli.Apply(patch) + patch = cliPatch.Apply(patch) logLevelValue := 3 if patch.LogLevel != nil { @@ -201,7 +201,7 @@ func lspAction(ctx context.Context, cmd *cli.Command) error { lspOpts := &lsp.Options{ Config: patch, ConfigPath: cmd.String("config"), - CLI: cli, + CLI: cliPatch, } ss := lsp.NewStreamServer(lspOpts) @@ -220,12 +220,12 @@ func lspAction(ctx context.Context, cmd *cli.Command) error { func formatAction(ctx context.Context, cmd *cli.Command) error { file := cmd.Args().First() - cli, err := formatPatch(cmd) + cliPatch, err := formatPatch(cmd) if err != nil { return err } - return formatFile(file, cmd.Writer, cmd.Bool("w"), cmd.Bool("d"), cmd.String("config"), cli) + return formatFile(file, cmd.Writer, cmd.Bool("w"), cmd.Bool("d"), cmd.String("config"), cliPatch) } // dumpAction prints the parse tree, and optionally the formatted document @@ -310,10 +310,13 @@ func checkAction(ctx context.Context, cmd *cli.Command) error { patch = cliPatch.Apply(patch) - root := path - if info, err := os.Stat(path); err != nil { + info, err := os.Stat(path) + if err != nil { return err - } else if !info.IsDir() { + } + + root := path + if !info.IsDir() { root = filepath.Dir(path) } diff --git a/options/options.go b/options/options.go index abd5ce7..56e229b 100644 --- a/options/options.go +++ b/options/options.go @@ -15,7 +15,6 @@ import ( "fmt" "os" "path/filepath" - "slices" "strings" "github.com/karitham/thrift-ls/formatter" @@ -66,29 +65,8 @@ func (p Patch) Apply(base Patch) Patch { out.Align = p.Align } - if p.Separators != nil { - if out.Separators == nil { - out.Separators = &Separators{} - } - - 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{} - } - - for _, c := range formatter.AllConstructs { - if v := p.Break.Get(c); v != nil { - out.Break.Set(c, v) - } - } - } + out.Separators = overlayPerConstruct(out.Separators, p.Separators) + out.Break = overlayPerConstruct(out.Break, p.Break) if p.IncludePaths != nil { out.IncludePaths = p.IncludePaths @@ -101,6 +79,26 @@ func (p Patch) Apply(base Patch) Patch { return out } +// overlayPerConstruct copies the set fields of src onto dst, creating dst +// when it is nil. +func overlayPerConstruct[T *E, E any](dst, src *formatter.PerConstruct[T]) *formatter.PerConstruct[T] { + if src == nil { + return dst + } + + if dst == nil { + dst = &formatter.PerConstruct[T]{} + } + + for _, c := range formatter.AllConstructs { + if v := src.Get(c); v != nil { + dst.Set(c, v) + } + } + + return dst +} + // Default returns the default options as a fully-set patch. func Default() Patch { printWidth := 80 @@ -135,22 +133,24 @@ func (p Patch) Validate() error { return errors.New("tabWidth must be positive") } - 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.Align != nil { + if _, ok := alignMode(*p.Align); !ok { + return fmt.Errorf("align must be one of \"field\", \"assign\", \"disable\", got %q", *p.Align) + } } if p.Separators != nil { 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 v := p.Separators.Get(c); v != nil { + if _, ok := separatorMode(*v); !ok { + 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") - } + if p.Indent != nil && (p.Indent.Width <= 0 || !isWhitespaceOnly(p.Indent.Value)) { + return errors.New("indent must be a string of spaces or tabs") } return nil @@ -177,20 +177,17 @@ func (p Patch) Formatter() (formatter.Options, error) { } if p.Align != nil { - switch *p.Align { - case "field": - o.Align = formatter.AlignField - case "assign": - o.Align = formatter.AlignAssign - case "disable": - o.Align = formatter.AlignDisable + if mode, ok := alignMode(*p.Align); ok { + o.Align = mode } } if p.Separators != nil { for _, c := range formatter.AllConstructs { if v := p.Separators.Get(c); v != nil { - o.Separator.Set(c, separatorMode(*v)) + if mode, ok := separatorMode(*v); ok { + o.Separator.Set(c, mode) + } } } } @@ -206,31 +203,46 @@ func (p Patch) Formatter() (formatter.Options, error) { return o, nil } +// alignMode maps a config value to a formatter align mode. The second +// result reports whether the value is a known align mode. +func alignMode(s string) (formatter.AlignMode, bool) { + switch s { + case "field": + return formatter.AlignField, true + case "assign": + return formatter.AlignAssign, true + case "disable": + return formatter.AlignDisable, true + default: + return 0, false + } +} + // separatorMode maps a config value to a formatter separator mode. The -// value is validated before this is called. -func separatorMode(s string) formatter.SeparatorMode { +// second result reports whether the value is a known separator mode. +func separatorMode(s string) (formatter.SeparatorMode, bool) { switch s { case "comma": - return formatter.SeparatorComma + return formatter.SeparatorComma, true case "semicolon": - return formatter.SeparatorSemicolon + return formatter.SeparatorSemicolon, true case "none": - return formatter.SeparatorNone - default: // "preserve" - return formatter.SeparatorPreserve + return formatter.SeparatorNone, true + case "preserve": + return formatter.SeparatorPreserve, true + default: + return 0, false } } // 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, or a number of spaces. +// its display width. It is set from a literal string of spaces or tabs. type Indent struct { Value string // the indentation string, spaces or tabs Width int // display width of one level } -// UnmarshalJSON accepts a number (spaces) or a literal string of spaces -// or tabs. +// UnmarshalJSON accepts a literal string of spaces or tabs. func (i *Indent) UnmarshalJSON(data []byte) error { var s string if err := json.Unmarshal(data, &s); err != nil { @@ -247,11 +259,10 @@ func (i *Indent) UnmarshalJSON(data []byte) error { return nil } -// ParseIndentValue resolves a friendly indent spec: +// ParseIndentValue resolves a literal indent string: // // " " literal spaces, used as written // "\t" literal tabs, used as written -// "8" a number of spaces // // An empty spec yields the default of four spaces. func ParseIndentValue(s string) (Indent, error) { @@ -350,9 +361,12 @@ func FindConfig(dir string) (string, error) { for d := dir; ; d = filepath.Dir(d) { path := filepath.Join(d, ConfigFileName) - if _, err := os.Stat(path); err == nil { + _, err := os.Stat(path) + if err == nil { return path, nil - } else if !os.IsNotExist(err) { + } + + if !os.IsNotExist(err) { return "", err } diff --git a/resolver/resolver.go b/resolver/resolver.go index 77bab6d..9943ec4 100644 --- a/resolver/resolver.go +++ b/resolver/resolver.go @@ -65,19 +65,15 @@ func NewWithFS(includePaths []string, fsys fs.FS) *Resolver { // Resolve resolves an include path relative to the current file. // It first tries relative to currentFile's directory, then tries each -// configured include path in order. Returns the resolved absolute file path, -// or the relative path as a fallback if not found. +// configured include path in order. Returns the resolved path, or the +// candidate relative to currentFile's directory as a fallback if not found. func (r *Resolver) Resolve(currentFile, includePath string) string { - // First try relative to current file's directory basePath := filepath.Dir(currentFile) resolvedPath := filepath.Join(basePath, includePath) - - // Check if file exists if r.exists(resolvedPath) { return resolvedPath } - // Try each configured include path for _, ip := range r.includePaths { candidatePath := filepath.Join(ip, includePath) if r.exists(candidatePath) { @@ -85,7 +81,6 @@ func (r *Resolver) Resolve(currentFile, includePath string) string { } } - // Return relative path as fallback return resolvedPath } diff --git a/syntax/lexer.go b/syntax/lexer.go index 9125d7b..2c41935 100644 --- a/syntax/lexer.go +++ b/syntax/lexer.go @@ -260,39 +260,29 @@ func (l *lexer) run() ([]Token, []Error) { // empty lines immediately before the next real token. func (l *lexer) scanTrivia() (blankLines int, comments []Token) { for l.off < len(l.src) { + var t Token switch c := l.src[l.off]; { case isWhitespace(c): blankLines += l.scanWhitespace() - case c == '/' && l.peekByte(1) == '/': - t := l.scanLineComment() - t.BlankLinesBefore = blankLines - blankLines = 0 - comments = append(comments, t) + continue + case c == '/' && l.peekByte(1) == '/', c == '#': + t = l.scanLineComment() case c == '/' && l.peekByte(1) == '*': - t := l.scanBlockComment() - t.BlankLinesBefore = blankLines - blankLines = 0 - - comments = append(comments, t) - case c == '#': - t := l.scanLineComment() - t.BlankLinesBefore = blankLines - blankLines = 0 - - comments = append(comments, t) + t = l.scanBlockComment() case c == '@': // Java-style annotations (@name{...}) are preserved as trivia, // like comments, so they round-trip without being part of the // grammar. - t := l.scanLineAnnotation() - t.BlankLinesBefore = blankLines - blankLines = 0 - - comments = append(comments, t) + t = l.scanLineAnnotation() default: return blankLines, comments } + + t.BlankLinesBefore = blankLines + blankLines = 0 + + comments = append(comments, t) } return blankLines, comments @@ -320,14 +310,17 @@ func (l *lexer) scanWhitespace() int { case ' ', '\t': l.advanceByte() default: - if newlines > 0 { - return newlines - 1 - } - - return 0 + return blankLinesBefore(newlines) } } + return blankLinesBefore(newlines) +} + +// blankLinesBefore converts a newline run into its blank-line count: n +// consecutive newlines yield n-1 blank lines, the last newline only ending +// the current line. +func blankLinesBefore(newlines int) int { if newlines > 0 { return newlines - 1 } @@ -369,7 +362,7 @@ func (l *lexer) scanBlockComment() Token { if l.off >= len(l.src) { l.errorfAt(start, "unterminated comment") - break + return l.finishTrivia(TokenBlockComment, start) } if l.src[l.off] == '*' && l.peekByte(1) == '/' { @@ -392,8 +385,6 @@ func (l *lexer) scanBlockComment() Token { l.advanceRune() } - - return l.finishTrivia(TokenBlockComment, start) } func (l *lexer) pos() srcPos { @@ -658,7 +649,7 @@ func (l *lexer) scanString() Token { if l.off >= len(l.src) { l.errorfAt(start, "unterminated string literal") - break + return Token{Kind: TokenStringLiteral, Text: l.src[start.offset:l.off], Offset: start.offset, Line: start.line, Col: start.col} } c := l.src[l.off] @@ -695,8 +686,6 @@ 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} } func (l *lexer) peekByte(ahead int) byte { diff --git a/syntax/parser.go b/syntax/parser.go index 228c08c..dcb197f 100644 --- a/syntax/parser.go +++ b/syntax/parser.go @@ -94,6 +94,17 @@ func (p *parser) synchronizeTo(kinds ...TokenKind) { } } +// recoverEntry skips past a malformed entry in a list-like construct: +// synchronize to a separator or the terminator, then consume one stray +// separator so the enclosing loop can continue with the next entry. +func (p *parser) recoverEntry(term TokenKind) { + p.synchronizeTo(TokenComma, TokenSemicolon, term) + + if !p.at(term) && !p.at(TokenEOF) { + p.advance() + } +} + func (p *parser) errorfCur(format string, args ...any) { t := p.cur() p.errs = append(p.errs, Error{ @@ -124,37 +135,21 @@ func (p *parser) parseDocument() *Document { case TokenEOF: return doc case TokenInclude: - if n := p.parseInclude(); n != nil { - doc.Nodes = append(doc.Nodes, n) - } + doc.appendNode(p.parseInclude()) case TokenCPPInclude: - if n := p.parseCPPInclude(); n != nil { - doc.Nodes = append(doc.Nodes, n) - } + doc.appendNode(p.parseCPPInclude()) case TokenNamespace: - if n := p.parseNamespace(); n != nil { - doc.Nodes = append(doc.Nodes, n) - } + doc.appendNode(p.parseNamespace()) case TokenConst: - if n := p.parseConst(); n != nil { - doc.Nodes = append(doc.Nodes, n) - } + doc.appendNode(p.parseConst()) case TokenTypedef: - if n := p.parseTypedef(); n != nil { - doc.Nodes = append(doc.Nodes, n) - } + doc.appendNode(p.parseTypedef()) case TokenEnum: - if n := p.parseEnum(); n != nil { - doc.Nodes = append(doc.Nodes, n) - } + doc.appendNode(p.parseEnum()) case TokenStruct, TokenUnion, TokenException: - if n := p.parseStruct(); n != nil { - doc.Nodes = append(doc.Nodes, n) - } + doc.appendNode(p.parseStruct()) case TokenService: - if n := p.parseService(); n != nil { - doc.Nodes = append(doc.Nodes, n) - } + doc.appendNode(p.parseService()) default: p.errorfCur("unexpected token %q at top level", p.cur().Text) p.synchronizeTo(TokenInclude, TokenCPPInclude, TokenNamespace, @@ -164,9 +159,18 @@ func (p *parser) parseDocument() *Document { } } +// appendNode appends n to the document's node list. The top-level parse +// functions return Node (not a concrete pointer), so a failed parse is a +// nil interface and this guard suffices. +func (d *Document) appendNode(n Node) { + if n != nil { + d.Nodes = append(d.Nodes, n) + } +} + // --- headers --------------------------------------------------------------- -func (p *parser) parseInclude() *Include { +func (p *parser) parseInclude() Node { n := &Include{nodeBase: nodeBase{first: p.nextReal(p.pos)}} p.advance() // include @@ -186,7 +190,7 @@ func (p *parser) parseInclude() *Include { return n } -func (p *parser) parseCPPInclude() *CPPInclude { +func (p *parser) parseCPPInclude() Node { n := &CPPInclude{nodeBase: nodeBase{first: p.nextReal(p.pos)}} p.advance() // cpp_include @@ -206,20 +210,19 @@ func (p *parser) parseCPPInclude() *CPPInclude { return n } -func (p *parser) parseNamespace() *Namespace { +func (p *parser) parseNamespace() Node { n := &Namespace{nodeBase: nodeBase{first: p.nextReal(p.pos)}} p.advance() // namespace - switch { - case p.at(TokenIdentifier), p.at(TokenStar): - n.Scope = p.advance() - default: + if !p.at(TokenIdentifier) && !p.at(TokenStar) { p.errorfCur("expected namespace scope, got %q", p.cur().Text) p.synchronizeTo(TokenEOF) return nil } + n.Scope = p.advance() + n.Name = p.expectIdentifier("namespace name") if n.Name == nil { p.synchronizeTo(TokenEOF) @@ -235,7 +238,7 @@ func (p *parser) parseNamespace() *Namespace { // --- definitions ----------------------------------------------------------- -func (p *parser) parseConst() *Const { +func (p *parser) parseConst() Node { n := &Const{nodeBase: nodeBase{first: p.nextReal(p.pos)}} p.advance() // const @@ -266,16 +269,14 @@ func (p *parser) parseConst() *Const { return nil } - if sep := p.acceptSeparator(); sep != 0 { - n.Sep = sep - } + n.Sep = p.acceptSeparator() n.last = p.pos - 1 return n } -func (p *parser) parseTypedef() *Typedef { +func (p *parser) parseTypedef() Node { n := &Typedef{nodeBase: nodeBase{first: p.nextReal(p.pos)}} p.advance() // typedef @@ -294,16 +295,14 @@ func (p *parser) parseTypedef() *Typedef { } n.Annotations = p.parseAnnotationsIfPresent() - if sep := p.acceptSeparator(); sep != 0 { - n.Sep = sep - } + n.Sep = p.acceptSeparator() n.last = p.pos - 1 return n } -func (p *parser) parseEnum() *Enum { +func (p *parser) parseEnum() Node { n := &Enum{nodeBase: nodeBase{first: p.nextReal(p.pos)}} p.advance() // enum @@ -326,11 +325,7 @@ func (p *parser) parseEnum() *Enum { n.Values = append(n.Values, v) } - if !p.at(TokenRBrace) { - p.errorfCur("expected '}' to close enum, got %q", p.cur().Text) - } else { - p.advance() - } + p.expect(TokenRBrace, "'}' to close enum") } else { p.errorfCur("expected '{' after enum name, got %q", p.cur().Text) } @@ -362,16 +357,14 @@ func (p *parser) parseEnumValue() *EnumValue { } v.Annotations = p.parseAnnotationsIfPresent() - if sep := p.acceptSeparator(); sep != 0 { - v.Sep = sep - } + v.Sep = p.acceptSeparator() v.last = p.pos - 1 return v } -func (p *parser) parseStruct() *Struct { +func (p *parser) parseStruct() Node { n := &Struct{nodeBase: nodeBase{first: p.nextReal(p.pos)}, Kind: StructKind(p.cur().Kind)} p.advance() // struct | union | exception @@ -384,11 +377,7 @@ func (p *parser) parseStruct() *Struct { if p.accept(TokenLBrace) != nil { n.Fields = p.parseFieldList(TokenRBrace) - if !p.at(TokenRBrace) { - p.errorfCur("expected '}' to close struct, got %q", p.cur().Text) - } else { - p.advance() - } + p.expect(TokenRBrace, "'}' to close struct") } else { p.errorfCur("expected '{' after struct name, got %q", p.cur().Text) } @@ -399,7 +388,7 @@ func (p *parser) parseStruct() *Struct { return n } -func (p *parser) parseService() *Service { +func (p *parser) parseService() Node { n := &Service{nodeBase: nodeBase{first: p.nextReal(p.pos)}} p.advance() // service @@ -435,11 +424,7 @@ func (p *parser) parseService() *Service { n.Functions = append(n.Functions, f) } - if !p.at(TokenRBrace) { - p.errorfCur("expected '}' to close service, got %q", p.cur().Text) - } else { - p.advance() - } + p.expect(TokenRBrace, "'}' to close service") } else { p.errorfCur("expected '{' after service name, got %q", p.cur().Text) } @@ -486,11 +471,7 @@ func (p *parser) parseFunction() *Function { } f.Args = p.parseFieldList(TokenRParen) - if !p.at(TokenRParen) { - p.errorfCur("expected ')' to close arguments, got %q", p.cur().Text) - } else { - p.advance() - } + p.expect(TokenRParen, "')' to close arguments") if p.at(TokenThrows) { p.advance() @@ -504,19 +485,13 @@ func (p *parser) parseFunction() *Function { 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() - } + p.expect(TokenRParen, "')' to close throws") f.Throws.last = p.pos - 1 } f.Annotations = p.parseAnnotationsIfPresent() - if sep := p.acceptSeparator(); sep != 0 { - f.Sep = sep - } + f.Sep = p.acceptSeparator() f.last = p.pos - 1 @@ -533,11 +508,7 @@ func (p *parser) parseFieldList(term TokenKind) []*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 - } + p.recoverEntry(term) continue } @@ -601,9 +572,7 @@ func (p *parser) parseField() (*Field, bool) { } f.Annotations = p.parseAnnotationsIfPresent() - if sep := p.acceptSeparator(); sep != 0 { - f.Sep = sep - } + f.Sep = p.acceptSeparator() f.last = p.pos - 1 @@ -751,11 +720,7 @@ func (p *parser) parseConstValue() *ConstValue { 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 - } + p.recoverEntry(TokenRBracket) continue } @@ -765,11 +730,7 @@ func (p *parser) parseConstValue() *ConstValue { p.acceptSeparator() } - if !p.at(TokenRBracket) { - p.errorfCur("expected ']' to close list constant, got %q", p.cur().Text) - } else { - p.advance() - } + p.expect(TokenRBracket, "']' to close list constant") case TokenLBrace: v.Kind = ValueMap @@ -779,32 +740,20 @@ func (p *parser) parseConstValue() *ConstValue { 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() - } + p.recoverEntry(TokenRBrace) continue } if !p.expect(TokenColon, "':' between map key and value") { - p.synchronizeTo(TokenComma, TokenSemicolon, TokenRBrace) - - if !p.at(TokenRBrace) && !p.at(TokenEOF) { - p.advance() - } + p.recoverEntry(TokenRBrace) continue } value := p.parseConstValue() if value == nil { - p.synchronizeTo(TokenComma, TokenSemicolon, TokenRBrace) - - if !p.at(TokenRBrace) && !p.at(TokenEOF) { - p.advance() - } + p.recoverEntry(TokenRBrace) continue } @@ -814,11 +763,7 @@ func (p *parser) parseConstValue() *ConstValue { p.acceptSeparator() } - if !p.at(TokenRBrace) { - p.errorfCur("expected '}' to close map constant, got %q", p.cur().Text) - } else { - p.advance() - } + p.expect(TokenRBrace, "'}' to close map constant") default: p.errorfCur("expected constant value, got %q", p.cur().Text) @@ -852,11 +797,7 @@ func (p *parser) parseAnnotations() *Annotations { for !p.at(TokenRParen) && !p.at(TokenEOF) { 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() - } + p.recoverEntry(TokenRParen) continue } @@ -875,19 +816,13 @@ func (p *parser) parseAnnotations() *Annotations { } } - if sep := p.acceptSeparator(); sep != 0 { - item.Sep = sep - } + item.Sep = p.acceptSeparator() item.last = p.pos - 1 a.Items = append(a.Items, item) } - if !p.at(TokenRParen) { - p.errorfCur("expected ')' to close annotations, got %q", p.cur().Text) - } else { - p.advance() - } + p.expect(TokenRParen, "')' to close annotations") a.last = p.pos - 1 diff --git a/vscode/package-lock.json b/vscode/package-lock.json index 74d1008..e518ad8 100644 --- a/vscode/package-lock.json +++ b/vscode/package-lock.json @@ -1,12 +1,12 @@ { "name": "thrift-ls", - "version": "0.1.1", + "version": "0.1.4", "lockfileVersion": 3, "requires": true, "packages": { "": { "name": "thrift-ls", - "version": "0.1.1", + "version": "0.1.4", "license": "MIT", "dependencies": { "vscode-languageclient": "^9.0.1" diff --git a/vscode/package.json b/vscode/package.json index e428f50..d9d9a8e 100644 --- a/vscode/package.json +++ b/vscode/package.json @@ -2,7 +2,7 @@ "name": "thrift-ls", "displayName": "Thrift Language Server", "description": "Language server and formatter for Apache Thrift IDL files", - "version": "0.1.3", + "version": "0.1.4", "publisher": "karitham", "license": "MIT", "engines": {