diff --git a/CLAUDE.md b/CLAUDE.md index 6df4a9c..2d338b1 100644 --- a/CLAUDE.md +++ b/CLAUDE.md @@ -240,9 +240,11 @@ Every other text field -- `class_names`, `method_names`, `base_types`, `field_ty | `imports` | -- | `using` / `import` / `include` directives | all except SQL | | `params` | METHOD | Parameter list for METHOD | C#, Python, JS, Rust, C++ | | `declarations` | NAME | The declaration(s) of NAME (narrow with `symbol_kind`) | all | -| `body` | NAME | Full source of NAME's declaration | C# only | +| `body` | NAME | Full source of NAME's declaration(s). Works in both `query_codebase` (returns bodies across every matching file) and `query_single_file` (one file only). Narrow with `symbol_kind`. | C# only | | `at` | LINE:COL | Deepest AST node at position + enclosing scope chain | C# only | | `calls` | METHOD | Call sites of METHOD. Qualify with `Type.Method` to restrict by receiver -- the qualifier matches both the literal receiver text (`Foo.Save()`) **and** any receiver whose declared/inferred type resolves to that name via the method-scoped var-type map (`store.Save()` where `store: IRepository` matches `IRepository.Save`). When the receiver's type is conflicted in its scope, the qualified match is skipped -- the bare name still finds the call. Pass a METHOD name only -- a variable/receiver name silently returns empty; use `all_refs` on the variable for that. | all | +| `caller_of` | METHOD | Like `calls`, but groups call sites by the **enclosing caller** -- one row per `(TypeName.MemberName)` caller with a count of how many sites it contains. Collapses noisy `calls METHOD` output into a unique-caller view. Useful for "who depends on this". | C# only | +| `callee_of` | METHOD | The inverse -- walk the body of the method named METHOD and emit one row per distinct callee with an invocation count. Constructor calls (`new T()`) are reported as `T (N invocations, ctor)`. Useful for "what does this method depend on" / "what could be slow here". | C# only | | `implements` | TYPE | Types that inherit/implement TYPE | all except SQL | | `uses` | TYPE | Type references; narrow with `uses_kind` (`field`/`param`/`return`/`cast`/`base`/`locals`) | C# only | | `casts` | TYPE | `(TYPE)expr` / `as TYPE` sites | C# only | @@ -250,7 +252,9 @@ Every other text field -- `class_names`, `method_names`, `base_types`, `field_ty | `accesses_of` | MEMBER | Access sites of property/field by name (`"Order.Status"` restricts) | C# only | | `accesses_on` | TYPE | `.Member` accesses on locals/params/fields typed as TYPE (plus `new T { ... }` and `with` mutations). Returns nothing when the variable is only assigned, returned, or forwarded as an argument -- no `.Member` exists. Fall back to `all_refs` on the variable name. | C# only | | `all_refs` | NAME | Every identifier occurrence (broadest -- AST-only, skips strings/comments). For SQL this is a plain substring scan over lines. | all | -| `var_type` | NAME | For each occurrence of NAME, report the resolved type from the method-scoped var-type map, or `(unresolved)` / `(conflicting)` when the resolver can't pin it down. Saves an `at LINE:COL` round-trip when you just want the type. | C# only | +| `var_type` | NAME | For each occurrence of NAME, report the resolved type from the method-scoped var-type map, or `(unresolved)` / `(conflicting)` when the resolver can't pin it down. Works in both `query_codebase` (every file that mentions NAME) and `query_single_file`. Saves an `at LINE:COL` round-trip when you just want the type. | C# only | + +**Enclosing-scope filter (pattern modes).** `calls`, `uses`, `casts`, `accesses_of`, `accesses_on`, `all_refs` accept `enclosing_method="WriteBack"` and/or `enclosing_class="OrderProcessor"`. The two compose as a logical AND. Useful for pinpointing call sites in a specific member, e.g. `calls("Save", enclosing_method="WriteBack")` returns only the `Save()` calls that happen inside `WriteBack` methods. (C# only.) **Visibility filter (declaration modes).** `declarations`, `classes`, `methods`, `fields` accept `visibility="public,internal,protected,private"` (comma-separated, any subset). The filter is applied AST-side per declaration using the same defaults the indexer uses: top-level types default to `internal`, nested types to `private`, interface members to `public`, enum members to `public`, class/struct/record members to `private`. Compound modifiers collapse to their dominant role (`protected internal` -> `protected`, `private protected` -> `private`). Languages other than C# currently don't capture visibility -- passing the filter against them returns nothing rather than over-matching. diff --git a/mcp_server.py b/mcp_server.py index 4ebbaca..6d0ddbd 100644 --- a/mcp_server.py +++ b/mcp_server.py @@ -184,6 +184,9 @@ def query_codebase( symbol_kind: str = "", uses_kind: str = "", visibility: str = "", + head_lines: int = 0, + enclosing_method: str = "", + enclosing_class: str = "", exclude_path: str = "", ) -> str: """Index pre-filter + tree-sitter AST. Returns one of three response shapes @@ -273,6 +276,19 @@ def query_codebase( nothing when this filter is set. (C# captures explicit modifiers plus interface-public / enum-public / nested type defaults.) + head_lines: For `body` and `declarations include_body=True` -- truncate + each emitted body to the first N source lines (signature + + body together), with a `... +K more lines` tail marker. + Default 0 = no truncation. Useful when scanning many bodies + at once. + enclosing_method: + For pattern modes (calls / uses / casts / accesses_of / + accesses_on / all_refs) -- only keep hits inside a member + with this exact name. Combine with `enclosing_class=` to + pinpoint a specific call-site context. (C# only.) + enclosing_class: + Same as above, narrowed to a type by name. Composes with + `enclosing_method=` as a logical AND. (C# only.) exclude_path: Comma-separated list of folder paths to exclude from results. Each value is matched as an exact ancestor folder, not a glob -- wildcards are not supported. Behavior: @@ -290,21 +306,25 @@ def query_codebase( query_codebase("uses", "IDataStore", uses_kind="param", sub="services") query_codebase("implements", "IRepository") query_codebase("declarations", "SaveChanges", symbol_kind="method") + query_codebase("body", "SaveChanges", symbol_kind="method") # full source of every match + query_codebase("var_type", "store", sub="services") # resolved type of every ``store`` occurrence query_codebase("calls", "SaveChanges", sub="services,vendor") query_codebase("calls", "SaveChanges", exclude_path="tests,generated") query_codebase("uses", "IRepo", sub="services", exclude_path="services/legacy")""" # File-targeted modes don't make sense for a codebase-wide search. - # `body`/`at`/`params`/`var_type` need an explicit file; listing modes + # `at`/`params` need an explicit file; listing modes # (methods/fields/classes/imports/capabilities) only describe a single # file's structure. Catch them here so the agent sees an actionable # redirect instead of the daemon's generic "unknown mode" error. + # ``body`` and ``var_type`` work codebase-wide now (index pre-filter + # + per-file AST), so they're no longer in this set. _FILE_ONLY = { "methods", "fields", "classes", "imports", "capabilities", - "body", "at", "params", "var_type", + "at", "params", } m = mode.lower().strip().replace("-", "_") if m in _FILE_ONLY: - if m in ("body", "at", "params", "var_type"): + if m in ("at", "params"): why = "needs a specific file" example_arg = pattern or ('"SaveChanges"' if m != "at" else '"42:10"') example = f' query_single_file("{m}", {example_arg}, file="path/to/File.cs")' @@ -321,6 +341,9 @@ def query_codebase( "include_body": include_body, "symbol_kind": symbol_kind or "", "uses_kind": uses_kind or "", "visibility": visibility or "", + "head_lines": int(head_lines) if head_lines else 0, + "enclosing_method": enclosing_method or "", + "enclosing_class": enclosing_class or "", "exclude_path": exclude_path or "", }) except Exception as e: @@ -429,6 +452,12 @@ def _qsf_call(file_rel: str) -> str: args.append(f'uses_kind="{uses_kind}"') if visibility: args.append(f'visibility="{visibility}"') + if head_lines: + args.append(f'head_lines={int(head_lines)}') + if enclosing_method: + args.append(f'enclosing_method="{enclosing_method}"') + if enclosing_class: + args.append(f'enclosing_class="{enclosing_class}"') return "query_single_file(" + ", ".join(args) + ")" # Tier 2 -- many files: filenames + counts only. @@ -459,11 +488,19 @@ def _qsf_call(file_rel: str) -> str: output = "\n".join(out_lines) if truncated_files: - notes = [f"\n\n{len(truncated_files)} file(s) had more than " - f"{_PER_FILE_DETAIL_LINES} hits -- showing first " - f"{_PER_FILE_DETAIL_LINES} of each. To see all hits in a file:"] + # Compact form: one short line per capped file showing how many + # extra hits are available. The agent already knows the mode and + # pattern; spelling out the full ``query_single_file(...)`` call + # for every capped file would burn ~100 chars per file in pure + # boilerplate. ``[+K capped] path`` is enough -- the agent can + # reissue a query_single_file call when it wants the rest. + notes = [f"\n\n{len(truncated_files)} file(s) capped at " + f"{_PER_FILE_DETAIL_LINES} hits each. Issue " + f"query_single_file(\"{m}\", \"{pattern}\", file=...) " + f"to see all hits."] for rel, total in truncated_files: - notes.append(f" {_qsf_call(rel)} # {total} total hits") + extra = total - _PER_FILE_DETAIL_LINES + notes.append(f" [+{extra} capped] {rel}") output += "\n".join(notes) output, truncated = _truncate(output) @@ -485,6 +522,9 @@ def query_single_file( symbol_kind: str = "", uses_kind: str = "", visibility: str = "", + head_lines: int = 0, + enclosing_method: str = "", + enclosing_class: str = "", head_limit: int = 250, offset: int = 0, ) -> str: @@ -515,6 +555,14 @@ def query_single_file( calls METHOD Call sites of METHOD. Pass a METHOD name, not a receiver -- `obj.Foo()` is matched by calls("Foo") not calls("obj"); for variable usage use all_refs. + caller_of METHOD Like ``calls``, but groups results by the enclosing + caller -- one row per ``(TypeName.MemberName)`` + with a count of how many call sites it contains. + (C# only.) + callee_of METHOD The inverse: walk the body of the method named + METHOD and emit one row per distinct callee with + invocation counts. Constructor calls show as + ``T (N invocations, ctor)``. (C# only.) implements TYPE Types that inherit/implement TYPE. uses TYPE Type references. Omit `uses_kind` (or "all") for the union of every role; narrow with `uses_kind` @@ -570,6 +618,14 @@ def query_single_file( language's defaults (interface members => public, nested types => private, top-level types => internal); other languages currently return nothing when this filter is set. + head_lines: For `body` and `declarations include_body=True` -- truncate + each body to the first N source lines (signature + body + together) with a `... +K more lines` tail marker. + Default 0 = no truncation. + enclosing_method / enclosing_class: + For pattern modes -- restrict hits to those inside a + member / type with the given name. Composes with both + filters as a logical AND. (C# only.) head_limit: Max results to return (default 250). Use with offset to page. offset: Skip first N results before applying head_limit (default 0). @@ -613,6 +669,9 @@ def query_single_file( symbol_kind=symbol_kind or None, uses_kind=uses_kind or None, visibility=visibility or None, + head_lines=int(head_lines) if head_lines else None, + enclosing_method=enclosing_method or None, + enclosing_class=enclosing_class or None, ) except ValueError as e: # Unknown extension or unsupported mode -- propagate the helpful diff --git a/query/cs.py b/query/cs.py index 372492c..760518d 100644 --- a/query/cs.py +++ b/query/cs.py @@ -227,6 +227,78 @@ def _cls_prefix(node, src) -> str: return f"[in {cls}] " if cls else "" +def _member_decl_name_node(member_node): + """Return the name node for a member declaration, or ``None``. + + Most member kinds expose their name via ``child_by_field_name("name")``; + ``field_declaration`` and ``event_field_declaration`` carry the name in + a nested ``variable_declarator`` and need the unwrap. When multiple + declarators share one ``int a, b, c`` field, the first is returned -- + the position-aware variant (``_field_declarator_name``) lives below + and is used by ``q_at``. + """ + nm = member_node.child_by_field_name("name") + if nm is not None: + return nm + if member_node.type in ("field_declaration", "event_field_declaration"): + vd = next((c for c in member_node.children + if c.type == "variable_declaration"), None) + if vd is not None: + decl = next((c for c in vd.children + if c.type == "variable_declarator"), None) + if decl is not None: + return decl.child_by_field_name("name") + return None + + +def _enclosing_member_name(node, src) -> str: + """Walk up to find the enclosing member declaration and return its name. + + Returns ``""`` when the node is at type-level (e.g. a static initialiser + sitting outside any member) or at namespace level. + """ + p = node.parent + while p is not None: + if p.type in _MEMBER_DECL_NODES: + name_node = _member_decl_name_node(p) + return _text(name_node, src).strip() if name_node is not None else "" + p = p.parent + return "" + + +def _scope_prefix(node, src) -> str: + """Return ``'[in TypeName.MemberName] '`` for a node inside a member, + ``'[in TypeName] '`` for a node at type-level, or ``''`` at namespace + level. Used to enrich AST hits in pattern modes (``calls`` / ``uses`` / + ``accesses_of`` / ``accesses_on`` / ``casts``) with the enclosing scope + so the agent doesn't have to issue a follow-up ``at LINE:COL`` query. + """ + type_name = _enclosing_type_name(node, src) + member_name = _enclosing_member_name(node, src) + if type_name and member_name: + return f"[in {type_name}.{member_name}] " + if type_name: + return f"[in {type_name}] " + return "" + + +def _passes_enclosing_filter(node, src, + enclosing_method: str | None, + enclosing_class: str | None) -> bool: + """True if ``node`` lives inside a member named ``enclosing_method`` and + a type named ``enclosing_class``. Either filter can be ``None`` / empty + to skip that dimension. Used to narrow pattern-mode hits to a specific + call-site context, e.g. ``calls("Save", enclosing_method="WriteBack")``. + """ + if enclosing_method: + if _enclosing_member_name(node, src) != enclosing_method: + return False + if enclosing_class: + if _enclosing_type_name(node, src) != enclosing_class: + return False + return True + + _LITERAL_NODES = { "comment", "string_literal", "verbatim_string_literal", "interpolated_string_expression", "character_literal", @@ -1340,28 +1412,33 @@ def q_fields(src, tree, lines, visibility=None): if _visibility_keep(getattr(_r, "visibility", ""), allowed)] -def q_calls(src, tree, lines, method_name): +def _iter_call_sites(tree, src, method_name): + """Yield ``(invocation_node, name_node)`` for every call site of + ``method_name`` in the tree. + + ``method_name`` may be qualified as ``Type.Method`` -- the qualifier + matches the receiver either by literal text or by resolved type via + the per-method var-type map. When no qualifier is supplied, + ``new T(...)`` object-creation expressions for a type named + ``method_name`` are also yielded. + + Shared by ``q_calls`` and ``q_caller_of`` so both modes have + identical matching semantics. + """ if "." in method_name: qualifier, bare_name = method_name.rsplit(".", 1) else: qualifier, bare_name = None, method_name - # Build the var-type map lazily -- only when a qualified pattern is - # supplied. Bare-method searches don't need receiver resolution. var_map = _build_var_type_map(tree, src) if qualifier else None + _qualifier_str = qualifier or "" - _qualifier_str: str = qualifier or "" - - def _qualifier_matches(expr_node) -> bool: - """True if ``expr_node`` (the receiver of a member access) matches - ``qualifier`` either by literal text or by resolved type.""" + def _qualifier_matches(expr_node): if expr_node is None or not _qualifier_str: return False expr_txt = _text(expr_node, src).strip() if expr_txt == _qualifier_str or expr_txt.endswith("." + _qualifier_str): return True - # Resolved-type fallback: look up an identifier receiver in the - # method-scoped var-type map and compare its unqualified type name. if expr_node.type == "identifier" and var_map is not None: rt = var_map.resolve_at(expr_txt, expr_node) if isinstance(rt, str) and rt: @@ -1370,27 +1447,6 @@ def _qualifier_matches(expr_node) -> bool: return True return False - def _report(node, name_node): - """Build a result tuple pinpointing the method-name token. - - Reports the line of ``name_node`` (where the matched identifier - actually appears) rather than the start of the surrounding - invocation -- for chained calls like ``a.B().Method(...)`` that - means the result lands on the line of ``Method``, not on - whichever line the chain begins. Source text is the single line - containing the name, not the multi-line invocation node. - """ - anchor = name_node if name_node is not None else node - row = anchor.start_point[0] - line_text = lines[row].strip() if 0 <= row < len(lines) else "" - if not line_text: - # Fall back to a truncated render of the node when the source - # row is empty/missing (defensive -- shouldn't happen for real - # call sites). - line_text = _truncate_raw(node, src) - return (_line(anchor), line_text) - - results = [] for node in _find_all(tree.root_node, lambda n: n.type == "invocation_expression"): if _in_literal(node): continue @@ -1429,8 +1485,8 @@ def _report(node, name_node): if nn: matched = _strip_generic(_text(nn, src)) match_name_node = nn - if matched == bare_name: - results.append(_report(node, match_name_node)) + if matched == bare_name and match_name_node is not None: + yield node, match_name_node if qualifier is None: for node in _find_all(tree.root_node, lambda n: n.type == "object_creation_expression"): @@ -1444,11 +1500,153 @@ def _report(node, name_node): continue if _strip_generic(_text(idents[-1], src)) == bare_name: # ``new T(...)`` -- anchor at the type name token. - results.append(_report(node, idents[-1])) + yield node, idents[-1] + + +def q_calls(src, tree, lines, method_name, + enclosing_method=None, enclosing_class=None): + """Reports the line of the matched-name token (e.g. ``Method`` in + ``a.B().Method(...)`` lands on ``Method``'s line, not the start of + the chain). Source text is the single line containing that token, + prefixed with the enclosing scope chain so agents don't need a + follow-up ``at LINE:COL`` query. + """ + results = [] + for node, name_node in _iter_call_sites(tree, src, method_name): + if not _passes_enclosing_filter(node, src, + enclosing_method, enclosing_class): + continue + row = name_node.start_point[0] + line_text = lines[row].strip() if 0 <= row < len(lines) else "" + if not line_text: + # Defensive fallback for empty/missing rows -- shouldn't + # happen for real call sites. + line_text = _truncate_raw(node, src) + results.append((_line(name_node), + f"{_scope_prefix(node, src)}{line_text}")) + return results + + +def q_caller_of(src, tree, lines, method_name): + """Group ``calls`` hits by the enclosing method. + + Where ``calls METHOD`` returns one row per call site (potentially many + inside the same caller), ``caller_of METHOD`` collapses those into one + row per (TypeName.MemberName) caller, with a count of how many call + sites that caller contains. Useful for impact analysis ("who depends + on this") without the per-site noise. + + Result text per caller: ``[in TypeName.MemberName] (N call sites)``. + Rows are emitted at the line of the *first* call site so the caller + line ordering matches source ordering. + """ + # Walk the AST directly via the shared call-site iterator -- the + # caller key comes straight from _scope_prefix's components, no + # text-format coupling to q_calls' output. + by_caller: dict[str, list[int]] = {} + for node, name_node in _iter_call_sites(tree, src, method_name): + type_name = _enclosing_type_name(node, src) + member_name = _enclosing_member_name(node, src) + if type_name and member_name: + caller = f"{type_name}.{member_name}" + elif type_name: + caller = type_name + else: + caller = "" + by_caller.setdefault(caller, []).append(_line(name_node)) + results = [] + for caller, line_nos in by_caller.items(): + first_line = min(line_nos) + n = len(line_nos) + plural = "site" if n == 1 else "sites" + results.append((first_line, f"[in {caller}] ({n} call {plural})")) + results.sort(key=lambda x: x[0]) + return results + + +def q_callee_of(src, tree, lines, method_name): + """List every callee invoked inside the method named ``method_name``. + + The inverse of ``caller_of``: given a method, walk its body and return + one row per distinct callee name, with a count of how many times that + callee is invoked inside this method. Useful for "what does this + method depend on" / "what could be slow here" analysis without having + to read the full body. + + Result text per callee: ``Callee (N invocations)``. Constructor calls + (``new T()``) are reported as ``T (N invocations, ctor)``. + """ + # Find every method-shaped declaration with the given name. + target_types = {"method_declaration", "constructor_declaration", + "local_function_statement"} + matches: list = [] + for node in _find_all(tree.root_node, lambda n: n.type in target_types): + nm = node.child_by_field_name("name") + if nm and _text(nm, src).strip() == method_name: + matches.append(node) + if not matches: + return [] + + results = [] + for method_node in matches: + body = method_node.child_by_field_name("body") + if body is None: + continue + counts: dict[tuple[str, bool], int] = {} + first_seen: dict[tuple[str, bool], int] = {} + for inv in _find_all(body, lambda n: n.type == "invocation_expression"): + if _in_literal(inv): + continue + fn = inv.child_by_field_name("function") + if fn is None: + continue + callee_name = "" + if fn.type == "member_access_expression": + nm = fn.child_by_field_name("name") + if nm: + callee_name = _strip_generic(_text(nm, src)) + elif fn.type in ("identifier", "generic_name"): + callee_name = _strip_generic(_text(fn, src)) + elif fn.type == "conditional_access_expression": + binding = next((c for c in fn.children + if c.type == "member_binding_expression"), None) + if binding: + nm = binding.child_by_field_name("name") + if nm: + callee_name = _strip_generic(_text(nm, src)) + if not callee_name: + continue + key = (callee_name, False) + counts[key] = counts.get(key, 0) + 1 + first_seen.setdefault(key, _line(inv)) + # Constructor calls -- ``new T(...)`` -- treated as callees of T. + for ctor in _find_all(body, lambda n: n.type == "object_creation_expression"): + if _in_literal(ctor): + continue + tn = ctor.child_by_field_name("type") + if tn is None: + continue + idents = _find_all(tn, lambda n: n.type == "identifier") + if not idents: + continue + callee_name = _strip_generic(_text(idents[-1], src)) + key = (callee_name, True) + counts[key] = counts.get(key, 0) + 1 + first_seen.setdefault(key, _line(ctor)) + # Per-method anchor: emit each callee at its first-invocation line. + for (callee, is_ctor), n in counts.items(): + anchor = first_seen[(callee, is_ctor)] + plural = "invocation" if n == 1 else "invocations" + suffix = ", ctor" if is_ctor else "" + results.append( + (anchor, + f"[in {method_name}] {callee} ({n} {plural}{suffix})")) + results.sort(key=lambda x: x[0]) return results -def q_accesses_of(src, tree, lines, member_name): +def q_accesses_of(src, tree, lines, member_name, + enclosing_method=None, enclosing_class=None): if "." in member_name: qualifier, bare_name = member_name.rsplit(".", 1) else: @@ -1457,6 +1655,11 @@ def q_accesses_of(src, tree, lines, member_name): results = [] seen_rows = set() + def _emit(containing_node, text): + prefix = _scope_prefix(containing_node, src) + results.append((_line(containing_node), + f"{prefix}{text}" if prefix else text)) + def _check_access(member_node, expr_node, containing_node): if not member_node: return @@ -1466,11 +1669,14 @@ def _check_access(member_node, expr_node, containing_node): expr_txt = _text(expr_node, src).strip() if expr_node else "" if not (expr_txt == qualifier or expr_txt.endswith("." + qualifier)): return + if not _passes_enclosing_filter(containing_node, src, + enclosing_method, enclosing_class): + return row = containing_node.start_point[0] if row in seen_rows: return seen_rows.add(row) - results.append((_line(containing_node), _truncate_raw(containing_node, src))) + _emit(containing_node, _truncate_raw(containing_node, src)) for node in _find_all(tree.root_node, lambda n: n.type == "member_access_expression"): if _in_literal(node): @@ -1494,11 +1700,14 @@ def _check_access(member_node, expr_node, containing_node): obj_type = _unqualify(_text(type_node, src).strip()) if type_node else None if qualifier and obj_type != qualifier: continue + if not _passes_enclosing_filter(assign, src, + enclosing_method, enclosing_class): + continue row = assign.start_point[0] if row in seen_rows: continue seen_rows.add(row) - results.append((_line(assign), _truncate_raw(assign, src))) + _emit(assign, _truncate_raw(assign, src)) # With-expression member mutations -- w with { Value = 10 } for wi, src_ident, prop in _iter_with_members(tree, src): @@ -1506,11 +1715,14 @@ def _check_access(member_node, expr_node, containing_node): continue if qualifier and _text(src_ident, src).strip() != qualifier: continue + if not _passes_enclosing_filter(wi, src, + enclosing_method, enclosing_class): + continue row = wi.start_point[0] if row in seen_rows: continue seen_rows.add(row) - results.append((_line(wi), _truncate_raw(wi, src))) + _emit(wi, _truncate_raw(wi, src)) results.sort(key=lambda x: x[0]) return results @@ -1534,7 +1746,8 @@ def q_implements(src, tree, lines, type_name): return results -def _q_uses_all(src, tree, lines, type_name): +def _q_uses_all(src, tree, lines, type_name, + enclosing_method=None, enclosing_class=None): results = [] seen_rows = set() @@ -1568,6 +1781,9 @@ def _is_invocation_target(node): continue if _is_invocation_target(node): continue + if not _passes_enclosing_filter(node, src, + enclosing_method, enclosing_class): + continue row = node.start_point[0] if row in seen_rows: continue @@ -1586,10 +1802,17 @@ def q_usings(src, tree, lines): def q_declarations(src, tree, lines, name, include_body=False, symbol_kind=None, - visibility=None): + visibility=None, head_lines=None): kind_nodes = SYMBOL_KIND_TO_NODES.get((symbol_kind or "").lower().strip()) target_nodes = kind_nodes if kind_nodes is not None else (_TYPE_DECL_NODES | _MEMBER_DECL_NODES) allowed = _parse_visibility_filter(visibility) + # head_lines normalisation: None / 0 / negative means "no truncation". + try: + head_n = int(head_lines) if head_lines is not None else None + except (TypeError, ValueError): + head_n = None + if head_n is not None and head_n <= 0: + head_n = None results = [] for node in _find_all(tree.root_node, lambda n: n.type in target_nodes): name_node = node.child_by_field_name("name") @@ -1611,9 +1834,19 @@ def q_declarations(src, tree, lines, name, include_body=False, symbol_kind=None, body_node = node.child_by_field_name("body") if body_node and not include_body: sig_end_row = body_node.start_point[0] - content = "\n".join(lines[start_row:sig_end_row]).rstrip() + content_lines = lines[start_row:sig_end_row] + content = "\n".join(content_lines).rstrip() else: - content = "\n".join(lines[start_row:end_row + 1]) + content_lines = lines[start_row:end_row + 1] + content = "\n".join(content_lines) + # head_lines truncation: clip the content (signature + body together) + # to the first N source lines, appending a "... +K more lines" tail + # marker so the agent knows there's more to fetch if it wants. + if head_n is not None and head_n < len(content_lines): + kept = content_lines[:head_n] + remaining = len(content_lines) - head_n + content = "\n".join(kept).rstrip() + content = f"{content}\n... +{remaining} more lines" header = f"[{kind}] {name} {start_row + 1}-{end_row + 1}:" results.append((_line(node), f"{header}\n{content}")) return results @@ -1740,15 +1973,21 @@ def _contains(n) -> bool: return [(out_line, text)] -def q_body(src, tree, lines, name, symbol_kind=None): +def q_body(src, tree, lines, name, symbol_kind=None, head_lines=None): """Return the full source of every member declaration named ``name``. Sugar for ``q_declarations(..., include_body=True)`` with a single-name intent: the agent says "give me the source of SaveChanges" and gets the whole method block (or every overload, if more than one matches). + + ``head_lines`` truncates each body to the first N source lines (header + line excluded; truncation indicator appended). Useful when scanning + many bodies at once and the signature plus the first few lines is + enough. """ return q_declarations(src, tree, lines, name, - include_body=True, symbol_kind=symbol_kind) + include_body=True, symbol_kind=symbol_kind, + head_lines=head_lines) def q_params(src, tree, lines, method_name): @@ -1784,12 +2023,16 @@ def q_params(src, tree, lines, method_name): return results -def _q_field_type(src, tree, lines, type_name): +def _q_field_type(src, tree, lines, type_name, + enclosing_method=None, enclosing_class=None): results = [] for node in _find_all(tree.root_node, lambda n: n.type in ("field_declaration", "event_field_declaration", "property_declaration")): + if not _passes_enclosing_filter(node, src, + enclosing_method, enclosing_class): + continue if node.type in ("field_declaration", "event_field_declaration"): type_txt = _field_type(node, src) if type_name not in _type_names(type_txt): @@ -1815,7 +2058,8 @@ def _q_field_type(src, tree, lines, type_name): return results -def _q_param_type(src, tree, lines, type_name): +def _q_param_type(src, tree, lines, type_name, + enclosing_method=None, enclosing_class=None): results = [] for mnode in _find_all(tree.root_node, lambda n: n.type in ("method_declaration", "constructor_declaration", @@ -1824,6 +2068,9 @@ def _q_param_type(src, tree, lines, type_name): params_node = mnode.child_by_field_name("parameters") if not params_node: continue + if not _passes_enclosing_filter(mnode, src, + enclosing_method, enclosing_class): + continue name_node = mnode.child_by_field_name("name") mname = _text(name_node, src).strip() if name_node else "" kind = mnode.type.replace("_declaration", "").replace("statement", "").replace("_", " ").strip() @@ -1843,13 +2090,17 @@ def _q_param_type(src, tree, lines, type_name): return results -def _q_return_type(src, tree, lines, type_name): +def _q_return_type(src, tree, lines, type_name, + enclosing_method=None, enclosing_class=None): results = [] for node in _find_all(tree.root_node, lambda n: n.type in ("method_declaration", "constructor_declaration", "local_function_statement", "delegate_declaration")): + if not _passes_enclosing_filter(node, src, + enclosing_method, enclosing_class): + continue # method_declaration exposes its return type via "returns"; all other # supported node types use "type". if node.type == "method_declaration": @@ -1875,21 +2126,28 @@ def _q_return_type(src, tree, lines, type_name): return results -def _q_local_type(src, tree, lines, type_name): +def _q_local_type(src, tree, lines, type_name, + enclosing_method=None, enclosing_class=None): results = [ - (_line(node), f"[local] {type_txt} {var_txt} {_cls_prefix(node, src)}") + (_line(node), f"{_scope_prefix(node, src)}[local] {type_txt} {var_txt}") for node, type_txt, var_txt in _iter_all_locals(tree, src, type_name) + if _passes_enclosing_filter(node, src, + enclosing_method, enclosing_class) ] results.sort(key=lambda x: x[0]) return results -def _q_base_uses(src, tree, lines, type_name): +def _q_base_uses(src, tree, lines, type_name, + enclosing_method=None, enclosing_class=None): results = [] for node in _find_all(tree.root_node, lambda n: n.type in _TYPE_DECL_NODES): bases = _base_type_names(node, src) if not any(type_name in _type_names(b) for b in bases): continue + if not _passes_enclosing_filter(node, src, + enclosing_method, enclosing_class): + continue name_node = node.child_by_field_name("name") if not name_node: continue @@ -1902,35 +2160,55 @@ def _q_base_uses(src, tree, lines, type_name): return results -def q_uses(src, tree, lines, type_name, uses_kind=None): +def q_uses(src, tree, lines, type_name, uses_kind=None, + enclosing_method=None, enclosing_class=None): + """Type-reference query, narrowable by ``uses_kind``. + + Each sub-helper applies the optional ``enclosing_method`` / + ``enclosing_class`` filter during its own walk -- there's no + post-pass. Note that for declaration-level kinds (``field``, + ``param``, ``return``, ``base``) the enclosing-method check runs + against the *declaration* node itself, which has no enclosing + method by definition, so passing ``enclosing_method`` against these + drops every row. Use ``locals`` / ``cast`` / ``all`` (or no + ``uses_kind``) when you want call-site context. + """ k = (uses_kind or "all").lower().strip() + kw = {"enclosing_method": enclosing_method, + "enclosing_class": enclosing_class} if k == "field": - return _q_field_type(src, tree, lines, type_name) - elif k == "param": - return _q_param_type(src, tree, lines, type_name) - elif k == "return": - return _q_return_type(src, tree, lines, type_name) - elif k == "cast": - return q_casts(src, tree, lines, type_name) - elif k == "base": - return _q_base_uses(src, tree, lines, type_name) - elif k == "locals": - return _q_local_type(src, tree, lines, type_name) - else: - return _q_uses_all(src, tree, lines, type_name) - - -def q_casts(src, tree, lines, type_name): - results = [ - (_line(node), lines[node.start_point[0]].strip() - if node.start_point[0] < len(lines) else "") - for node in _iter_cast_nodes(tree, src, type_name) - ] + return _q_field_type(src, tree, lines, type_name, **kw) + if k == "param": + return _q_param_type(src, tree, lines, type_name, **kw) + if k == "return": + return _q_return_type(src, tree, lines, type_name, **kw) + if k == "cast": + return q_casts(src, tree, lines, type_name, **kw) + if k == "base": + return _q_base_uses(src, tree, lines, type_name, **kw) + if k == "locals": + return _q_local_type(src, tree, lines, type_name, **kw) + return _q_uses_all(src, tree, lines, type_name, **kw) + + +def q_casts(src, tree, lines, type_name, + enclosing_method=None, enclosing_class=None): + results = [] + for node in _iter_cast_nodes(tree, src, type_name): + if not _passes_enclosing_filter(node, src, + enclosing_method, enclosing_class): + continue + row = node.start_point[0] + snippet = lines[row].strip() if row < len(lines) else "" + prefix = _scope_prefix(node, src) + results.append((_line(node), + f"{prefix}{snippet}" if prefix else snippet)) results.sort(key=lambda x: x[0]) return results -def q_accesses_on(src, tree, lines, type_name): +def q_accesses_on(src, tree, lines, type_name, + enclosing_method=None, enclosing_class=None): var_map = _build_var_type_map(tree, src) def _matches(name: str, node) -> bool: @@ -1941,12 +2219,19 @@ def _matches(name: str, node) -> bool: seen_rows = set() def _emit(node, member_name): + if not _passes_enclosing_filter(node, src, + enclosing_method, enclosing_class): + return row = node.start_point[0] if row in seen_rows: return seen_rows.add(row) line_text = lines[row].strip() if row < len(lines) else "" - results.append((_line(node), f".{member_name} <- {line_text}")) + prefix = _scope_prefix(node, src) + results.append( + (_line(node), + f"{prefix}.{member_name} <- {line_text}" + if prefix else f".{member_name} <- {line_text}")) # Direct member access: var.Member for node in _find_all(tree.root_node, lambda n: n.type == "member_access_expression"): @@ -1984,24 +2269,35 @@ def _emit(node, member_name): for assign, type_node, lhs in _iter_initializer_members(tree, src): if not type_node or type_name not in _type_names(_text(type_node, src).strip()): continue + if not _passes_enclosing_filter(assign, src, + enclosing_method, enclosing_class): + continue row = assign.start_point[0] line_text = lines[row].strip() if row < len(lines) else "" - results.append((_line(assign), f".{_text(lhs, src).strip()} <- {line_text}")) + prefix = _scope_prefix(assign, src) + body = f".{_text(lhs, src).strip()} <- {line_text}" + results.append((_line(assign), f"{prefix}{body}" if prefix else body)) # With-expression member mutations (C# 9 records) -- obj with { Prop = val } # Each member is emitted independently for the same reason as above. for wi, src_ident, prop in _iter_with_members(tree, src): if not _matches(_text(src_ident, src).strip(), wi): continue + if not _passes_enclosing_filter(wi, src, + enclosing_method, enclosing_class): + continue row = wi.start_point[0] line_text = lines[row].strip() if row < len(lines) else "" - results.append((_line(wi), f".{_text(prop, src).strip()} <- {line_text}")) + prefix = _scope_prefix(wi, src) + body = f".{_text(prop, src).strip()} <- {line_text}" + results.append((_line(wi), f"{prefix}{body}" if prefix else body)) results.sort(key=lambda x: x[0]) return results -def q_all_refs(src, tree, lines, name): +def q_all_refs(src, tree, lines, name, + enclosing_method=None, enclosing_class=None): results = [] seen_rows = set() for node in _find_all(tree.root_node, lambda n: n.type == "identifier"): @@ -2009,6 +2305,9 @@ def q_all_refs(src, tree, lines, name): continue if _in_literal(node): continue + if not _passes_enclosing_filter(node, src, + enclosing_method, enclosing_class): + continue row = node.start_point[0] if row in seen_rows: continue @@ -2065,7 +2364,9 @@ def q_var_type(src, tree, lines, name): # -- Process function ---------------------------------------------------------- def query_cs_bytes(src_bytes: bytes, mode: str, mode_arg: str, include_body=False, - symbol_kind=None, uses_kind=None, visibility=None, **kwargs): + symbol_kind=None, uses_kind=None, visibility=None, + head_lines=None, enclosing_method=None, + enclosing_class=None, **kwargs): """Parse C# bytes and return list[{"line": N, "text": "..."}] for the given mode.""" src_bytes = _strip_else_branches(src_bytes) try: @@ -2080,19 +2381,34 @@ def query_cs_bytes(src_bytes: bytes, mode: str, mode_arg: str, include_body=Fals "classes": lambda: q_classes(src_bytes, tree, lines, visibility=visibility), "methods": lambda: q_methods(src_bytes, tree, lines, visibility=visibility), "fields": lambda: q_fields(src_bytes, tree, lines, visibility=visibility), - "calls": lambda: q_calls(src_bytes, tree, lines, mode_arg), + "calls": lambda: q_calls(src_bytes, tree, lines, mode_arg, + enclosing_method=enclosing_method, + enclosing_class=enclosing_class), + "caller_of": lambda: q_caller_of(src_bytes, tree, lines, mode_arg), + "callee_of": lambda: q_callee_of(src_bytes, tree, lines, mode_arg), "implements": lambda: q_implements(src_bytes, tree, lines, mode_arg), - "uses": lambda: q_uses(src_bytes, tree, lines, mode_arg, uses_kind=uses_kind), - "accesses_on": lambda: q_accesses_on(src_bytes, tree, lines, mode_arg), - "all_refs": lambda: q_all_refs(src_bytes, tree, lines, mode_arg), - "casts": lambda: q_casts(src_bytes, tree, lines, mode_arg), + "uses": lambda: q_uses(src_bytes, tree, lines, mode_arg, uses_kind=uses_kind, + enclosing_method=enclosing_method, + enclosing_class=enclosing_class), + "accesses_on": lambda: q_accesses_on(src_bytes, tree, lines, mode_arg, + enclosing_method=enclosing_method, + enclosing_class=enclosing_class), + "all_refs": lambda: q_all_refs(src_bytes, tree, lines, mode_arg, + enclosing_method=enclosing_method, + enclosing_class=enclosing_class), + "casts": lambda: q_casts(src_bytes, tree, lines, mode_arg, + enclosing_method=enclosing_method, + enclosing_class=enclosing_class), "attrs": lambda: q_attrs(src_bytes, tree, lines, mode_arg), - "accesses_of": lambda: q_accesses_of(src_bytes, tree, lines, mode_arg), + "accesses_of": lambda: q_accesses_of(src_bytes, tree, lines, mode_arg, + enclosing_method=enclosing_method, + enclosing_class=enclosing_class), "imports": lambda: q_usings(src_bytes, tree, lines), "declarations": lambda: q_declarations(src_bytes, tree, lines, mode_arg, include_body=include_body, symbol_kind=symbol_kind, - visibility=visibility), - "body": lambda: q_body(src_bytes, tree, lines, mode_arg, symbol_kind=symbol_kind), + visibility=visibility, head_lines=head_lines), + "body": lambda: q_body(src_bytes, tree, lines, mode_arg, + symbol_kind=symbol_kind, head_lines=head_lines), "at": lambda: q_at(src_bytes, tree, lines, mode_arg), "params": lambda: q_params(src_bytes, tree, lines, mode_arg), "var_type": lambda: q_var_type(src_bytes, tree, lines, mode_arg), diff --git a/query/dispatch.py b/query/dispatch.py index 608c735..016d3e0 100644 --- a/query/dispatch.py +++ b/query/dispatch.py @@ -71,7 +71,8 @@ def _fn(src_bytes, mode, mode_arg, **kwargs): def query_file(src_bytes: bytes, ext: str, mode: str, mode_arg: str = "", include_body=False, symbol_kind=None, uses_kind=None, - visibility=None, **kwargs): + visibility=None, head_lines=None, + enclosing_method=None, enclosing_class=None, **kwargs): """Query src_bytes using the given mode. Returns ``list[{"line": N, "text": "..."}]`` on success. @@ -84,6 +85,16 @@ def query_file(src_bytes: bytes, ext: str, mode: str, mode_arg: str = "", ``visibility`` is an optional comma-separated filter for declaration modes (``classes``/``methods``/``fields``/``declarations``). Languages that don't capture visibility silently ignore it. + + ``head_lines`` truncates each ``body`` / ``declarations include_body=True`` + match to the first N source lines (signature + body together) with a + ``... +K more lines`` tail marker. Other modes ignore it. + + ``enclosing_method`` / ``enclosing_class`` narrow pattern-mode hits to + those that occur inside a member / type with the given name. Useful + for call-site context queries like ``calls("Save", + enclosing_method="WriteBack")``. Languages that don't capture + enclosing-scope structure silently ignore these. """ fn = _EXT_TO_QUERY_BYTES.get(ext) if fn is None: @@ -93,7 +104,10 @@ def query_file(src_bytes: bytes, ext: str, mode: str, mode_arg: str = "", ) return fn(src_bytes, mode, mode_arg, include_body=include_body, symbol_kind=symbol_kind, - uses_kind=uses_kind, visibility=visibility, **kwargs) + uses_kind=uses_kind, visibility=visibility, + head_lines=head_lines, + enclosing_method=enclosing_method, + enclosing_class=enclosing_class, **kwargs) def describe_file(src_bytes: bytes, ext: str) -> FileDescription: diff --git a/query/tests/test_cs_at_and_body.py b/query/tests/test_cs_at_and_body.py index f8cc23e..839820f 100644 --- a/query/tests/test_cs_at_and_body.py +++ b/query/tests/test_cs_at_and_body.py @@ -214,5 +214,67 @@ def test_body_with_symbol_kind_restricts_results(self): self.assertIn("[method]", text) +# ── q_body / q_declarations head_lines truncation ──────────────────────────── + + +_LONG_BODY_SRC = """\ +namespace Acme { + public class C { + public void LongMethod() + { + // line 1 of body + int a = 1; + int b = 2; + int c = 3; + int d = 4; + int e = 5; + int f = 6; + int g = 7; + } + } +} +""" + + +class TestQBodyHeadLines(unittest.TestCase): + """``q_body`` with ``head_lines=N`` returns only the first N source + lines of each match (signature + body together), with a + ``... +K more lines`` tail marker so the caller knows there's more.""" + + def setUp(self): + b = _LONG_BODY_SRC.encode() + tree = _PARSER.parse(b) + self.fx = (b, tree, _LONG_BODY_SRC.splitlines()) + + def test_no_truncation_by_default(self): + results = q_body(*self.fx, "LongMethod") + assert len(results) == 1 + text = results[0][-1] + # All 13 source lines (sig + 12 body lines incl. braces) present. + assert "int g = 7;" in text + assert "+more lines" not in text + + def test_head_lines_truncates(self): + results = q_body(*self.fx, "LongMethod", head_lines=4) + assert len(results) == 1 + text = results[0][-1] + # First few lines of the method (signature + body opening) kept; + # later lines dropped; tail marker appended. + assert "int g = 7;" not in text + assert "more lines" in text + + def test_head_lines_zero_means_no_truncation(self): + results = q_body(*self.fx, "LongMethod", head_lines=0) + text = results[0][-1] + assert "int g = 7;" in text + assert "more lines" not in text + + def test_head_lines_larger_than_body_is_noop(self): + results = q_body(*self.fx, "LongMethod", head_lines=1000) + text = results[0][-1] + assert "int g = 7;" in text + assert "more lines" not in text + + if __name__ == "__main__": unittest.main() diff --git a/query/tests/test_cs_call_graph.py b/query/tests/test_cs_call_graph.py new file mode 100644 index 0000000..6511329 --- /dev/null +++ b/query/tests/test_cs_call_graph.py @@ -0,0 +1,145 @@ +""" +Tests for ``caller_of`` / ``callee_of`` -- the call-graph neighbour modes. + +``caller_of METHOD`` groups every call site of METHOD by the enclosing +method, returning one row per (TypeName.MemberName) caller with a count +of how many call sites that caller contains. + +``callee_of METHOD`` walks the body of the method named METHOD and +returns one row per distinct callee with a count of invocations. Object- +creation expressions (``new T(...)``) are reported as callees of T. +""" +from __future__ import annotations + +import unittest + +import tree_sitter_c_sharp as tscsharp +from tree_sitter import Language, Parser + +from query.cs import q_caller_of, q_callee_of + +_CS = Language(tscsharp.language()) +_PARSER = Parser(_CS) + + +def _parse(src: str): + b = src.encode() + tree = _PARSER.parse(b) + return b, tree, src.splitlines() + + +# --------------------------------------------------------------------------- +# caller_of +# --------------------------------------------------------------------------- + + +class TestCallerOf(unittest.TestCase): + def test_groups_multiple_call_sites_into_one_caller(self): + src = """class C { + void A() { + Foo(); + Foo(); + Foo(); + } + }""" + b, tree, lines = _parse(src) + r = q_caller_of(b, tree, lines, "Foo") + assert len(r) == 1, r + _, text = r[0] + assert text.startswith("[in C.A]") + assert "3 call sites" in text + + def test_separate_callers_emit_separate_rows(self): + src = """class C { + void A() { Foo(); } + void B() { Foo(); Foo(); } + }""" + b, tree, lines = _parse(src) + r = q_caller_of(b, tree, lines, "Foo") + assert len(r) == 2, r + texts = [t for _, t in r] + assert any("[in C.A]" in t and "1 call site)" in t for t in texts), texts + assert any("[in C.B]" in t and "2 call sites)" in t for t in texts), texts + + def test_empty_when_no_calls(self): + src = "class C { void A() {} }" + b, tree, lines = _parse(src) + assert q_caller_of(b, tree, lines, "NoSuchMethod") == [] + + def test_results_sorted_by_line(self): + src = """class C { + void Z() { Foo(); } + void A() { Foo(); } + }""" + b, tree, lines = _parse(src) + r = q_caller_of(b, tree, lines, "Foo") + lns = [ln for ln, _ in r] + assert lns == sorted(lns), lns + + +# --------------------------------------------------------------------------- +# callee_of +# --------------------------------------------------------------------------- + + +class TestCalleeOf(unittest.TestCase): + def test_groups_distinct_callees_with_counts(self): + src = """class C { + void Driver() { + Repo.Save(); + Repo.Save(); + Logger.Info("hi"); + } + }""" + b, tree, lines = _parse(src) + r = q_callee_of(b, tree, lines, "Driver") + # Two distinct callees inside Driver: Save (2x), Info (1x). + texts = [t for _, t in r] + assert any("Save" in t and "2 invocations" in t for t in texts), texts + assert any("Info" in t and "1 invocation)" in t for t in texts), texts + + def test_constructor_calls_reported_as_ctor(self): + src = """class C { + void Make() { + var w = new Widget(); + var g = new Gadget(); + var w2 = new Widget(); + } + }""" + b, tree, lines = _parse(src) + r = q_callee_of(b, tree, lines, "Make") + texts = [t for _, t in r] + assert any("Widget" in t and "2 invocations, ctor" in t for t in texts), texts + assert any("Gadget" in t and "1 invocation, ctor" in t for t in texts), texts + + def test_only_searches_the_named_method_body(self): + src = """class C { + void Target() { Inside(); } + void Other() { NotInTarget(); } + }""" + b, tree, lines = _parse(src) + r = q_callee_of(b, tree, lines, "Target") + texts = " ".join(t for _, t in r) + assert "Inside" in texts + assert "NotInTarget" not in texts + + def test_empty_when_method_not_found(self): + src = "class C { void M() { Foo(); } }" + b, tree, lines = _parse(src) + assert q_callee_of(b, tree, lines, "DoesNotExist") == [] + + def test_handles_conditional_access_call(self): + # ``obj?.Method(...)`` should still be picked up as a callee. + src = """class C { + void Run(Repo r) { + r?.Save(); + } + }""" + b, tree, lines = _parse(src) + r = q_callee_of(b, tree, lines, "Run") + texts = " ".join(t for _, t in r) + assert "Save" in texts, texts + + +if __name__ == "__main__": + unittest.main() diff --git a/query/tests/test_cs_enclosing_filter.py b/query/tests/test_cs_enclosing_filter.py new file mode 100644 index 0000000..3aa6b8c --- /dev/null +++ b/query/tests/test_cs_enclosing_filter.py @@ -0,0 +1,175 @@ +""" +Tests for the ``enclosing_method=`` and ``enclosing_class=`` filters on +pattern modes. + +The filters narrow per-hit results to those inside a member / type with +the given name. Cross-cuts every pattern mode that emits from inside +method bodies: ``calls``, ``uses``, ``casts``, ``accesses_of``, +``accesses_on``, ``all_refs``. +""" +from __future__ import annotations + +import unittest + +import tree_sitter_c_sharp as tscsharp +from tree_sitter import Language, Parser + +from query.cs import ( + q_calls, q_casts, q_accesses_of, q_accesses_on, q_all_refs, q_uses, +) + +_CS = Language(tscsharp.language()) +_PARSER = Parser(_CS) + + +def _parse(src: str): + b = src.encode() + tree = _PARSER.parse(b) + return b, tree, src.splitlines() + + +_SRC = """\ +namespace Acme { + public class WriteBack { + public void Run() { + Save(); + Save(); + } + public void Other() { + Save(); + Save(); + Save(); + } + } + public class ReadOnly { + public void Run() { + Save(); + } + } +} +""" + + +class TestEnclosingMethodFilter(unittest.TestCase): + def test_no_filter_returns_all_hits(self): + b, tree, lines = _parse(_SRC) + r = q_calls(b, tree, lines, "Save") + assert len(r) == 6, r + + def test_enclosing_method_narrows(self): + b, tree, lines = _parse(_SRC) + r = q_calls(b, tree, lines, "Save", enclosing_method="Run") + # Both Run methods qualify (WriteBack.Run has 2 calls, ReadOnly.Run has 1). + assert len(r) == 3, r + for _, t in r: + assert "[in WriteBack.Run]" in t or "[in ReadOnly.Run]" in t, t + + def test_enclosing_class_narrows(self): + b, tree, lines = _parse(_SRC) + r = q_calls(b, tree, lines, "Save", enclosing_class="WriteBack") + # All 5 Save() calls in WriteBack (2 in Run + 3 in Other). + assert len(r) == 5, r + for _, t in r: + assert "[in WriteBack." in t, t + + def test_both_filters_compose(self): + b, tree, lines = _parse(_SRC) + r = q_calls(b, tree, lines, "Save", + enclosing_method="Run", + enclosing_class="WriteBack") + # Only WriteBack.Run -- 2 calls. + assert len(r) == 2, r + for _, t in r: + assert "[in WriteBack.Run]" in t, t + + def test_no_match_when_method_name_wrong(self): + b, tree, lines = _parse(_SRC) + r = q_calls(b, tree, lines, "Save", enclosing_method="NoSuchMethod") + assert r == [], r + + +class TestEnclosingFilterOnAccessesOn(unittest.TestCase): + def test_filter_applies_to_accesses_on(self): + # ``q_accesses_on`` dedupes by source row, so each access must be on + # its own line for the test to count distinct hits. + src = """class C { + void A(Repo r) { + r.Save(); + r.Touch(); + } + void B(Repo r) { + r.Save(); + } + }""" + b, tree, lines = _parse(src) + all_r = q_accesses_on(b, tree, lines, "Repo") + assert len(all_r) == 3, all_r + a_only = q_accesses_on(b, tree, lines, "Repo", enclosing_method="A") + assert len(a_only) == 2, a_only + + +class TestEnclosingFilterOnCasts(unittest.TestCase): + def test_filter_applies_to_casts(self): + src = """class C { + void A() { int x = (int)42L; } + void B() { int y = (int)99L; } + }""" + b, tree, lines = _parse(src) + all_c = q_casts(b, tree, lines, "int") + assert len(all_c) == 2, all_c + a_only = q_casts(b, tree, lines, "int", enclosing_method="A") + assert len(a_only) == 1, a_only + + +class TestEnclosingFilterOnAccessesOf(unittest.TestCase): + def test_filter_applies_to_accesses_of(self): + # Same one-access-per-line constraint as accesses_on. + src = """class C { + void A(Foo f) { + var v = f.Value; + } + void B(Foo f) { + var v = f.Value; + var w = f.Value; + } + }""" + b, tree, lines = _parse(src) + all_r = q_accesses_of(b, tree, lines, "Value") + assert len(all_r) == 3, all_r + b_only = q_accesses_of(b, tree, lines, "Value", enclosing_method="B") + assert len(b_only) == 2, b_only + + +class TestEnclosingFilterOnAllRefs(unittest.TestCase): + def test_filter_applies_to_all_refs(self): + src = """class C { + void A() { var foo = 1; foo = 2; } + void B() { var foo = 3; } + }""" + b, tree, lines = _parse(src) + all_r = q_all_refs(b, tree, lines, "foo") + # 3 unique row hits: line 2 (decl+assign on different rows), line 2 assign? Actually: + # Line 2: ``var foo = 1; foo = 2;`` -- single row, so 1 hit + # Line 3: ``var foo = 3;`` -- 1 hit + # all_refs dedupes by row, so 2 hits total. + assert len(all_r) == 2, all_r + a_only = q_all_refs(b, tree, lines, "foo", enclosing_method="A") + assert len(a_only) == 1, a_only + + +class TestEnclosingFilterOnUsesLocals(unittest.TestCase): + def test_filter_applies_to_uses_locals(self): + src = """class C { + void A() { Repo r = null; r.Save(); } + void B() { Repo r = null; r.Save(); } + }""" + b, tree, lines = _parse(src) + all_l = q_uses(b, tree, lines, "Repo", uses_kind="locals") + assert len(all_l) == 2, all_l + a_only = q_uses(b, tree, lines, "Repo", uses_kind="locals", + enclosing_method="A") + assert len(a_only) == 1, a_only + + +if __name__ == "__main__": + unittest.main() diff --git a/query/tests/test_cs_object_initializer_accesses.py b/query/tests/test_cs_object_initializer_accesses.py index d229f12..bb28b26 100644 --- a/query/tests/test_cs_object_initializer_accesses.py +++ b/query/tests/test_cs_object_initializer_accesses.py @@ -44,7 +44,15 @@ def _lns(results): return {ln for ln, _ in results} def _members(results): - return {txt.split(" <-")[0].lstrip(".") for _, txt in results} + # Result text is ``[in Class.Method] .MemberName <- source`` (or just + # ``.MemberName <- source`` when the access is outside any member). + # The member is the token between the last dot and the `` <-`` marker. + out = set() + for _, txt in results: + before = txt.split(" <-", 1)[0] + if "." in before: + out.add(before.rsplit(".", 1)[1].strip()) + return out def _line_no(fragment): for i, ln in enumerate(_LINES): @@ -76,7 +84,7 @@ def test_multi_member_same_line_both_reported(self): r = self._accesses("Widget") line = _line_no("Value = 1, Name") line_results = [(ln, txt) for ln, txt in r if ln == line] - member_names = {txt.split(" <-")[0].lstrip(".") for _, txt in line_results} + member_names = _members(line_results) assert "Value" in member_names, f"'Value' missing from same-line initializer: {r}" assert "Name" in member_names, f"'Name' missing from same-line initializer: {r}" diff --git a/query/tests/test_cs_scope_prefix.py b/query/tests/test_cs_scope_prefix.py new file mode 100644 index 0000000..76ecaa7 --- /dev/null +++ b/query/tests/test_cs_scope_prefix.py @@ -0,0 +1,143 @@ +""" +Tests for the enclosing-scope prefix on pattern-mode AST hits. + +Each result emitted by ``calls`` / ``uses`` / ``accesses_of`` / +``accesses_on`` / ``casts`` is prepended with ``[in TypeName.MemberName] `` +(or ``[in TypeName] `` at type-level, or nothing at namespace level) so +agents can tell which class/method a hit lives in without a follow-up +``at LINE:COL`` query. +""" +from __future__ import annotations + +import unittest + +import tree_sitter_c_sharp as tscsharp +from tree_sitter import Language, Parser + +from query.cs import ( + q_calls, q_accesses_of, q_accesses_on, q_casts, q_uses, + _scope_prefix, _enclosing_member_name, _find_all, +) + +_CS = Language(tscsharp.language()) +_PARSER = Parser(_CS) + + +def _parse(src: str): + b = src.encode() + tree = _PARSER.parse(b) + return b, tree, src.splitlines() + + +_SRC = """\ +namespace Acme { + public class Widget { + private int _count; + public void DoWork() { + Logger.Info("hi"); + _count = (int)42L; + } + public int Value { get { return _count; } } + } + public class Caller { + public void Run(Widget w) { + w.DoWork(); + var v = w.Value; + } + } +} +""" + + +class TestScopePrefixHelpers(unittest.TestCase): + """Unit tests on the helper that builds the prefix.""" + + def test_namespace_level_has_no_prefix(self): + # No class around the node -> empty string. + b, tree, _ = _parse("namespace N { }") + root = tree.root_node + assert _scope_prefix(root, b) == "" + + def test_inside_method_yields_class_dot_member(self): + b, tree, _ = _parse(_SRC) + # Find the invocation_expression for ``Logger.Info(...)`` -- it + # lives inside ``Widget.DoWork``. + calls = _find_all(tree.root_node, lambda n: n.type == "invocation_expression") + assert calls, "expected at least one invocation in fixture" + prefix = _scope_prefix(calls[0], b) + assert prefix == "[in Widget.DoWork] ", prefix + + def test_field_declaration_name_resolves_via_declarator(self): + # field_declaration has no direct ``name`` field; the enclosing- + # member helper should still recover the variable name. + b, tree, _ = _parse("class C { private int Foo = 0; }") + fields = _find_all(tree.root_node, lambda n: n.type == "field_declaration") + # Pick any descendant of the field to walk up from. + descendant = next((c for c in fields[0].children if c.is_named), fields[0]) + assert _enclosing_member_name(descendant, b) == "Foo" + + +class TestQCallsScopePrefix(unittest.TestCase): + def test_call_inside_method_carries_scope(self): + b, tree, lines = _parse(_SRC) + r = q_calls(b, tree, lines, "Info") + assert len(r) == 1, r + _, text = r[0] + assert text.startswith("[in Widget.DoWork] "), text + + def test_call_in_different_method_carries_its_own_scope(self): + b, tree, lines = _parse(_SRC) + r = q_calls(b, tree, lines, "DoWork") + # Two hits: declaration matches no, the call site ``w.DoWork()`` is + # inside Caller.Run. + assert any("[in Caller.Run] " in t for _, t in r), r + + +class TestQAccessesOnScopePrefix(unittest.TestCase): + def test_member_access_via_typed_local_includes_scope(self): + b, tree, lines = _parse(_SRC) + r = q_accesses_on(b, tree, lines, "Widget") + # ``w.DoWork()`` and ``w.Value`` both inside Caller.Run. + assert r, r + for _, txt in r: + assert txt.startswith("[in Caller.Run] "), txt + + +class TestQAccessesOfScopePrefix(unittest.TestCase): + def test_property_read_inside_method_carries_scope(self): + b, tree, lines = _parse(_SRC) + r = q_accesses_of(b, tree, lines, "Value") + # ``w.Value`` access happens inside Caller.Run. + assert r, r + assert any("[in Caller.Run] " in t for _, t in r), r + + +class TestQCastsScopePrefix(unittest.TestCase): + def test_cast_inside_method_carries_scope(self): + b, tree, lines = _parse(_SRC) + # ``(int)42L`` lives inside Widget.DoWork. + r = q_casts(b, tree, lines, "int") + assert r, r + _, txt = r[0] + assert txt.startswith("[in Widget.DoWork] "), txt + + +class TestQUsesLocalsScopePrefix(unittest.TestCase): + def test_typed_local_inside_method_carries_scope(self): + # ``Widget w`` would normally be a parameter (handled separately by + # the param uses_kind); use an explicit local for clarity. + src = """class C { + void M() { + Widget local = null; + local.DoWork(); + } + }""" + b, tree, lines = _parse(src) + r = q_uses(b, tree, lines, "Widget", uses_kind="locals") + assert r, r + _, text = r[0] + assert text.startswith("[in C.M] "), text + + +if __name__ == "__main__": + unittest.main() diff --git a/query/tests/test_cs_throttle.py b/query/tests/test_cs_throttle.py index bcfa659..0ef35ef 100644 --- a/query/tests/test_cs_throttle.py +++ b/query/tests/test_cs_throttle.py @@ -59,6 +59,22 @@ def _texts(results): return " ".join(t for t in (row[-1] for row in results)) +def _stripped_texts(results): + """Like ``_texts`` but drops the ``[in TypeName.MemberName] `` enclosing- + scope prefix that pattern modes prepend to each hit. Use for assertions + that care about the source-snippet content only, e.g. checking that a + different member name isn't accidentally matched -- the prefix itself + legitimately contains the enclosing method name and would spoil a naive + substring check.""" + out = [] + for row in results: + t = row[-1] + if t.startswith("[in ") and "] " in t: + t = t.split("] ", 1)[1] + out.append(t) + return " ".join(out) + + # =========================================================================== # declarations # =========================================================================== @@ -214,9 +230,15 @@ def test_finds_record_attempt_call_site(self): assert r, "_policy.RecordAttempt must be found" def test_different_member_not_returned(self): + # Sanity check: an ``accesses_of("TotalMilliseconds")`` search must + # not match call sites of an unrelated member like ``RecordAttempt``. + # Strip the ``[in TypeName.MemberName] `` scope prefix before the + # substring check -- the prefix legitimately contains the enclosing + # method's name (which can be ``RecordAttempt``) without that being + # a mismatch on the accessed member. r = self._of("TotalMilliseconds") - texts = _texts(r) - assert "RecordAttempt" not in texts + snippets = _stripped_texts(r) + assert "RecordAttempt" not in snippets def test_nonexistent_member_returns_empty(self): assert self._of("NoSuchMember") == [] diff --git a/query/tests/test_cs_with_expression_accesses.py b/query/tests/test_cs_with_expression_accesses.py index 7ad7b23..0607860 100644 --- a/query/tests/test_cs_with_expression_accesses.py +++ b/query/tests/test_cs_with_expression_accesses.py @@ -42,7 +42,15 @@ def _lns(results): return {ln for ln, _ in results} def _members(results): - return {txt.split(" <-")[0].lstrip(".") for _, txt in results} + # Same logic as the object-initializer test: extract the member token + # between the last ``.`` and the `` <-`` source separator. Tolerates + # an optional leading ``[in TypeName.MemberName] `` scope prefix. + out = set() + for _, txt in results: + before = txt.split(" <-", 1)[0] + if "." in before: + out.add(before.rsplit(".", 1)[1].strip()) + return out def _line_no(fragment): for i, ln in enumerate(_LINES): @@ -74,7 +82,7 @@ def test_multi_member_same_line_both_reported(self): r = self._accesses("Coord") line = _line_no("X = 0, Y = 0") line_results = [(ln, txt) for ln, txt in r if ln == line] - names = {txt.split(" <-")[0].lstrip(".") for _, txt in line_results} + names = _members(line_results) assert "X" in names and "Y" in names, \ f"Both X and Y must appear on same-line with result: {line_results}" diff --git a/tests/unit/test_mcp_server.py b/tests/unit/test_mcp_server.py index 9fc7af2..3f6297d 100644 --- a/tests/unit/test_mcp_server.py +++ b/tests/unit/test_mcp_server.py @@ -576,11 +576,12 @@ def test_caps_at_10_lines_per_file(self): # Lines 11-25 absent for i in range(11, 26): assert f"src/Big.cs:{i}:" not in result - # Per-file suggestion appended -- uses bare relative path, not the - # legacy $SRC_ROOT/ placeholder. - assert "25 total hits" in result + # Per-file suggestion appended in the compact ``[+K capped] path`` form. + # 25 hits - 10 shown = 15 capped. + assert "[+15 capped] src/Big.cs" in result + # The general "issue query_single_file(...)" reminder is at the top + # of the suggestion block (one line shared across all capped files). assert "query_single_file" in result - assert 'file="src/Big.cs"' in result def test_no_suggestion_when_under_cap(self): """File with <=10 hits shouldn't trigger a per-file suggestion.""" @@ -631,9 +632,10 @@ def test_multi_file_with_some_capped(self): # Small.cs shown in full assert "src/Small.cs:1: hit 1" in result assert "src/Small.cs:2: hit 2" in result - # Suggestion only for Big.cs - assert "src/Big.cs" in result - assert "15 total hits" in result + # Suggestion only for Big.cs: 15 hits - 10 shown = 5 capped. + assert "[+5 capped] src/Big.cs" in result + # Small.cs is fully shown so it should NOT appear in the capped list. + assert "capped] src/Small.cs" not in result # -- _sync_state --------------------------------------------------------------- diff --git a/tsquery_server.py b/tsquery_server.py index 37962eb..f06eb80 100644 --- a/tsquery_server.py +++ b/tsquery_server.py @@ -106,7 +106,19 @@ def _get_query_module(): def _run_query(mode: str, pattern: str, files: list[Path], include_body: bool = False, symbol_kind: str = "", - uses_kind: str = "", visibility: str = "") -> list: + uses_kind: str = "", visibility: str = "", + head_lines: int | None = None, + enclosing_method: str = "", + enclosing_class: str = "") -> list: + """Per-file AST pass for the index-pre-filtered candidate set. + + Files whose language doesn't support ``mode`` are silently skipped + rather than crashing the whole query -- a codebase-wide query like + ``body SaveChanges`` can match files in many languages, and only the + ones whose extractor knows ``body`` should contribute results. The + ValueError that ``query_file`` raises for unsupported modes is + treated as a no-op for that file. + """ _q = _get_query_module() results = [] for path in files: @@ -117,11 +129,18 @@ def _run_query(mode: str, pattern: str, files: list[Path], except OSError as e: print(f"ERROR reading {native}: {e}", file=sys.stderr) continue - matches = _q.query_file(src_bytes, ext, mode, pattern, - include_body=include_body, - symbol_kind=symbol_kind, - uses_kind=uses_kind, - visibility=visibility) + try: + matches = _q.query_file(src_bytes, ext, mode, pattern, + include_body=include_body, + symbol_kind=symbol_kind, + uses_kind=uses_kind, + visibility=visibility, + head_lines=head_lines, + enclosing_method=enclosing_method or None, + enclosing_class=enclosing_class or None) + except ValueError: + # Mode unsupported for this file's language - skip silently. + continue if matches: results.append({"file": str(native), "matches": matches}) return results @@ -131,7 +150,16 @@ def _run_query(mode: str, pattern: str, files: list[Path], _EXT_TO_TS_AND_AST: dict[str, tuple[str, str]] = { "declarations": ("symbols", "declarations"), + "body": ("symbols", "body"), "calls": ("calls", "calls"), + # ``caller_of`` shares the index pre-filter with ``calls`` -- same set + # of candidate files contain the method invocations -- but the AST + # post-pass groups results by the enclosing method. + "caller_of": ("calls", "caller_of"), + # ``callee_of`` is "what does THIS method call". Pre-filter on the + # method's declaration field (``method_names``) -- the file we want is + # the one that *declares* the method, not the ones that call it. + "callee_of": ("symbols", "callee_of"), "implements": ("implements", "implements"), "uses": ("uses", "uses"), "casts": ("casts", "casts"), @@ -139,6 +167,7 @@ def _run_query(mode: str, pattern: str, files: list[Path], "accesses_of": ("accesses_of", "accesses_of"), "accesses_on": ("uses", "accesses_on"), "all_refs": ("all_refs", "all_refs"), + "var_type": ("all_refs", "var_type"), } @@ -486,7 +515,15 @@ def _handle(self) -> None: symbol_kind = str(body.get("symbol_kind", "") or "") uses_kind = str(body.get("uses_kind", "") or "") visibility = str(body.get("visibility", "") or "") + enclosing_method = str(body.get("enclosing_method", "") or "") + enclosing_class = str(body.get("enclosing_class", "") or "") exclude_path = str(body.get("exclude_path", "") or "") + try: + head_lines_raw = body.get("head_lines", None) + head_lines = (int(head_lines_raw) + if head_lines_raw not in (None, "") else None) + except (TypeError, ValueError): + head_lines = None if mode not in _EXT_TO_TS_AND_AST: self._send_json(400, {"error": f"unknown mode: {mode!r}"}) @@ -559,7 +596,10 @@ def _handle(self) -> None: ast_results = _run_query(ast_mode, pattern, file_list, include_body=include_body, symbol_kind=symbol_kind, uses_kind=uses_kind, - visibility=visibility) + visibility=visibility, + head_lines=head_lines, + enclosing_method=enclosing_method, + enclosing_class=enclosing_class) response_hits = [] for ast_item in ast_results: