diff --git a/README.md b/README.md index 4ca6181..32d1617 100644 --- a/README.md +++ b/README.md @@ -154,6 +154,7 @@ vim.lsp.config('dexter', { filetypes = { 'elixir', 'eelixir', 'heex' }, init_options = { followDelegates = true, -- jump through defdelegate to the target function + -- definitionStyle = "all", -- "all" returns all function heads; "first" jumps to the first one -- stdlibPath = "", -- override Elixir stdlib path (auto-detected) -- debug = false, -- verbose logging to stderr (view with :LspLog) }, @@ -271,6 +272,21 @@ If Zed shows a *"could not detect Elixir stdlib"* warning on startup — common Equivalently, set the `DEXTER_ELIXIR_LIB_ROOT` environment variable via `lsp.dexter.binary.env`. +To configure other LSP options, such as returning only the first matching function head, add them to the same `initialization_options` object (see [LSP options](#lsp-options)): + +```json +{ + "lsp": { + "dexter": { + "initialization_options": { + "followDelegates": true, + "definitionStyle": "first" + } + } + } +} +``` + ### Emacs The emacs instructions assume you're using **use-package**. @@ -592,6 +608,7 @@ If the persistent process can't start, dexter falls back to running `mix format` Dexter reads `initializationOptions` from your editor configuration: - **`followDelegates`** (boolean, default: `true`): follow `defdelegate` targets on lookup. +- **`definitionStyle`** (string, default: `"all"`): controls how many locations are returned when a function has multiple heads (clauses). `"all"` returns every definition site; `"first"` returns only the first one, which makes editors like Zed jump directly instead of showing a picker. - **`stdlibPath`** (string): override the Elixir stdlib directory to index. Defaults to auto-detection; use this if your install is non-standard. - **`debug`** (boolean, default: `false`): enable verbose logging for this editor session. Logs timing and resolution details for every definition, hover, references, and rename request to your editor's LSP log and to the workspace daemon's log (see [Debugging](#debugging)). Can also be enabled via the `DEXTER_DEBUG=true` environment variable. - **`maxTransientDocuments`** (integer, default: `50`): cap on how many lazily-loaded buffers the server retains in memory. When an LSP client (e.g. Claude Code) queries a file it never opened via `didOpen`, dexter reads it from disk and caches it. Editor-owned buffers are unaffected; only disk-loaded entries are subject to LRU eviction. Set to `0` to disable transient caching. diff --git a/internal/lsp/elixir.go b/internal/lsp/elixir.go index 6e3f525..a7b9962 100644 --- a/internal/lsp/elixir.go +++ b/internal/lsp/elixir.go @@ -68,6 +68,267 @@ func (tf *TokenizedFile) FullExpressionAtCursor(line, col int) CursorContext { return ctx } +// ArityAtCallsite returns the call arity at the given expression position, or +// -1 when arity can't be determined. Handles: +// - Foo.bar(a, b) → 2 +// - Foo.bar() → 0 +// - &Foo.bar/2 → 2 (capture syntax) +// - x |> Foo.bar(y) → 2 (pipe injects one implicit arg) +// - Foo.bar → -1 (no call suffix, arity unknown) +// +// line is 0-based. startCol/endCol are the expression's 0-based column bounds +// (as returned in CursorContext.ExprStart/ExprEnd). +func (tf *TokenizedFile) ArityAtCallsite(line, startCol, endCol int) int { + endOffset := parser.LineColToOffset(tf.lineStarts, line, endCol-1) + if endOffset >= 0 && parser.TokenAtOffset(tf.interp, endOffset) >= 0 { + // The interpolation stream intentionally contains only references, not + // the delimiters required to determine a call's arity. + return -1 + } + return arityAtCallsite(tf.tokens, tf.source, tf.lineStarts, line, startCol, endCol) +} + +func arityAtCallsite(tokens []parser.Token, source []byte, lineStarts []int, line, startCol, endCol int) int { + n := len(tokens) + if n == 0 || endCol <= 0 { + return -1 + } + + // Locate the last token of the expression (index of the char at endCol-1). + endOffset := parser.LineColToOffset(lineStarts, line, endCol-1) + if endOffset < 0 { + return -1 + } + endIdx := parser.TokenAtOffset(tokens, endOffset) + if endIdx < 0 { + return -1 + } + + w := parser.NewTokenWalker(source, tokens) + w.SetPos(endIdx + 1) + w.SkipToNextSig() + j := w.Pos() + + arity := -1 + switch { + case j < n && tokens[j].Kind == parser.TokOpenParen: + var closeIdx int + arity, closeIdx = countCallArgs(source, tokens, j) + if arity >= 0 { + w.SetPos(closeIdx + 1) + w.SkipToNextSig() + if w.CurrentKind() == parser.TokDo { + startOffset := parser.LineColToOffset(lineStarts, line, startCol) + startIdx := parser.TokenAtOffset(tokens, startOffset) + prev := w.PreviousSigPos(startIdx) + switch { + case prev >= 0 && isDefinitionKeyword(tokens[prev].Kind): + // `def name(a, b) do` opens the body, not a keyword argument. + case prev >= 0 && tokenCanOwnFollowingExpression(tokens[prev].Kind): + return -1 + default: + arity++ + } + } + } + case j < n && tokens[j].Kind == parser.TokOther && + tokens[j].End-tokens[j].Start == 1 && source[tokens[j].Start] == '/': + // Capture syntax: &Foo.bar/2 + startOffset := parser.LineColToOffset(lineStarts, line, startCol) + startIdx := parser.TokenAtOffset(tokens, startOffset) + prev := w.PreviousSigPos(startIdx) + if prev < 0 || tokens[prev].Kind != parser.TokOther || + tokens[prev].End-tokens[prev].Start != 1 || source[tokens[prev].Start] != '&' { + return -1 + } + w.SetPos(j + 1) + w.SkipToNextSig() + k := w.Pos() + if k < n && tokens[k].Kind == parser.TokNumber { + if a, ok := parseNumberTokenArity(source, tokens[k]); ok { + arity = a + } + } + } + + if arity < 0 { + return -1 + } + + // Pipe adjustment: if the expression is the RHS of a |>, add one for the + // implicit first argument. + startOffset := parser.LineColToOffset(lineStarts, line, startCol) + if startOffset >= 0 { + startIdx := parser.TokenAtOffset(tokens, startOffset) + if prev := w.PreviousSigPos(startIdx); prev >= 0 && tokens[prev].Kind == parser.TokPipe { + return arity + 1 + } + } + + return arity +} + +// countCallArgs counts top-level arguments inside a parenthesized call, +// starting at openIdx which must be a TokOpenParen. It returns the arity and +// matching close-token index, or -1, -1 when the expression is unbalanced. +func countCallArgs(source []byte, tokens []parser.Token, openIdx int) (int, int) { + if openIdx >= len(tokens) || tokens[openIdx].Kind != parser.TokOpenParen { + return -1, -1 + } + w := parser.NewTokenWalker(source, tokens) + w.SetPos(openIdx) + w.Advance() + args := 0 + hasContent := false + keywordTail := false + // innerCall marks a parenthesis-free call such as `if ready, do: x` or + // `fetch user, opts` inside the argument list: Elixir gives it every + // following top-level comma until a do block ends it. + innerCall := false + for w.More() { + pos := w.Pos() + kind := w.CurrentKind() + switch kind { + case parser.TokCloseParen, parser.TokCloseBracket, parser.TokCloseBrace, parser.TokCloseAngle: + if w.Depth() == 1 && w.BlockDepth() == 0 { + if hasContent { + return args + 1, pos + } + return 0, pos + } + case parser.TokDo: + if w.Depth() == 1 && w.BlockDepth() == 0 { + innerCall = false + } + hasContent = true + case parser.TokIdent: + if w.Depth() == 1 && w.BlockDepth() == 0 && !innerCall && startsParenFreeCall(source, tokens, pos) { + innerCall = true + } + hasContent = true + case parser.TokComma: + if w.Depth() == 1 && w.BlockDepth() == 0 { + if innerCall { + // Only a sole argument may be a parenthesis-free call + // followed by commas; anything else is a syntax error. + if args > 0 { + return -1, -1 + } + w.Advance() + continue + } + // Elixir's trailing keyword syntax is one list argument even + // though its entries are separated by top-level commas. + if keywordTail { + w.Advance() + continue + } + args++ + hasContent = false + w.Advance() + continue + } + hasContent = true + case parser.TokColon: + prev := w.PreviousSigPos(pos) + if w.Depth() == 1 && w.BlockDepth() == 0 && prev > openIdx && tokens[prev].Kind == parser.TokIdent { + keywordTail = true + } + hasContent = true + case parser.TokEOL, parser.TokComment: + // skip + default: + hasContent = true + } + w.Advance() + } + return -1, -1 +} + +// startsParenFreeCall reports whether the identifier at pos is called without +// parentheses, as in `if ready, ...` or `fetch user`: Elixir parses a name +// followed by whitespace and the start of an expression as a call. Word +// operators (`a in b`, `not c`) and binary minus (`a - 1`) are not calls. +func startsParenFreeCall(source []byte, tokens []parser.Token, pos int) bool { + if pos+1 >= len(tokens) { + return false + } + name := string(source[tokens[pos].Start:tokens[pos].End]) + if isWordOperator(name) { + return false + } + next := tokens[pos+1] + if next.Start <= tokens[pos].End { + // `foo(`, `foo.`, `foo[` and `foo:` are not parenthesis-free calls. + return false + } + switch next.Kind { + case parser.TokIdent: + return !isWordOperator(string(source[next.Start:next.End])) + case parser.TokModule, parser.TokNumber, parser.TokString, parser.TokHeredoc, + parser.TokSigil, parser.TokCharLiteral, parser.TokAtom, parser.TokOpenBracket, + parser.TokOpenBrace, parser.TokPercent, parser.TokFn, parser.TokAttr, + parser.TokAttrDoc, parser.TokAttrSpec, parser.TokAttrType, + parser.TokAttrBehaviour, parser.TokAttrCallback: + return true + case parser.TokOther: + if next.End-next.Start != 1 { + return false + } + switch source[next.Start] { + case '&', '^', '!': + return true + case '-', '+': + // `foo -1` is a call; `a - 1` is subtraction. + return pos+2 < len(tokens) && tokens[pos+2].Start == next.End + } + } + return false +} + +func isWordOperator(name string) bool { + switch name { + case "in", "and", "or", "not", "when": + return true + } + return false +} + +func isDefinitionKeyword(kind parser.TokenKind) bool { + switch kind { + case parser.TokDef, parser.TokDefp, parser.TokDefmacro, parser.TokDefmacrop, + parser.TokDefguard, parser.TokDefguardp: + return true + } + return false +} + +func tokenCanOwnFollowingExpression(kind parser.TokenKind) bool { + switch kind { + case parser.TokIdent, parser.TokModule, parser.TokNumber, parser.TokString, + parser.TokHeredoc, parser.TokSigil, parser.TokCharLiteral, parser.TokAtom, + parser.TokCloseParen, parser.TokCloseBracket, parser.TokCloseBrace, parser.TokCloseAngle: + return true + default: + return false + } +} + +func parseNumberTokenArity(source []byte, t parser.Token) (int, bool) { + text := source[t.Start:t.End] + n := 0 + for _, b := range text { + if b < '0' || b > '9' { + return 0, false + } + n = n*10 + int(b-'0') + if n > 255 { // arity fits in a byte in practice + return 0, false + } + } + return n, true +} + // FirstDefmodule returns the first defmodule name found, or "". func (tf *TokenizedFile) FirstDefmodule() string { for i := 0; i < tf.n; i++ { @@ -116,6 +377,86 @@ func (tf *TokenizedFile) FindTypeDefinition(functionName string) (int, bool) { return tf.findDefinition(functionName, true) } +// FindDefinitionLines returns the callable or type definition lines declared +// directly in module; declarations in nested or sibling modules belong to +// those modules. An arity below zero keeps every arity. preferType selects the +// namespace to prefer when a type and callable share a name. +func (tf *TokenizedFile) FindDefinitionLines(module, functionName string, arity int, preferType bool) []int { + var functionLines, typeLines []int + type moduleFrame struct { + name string + blockDepth int + } + var stack []moduleFrame + w := parser.NewTokenWalker(tf.source, tf.tokens) + for w.More() { + i := w.Pos() + tok := w.Current() + blockDepth := w.BlockDepth() + w.Advance() + switch tok.Kind { + case parser.TokDefmodule, parser.TokDefprotocol, parser.TokDefimpl: + parent := "" + if len(stack) > 0 { + parent = stack[len(stack)-1].name + } + if name, _, hasDo := tokParseModuleDef(tf.source, tf.tokens, i+1, parent); name != "" && hasDo { + // The walker counts the module's do when it reaches it. + stack = append(stack, moduleFrame{name: name, blockDepth: blockDepth + 1}) + } + continue + case parser.TokEnd: + if len(stack) > 0 && stack[len(stack)-1].blockDepth == blockDepth { + stack = stack[:len(stack)-1] + } + continue + } + if len(stack) == 0 || stack[len(stack)-1].name != module { + continue + } + switch tok.Kind { + case parser.TokDef, parser.TokDefp, parser.TokDefmacro, parser.TokDefmacrop, + parser.TokDefguard, parser.TokDefguardp, parser.TokDefdelegate: + name, j, ok := parser.StaticDeclarationName(tf.source, tf.tokens, tf.n, i) + if !ok || name != functionName { + continue + } + maxArity, defaultCount := 0, 0 + pj := tokNextSig(tf.tokens, tf.n, j+1) + if pj < tf.n && tf.tokens[pj].Kind == parser.TokOpenParen { + maxArity, defaultCount, _, _ = parser.CollectParams(tf.source, tf.tokens, tf.n, pj) + } + if arity < 0 || (arity >= maxArity-defaultCount && arity <= maxArity) { + functionLines = append(functionLines, tok.Line) + } + + case parser.TokAttrType: + name, j, ok := parser.StaticDeclarationName(tf.source, tf.tokens, tf.n, i) + if !ok || name != functionName { + continue + } + typeArity := 0 + pj := tokNextSig(tf.tokens, tf.n, j+1) + if pj < tf.n && tf.tokens[pj].Kind == parser.TokOpenParen { + typeArity, _, _, _ = parser.CollectParams(tf.source, tf.tokens, tf.n, pj) + } + if arity < 0 || arity == typeArity { + typeLines = append(typeLines, tok.Line) + } + } + } + if preferType { + if len(typeLines) > 0 { + return typeLines + } + return functionLines + } + if len(functionLines) > 0 { + return functionLines + } + return typeLines +} + // findDefinition returns the line of the first matching definition. A module // may declare both a type and a function under one name — Ecto.Schema has // `@type schema` above `defmacro schema/2` — so file order alone cannot decide diff --git a/internal/lsp/elixir_test.go b/internal/lsp/elixir_test.go index 1ed4308..9372c71 100644 --- a/internal/lsp/elixir_test.go +++ b/internal/lsp/elixir_test.go @@ -437,6 +437,143 @@ func TestExpressionAtCursor_ExprBounds(t *testing.T) { } } +func TestArityAtCallsite_KeywordTailCountsAsOneArgument(t *testing.T) { + code := "SharedLib.Repo.insert(changeset, returning: true, on_conflict: :replace)" + tf := NewTokenizedFile(code) + ctx := tf.ExpressionAtCursor(0, strings.Index(code, "insert")+2) + if got := tf.ArityAtCallsite(0, ctx.ExprStart, ctx.ExprEnd); got != 2 { + t.Fatalf("ArityAtCallsite() = %d, want 2", got) + } +} + +func TestArityAtCallsite_ComplexForms(t *testing.T) { + tests := []struct { + name string + code string + want int + }{ + { + name: "commas in fn body do not add arguments", + code: "SharedLib.Worker.run(fn left, right -> {left, right} end)", + want: 1, + }, + { + name: "trailing do block is a keyword list argument", + code: "SharedLib.Worker.run(:value) do\n :ok\nend", + want: 2, + }, + { + name: "inline do keyword tail is one argument", + code: "SharedLib.Worker.run(:value, do: :ok, else: :error)", + want: 2, + }, + { + name: "slash without capture is ambiguous", + code: "SharedLib.Worker.run / 2", + want: -1, + }, + { + name: "capture slash supplies arity", + code: "&SharedLib.Worker.run/2", + want: 2, + }, + { + name: "outer block ownership is ambiguous", + code: "if SharedLib.Worker.run(:value) do\n :ok\nend", + want: -1, + }, + { + name: "parenthesis-free call is ambiguous", + code: "SharedLib.Worker.run :value, mode: :fast", + want: -1, + }, + { + name: "unparenthesized if owns the following commas", + code: "SharedLib.Worker.run(if ready, do: :ok, else: :error)", + want: 1, + }, + { + name: "unparenthesized for owns its generators", + code: "SharedLib.Worker.run(for x <- xs, y <- ys, do: {x, y})", + want: 1, + }, + { + name: "unparenthesized with owns its clauses", + code: "SharedLib.Worker.run(with {:ok, a} <- fetch(), {:ok, b} <- load(a), do: b)", + want: 1, + }, + { + name: "parenthesis-free remote call owns the following commas", + code: "SharedLib.Worker.run(MyApp.Accounts.get user, opts)", + want: 1, + }, + { + name: "parenthesis-free call after a match owns the following commas", + code: "SharedLib.Worker.run(result = fetch user, opts)", + want: 1, + }, + { + name: "parenthesis-free call as the last argument", + code: "SharedLib.Worker.run(:value, fetch user)", + want: 2, + }, + { + name: "parenthesized if keeps outer arguments", + code: "SharedLib.Worker.run(if(ready, do: :ok), :value)", + want: 2, + }, + { + name: "word operators are not calls", + code: "SharedLib.Worker.run(a in b, not c, d and e, f or g)", + want: 4, + }, + { + name: "binary minus is not a call", + code: "SharedLib.Worker.run(a - 1, b - c, d)", + want: 3, + }, + { + name: "unary minus after a space starts a call", + code: "SharedLib.Worker.run(fetch -1, d)", + want: 1, + }, + { + name: "definition head do block is not an argument", + code: "def run(left, right) do\n :ok\nend", + want: 2, + }, + { + name: "private macro head do block is not an argument", + code: "defmacrop run(left) do\n :ok\nend", + want: 1, + }, + { + name: "do block ends a parenthesis-free call", + code: "SharedLib.Worker.run(case x do\n _ -> {1, 2}\nend, y)", + want: 2, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + tf := NewTokenizedFile(tt.code) + ctx := tf.ExpressionAtCursor(0, strings.Index(tt.code, "run")+1) + if got := tf.ArityAtCallsite(0, ctx.ExprStart, ctx.ExprEnd); got != tt.want { + t.Fatalf("ArityAtCallsite() = %d, want %d", got, tt.want) + } + }) + } +} + +func TestArityAtCallsite_InterpolationIsAmbiguous(t *testing.T) { + code := `"#{SharedLib.Worker.run(:value)}"` + tf := NewTokenizedFile(code) + ctx := tf.ExpressionAtCursor(0, strings.Index(code, "run")+1) + if got := tf.ArityAtCallsite(0, ctx.ExprStart, ctx.ExprEnd); got != -1 { + t.Fatalf("ArityAtCallsite() = %d, want -1", got) + } +} + func TestCursorContext_Expr(t *testing.T) { tests := []struct { mod, fn, want string diff --git a/internal/lsp/generated_definition_test.go b/internal/lsp/generated_definition_test.go index 1ac136b..2d0daf1 100644 --- a/internal/lsp/generated_definition_test.go +++ b/internal/lsp/generated_definition_test.go @@ -828,3 +828,51 @@ func TestDefinitionGeneratedFunctionWithoutDeclarationKeepsModuleLine(t *testing }) } } + +// splitArityDefinitions records each arity of get_room_by_slug! on its own +// line, so a result shows which arity it came from. +func splitArityDefinitions() []dbgiDefinition { + const generator = "deps/ash/lib/ash/code_interface.ex" + return []dbgiDefinition{ + {name: "list_rooms", arity: 0, line: 7, keepFile: generator, keepLine: 1112}, + {name: "get_room_by_slug!", arity: 1, line: generatedDefineLine, keepFile: generator, keepLine: 1112}, + {name: "get_room_by_slug!", arity: 2, line: 7, keepFile: generator, keepLine: 1112}, + } +} + +func TestDefinitionBareGeneratedFunctionFiltersByCallArity(t *testing.T) { + server, domainPath := newGeneratedDefinitionFixture(t, generatedDomainRel, splitArityDefinitions()...) + locations := generatedDefinitionAt(t, server, generatedDomainRel, generatedDomainSource, 11, 24) + expectSingleLocation(t, locations, domainPath, generatedDefineLine) +} + +func TestDefinitionBareGeneratedFunctionAppliesDefinitionStyle(t *testing.T) { + server, domainPath := newGeneratedDefinitionFixture(t, generatedDomainRel, splitArityDefinitions()...) + // No generated arity matches, so every arity remains a candidate. + source := strings.Replace(generatedDomainSource, `get_room_by_slug!("lounge")`, `get_room_by_slug!("lounge", 1, 2)`, 1) + + if locations := generatedDefinitionAt(t, server, generatedDomainRel, source, 11, 24); len(locations) != 2 { + t.Fatalf("expected both arities with definitionStyle all, got %#v", locations) + } + + server.definitionStyle = "first" + locations := generatedDefinitionAt(t, server, generatedDomainRel, source, 11, 24) + if len(locations) != 1 || uriToPath(locations[0].URI) != domainPath { + t.Fatalf("expected one location with definitionStyle first, got %#v", locations) + } +} + +func TestLookupNameGeneratedArityMissFallsBackToModule(t *testing.T) { + server, domainPath := newGeneratedDefinitionFixture(t, generatedDomainRel, lineDefinitions()...) + locations, err := server.LookupName("MyApp.Chat", "get_room_by_slug!", NameLookupOptions{ + Arity: 5, + ExactArity: true, + FallbackToModule: true, + }) + if err != nil { + t.Fatal(err) + } + if len(locations) != 1 || locations[0].FilePath != domainPath || locations[0].Line != generatedModuleLine { + t.Fatalf("expected module fallback %s:%d, got %#v", domainPath, generatedModuleLine, locations) + } +} diff --git a/internal/lsp/name_navigation.go b/internal/lsp/name_navigation.go index 3719bcf..7638e96 100644 --- a/internal/lsp/name_navigation.go +++ b/internal/lsp/name_navigation.go @@ -1,6 +1,9 @@ package lsp -import "github.com/remoteoss/dexter/internal/store" +import ( + "github.com/remoteoss/dexter/internal/beam" + "github.com/remoteoss/dexter/internal/store" +) // NameKind selects which Elixir namespace a canonical name refers to. type NameKind uint8 @@ -26,6 +29,8 @@ type NameLookupOptions struct { External bool FallbackToModule bool ExcludeStdlib bool + Arity int + ExactArity bool // ExactModule places a generated function only at its own module's // definition. Without it, a module that exists only as a BEAM, such as // Phoenix route helpers, resolves to the nearest lexical parent with source; @@ -57,8 +62,12 @@ func (s *Server) LookupName(module, function string, opts NameLookupOptions) ([] var results []store.LookupResult var err error - if opts.FollowDelegates { + if opts.FollowDelegates && opts.ExactArity { + results, err = s.store.LookupFollowDelegateByArity(module, function, opts.Arity) + } else if opts.FollowDelegates { results, err = s.store.LookupFollowDelegate(module, function) + } else if opts.ExactArity { + results, err = s.store.LookupFunctionByArity(module, function, opts.Arity) } else { results, err = s.store.LookupFunction(module, function) } @@ -70,22 +79,32 @@ func (s *Server) LookupName(module, function string, opts NameLookupOptions) ([] results = filterOutPrivate(results) } if len(results) == 0 { - results = filterLookupKind(s.lookupThroughUseOfWithFollow(module, function, opts.FollowDelegates), opts.Kind) + arity := -1 + if opts.ExactArity { + arity = opts.Arity + } + results = filterLookupKind(s.lookupThroughUseOfWithFollow(module, function, opts.FollowDelegates, arity), opts.Kind) } if len(results) == 0 && opts.Kind != NameKindType { if generated, found := s.generatedSymbol(module, "", function); found && len(generated) > 0 { - // A line the compiled module records for the function is its own - // definition, so even an exact lookup takes it. - var precise bool - results, precise = s.generatedDefinitionResultsFor(module, "", generated) - if !precise && opts.ExactModule { - if results, err = s.store.LookupModule(module); err != nil { - return nil, err - } + if opts.ExactArity { + generated = filterGeneratedFunctionsByArity(generated, opts.Arity) } - if !precise && len(results) > 0 { - results[0].Arity = generated[0].Arity - results[0].Kind = generated[0].Kind + // No generated arity matching the call leaves the module fallback. + if len(generated) > 0 { + // A line the compiled module records for the function is its own + // definition, so even an exact lookup takes it. + var precise bool + results, precise = s.generatedDefinitionResultsFor(module, "", generated) + if !precise && opts.ExactModule { + if results, err = s.store.LookupModule(module); err != nil { + return nil, err + } + } + if !precise && len(results) > 0 { + results[0].Arity = generated[0].Arity + results[0].Kind = generated[0].Kind + } } } } @@ -98,6 +117,29 @@ func (s *Server) LookupName(module, function string, opts NameLookupOptions) ([] return s.lookupLocations(results, opts.ExcludeStdlib), nil } +// generatedFunctionsForCall narrows generated functions to the call's arity, +// keeping every arity when it is unknown or none matches: a generated function +// is still the best target the call has. +func generatedFunctionsForCall(functions []beam.Function, arity int) []beam.Function { + if arity < 0 { + return functions + } + if filtered := filterGeneratedFunctionsByArity(functions, arity); len(filtered) > 0 { + return filtered + } + return functions +} + +func filterGeneratedFunctionsByArity(functions []beam.Function, arity int) []beam.Function { + filtered := make([]beam.Function, 0, len(functions)) + for _, function := range functions { + if function.Arity == arity { + filtered = append(filtered, function) + } + } + return filtered +} + // ReferenceNames finds references after a frontend has resolved a canonical // module/function name. Cursor-specific alias and variable resolution stays in // the LSP adapter. diff --git a/internal/lsp/server.go b/internal/lsp/server.go index b375fb5..85c9de2 100644 --- a/internal/lsp/server.go +++ b/internal/lsp/server.go @@ -199,6 +199,7 @@ type Server struct { clientLog *clientLog // forwards this session's log lines to its editor followDelegates bool debug bool + definitionStyle string // "all" (default) or "first": controls multi-head definition results mixBin string // resolved path to the mix binary beams map[string]*beamProcess // build root → persistent BEAM process @@ -279,6 +280,7 @@ func NewServerWithOptions(s *store.Store, projectRoot string, opts ServerOptions projectRoot: projectRoot, explicitRoot: projectRoot != "", followDelegates: true, + definitionStyle: "all", // Read here as well as in Initialize: the daemon's headless service // answers CLI and MCP calls and never receives an initialize request. debug: os.Getenv("DEXTER_DEBUG") == "true", @@ -908,6 +910,11 @@ func (s *Server) Initialize(ctx context.Context, params *protocol.InitializePara if v, ok := opts["maxTransientDocuments"].(float64); ok { s.docs.SetMaxTransient(int(v)) } + if v, ok := opts["definitionStyle"].(string); ok { + if v == "all" || v == "first" { + s.definitionStyle = v + } + } } if os.Getenv("DEXTER_DEBUG") == "true" { s.debug = true @@ -1248,6 +1255,7 @@ func (s *Server) Definition(ctx context.Context, params *protocol.DefinitionPara expr := tf.ResolveModuleExpr(exprCtx.Expr(), lineNum) moduleRef, functionName := ExtractModuleAndFunction(expr) + callArity := tf.ArityAtCallsite(lineNum, exprCtx.ExprStart, exprCtx.ExprEnd) if moduleRef != "" { if aliasParent, inBlock := tf.ExtractAliasBlockParent(lineNum); inBlock { @@ -1257,7 +1265,7 @@ func (s *Server) Definition(ctx context.Context, params *protocol.DefinitionPara aliases := tf.ExtractAliasesInScope(lineNum) s.mergeAliasesFromUseTokenized(tf, aliases) - s.debugf("Definition: expr=%q module=%q function=%q", expr, moduleRef, functionName) + s.debugf("Definition: expr=%q module=%q function=%q arity=%d", expr, moduleRef, functionName, callArity) // Bare identifier — check variable first (cheap tree-sitter lookup), then functions if moduleRef == "" { @@ -1286,9 +1294,9 @@ func (s *Server) Definition(ctx context.Context, params *protocol.DefinitionPara if provider, functions, found := s.generatedSymbolInScope(currentModule, func() []string { return s.enclosingBlockPath(docURI, lineNum, col) }, functionName); found { - if results, precise := s.generatedDefinitionResultsFor(provider.module, provider.beamPath, functions); len(results) > 0 { + if results, precise := s.generatedDefinitionResultsFor(provider.module, provider.beamPath, generatedFunctionsForCall(functions, callArity)); len(results) > 0 { s.debugf("Definition: generated bare %q provider=%s precise=%t", functionName, provider.module, precise) - return storeResultsToLocations(results), nil + return s.applyDefinitionStyle(storeResultsToLocations(results)), nil } } s.debugf("Definition: could not resolve bare function %q", functionName) @@ -1298,15 +1306,16 @@ func (s *Server) Definition(ctx context.Context, params *protocol.DefinitionPara // Current module — return buffer location directly (works before indexing). // In a typespec the bare name is the type, everywhere else the function. if fullModule == currentModule { - find := tf.FindFunctionDefinition - if tf.InTypespec(lineNum) { - find = tf.FindTypeDefinition - } - if line, found := find(functionName); found { - return []protocol.Location{{ - URI: params.TextDocument.URI, - Range: lineRange(line - 1), - }}, nil + lines := tf.FindDefinitionLines(fullModule, functionName, callArity, tf.InTypespec(lineNum)) + if len(lines) > 0 { + locations := make([]protocol.Location, 0, len(lines)) + for _, line := range lines { + locations = append(locations, protocol.Location{ + URI: params.TextDocument.URI, + Range: lineRange(line - 1), + }) + } + return s.applyDefinitionStyle(locations), nil } } @@ -1318,25 +1327,27 @@ func (s *Server) Definition(ctx context.Context, params *protocol.DefinitionPara Kind: kind, FollowDelegates: s.followDelegates, External: fullModule != extractEnclosingModuleFromTokens(tf.source, tf.tokens, lineNum), + Arity: callArity, + ExactArity: callArity >= 0, }) if err == nil && len(results) > 0 { s.debugf("Definition: found %d semantic result(s) for %s.%s", len(results), fullModule, functionName) - return nameLocationsToProtocol(results), nil + return s.applyDefinitionStyle(nameLocationsToProtocol(results)), nil } // Fallback for use-chain inline defs (not stored as module definitions) - if results := s.lookupThroughUse(text, functionName, aliases); len(results) > 0 { + if results := s.lookupThroughUseWithFollow(text, functionName, aliases, s.followDelegates, callArity); len(results) > 0 { s.debugf("Definition: found %d result(s) via current file use chain for %s", len(results), functionName) - return storeResultsToLocations(byKindForContext(tf, lineNum, results)), nil + return s.applyDefinitionStyle(storeResultsToLocations(byKindForContext(tf, lineNum, results))), nil } currentModule = s.store.LookupEnclosingModule(uriToPath(protocol.DocumentURI(docURI)), lineNum+1) if provider, functions, found := s.generatedSymbolInScope(currentModule, func() []string { return s.enclosingBlockPath(docURI, lineNum, col) }, functionName); found { - if results, precise := s.generatedDefinitionResultsFor(provider.module, provider.beamPath, functions); len(results) > 0 { + if results, precise := s.generatedDefinitionResultsFor(provider.module, provider.beamPath, generatedFunctionsForCall(functions, callArity)); len(results) > 0 { s.debugf("Definition: generated fallback for bare %q provider=%s precise=%t", functionName, provider.module, precise) - return storeResultsToLocations(results), nil + return s.applyDefinitionStyle(storeResultsToLocations(results)), nil } } @@ -1359,11 +1370,13 @@ func (s *Server) Definition(ctx context.Context, params *protocol.DefinitionPara FollowDelegates: s.followDelegates, External: fullModule != extractEnclosingModuleFromTokens(tf.source, tf.tokens, lineNum), FallbackToModule: true, + Arity: callArity, + ExactArity: callArity >= 0, }) if err != nil { return nil, nil } - return nameLocationsToProtocol(results), nil + return s.applyDefinitionStyle(nameLocationsToProtocol(results)), nil } results, err := s.LookupName(fullModule, "", NameLookupOptions{}) @@ -1390,6 +1403,13 @@ func nameLocationsToProtocol(results []NameLocation) []protocol.Location { return locations } +func (s *Server) applyDefinitionStyle(locations []protocol.Location) []protocol.Location { + if s.definitionStyle == "first" && len(locations) > 1 { + return locations[:1] + } + return locations +} + func storeResultsToLocations(results []store.LookupResult) []protocol.Location { type locKey struct { filePath string @@ -2612,10 +2632,10 @@ func (s *Server) parseUsingFile(filePath, moduleName string) *usingCacheEntry { // fullModule's source file. This handles qualified calls like M.func() where // func is not defined directly in M but is injected by a macro M uses. func (s *Server) lookupThroughUseOf(fullModule, functionName string) []store.LookupResult { - return s.lookupThroughUseOfWithFollow(fullModule, functionName, s.followDelegates) + return s.lookupThroughUseOfWithFollow(fullModule, functionName, s.followDelegates, -1) } -func (s *Server) lookupThroughUseOfWithFollow(fullModule, functionName string, followDelegates bool) []store.LookupResult { +func (s *Server) lookupThroughUseOfWithFollow(fullModule, functionName string, followDelegates bool, arity int) []store.LookupResult { modResults, err := s.store.LookupModule(fullModule) if err != nil || len(modResults) == 0 { return nil @@ -2624,7 +2644,7 @@ func (s *Server) lookupThroughUseOfWithFollow(fullModule, functionName string, f if !ok { return nil } - return s.lookupThroughUseWithFollow(fileText, functionName, ExtractAliases(fileText), followDelegates) + return s.lookupThroughUseWithFollow(fileText, functionName, ExtractAliases(fileText), followDelegates, arity) } // lookupThroughUse searches for functionName in definitions injected by `use` @@ -2632,15 +2652,15 @@ func (s *Server) lookupThroughUseOfWithFollow(fullModule, functionName string, f // priority over imported ones. Later `use` declarations shadow earlier ones. // Transitive use chains (use inside __using__ body) are followed recursively. func (s *Server) lookupThroughUse(text, functionName string, aliases map[string]string) []store.LookupResult { - return s.lookupThroughUseWithFollow(text, functionName, aliases, s.followDelegates) + return s.lookupThroughUseWithFollow(text, functionName, aliases, s.followDelegates, -1) } -func (s *Server) lookupThroughUseWithFollow(text, functionName string, aliases map[string]string, followDelegates bool) []store.LookupResult { +func (s *Server) lookupThroughUseWithFollow(text, functionName string, aliases map[string]string, followDelegates bool, arity int) []store.LookupResult { useCalls := ExtractUsesWithOpts(text, aliases) visited := make(map[string]bool) for i := len(useCalls) - 1; i >= 0; i-- { - if result := s.lookupInUsingEntryForWithFollow(useCalls[i].Module, functionName, useCalls[i].dispatchAtom(), useCalls[i].Opts, visited, followDelegates); result != nil { + if result := s.lookupInUsingEntryForWithFollow(useCalls[i].Module, functionName, useCalls[i].dispatchAtom(), useCalls[i].Opts, visited, followDelegates, arity); result != nil { return result } } @@ -2652,7 +2672,7 @@ func (s *Server) lookupThroughUseWithFollow(text, functionName string, aliases m // consumerOpts are the keyword args from the `use Module, key: Val` call and // are used to resolve dynamic imports like `import unquote(mod)`. func (s *Server) lookupInUsingEntry(moduleName, functionName string, consumerOpts map[string]string, visited map[string]bool) []store.LookupResult { - return s.lookupInUsingEntryForWithFollow(moduleName, functionName, "", consumerOpts, visited, s.followDelegates) + return s.lookupInUsingEntryForWithFollow(moduleName, functionName, "", consumerOpts, visited, s.followDelegates, -1) } // bodyFor picks the injected body for a `use` call. An ordinary __using__ has a @@ -2682,14 +2702,38 @@ func usingVisitKey(moduleName, which string) string { return moduleName + "\x00" + which } +// usingVisitKeyWithOpts also keys on the consumer opts: one injector used twice +// with different opts selects different providers, so visiting it for one +// `use` must not skip it for the other. +func usingVisitKeyWithOpts(moduleName, which string, opts map[string]string) string { + key := usingVisitKey(moduleName, which) + if len(opts) == 0 { + return key + } + keys := make([]string, 0, len(opts)) + for k := range opts { + keys = append(keys, k) + } + sort.Strings(keys) + var b strings.Builder + b.WriteString(key) + for _, k := range keys { + b.WriteByte(0) + b.WriteString(k) + b.WriteByte('=') + b.WriteString(opts[k]) + } + return b.String() +} + // lookupInUsingEntryFor is lookupInUsingEntry with the dispatch atom from the // `use` site (empty for an ordinary `use Module`). func (s *Server) lookupInUsingEntryFor(moduleName, functionName, which string, consumerOpts map[string]string, visited map[string]bool) []store.LookupResult { - return s.lookupInUsingEntryForWithFollow(moduleName, functionName, which, consumerOpts, visited, s.followDelegates) + return s.lookupInUsingEntryForWithFollow(moduleName, functionName, which, consumerOpts, visited, s.followDelegates, -1) } -func (s *Server) lookupInUsingEntryForWithFollow(moduleName, functionName, which string, consumerOpts map[string]string, visited map[string]bool, followDelegates bool) []store.LookupResult { - visitKey := usingVisitKey(moduleName, which) +func (s *Server) lookupInUsingEntryForWithFollow(moduleName, functionName, which string, consumerOpts map[string]string, visited map[string]bool, followDelegates bool, arity int) []store.LookupResult { + visitKey := usingVisitKeyWithOpts(moduleName, which, consumerOpts) if visited[visitKey] { return nil } @@ -2708,18 +2752,34 @@ func (s *Server) lookupInUsingEntryForWithFollow(moduleName, functionName, which if defs, ok := body.inlineDefs[functionName]; ok { var results []store.LookupResult for _, d := range defs { - results = append(results, store.LookupResult{FilePath: entry.filePath, Line: d.line}) + if arity >= 0 && d.arity != arity { + continue + } + results = append(results, store.LookupResult{ + FilePath: entry.filePath, + Line: d.line, + Kind: d.kind, + Arity: d.arity, + }) + } + if len(results) > 0 { + return results } - return results } // Static imports for j := len(body.imports) - 1; j >= 0; j-- { var results []store.LookupResult var err error - if followDelegates { + if followDelegates && arity >= 0 { + results, err = s.store.LookupFollowDelegateByArity(body.imports[j], functionName, arity) + results = publicOnly(results) + } else if followDelegates { results, err = s.store.LookupFollowDelegate(body.imports[j], functionName) results = publicOnly(results) + } else if arity >= 0 { + results, err = s.store.LookupFunctionByArity(body.imports[j], functionName, arity) + results = publicOnly(results) } else { results, err = s.store.LookupPublicFunction(body.imports[j], functionName) } @@ -2742,9 +2802,15 @@ func (s *Server) lookupInUsingEntryForWithFollow(moduleName, functionName, which case "import": var results []store.LookupResult var err error - if followDelegates { + if followDelegates && arity >= 0 { + results, err = s.store.LookupFollowDelegateByArity(mod, functionName, arity) + results = publicOnly(results) + } else if followDelegates { results, err = s.store.LookupFollowDelegate(mod, functionName) results = publicOnly(results) + } else if arity >= 0 { + results, err = s.store.LookupFunctionByArity(mod, functionName, arity) + results = publicOnly(results) } else { results, err = s.store.LookupPublicFunction(mod, functionName) } @@ -2752,7 +2818,7 @@ func (s *Server) lookupInUsingEntryForWithFollow(moduleName, functionName, which return results } case "use": - if result := s.lookupInUsingEntryForWithFollow(mod, functionName, "", nil, visited, followDelegates); result != nil { + if result := s.lookupInUsingEntryForWithFollow(mod, functionName, "", nil, visited, followDelegates, arity); result != nil { return result } } @@ -2761,7 +2827,7 @@ func (s *Server) lookupInUsingEntryForWithFollow(moduleName, functionName, which // Transitive uses: use Module inside the __using__ body (double-use chains) for k := len(body.transCalls) - 1; k >= 0; k-- { call := body.transCalls[k] - if result := s.lookupInUsingEntryForWithFollow(call.Module, functionName, call.dispatchAtom(), call.Opts, visited, followDelegates); result != nil { + if result := s.lookupInUsingEntryForWithFollow(call.Module, functionName, call.dispatchAtom(), call.Opts, visited, followDelegates, arity); result != nil { return result } } @@ -2769,7 +2835,7 @@ func (s *Server) lookupInUsingEntryForWithFollow(moduleName, functionName, which if body.hasTransCall(body.transUses[k]) { continue } - if result := s.lookupInUsingEntryForWithFollow(body.transUses[k], functionName, "", nil, visited, followDelegates); result != nil { + if result := s.lookupInUsingEntryForWithFollow(body.transUses[k], functionName, "", nil, visited, followDelegates, arity); result != nil { return result } } diff --git a/internal/lsp/server_test.go b/internal/lsp/server_test.go index d70db75..04b706c 100644 --- a/internal/lsp/server_test.go +++ b/internal/lsp/server_test.go @@ -245,6 +245,9 @@ func TestServer_InitializationOptions(t *testing.T) { if server.debug { t.Error("debug should default to false") } + if server.definitionStyle != "all" { + t.Errorf("definitionStyle: got %q, want %q", server.definitionStyle, "all") + } }) // Claude Code plugin template substitution yields strings, not booleans. @@ -254,12 +257,13 @@ func TestServer_InitializationOptions(t *testing.T) { opts map[string]interface{} wantFollowDel bool wantDebug bool + wantStyle string }{ - {"bool true/false", map[string]interface{}{"followDelegates": false, "debug": true}, false, true}, - {"string true/false", map[string]interface{}{"followDelegates": "false", "debug": "true"}, false, true}, - {"string 1/0", map[string]interface{}{"followDelegates": "0", "debug": "1"}, false, true}, - {"empty string leaves default", map[string]interface{}{"followDelegates": "", "debug": ""}, true, false}, - {"unsupported type leaves default", map[string]interface{}{"followDelegates": 1, "debug": 0}, true, false}, + {"bool true/false", map[string]interface{}{"followDelegates": false, "debug": true, "definitionStyle": "first"}, false, true, "first"}, + {"string true/false", map[string]interface{}{"followDelegates": "false", "debug": "true", "definitionStyle": "all"}, false, true, "all"}, + {"string 1/0", map[string]interface{}{"followDelegates": "0", "debug": "1"}, false, true, "all"}, + {"empty string leaves default", map[string]interface{}{"followDelegates": "", "debug": "", "definitionStyle": ""}, true, false, "all"}, + {"unsupported values leave default", map[string]interface{}{"followDelegates": 1, "debug": 0, "definitionStyle": "bogus"}, true, false, "all"}, } for _, tc := range cases { @@ -280,10 +284,398 @@ func TestServer_InitializationOptions(t *testing.T) { if server.debug != tc.wantDebug { t.Errorf("debug: got %v, want %v", server.debug, tc.wantDebug) } + if server.definitionStyle != tc.wantStyle { + t.Errorf("definitionStyle: got %q, want %q", server.definitionStyle, tc.wantStyle) + } + }) + } +} + +func TestServer_ApplyDefinitionStyle(t *testing.T) { + server, cleanup := setupTestServer(t) + defer cleanup() + + locs := []protocol.Location{ + {URI: "file:///a.ex", Range: lineRange(0)}, + {URI: "file:///a.ex", Range: lineRange(5)}, + {URI: "file:///a.ex", Range: lineRange(9)}, + } + + // Default "all" returns everything + got := server.applyDefinitionStyle(locs) + if len(got) != 3 { + t.Errorf("expected 3 locations with style %q, got %d", "all", len(got)) + } + + // "first" returns only the first + server.definitionStyle = "first" + got = server.applyDefinitionStyle(locs) + if len(got) != 1 { + t.Errorf("expected 1 location with style %q, got %d", "first", len(got)) + } + if got[0].Range.Start.Line != 0 { + t.Errorf("expected first location (line 0), got line %d", got[0].Range.Start.Line) + } + + // Single location is unaffected by "first" + got = server.applyDefinitionStyle(locs[:1]) + if len(got) != 1 { + t.Errorf("expected 1 location with style %q and single input, got %d", "first", len(got)) + } + + // Empty slice is unaffected + got = server.applyDefinitionStyle(nil) + if len(got) != 0 { + t.Errorf("expected 0 locations for nil input, got %d", len(got)) + } +} + +// A call with a single matching definition returns exactly one location. +func TestDefinition_SingleDef_ReturnsOneLocation(t *testing.T) { + server, cleanup := setupTestServer(t) + defer cleanup() + + indexFile(t, server.store, server.projectRoot, "lib/math.ex", `defmodule MyApp.Math do + def add(a, b), do: a + b +end +`) + + callerPath := filepath.Join(server.projectRoot, "lib", "caller.ex") + callerContent := `defmodule MyApp.Caller do + def run, do: MyApp.Math.add(1, 2) +end +` + indexFile(t, server.store, server.projectRoot, "lib/caller.ex", callerContent) + callerURI := "file://" + callerPath + server.docs.Set(callerURI, callerContent) + + // Cursor on "add" in MyApp.Math.add(1, 2) + locs := definitionAt(t, server, callerURI, 1, 25) + if len(locs) != 1 { + t.Fatalf("expected exactly 1 location for single def, got %d: %+v", len(locs), locs) + } +} + +// A call returns only the definition matching its arity. +func TestDefinition_MultiArity_ReturnsOneLocation(t *testing.T) { + server, cleanup := setupTestServer(t) + defer cleanup() + + indexFile(t, server.store, server.projectRoot, "lib/math.ex", `defmodule MyApp.Math do + def square(x), do: x * x + + def square(x, factor), do: (x * x) * factor +end +`) + + callerPath := filepath.Join(server.projectRoot, "lib", "caller.ex") + callerContent := `defmodule MyApp.Caller do + def run, do: MyApp.Math.square(3) +end +` + indexFile(t, server.store, server.projectRoot, "lib/caller.ex", callerContent) + callerURI := "file://" + callerPath + server.docs.Set(callerURI, callerContent) + + // Cursor on "square" in MyApp.Math.square(3) + locs := definitionAt(t, server, callerURI, 1, 28) + if len(locs) != 1 { + t.Fatalf("expected 1 location for square/1, got %d: %+v", len(locs), locs) + } +} + +func TestDefinition_CurrentModuleBareCallUsesArityAndStyle(t *testing.T) { + t.Run("selects matching arity", func(t *testing.T) { + server, cleanup := setupTestServer(t) + defer cleanup() + + content := `defmodule MyApp.Current do + def calculate(value), do: value + def calculate(left, right), do: left + right + def run, do: calculate(1, 2) +end +` + path := filepath.Join(server.projectRoot, "lib", "current.ex") + indexFile(t, server.store, server.projectRoot, "lib/current.ex", content) + uri := "file://" + path + server.docs.Set(uri, content) + + locs := definitionAt(t, server, uri, 3, 17) + if len(locs) != 1 || locs[0].Range.Start.Line != 2 { + t.Fatalf("expected calculate/2 on line 2, got %+v", locs) + } + }) + + t.Run("returns all same-arity heads by default", func(t *testing.T) { + server, cleanup := setupTestServer(t) + defer cleanup() + + content := `defmodule MyApp.Current do + def calculate(:first), do: 1 + def calculate(:second), do: 2 + def run, do: calculate(:first) +end +` + path := filepath.Join(server.projectRoot, "lib", "current.ex") + indexFile(t, server.store, server.projectRoot, "lib/current.ex", content) + uri := "file://" + path + server.docs.Set(uri, content) + + locs := definitionAt(t, server, uri, 3, 17) + if len(locs) != 2 { + t.Fatalf("expected both calculate/1 heads, got %+v", locs) + } + }) + + t.Run("ignores declarations in nested and later modules", func(t *testing.T) { + server, cleanup := setupTestServer(t) + defer cleanup() + + content := `defmodule MyApp.Current do + def calculate, do: :current + def run, do: calculate() + + defmodule Inner do + def calculate, do: :inner + end +end + +defmodule MyApp.Other do + def calculate, do: :other +end +` + path := filepath.Join(server.projectRoot, "lib", "current.ex") + indexFile(t, server.store, server.projectRoot, "lib/current.ex", content) + uri := "file://" + path + server.docs.Set(uri, content) + + locs := definitionAt(t, server, uri, 2, 17) + if len(locs) != 1 || locs[0].Range.Start.Line != 1 { + t.Fatalf("expected only MyApp.Current.calculate/0 on line 1, got %+v", locs) + } + }) +} + +func TestDefinition_UseChainSelectsProviderByArity(t *testing.T) { + server, cleanup := setupTestServer(t) + defer cleanup() + + indexFile(t, server.store, server.projectRoot, "lib/one.ex", `defmodule SharedLib.One do + def execute(value), do: value +end +`) + indexFile(t, server.store, server.projectRoot, "lib/two.ex", `defmodule SharedLib.Two do + def execute(left, right), do: {left, right} +end +`) + indexFile(t, server.store, server.projectRoot, "lib/injector.ex", `defmodule SharedLib.Injector do + defmacro __using__(_opts) do + quote do + import SharedLib.One + import SharedLib.Two + end + end +end +`) + + callerPath := filepath.Join(server.projectRoot, "lib", "consumer.ex") + callerContent := `defmodule MyApp.Consumer do + use SharedLib.Injector + def run, do: execute(:value) +end +` + indexFile(t, server.store, server.projectRoot, "lib/consumer.ex", callerContent) + callerURI := "file://" + callerPath + server.docs.Set(callerURI, callerContent) + + locs := definitionAt(t, server, callerURI, 2, 16) + if len(locs) != 1 || !strings.HasSuffix(string(locs[0].URI), "/lib/one.ex") { + t.Fatalf("expected SharedLib.One.execute/1, got %+v", locs) + } +} + +func TestDefinition_UseChainSameInjectorDifferentOptsSelectsProviderByArity(t *testing.T) { + server, cleanup := setupTestServer(t) + defer cleanup() + + indexFile(t, server.store, server.projectRoot, "lib/one.ex", `defmodule SharedLib.One do + def execute(value), do: value +end +`) + indexFile(t, server.store, server.projectRoot, "lib/two.ex", `defmodule SharedLib.Two do + def execute(left, right), do: {left, right} +end +`) + indexFile(t, server.store, server.projectRoot, "lib/injector.ex", `defmodule SharedLib.Injector do + defmacro __using__(opts) do + provider = Keyword.get(opts, :provider, SharedLib.Two) + quote do + import unquote(provider) + end + end +end +`) + + callerPath := filepath.Join(server.projectRoot, "lib", "consumer.ex") + callerContent := `defmodule MyApp.Consumer do + use SharedLib.Injector, provider: SharedLib.One + use SharedLib.Injector, provider: SharedLib.Two + def run, do: execute(:value) +end +` + indexFile(t, server.store, server.projectRoot, "lib/consumer.ex", callerContent) + callerURI := "file://" + callerPath + server.docs.Set(callerURI, callerContent) + + locs := definitionAt(t, server, callerURI, 3, 16) + if len(locs) != 1 || !strings.HasSuffix(string(locs[0].URI), "/lib/one.ex") { + t.Fatalf("expected SharedLib.One.execute/1, got %+v", locs) + } +} + +func TestDefinition_KeywordTailCountsAsOneArgument(t *testing.T) { + server, cleanup := setupTestServer(t) + defer cleanup() + + indexFile(t, server.store, server.projectRoot, "lib/repo.ex", `defmodule SharedLib.Repo do + def insert(changeset, opts), do: {changeset, opts} + def insert(changeset, opts, metadata), do: {changeset, opts, metadata} +end +`) + + callerPath := filepath.Join(server.projectRoot, "lib", "caller.ex") + callerContent := `defmodule MyApp.Caller do + def run(changeset), do: SharedLib.Repo.insert(changeset, returning: true, on_conflict: :replace) +end +` + indexFile(t, server.store, server.projectRoot, "lib/caller.ex", callerContent) + callerURI := "file://" + callerPath + server.docs.Set(callerURI, callerContent) + + locs := definitionAt(t, server, callerURI, 1, 44) + if len(locs) != 1 { + t.Fatalf("expected exactly 1 location for insert/2, got %d: %+v", len(locs), locs) + } + if got := locs[0].Range.Start.Line; got != 1 { + t.Fatalf("expected keyword tail to resolve insert/2 on line 1, got line %d", got) + } +} + +// A defdelegate and a def sharing a name resolve to the one matching the call. +func TestDefinition_DelegateAndDefSameName_ReturnsOneLocation(t *testing.T) { + server, cleanup := setupTestServer(t) + defer cleanup() + + indexFile(t, server.store, server.projectRoot, "lib/worker.ex", `defmodule MyApp.Worker do + def call(attrs), do: {:ok, attrs} +end +`) + + indexFile(t, server.store, server.projectRoot, "lib/api.ex", `defmodule MyApp.Api do + defdelegate do_thing(x), to: MyApp.Worker, as: :call + + def do_thing(x, y), do: {x, y} +end +`) + + callerPath := filepath.Join(server.projectRoot, "lib", "caller.ex") + callerContent := `defmodule MyApp.Caller do + def run, do: MyApp.Api.do_thing("hello") +end +` + indexFile(t, server.store, server.projectRoot, "lib/caller.ex", callerContent) + callerURI := "file://" + callerPath + server.docs.Set(callerURI, callerContent) + + // Cursor on "do_thing" in MyApp.Api.do_thing("hello") + locs := definitionAt(t, server, callerURI, 1, 27) + if len(locs) != 1 { + t.Fatalf("expected 1 location for do_thing/1, got %d: %+v", len(locs), locs) + } +} + +// A delegate with default arguments calls its target at the declared full arity. +func TestDefinition_DefaultArgumentDelegateFollowsDeclaredArity(t *testing.T) { + for _, tc := range []struct { + name string + worker string + }{ + {"skips a target overload at the call arity", `defmodule SharedLib.Worker do + def run(x), do: x + def run(x, opts), do: {x, opts} +end +`}, + {"reaches a target with only the full arity", `defmodule SharedLib.Worker do + def run(x, opts), do: {x, opts} +end +`}, + } { + t.Run(tc.name, func(t *testing.T) { + server, cleanup := setupTestServer(t) + defer cleanup() + + indexFile(t, server.store, server.projectRoot, "lib/worker.ex", tc.worker) + indexFile(t, server.store, server.projectRoot, "lib/api.ex", `defmodule SharedLib.Api do + defdelegate run(x, opts \\ []), to: SharedLib.Worker +end +`) + callerPath := filepath.Join(server.projectRoot, "lib", "caller.ex") + callerContent := `defmodule MyApp.Caller do + def call, do: SharedLib.Api.run(:value) +end +` + indexFile(t, server.store, server.projectRoot, "lib/caller.ex", callerContent) + callerURI := "file://" + callerPath + server.docs.Set(callerURI, callerContent) + + workerLines, _ := server.store.LookupFunctionByArity("SharedLib.Worker", "run", 2) + if len(workerLines) != 1 { + t.Fatalf("expected one indexed SharedLib.Worker.run/2, got %+v", workerLines) + } + locs := definitionAt(t, server, callerURI, 1, 30) + if len(locs) != 1 || !strings.HasSuffix(string(locs[0].URI), "/lib/worker.ex") || + int(locs[0].Range.Start.Line) != workerLines[0].Line-1 { + t.Fatalf("expected SharedLib.Worker.run/2, got %+v", locs) + } }) } } +// Same-arity heads: "all" returns every head, "first" only the first. +func TestDefinition_MultipleHeadsSameArity_StyleControlled(t *testing.T) { + server, cleanup := setupTestServer(t) + defer cleanup() + + indexFile(t, server.store, server.projectRoot, "lib/accounts.ex", `defmodule MyApp.Accounts do + def fetch_user(%{id: id}), do: id + def fetch_user(email) when is_binary(email), do: email + def fetch_user(id) when is_integer(id), do: id +end +`) + + callerPath := filepath.Join(server.projectRoot, "lib", "caller.ex") + callerContent := `defmodule MyApp.Caller do + def run, do: MyApp.Accounts.fetch_user("nick@example.com") +end +` + indexFile(t, server.store, server.projectRoot, "lib/caller.ex", callerContent) + callerURI := "file://" + callerPath + server.docs.Set(callerURI, callerContent) + + // Cursor on "fetch_user" in MyApp.Accounts.fetch_user("...") + locs := definitionAt(t, server, callerURI, 1, 32) + if len(locs) < 2 { + t.Fatalf("expected multiple locations for 3 function heads with style=all, got %d", len(locs)) + } + + // With style="first", caller should see only the first head. + server.definitionStyle = "first" + locs = definitionAt(t, server, callerURI, 1, 32) + if len(locs) != 1 { + t.Fatalf("expected exactly 1 location with definitionStyle=first, got %d", len(locs)) + } +} + func definitionAt(t *testing.T, server *Server, uri string, line, col uint32) []protocol.Location { t.Helper() result, err := server.Definition(context.Background(), &protocol.DefinitionParams{ diff --git a/internal/parser/token_walk_test.go b/internal/parser/token_walk_test.go index 42db635..1f16039 100644 --- a/internal/parser/token_walk_test.go +++ b/internal/parser/token_walk_test.go @@ -35,6 +35,26 @@ func TestTrackBlockDepth(t *testing.T) { } } +func TestTokenWalker_PreviousSigPos(t *testing.T) { + source := []byte("left\n# comment\n|> right") + tokens := Tokenize(source) + w := NewTokenWalker(source, tokens) + right := -1 + for i, token := range tokens { + if TokenText(source, token) == "right" { + right = i + break + } + } + if right < 0 { + t.Fatal("right token not found") + } + prev := w.PreviousSigPos(right) + if prev < 0 || tokens[prev].Kind != TokPipe { + t.Fatalf("previous significant token = %d, want pipe", prev) + } +} + func TestAliasShortName(t *testing.T) { tests := []struct { in string diff --git a/internal/parser/token_walker.go b/internal/parser/token_walker.go index 1ab2538..ebf2a6b 100644 --- a/internal/parser/token_walker.go +++ b/internal/parser/token_walker.go @@ -124,6 +124,18 @@ func (w *TokenWalker) NextSigPos() int { return NextSigToken(w.Tokens, w.N, w.pos) } +// PreviousSigPos returns the previous significant token before before, or -1. +// EOL and comment tokens are skipped. +func (w *TokenWalker) PreviousSigPos(before int) int { + for i := before - 1; i >= 0; i-- { + kind := w.Tokens[i].Kind + if kind != TokEOL && kind != TokComment { + return i + } + } + return -1 +} + // Depth returns the current bracket depth. func (w *TokenWalker) Depth() int { return w.depth diff --git a/internal/store/store.go b/internal/store/store.go index ff6c71b..e026bfe 100644 --- a/internal/store/store.go +++ b/internal/store/store.go @@ -1456,6 +1456,21 @@ func (s *Store) LookupPublicFunction(module, function string) ([]LookupResult, e ) } +// LookupFunctionByArity returns definitions for a function name, optionally +// filtered by arity. When arity < 0, behavior matches LookupFunction (all +// arities returned). When arity >= 0, results are filtered to exact matches, +// preventing callers from getting e.g. `foo/1` and `foo/2` rows for a call +// that only uses one of them. +func (s *Store) LookupFunctionByArity(module, function string, arity int) ([]LookupResult, error) { + if arity < 0 { + return s.LookupFunction(module, function) + } + return s.queryLookup( + "SELECT f.path, d.line, d.kind, d.arity, d.delegate_to, d.delegate_as FROM definitions d JOIN files f ON f.id = d.file_id WHERE d.module = ? AND d.function = ? AND d.arity = ? AND d.kind NOT IN ('module', 'defprotocol', 'defimpl', 'callback', 'macrocallback') ORDER BY CASE WHEN d.kind IN ('type', 'opaque') THEN 1 ELSE 0 END, d.line", + module, function, arity, + ) +} + // CallbackResult holds a @callback or @macrocallback definition with its arity. type CallbackResult struct { FilePath string @@ -2014,42 +2029,92 @@ func (s *Store) NextFunctionLine(filePath string, startLine int) int { } func (s *Store) LookupFollowDelegate(module, function string) ([]LookupResult, error) { - return s.lookupFollowDelegate(module, function, 0) + return s.LookupFollowDelegateByArity(module, function, -1) +} + +// LookupFollowDelegateByArity is like LookupFollowDelegate but filters by +// arity when arity >= 0. Delegate-following is partitioned per arity, so a +// module that mixes `defdelegate foo/1, to: X` with `def foo/2` still follows +// the /1 delegate for a /1 call without being blocked by the non-delegate /2 +// row. +func (s *Store) LookupFollowDelegateByArity(module, function string, arity int) ([]LookupResult, error) { + return s.lookupFollowDelegate(module, function, arity, 0) } -func (s *Store) lookupFollowDelegate(module, function string, depth int) ([]LookupResult, error) { +// declaredDelegateArity is the full arity of the declaration delegate belongs +// to. Default arguments index one row per arity, but each calls the target +// with every argument. +func declaredDelegateArity(rows []LookupResult, delegate LookupResult) int { + arity := delegate.Arity + for _, r := range rows { + if r.Kind == "defdelegate" && r.FilePath == delegate.FilePath && r.Line == delegate.Line && r.Arity > arity { + arity = r.Arity + } + } + return arity +} + +func (s *Store) lookupFollowDelegate(module, function string, arity, depth int) ([]LookupResult, error) { if depth > 5 { return nil, nil } - results, err := s.LookupFunction(module, function) + // Every arity is read so a delegate's declared arity is known: a + // `defdelegate run(x, opts \\ [])` row for run/1 still calls run/2. + all, err := s.LookupFunction(module, function) if err != nil { return nil, err } + results := all + if arity >= 0 { + results = make([]LookupResult, 0, len(all)) + for _, r := range all { + if r.Arity == arity { + results = append(results, r) + } + } + } + if len(results) == 0 { + return nil, nil + } - // If all results are defdelegates, follow them to the target - allDelegates := len(results) > 0 + // Group by arity. Within each arity group, follow if every row is a + // delegate. This handles the defdelegate/1 + def/2 mix correctly. + byArity := make(map[int][]LookupResult, len(results)) + order := make([]int, 0, len(results)) for _, r := range results { - if r.Kind != "defdelegate" || r.DelegateTo == "" { - allDelegates = false - break + if _, seen := byArity[r.Arity]; !seen { + order = append(order, r.Arity) } + byArity[r.Arity] = append(byArity[r.Arity], r) } - if allDelegates { - targetModule := results[0].DelegateTo - targetFunc := function - if results[0].DelegateAs != "" { - targetFunc = results[0].DelegateAs - } - targetResults, err := s.lookupFollowDelegate(targetModule, targetFunc, depth+1) - if err != nil { - return nil, err + var out []LookupResult + for _, a := range order { + group := byArity[a] + allDelegates := true + for _, r := range group { + if r.Kind != "defdelegate" || r.DelegateTo == "" { + allDelegates = false + break + } } - if len(targetResults) > 0 { - return targetResults, nil + if allDelegates { + targetModule := group[0].DelegateTo + targetFunc := function + if group[0].DelegateAs != "" { + targetFunc = group[0].DelegateAs + } + targetResults, err := s.lookupFollowDelegate(targetModule, targetFunc, declaredDelegateArity(all, group[0]), depth+1) + if err != nil { + return nil, err + } + if len(targetResults) > 0 { + out = append(out, targetResults...) + continue + } } + out = append(out, group...) } - - return results, nil + return out, nil }