From 6a9a1942f80f23b17899b8143510e478971b1ba7 Mon Sep 17 00:00:00 2001 From: Divyam Talwar Date: Thu, 24 Sep 2026 01:34:31 +0530 Subject: [PATCH] fix(graph): retain Rust implementation method identities --- src/graph/extract/rust.ts | 166 ++++++++++++----- tests/shared/graph/rust.test.ts | 306 +++++++++++++++++++++++++++++++- 2 files changed, 423 insertions(+), 49 deletions(-) diff --git a/src/graph/extract/rust.ts b/src/graph/extract/rust.ts index e9cd3d8e5..41f41d81f 100644 --- a/src/graph/extract/rust.ts +++ b/src/graph/extract/rust.ts @@ -41,12 +41,63 @@ export function extractRust( result.nodes.push(moduleNode); const declByName = new Map(); - collectDecls(root, relativePath, result, declByName, moduleNode); - collectCalls(root, result, declByName); + const impls: ImplState = { fnByDecl: new Map(), ownerTypes: new Map(), pendingOwners: [] }; + collectDecls(root, relativePath, result, declByName, moduleNode, impls, ""); + linkImplOwners(result, impls); + collectCalls(root, result, declByName, impls); return result; } +/** + * Impl bookkeeping shared by the declaration and call passes. + * - fnByDecl: function_item start position → the node declared for it, so a + * caller resolves to its own declaration instead of a same-named function. + * - ownerTypes: struct/enum declared in this file, keyed by inline-module + * scope + name, so an impl only links to a type in its own module. + * - pendingOwners: method_of edges awaiting their owner's local type node, + * resolved after all decls so a type declared below its impl still links. + */ +interface ImplState { + fnByDecl: Map; + ownerTypes: Map; + pendingOwners: { owner: string; method: string }[]; +} + +function declPos(node: TSNode): string { + return `${node.startPosition.row}:${node.startPosition.column}`; +} + +function scopedName(scope: string, name: string): string { + return `${scope}::${name}`; +} + +function pushFn( + result: FileExtraction, + declByName: Map, + impls: ImplState, + decl: TSNode, + node: GraphNode, + lookupKey?: string, +): void { + // On an id collision pushNode keeps the first node; attribute calls to it. + const existing = result.nodes.find((n) => n.id === node.id); + pushNode(result, declByName, node, lookupKey); + impls.fnByDecl.set(declPos(decl), existing ?? node); +} + +function pushOwnerType( + result: FileExtraction, + declByName: Map, + impls: ImplState, + scope: string, + node: GraphNode, +): void { + pushNode(result, declByName, node); + const key = scopedName(scope, node.label); + if (!impls.ownerTypes.has(key)) impls.ownerTypes.set(key, node.id); +} + // ─── Pass 1 + 2 ──────────────────────────────────────────────────────────── function collectDecls( @@ -55,6 +106,8 @@ function collectDecls( result: FileExtraction, declByName: Map, moduleNode: GraphNode, + impls: ImplState, + scope: string, ): void { for (let i = 0; i < node.namedChildCount; i++) { const child = node.namedChild(i); @@ -66,24 +119,24 @@ function collectDecls( /* c8 ignore next */ if (name === null) continue; const exported = isRustPub(child); - pushNode(result, declByName, makeNode(relativePath, name, "function", child, exported, LANG)); + pushFn(result, declByName, impls, child, makeNode(relativePath, name, "function", child, exported, LANG)); } else if (child.type === "struct_item") { const name = textOfField(child, "name"); /* c8 ignore next */ if (name === null) continue; - pushNode(result, declByName, makeNode(relativePath, name, "class", child, isRustPub(child), LANG)); + pushOwnerType(result, declByName, impls, scope, makeNode(relativePath, name, "class", child, isRustPub(child), LANG)); } else if (child.type === "enum_item") { const name = textOfField(child, "name"); /* c8 ignore next */ if (name === null) continue; - pushNode(result, declByName, makeNode(relativePath, name, "enum", child, isRustPub(child), LANG)); + pushOwnerType(result, declByName, impls, scope, makeNode(relativePath, name, "enum", child, isRustPub(child), LANG)); } else if (child.type === "trait_item") { const name = textOfField(child, "name"); /* c8 ignore next */ if (name === null) continue; pushNode(result, declByName, makeNode(relativePath, name, "interface", child, isRustPub(child), LANG)); } else if (child.type === "impl_item") { - collectImplMethods(child, relativePath, result, declByName); + collectImplMethods(child, relativePath, result, declByName, impls, scope); } else if (child.type === "mod_item") { const name = textOfField(child, "name"); /* c8 ignore next */ @@ -93,7 +146,7 @@ function collectDecls( const body = child.childForFieldName("body"); /* c8 ignore next */ if (body !== null) { - collectDecls(body, relativePath, result, declByName, moduleNode); + collectDecls(body, relativePath, result, declByName, moduleNode, impls, scopedName(scope, name)); } } else if (child.type === "use_declaration") { collectUseDecl(child, result, moduleNode); @@ -114,30 +167,47 @@ function isRustPub(node: TSNode): boolean { return false; } +/** + * Method identity is the impl's own spelling: `Rect::area` for an inherent + * impl, `Box::get` vs `Box::get` for specialized impls, and + * `::fmt` vs `::fmt` for trait impls so neither + * definition is lost. Only the base type (`Wrapper` → `Wrapper`) is used + * to link method_of, and only to a struct/enum declared in the same inline + * module of this file. Scoped, reference and other types get no owner. + * + * Known limitations (AST only, no name resolution): spelling is textual, so + * `Self`, type aliases and differently-qualified paths to one type are not + * unified; inline-module paths are not part of ids, so same-named items in + * sibling `mod` blocks share one node (as free fns already do); an impl of a + * type brought in with `use` is not linked to that type. + */ function collectImplMethods( impl: TSNode, relativePath: string, result: FileExtraction, declByName: Map, + impls: ImplState, + scope: string, ): void { - // impl_item → type field (the type being implemented) + declaration_list body + // impl_item → optional trait field, type field (the type being + // implemented) + declaration_list body const typeNode = impl.childForFieldName("type"); - /* c8 ignore next */ - const implTypeName = typeNode !== null ? typeNode.text.trim() : null; - const body = impl.childForFieldName("body"); /* c8 ignore next */ - if (body === null) return; + if (typeNode === null || body === null) return; + + const typeText = implSpelling(typeNode); + const traitNode = impl.childForFieldName("trait"); + const keyPrefix = traitNode === null ? typeText : `<${typeText} as ${implSpelling(traitNode)}>`; + const ownerName = implBaseTypeName(typeNode); for (let i = 0; i < body.namedChildCount; i++) { const member = body.namedChild(i); - /* c8 ignore next */ - if (member === null || member.type !== "function_item") continue; + if (member?.type !== "function_item") continue; const name = textOfField(member, "name"); /* c8 ignore next */ if (name === null) continue; - /* c8 ignore next */ - const key = implTypeName !== null ? `${implTypeName}::${name}` : name; + const key = `${keyPrefix}::${name}`; const methodNode: GraphNode = { id: nodeId(relativePath, key, "method"), label: name, @@ -147,19 +217,34 @@ function collectImplMethods( language: LANG, exported: isRustPub(member), }; - pushNode(result, declByName, methodNode, key); - /* c8 ignore next */ - if (implTypeName !== null) { - result.edges.push({ - source: nodeId(relativePath, implTypeName, "class"), - target: methodNode.id, - relation: "method_of", - confidence: "EXTRACTED", - }); + pushFn(result, declByName, impls, member, methodNode, key); + if (ownerName !== null) { + impls.pendingOwners.push({ owner: scopedName(scope, ownerName), method: methodNode.id }); } } } +/** Source spelling of an impl type/trait; preserve literal contents exactly. */ +function implSpelling(node: TSNode): string { + return node.text.trim(); +} + +function implBaseTypeName(typeNode: TSNode): string | null { + const base = typeNode.type === "generic_type" ? typeNode.childForFieldName("type") : typeNode; + return base?.type === "type_identifier" ? base.text : null; +} + +/** Emit method_of only for owners declared as a struct or enum in the impl's module. */ +function linkImplOwners(result: FileExtraction, impls: ImplState): void { + const seen = new Set(); + for (const { owner, method } of impls.pendingOwners) { + const source = impls.ownerTypes.get(owner); + if (source === undefined || seen.has(method)) continue; + seen.add(method); + result.edges.push({ source, target: method, relation: "method_of", confidence: "EXTRACTED" }); + } +} + function collectUseDecl( node: TSNode, result: FileExtraction, @@ -205,13 +290,14 @@ function collectCalls( node: TSNode, result: FileExtraction, declByName: Map, + impls: ImplState, ): void { if (node.type === "call_expression") { const fn = node.childForFieldName("function"); /* c8 ignore next */ if (fn !== null && fn.type === "identifier") { const target = declByName.get(fn.text); - const caller = findEnclosingFn(node, declByName); + const caller = findEnclosingFn(node, impls); /* c8 ignore next */ if (target !== undefined && caller !== null) { result.edges.push({ @@ -226,35 +312,23 @@ function collectCalls( for (let i = 0; i < node.namedChildCount; i++) { const child = node.namedChild(i); /* c8 ignore next */ - if (child !== null) collectCalls(child, result, declByName); + if (child !== null) collectCalls(child, result, declByName, impls); } } function findEnclosingFn( node: TSNode, - declByName: Map, + impls: ImplState, ): GraphNode | null { + // The nearest enclosing callable owns the call, resolved by declaration + // position, not name: `A::new`, `B::new` and a free `new` are distinct + // callers. Nested fns, closures and trait default bodies have no node of + // their own, so their calls are dropped rather than credited to an outer fn. let cur: TSNode | null = node.parent; while (cur !== null) { - if (cur.type === "function_item") { - const name = textOfField(cur, "name"); - /* c8 ignore next */ - if (name !== null) { - // check bare name first, then impl-qualified name - const found = declByName.get(name) ?? (() => { - for (const [k, v] of declByName) { - /* c8 ignore next */ - if (k.endsWith(`::${name}`) || k === name) return v; - } - /* c8 ignore next */ - return undefined; - })(); - /* c8 ignore next */ - if (found !== undefined) return found; - } - } + if (cur.type === "function_item") return impls.fnByDecl.get(declPos(cur)) ?? null; + if (cur.type === "closure_expression") return null; cur = cur.parent; } - /* c8 ignore next */ return null; } diff --git a/tests/shared/graph/rust.test.ts b/tests/shared/graph/rust.test.ts index d6a44565a..b43ac9a49 100644 --- a/tests/shared/graph/rust.test.ts +++ b/tests/shared/graph/rust.test.ts @@ -97,9 +97,7 @@ describe("Rust extraction", () => { expect(ex.edges.some(e => e.relation === "imports" && e.target.startsWith("external:"))).toBe(true); }); - it("resolves call from an impl method to a free function (triggers impl-qualified findEnclosingFn search)", () => { - // Covers lines 221-224: findEnclosingFn walks up to function_item inside an impl block; - // tries declByName.get(name) first then searches k.endsWith(::name) to find the impl-qualified key. + it("resolves call from an impl method to a free function", () => { const ex = extractRust( `fn setup() {}\nstruct Worker {}\nimpl Worker {\n pub fn run(&self) { setup(); }\n}\n`, "src/worker.rs", @@ -125,3 +123,305 @@ describe("Rust extraction", () => { expect(ex.parse_errors).toHaveLength(0); }); }); + +describe("Rust impl method ownership", () => { + const calls = (ex: ReturnType) => + ex.edges.filter(e => e.relation === "calls").map(e => `${e.source} -> ${e.target}`).sort(); + const methodOf = (ex: ReturnType) => + ex.edges.filter(e => e.relation === "method_of").map(e => `${e.source} -> ${e.target}`).sort(); + + it("attributes calls from A::new, B::new and a free new to their own declarations", () => { + const ex = extractRust( + [ + "fn helper_a() {}", + "fn helper_b() {}", + "fn helper_free() {}", + "fn new() { helper_free(); }", + "struct A;", + "impl A { fn new() -> A { helper_a(); A } }", + "struct B;", + "impl B { fn new() -> B { helper_b(); B } }", + "", + ].join("\n"), + "src/lib.rs", + ); + expect(ex.parse_errors).toHaveLength(0); + expect(ex.nodes.filter(n => n.label === "new").map(n => n.id).sort()).toEqual([ + "src/lib.rs:A::new:method", + "src/lib.rs:B::new:method", + "src/lib.rs:new:function", + ]); + expect(calls(ex)).toEqual([ + "src/lib.rs:A::new:method -> src/lib.rs:helper_a:function", + "src/lib.rs:B::new:method -> src/lib.rs:helper_b:function", + "src/lib.rs:new:function -> src/lib.rs:helper_free:function", + ]); + }); + + it("does not misattribute same-named impl methods when there is no free fn", () => { + const ex = extractRust( + "fn helper() {}\nstruct A;\nstruct B;\nimpl A { fn run(&self) {} }\nimpl B { fn run(&self) { helper(); } }\n", + "src/lib.rs", + ); + expect(calls(ex)).toEqual(["src/lib.rs:B::run:method -> src/lib.rs:helper:function"]); + }); + + it("links enum impl methods to the enum node", () => { + const ex = extractRust( + "pub enum Color { Red, Green }\nimpl Color { pub fn is_red(&self) -> bool { true } }\n", + "src/color.rs", + ); + expect(methodOf(ex)).toEqual(["src/color.rs:Color:enum -> src/color.rs:Color::is_red:method"]); + }); + + it("links generic struct impl methods to the base struct node", () => { + const ex = extractRust( + "fn helper() {}\npub struct Wrapper { v: T }\nimpl Wrapper { pub fn get(&self) -> &T { helper(); &self.v } }\n", + "src/wrap.rs", + ); + expect(ex.parse_errors).toHaveLength(0); + expect(ex.nodes.some(n => n.id === "src/wrap.rs:Wrapper::get:method")).toBe(true); + expect(methodOf(ex)).toEqual(["src/wrap.rs:Wrapper:class -> src/wrap.rs:Wrapper::get:method"]); + expect(calls(ex)).toEqual(["src/wrap.rs:Wrapper::get:method -> src/wrap.rs:helper:function"]); + }); + + it("keeps specialized generic impls distinct while linking both to the base type", () => { + const ex = extractRust( + [ + "pub struct Box(T);", + "fn for_i32() {}", + "fn for_u32() {}", + "impl Box {", + " pub fn get(&self) { for_i32(); }", + "}", + "impl Box {", + " fn get(&self) {", + " for_u32();", + " }", + "}", + "", + ].join("\n"), + "src/lib.rs", + ); + expect(ex.parse_errors).toHaveLength(0); + const gets = ex.nodes.filter(n => n.label === "get"); + expect(gets.map(n => [n.id, n.kind, n.source_location, n.exported])).toEqual([ + ["src/lib.rs:Box::get:method", "method", "L5", true], + ["src/lib.rs:Box::get:method", "method", "L8-10", false], + ]); + expect(methodOf(ex)).toEqual([ + "src/lib.rs:Box:class -> src/lib.rs:Box::get:method", + "src/lib.rs:Box:class -> src/lib.rs:Box::get:method", + ]); + expect(calls(ex)).toEqual([ + "src/lib.rs:Box::get:method -> src/lib.rs:for_i32:function", + "src/lib.rs:Box::get:method -> src/lib.rs:for_u32:function", + ]); + }); + + it("resolves owners declared after the impl block", () => { + const ex = extractRust( + "impl Late { fn a(&self) {} }\nimpl Shape { fn b(&self) {} }\nstruct Late;\nenum Shape { Sq }\n", + "src/lib.rs", + ); + expect(methodOf(ex)).toEqual([ + "src/lib.rs:Late:class -> src/lib.rs:Late::a:method", + "src/lib.rs:Shape:enum -> src/lib.rs:Shape::b:method", + ]); + }); + + it("resolves specialized and trait impl owners declared after the impl block", () => { + const ex = extractRust( + [ + "impl Tr for Late { fn a(&self) {} }", + "impl Late { fn a(&self) {} }", + "struct Late(T);", + "trait Tr { fn a(&self); }", + "", + ].join("\n"), + "src/lib.rs", + ); + expect(ex.parse_errors).toHaveLength(0); + expect(methodOf(ex)).toEqual([ + "src/lib.rs:Late:class -> src/lib.rs: as Tr>::a:method", + "src/lib.rs:Late:class -> src/lib.rs:Late::a:method", + ]); + }); + + it("emits no method_of edge for foreign or scoped impl types", () => { + const ex = extractRust( + "impl Foreign { fn a(&self) {} }\nimpl other::Scoped { fn b(&self) {} }\nimpl<'x> Local for &'x str { fn c(&self) {} }\nimpl other::Gen { fn d(&self) {} }\ntrait Local { fn c(&self); }\n", + "src/lib.rs", + ); + expect(ex.nodes.filter(n => n.kind === "method").map(n => n.id)).toEqual([ + "src/lib.rs:Foreign::a:method", + "src/lib.rs:other::Scoped::b:method", + "src/lib.rs:<&'x str as Local>::c:method", + "src/lib.rs:other::Gen::d:method", + ]); + expect(methodOf(ex)).toEqual([]); + }); + + it("links owners only within the impl's own inline module", () => { + const ex = extractRust( + [ + "struct A;", + "mod m {", + " struct B;", + " impl B { fn g(&self) {} }", + " impl A { fn f(&self) {} }", + " impl super::A { fn h(&self) {} }", + "}", + "impl B { fn top(&self) {} }", + "", + ].join("\n"), + "src/lib.rs", + ); + expect(ex.parse_errors).toHaveLength(0); + expect(ex.nodes.filter(n => n.kind === "method").map(n => n.id)).toEqual([ + "src/lib.rs:B::g:method", + "src/lib.rs:A::f:method", + "src/lib.rs:super::A::h:method", + "src/lib.rs:B::top:method", + ]); + expect(methodOf(ex)).toEqual(["src/lib.rs:B:class -> src/lib.rs:B::g:method"]); + }); + + it("keeps same-named methods from distinct trait impls and the inherent impl on one type", () => { + const ex = extractRust( + [ + "fn h1() {}", + "fn h2() {}", + "fn h3() {}", + "struct A;", + "trait T1 { fn go(&self); }", + "trait T2 { fn go(&self); }", + "impl T1 for A { fn go(&self) { h1(); } }", + "impl T2 for A {", + " fn go(&self) { h2(); }", + "}", + "impl A { pub fn go(&self) { h3(); } }", + "", + ].join("\n"), + "src/lib.rs", + ); + expect(ex.parse_errors).toHaveLength(0); + expect(ex.nodes.filter(n => n.kind === "method").map(n => [n.id, n.label, n.source_location])).toEqual([ + ["src/lib.rs:::go:method", "go", "L7"], + ["src/lib.rs:::go:method", "go", "L9"], + ["src/lib.rs:A::go:method", "go", "L11"], + ]); + expect(methodOf(ex)).toEqual([ + "src/lib.rs:A:class -> src/lib.rs:::go:method", + "src/lib.rs:A:class -> src/lib.rs:::go:method", + "src/lib.rs:A:class -> src/lib.rs:A::go:method", + ]); + expect(calls(ex)).toEqual([ + "src/lib.rs:::go:method -> src/lib.rs:h1:function", + "src/lib.rs:::go:method -> src/lib.rs:h2:function", + "src/lib.rs:A::go:method -> src/lib.rs:h3:function", + ]); + }); + + it("keeps generic trait impls with different arguments distinct", () => { + const ex = extractRust( + [ + "struct N(i64);", + "impl From for N { fn from(v: i32) -> Self { N(v as i64) } }", + "impl From for N { fn from(v: u32) -> Self { N(v as i64) } }", + "", + ].join("\n"), + "src/n.rs", + ); + expect(ex.parse_errors).toHaveLength(0); + expect(methodOf(ex)).toEqual([ + "src/n.rs:N:class -> src/n.rs:>::from:method", + "src/n.rs:N:class -> src/n.rs:>::from:method", + ]); + }); + + it("skips non-fn impl members", () => { + const ex = extractRust( + "struct A;\ntrait Tr { type Out; const N: u8; fn f(&self); }\nimpl Tr for A { type Out = u8; const N: u8 = 1; fn f(&self) {} }\n", + "src/lib.rs", + ); + expect(ex.parse_errors).toHaveLength(0); + expect(ex.nodes.filter(n => n.kind === "method").map(n => n.id)).toEqual(["src/lib.rs:::f:method"]); + }); + + it("merges cfg-duplicated declarations into the first node without losing calls", () => { + const ex = extractRust( + [ + "fn x() {}", + "fn y() {}", + "#[cfg(unix)]", + "fn a() { x(); }", + "#[cfg(not(unix))]", + "fn a() { y(); }", + "#[cfg(unix)]", + "struct S;", + "#[cfg(not(unix))]", + "struct S;", + "#[cfg(unix)]", + "impl S { fn m(&self) { x(); } }", + "#[cfg(not(unix))]", + "impl S { fn m(&self) { y(); } }", + "", + ].join("\n"), + "src/lib.rs", + ); + expect(ex.parse_errors).toHaveLength(0); + expect(ex.nodes.filter(n => n.label === "a").map(n => [n.id, n.source_location])).toEqual([ + ["src/lib.rs:a:function", "L4"], + ]); + expect(ex.nodes.filter(n => n.label === "m").map(n => [n.id, n.source_location])).toEqual([ + ["src/lib.rs:S::m:method", "L12"], + ]); + expect(methodOf(ex)).toEqual(["src/lib.rs:S:class -> src/lib.rs:S::m:method"]); + expect(calls(ex)).toEqual([ + "src/lib.rs:S::m:method -> src/lib.rs:x:function", + "src/lib.rs:S::m:method -> src/lib.rs:y:function", + "src/lib.rs:a:function -> src/lib.rs:x:function", + "src/lib.rs:a:function -> src/lib.rs:y:function", + ]); + }); + + it("does not credit calls in nested fns, closures or trait default bodies to an outer fn", () => { + const ex = extractRust( + [ + "fn helper() {}", + "fn outer() { fn inner() { helper(); } inner(); }", + "fn with_closure() { let f = || helper(); f(); helper(); }", + "trait Tr { fn d(&self) { helper(); } }", + "struct W;", + "impl W {", + " fn m(&self) {", + " fn nested() { helper(); }", + " let c = |v: u8| { helper(); v };", + " }", + "}", + "", + ].join("\n"), + "src/lib.rs", + ); + expect(ex.parse_errors).toHaveLength(0); + expect(calls(ex)).toEqual(["src/lib.rs:with_closure:function -> src/lib.rs:helper:function"]); + }); +}); + + +describe("Rust impl identity literal preservation", () => { + it("preserves literal bytes inside const-generic type expressions", () => { + const source = `struct Width; +impl Width<{ b"a b".len() }> { fn get(&self) {} } +impl Width<{ b"a b".len() }> { fn get(&self) {} } +`; + const result = extractRust(source, "width.rs"); + expect(result.parse_errors).toEqual([]); + const methods = result.nodes.filter((node) => node.kind === "method"); + expect(methods).toHaveLength(2); + expect(new Set(methods.map((node) => node.id)).size).toBe(2); + expect(methods.some((node) => node.id.includes('b"a b"'))).toBe(true); + expect(methods.some((node) => node.id.includes('b"a b"'))).toBe(true); + }); +});