diff --git a/__tests__/rust-path-super-glob.test.ts b/__tests__/rust-path-super-glob.test.ts new file mode 100644 index 0000000000..f52f8c494a --- /dev/null +++ b/__tests__/rust-path-super-glob.test.ts @@ -0,0 +1,85 @@ +/** + * Regression: #[path]-included test modules with `use super::*` must resolve + * bare callees through the declaring parent module, not a same-named symbol + * elsewhere in the workspace (ContextualWisdomLab/fast-mlsirm#1837). + */ +import { afterEach, beforeEach, expect, it } from 'vitest'; +import * as fs from 'node:fs'; +import * as os from 'node:os'; +import * as path from 'node:path'; +import { CodeGraph } from '../src'; + +let root: string; +let cg: CodeGraph | undefined; + +beforeEach(() => { + root = fs.mkdtempSync(path.join(os.tmpdir(), 'cg-rust-path-super-')); +}); + +afterEach(() => { + cg?.close(); + cg = undefined; + fs.rmSync(root, { recursive: true, force: true }); +}); + +async function indexWorkspace(layout: Record) { + fs.writeFileSync( + path.join(root, 'Cargo.toml'), + '[workspace]\nmembers = ["crates/core", "crates/py"]\nresolver = "2"\n' + ); + fs.mkdirSync(path.join(root, 'crates/core/src'), { recursive: true }); + fs.mkdirSync(path.join(root, 'crates/py/src'), { recursive: true }); + fs.mkdirSync(path.join(root, 'tests/unit'), { recursive: true }); + fs.writeFileSync( + path.join(root, 'crates/core/Cargo.toml'), + '[package]\nname = "core-crate"\nversion = "0.1.0"\nedition = "2021"\n' + ); + fs.writeFileSync( + path.join(root, 'crates/py/Cargo.toml'), + '[package]\nname = "py-crate"\nversion = "0.1.0"\nedition = "2021"\n' + ); + for (const [file, text] of Object.entries(layout)) { + const dest = path.join(root, file); + fs.mkdirSync(path.dirname(dest), { recursive: true }); + fs.writeFileSync(dest, text); + } + cg = await CodeGraph.init(root, { index: true }); +} + +function callTargets(testFn: string): string[] { + const caller = cg! + .getNodesByKind('function') + .find((n) => n.filePath === 'tests/unit/scaling_tests.rs' && n.name === testFn); + expect(caller).toBeDefined(); + return cg! + .getOutgoingEdges(caller!.id) + .filter((e) => e.kind === 'calls') + .map((e) => { + const target = cg!.getNode(e.target)!; + return `${target.filePath}:${target.qualifiedName ?? target.name}`; + }); +} + +it('resolves use super::* callee through #[path] parent, not a same-named decoy (#1837)', async () => { + await indexWorkspace({ + 'crates/core/src/lib.rs': 'pub mod scaling;', + 'crates/core/src/scaling.rs': `pub fn predict_rating_multi(x: i32) -> i32 { x + 1 } + +#[cfg(test)] +#[path = "../../../tests/unit/scaling_tests.rs"] +mod tests; +`, + 'crates/py/src/lib.rs': 'pub fn predict_rating_multi(x: i32) -> i32 { x + 99 }', + 'tests/unit/scaling_tests.rs': `use super::*; + +#[test] +fn pr_elom_rowmean() { + let _ = predict_rating_multi(1); +} +`, + }); + + expect(callTargets('pr_elom_rowmean')).toEqual([ + 'crates/core/src/scaling.rs:predict_rating_multi', + ]); +}); diff --git a/scripts/repro-rust-path-callee.sh b/scripts/repro-rust-path-callee.sh new file mode 100755 index 0000000000..372d3cd36f --- /dev/null +++ b/scripts/repro-rust-path-callee.sh @@ -0,0 +1,41 @@ +#!/usr/bin/env bash +# RED/GREEN repro for #[path] + use super::* callee resolution (fast-mlsirm#1837). +set -euo pipefail +ROOT="$(cd "$(dirname "$0")/.." && pwd)" +cd "$ROOT" +npm run build --silent +node --input-type=module -e " +import { CodeGraph } from './dist/index.js'; +import fs from 'node:fs'; +import os from 'node:os'; +import path from 'node:path'; + +const root = fs.mkdtempSync(path.join(os.tmpdir(), 'cg-repro-1837-')); +const write = (rel, text) => { + const dest = path.join(root, rel); + fs.mkdirSync(path.dirname(dest), { recursive: true }); + fs.writeFileSync(dest, text); +}; +write('Cargo.toml', '[workspace]\\nmembers = [\"crates/core\", \"crates/py\"]\\nresolver = \"2\"\\n'); +write('crates/core/Cargo.toml', '[package]\\nname = \"core-crate\"\\nversion = \"0.1.0\"\\nedition = \"2021\"\\n'); +write('crates/py/Cargo.toml', '[package]\\nname = \"py-crate\"\\nversion = \"0.1.0\"\\nedition = \"2021\"\\n'); +write('crates/core/src/lib.rs', 'pub mod scaling;'); +write('crates/core/src/scaling.rs', 'pub fn predict_rating_multi(x: i32) -> i32 { x + 1 }\\n\\n#[cfg(test)]\\n#[path = \"../../../tests/unit/scaling_tests.rs\"]\\nmod tests;\\n'); +write('crates/py/src/lib.rs', 'pub fn predict_rating_multi(x: i32) -> i32 { x + 99 }'); +write('tests/unit/scaling_tests.rs', 'use super::*;\\n\\n#[test]\\nfn pr_elom_rowmean() { let _ = predict_rating_multi(1); }\\n'); + +const cg = await CodeGraph.init(root, { index: true }); +const caller = cg.getNodesByKind('function').find(n => n.filePath === 'tests/unit/scaling_tests.rs' && n.name === 'pr_elom_rowmean'); +const edges = cg.getOutgoingEdges(caller.id).filter(e => e.kind === 'calls').map(e => { + const t = cg.getNode(e.target); + return t.filePath + ':' + (t.qualifiedName ?? t.name); +}); +const want = 'crates/core/src/scaling.rs:predict_rating_multi'; +if (!edges.includes(want)) { + console.error('FAIL: expected callee', want, 'got', edges); + process.exit(1); +} +console.log('OK:', edges.join(', ')); +cg.close(); +fs.rmSync(root, { recursive: true, force: true }); +" diff --git a/src/resolution/import-resolver.ts b/src/resolution/import-resolver.ts index cf1740f9af..3c40f6cf07 100644 --- a/src/resolution/import-resolver.ts +++ b/src/resolution/import-resolver.ts @@ -17,6 +17,12 @@ import { localReceiverTypePatterns, normalizeInferredTypeName, } from './name-matcher'; +import { + clearRustModulePathMemos, + getRustPathInclusionParent, + resolveRustSuperModuleFile, + rustSelfModuleDirAbs, +} from './rust-module-paths'; /** * Extension resolution order by language @@ -153,6 +159,7 @@ export function clearImportResolverMemos(context: ResolutionContext): void { fileExportIndexes.delete(context); luaFileBasenameIndexes.delete(context); cobolCopybookIndexes.delete(context); + clearRustModulePathMemos(context); } export function resolveImportPath( @@ -887,6 +894,8 @@ export function extractImportMappings( mappings.push(...extractPHPImports(content)); } else if (language === 'c' || language === 'cpp') { mappings.push(...extractCppImports(content)); + } else if (language === 'rust') { + mappings.push(...extractRustImports(content)); } return mappings; @@ -1194,6 +1203,87 @@ function extractCppImports(content: string): ImportMapping[] { return mappings; } +/** + * Extract Rust `use` declarations into import mappings. Named `use` items bind + * a local name to a module path; `use path::*` glob imports bind every public + * item from that module (Rust Reference, Use declarations). + */ +function extractRustImports(content: string): ImportMapping[] { + const mappings: ImportMapping[] = []; + + const expand = (spec: string): string[] => { + const open = spec.indexOf('{'); + if (open === -1) return [spec.trim()]; + const prefix = spec.slice(0, open); + let depth = 0; + let close = -1; + for (let i = open; i < spec.length; i++) { + if (spec[i] === '{') depth++; + else if (spec[i] === '}') { + depth--; + if (depth === 0) { + close = i; + break; + } + } + } + if (close === -1) return []; + const suffix = spec.slice(close + 1); + const inner = spec.slice(open + 1, close); + const parts: string[] = []; + let depth2 = 0; + let start = 0; + for (let i = 0; i <= inner.length; i++) { + const ch = inner[i]; + if (ch === '{') depth2++; + else if (ch === '}') depth2--; + if (i === inner.length || (ch === ',' && depth2 === 0)) { + const seg = inner.slice(start, i).trim(); + if (seg) parts.push(seg); + start = i + 1; + } + } + return parts.flatMap((p) => expand(prefix + p + suffix)); + }; + + const useRe = /(^|\n)\s*(?:pub(?:\([^)]*\))?\s+)?use\s+([^;]+);/g; + let m: RegExpExecArray | null; + while ((m = useRe.exec(content)) !== null) { + for (const spec of expand(m[2]!.replace(/\s+/g, ' '))) { + const aliasMatch = /^(.*?)\s+as\s+([A-Za-z_]\w*)$/.exec(spec); + const rawPath = (aliasMatch ? aliasMatch[1]! : spec).trim(); + if (!rawPath) continue; + + if (rawPath.endsWith('::*')) { + const modPath = rawPath.slice(0, -3); + if (!modPath) continue; + mappings.push({ + localName: '*', + exportedName: '*', + source: modPath, + isDefault: false, + isNamespace: true, + }); + continue; + } + + const segments = rawPath.split('::').map((s) => s.trim()).filter(Boolean); + const leaf = segments[segments.length - 1]; + if (!leaf || leaf === 'self' || leaf === 'super' || leaf === 'crate') continue; + const local = aliasMatch ? aliasMatch[2]! : leaf; + mappings.push({ + localName: local, + exportedName: leaf, + source: segments.join('::'), + isDefault: false, + isNamespace: false, + }); + } + } + + return mappings; +} + // Cache import mappings per file to avoid re-reading and re-parsing const importMappingCache = new Map(); @@ -1616,6 +1706,14 @@ export function resolveViaImport( if (rustResult) return rustResult; } + // Rust bare names brought in by `use` (named or `path::*` glob). Without this, + // `use super::*` in a #[path]-included test module falls through to global + // name-matching and can land on a same-named symbol in another crate (#1837). + if (ref.language === 'rust' && !ref.referenceName.includes('::')) { + const rustUse = resolveRustUseImport(ref, imports, context); + if (rustUse) return rustUse; + } + // Lua / Luau `require(...)`: a dotted module path (`a.b.c` from // `require("a.b.c")`) or an instance-path leaf (`Signal` from // `require(script.Parent.Signal)`) — map it to a module file. There's no static @@ -2043,6 +2141,86 @@ function resolvePythonAbsoluteModule( return hit ? { original: ref, targetNodeId: hit.id, confidence: 0.9, resolvedBy: 'import' } : null; } +const RUST_SYMBOL_KINDS = new Set([ + 'function', + 'struct', + 'union', + 'enum', + 'trait', + 'type_alias', + 'constant', + 'method', + 'class', + 'interface', +]); + +function findRustSymbolInFile( + file: string, + name: string, + context: ResolutionContext, + callableOnly: boolean +): Node | undefined { + return context.getNodesInFile(file).find((n) => { + if (n.name !== name) return false; + if (!RUST_SYMBOL_KINDS.has(n.kind)) return false; + if (callableOnly && n.kind !== 'function' && n.kind !== 'method') return false; + return true; + }); +} + +/** + * Resolve a bare Rust reference through `use` bindings: named imports and + * `use module::*` globs (Rust Reference, Use declarations). + */ +function resolveRustUseImport( + ref: UnresolvedRef, + imports: ImportMapping[], + context: ResolutionContext +): ResolvedRef | null { + const callableOnly = ref.referenceKind === 'calls' || ref.referenceKind === 'function_ref'; + + // Named `use path::item` / `use path::item as alias` — map via stored source path. + for (const imp of imports) { + if (imp.isNamespace || imp.localName !== ref.referenceName) continue; + const qualified = imp.source; + if (!qualified.includes('::')) continue; + const fakeRef: UnresolvedRef = { ...ref, referenceName: qualified }; + const hit = resolveRustPathReference(fakeRef, context); + if (hit) return hit; + } + + // `use super::*` / `use crate::m::*` — search the imported module file. + for (const imp of imports) { + if (!imp.isNamespace || imp.exportedName !== '*') continue; + const file = resolveRustModulePathToFile(imp.source, ref.filePath, context); + if (!file || file === ref.filePath) continue; + const target = findRustSymbolInFile(file, ref.referenceName, context, callableOnly); + if (target) { + return { original: ref, targetNodeId: target.id, confidence: 0.9, resolvedBy: 'import' }; + } + } + + return null; +} + +/** Map a Rust module path string (`super`, `crate::m`, `self::sub`) to a file. */ +function resolveRustModulePathToFile( + modPath: string, + fromFile: string, + context: ResolutionContext +): string | null { + const segments = modPath.split('::').filter((s) => s.length > 0); + if (segments.length === 0) return null; + + let supers = 0; + while (supers < segments.length && segments[supers] === 'super') supers++; + if (supers > 0 && supers === segments.length) { + return resolveRustSuperModuleFile(supers, fromFile, context); + } + + return resolveRustModuleFile(segments, fromFile, context); +} + /** * Resolve a Rust qualified reference `A::B::C` by mapping the MODULE prefix * (`A::B`) to a file and finding the leaf symbol (`C`) in it. This is the Rust @@ -2064,20 +2242,7 @@ function resolveRustPathReference( const file = resolveRustModuleFile(modSegs, ref.filePath, context); if (!file || file === ref.filePath) return null; - const target = context.getNodesInFile(file).find( - (n) => - n.name === leaf && - (n.kind === 'function' || - n.kind === 'struct' || - n.kind === 'union' || - n.kind === 'enum' || - n.kind === 'trait' || - n.kind === 'type_alias' || - n.kind === 'constant' || - n.kind === 'method' || - n.kind === 'class' || - n.kind === 'interface') - ); + const target = findRustSymbolInFile(file, leaf, context, false); if (target) { return { original: ref, targetNodeId: target.id, confidence: 0.9, resolvedBy: 'import' }; } @@ -2102,12 +2267,17 @@ function rustCrateRootDir(fromFileAbs: string, context: ResolutionContext): stri } /** Directory under which the current file's module declares its SUBMODULES. */ -function rustSelfModuleDir(fromFileAbs: string): string { - const base = path.basename(fromFileAbs); - const dir = path.dirname(fromFileAbs); - // mod.rs / lib.rs / main.rs own their directory; `foo.rs`'s submodules live in `foo/`. - if (base === 'mod.rs' || base === 'lib.rs' || base === 'main.rs') return dir; - return path.join(dir, base.replace(/\.rs$/, '')); +function rustSelfModuleDir( + fromFileAbs: string, + fromFileRel: string, + context: ResolutionContext +): string { + const parentViaPath = getRustPathInclusionParent(fromFileRel, context); + if (parentViaPath) { + const projectRoot = context.getProjectRoot(); + return path.dirname(path.join(projectRoot, parentViaPath)); + } + return rustSelfModuleDirAbs(fromFileAbs); } /** @@ -2149,14 +2319,19 @@ function resolveRustModuleFile( return resolveUnder(rustCrateRootDir(fromAbs, context), segments.slice(1)); } if (first === 'self') { - return resolveUnder(rustSelfModuleDir(fromAbs), segments.slice(1)); + return resolveUnder(rustSelfModuleDir(fromAbs, fromFile, context), segments.slice(1)); } if (first === 'super') { let supers = 0; while (segments[supers] === 'super') supers++; - let dir: string | null = rustSelfModuleDir(fromAbs); - for (let s = 0; s < supers && dir; s++) dir = path.dirname(dir); - return resolveUnder(dir, segments.slice(supers)); + const rest = segments.slice(supers); + if (rest.length === 0) { + return resolveRustSuperModuleFile(supers, fromFile, context); + } + const parentFile = resolveRustSuperModuleFile(supers, fromFile, context); + if (!parentFile) return null; + const parentAbs = path.join(projectRoot, parentFile); + return resolveUnder(rustSelfModuleDir(parentAbs, parentFile, context), rest); } // Bare path. In expression position (`submodule::item()` — the router-assembly // and general cross-module-call pattern) the prefix is a SUBMODULE of the @@ -2164,7 +2339,7 @@ function resolveRustModuleFile( // Fall back to crate-relative for 2015-edition / crate-root items. External // crate paths (`serde::de::Error`) miss both and fall through to name-matching. return ( - resolveUnder(rustSelfModuleDir(fromAbs), segments) ?? + resolveUnder(rustSelfModuleDir(fromAbs, fromFile, context), segments) ?? resolveUnder(rustCrateRootDir(fromAbs, context), segments) ); } diff --git a/src/resolution/rust-module-paths.ts b/src/resolution/rust-module-paths.ts new file mode 100644 index 0000000000..4a235cbfee --- /dev/null +++ b/src/resolution/rust-module-paths.ts @@ -0,0 +1,121 @@ +/** + * Rust #[path] module inclusion and effective parent-module resolution. + * + * Per the Rust Reference (Modules chapter), a `mod` item with a `path` + * attribute loads its body from an external file while remaining a lexical + * submodule of the declaring module; `super` in that file refers to the + * declaring module, not to a path derived from the included file's location + * on disk (Rust Reference, n.d., https://doc.rust-lang.org/reference/items/modules.html). + */ + +import * as path from 'path'; +import { ResolutionContext } from './types'; + +const rustPathInclusionMemos = new WeakMap>(); + +/** `#[path = "..."] mod name;` — attribute may share a line with other attrs. */ +const RUST_PATH_ATTR_MOD_RE = /#\[\s*path\s*=\s*"([^"]+)"\s*\]\s*mod\s+\w+\s*;/g; + +function normalizeRelPath(p: string): string { + return p.replace(/\\/g, '/'); +} + +export function clearRustModulePathMemos(context: ResolutionContext): void { + rustPathInclusionMemos.delete(context); +} + +/** + * Map from included child file (repo-relative) to its declaring parent module + * file (repo-relative). + */ +export function buildRustPathInclusionMap(context: ResolutionContext): Map { + const cached = rustPathInclusionMemos.get(context); + if (cached) return cached; + + const map = new Map(); + const projectRoot = context.getProjectRoot(); + const toRel = (abs: string) => normalizeRelPath(path.relative(projectRoot, abs)); + + for (const file of context.getAllFiles()) { + if (!file.endsWith('.rs')) continue; + const content = context.readFile(file); + if (!content?.includes('#[path')) continue; + + const parentDir = path.dirname(path.join(projectRoot, file)); + RUST_PATH_ATTR_MOD_RE.lastIndex = 0; + let m: RegExpExecArray | null; + while ((m = RUST_PATH_ATTR_MOD_RE.exec(content)) !== null) { + const child = toRel(path.normalize(path.join(parentDir, m[1]!))); + map.set(child, normalizeRelPath(file)); + } + } + + rustPathInclusionMemos.set(context, map); + return map; +} + +/** Parent module file when `childFile` is loaded via `#[path]`, else null. */ +export function getRustPathInclusionParent( + childFile: string, + context: ResolutionContext +): string | null { + return buildRustPathInclusionMap(context).get(normalizeRelPath(childFile)) ?? null; +} + +/** + * Resolve a lone `super` (or `super::super::…` with no trailing module + * segments) to the parent module's source file. `#[path]`-included modules + * treat one `super` as the declaring file (Rust Reference, Paths). + */ +export function resolveRustSuperModuleFile( + superCount: number, + fromFile: string, + context: ResolutionContext +): string | null { + if (superCount <= 0) return null; + + const projectRoot = context.getProjectRoot(); + const parentViaPath = getRustPathInclusionParent(fromFile, context); + if (parentViaPath) { + if (superCount === 1) return parentViaPath; + // Additional `super` segments walk up from the declaring module file. + let dir = path.dirname(path.join(projectRoot, parentViaPath)); + for (let s = 1; s < superCount && dir; s++) { + dir = path.dirname(dir); + } + const base = path.basename(dir); + const lib = normalizeRelPath(path.join(dir, 'lib.rs')); + const main = normalizeRelPath(path.join(dir, 'main.rs')); + const modRs = normalizeRelPath(path.join(dir, 'mod.rs')); + if (context.fileExists(lib)) return lib; + if (context.fileExists(main)) return main; + if (context.fileExists(modRs)) return modRs; + return normalizeRelPath(path.join(dir, base + '.rs')); + } + + const fromAbs = path.join(projectRoot, fromFile); + let dir = rustSelfModuleDirAbs(fromAbs); + for (let s = 0; s < superCount && dir; s++) { + dir = path.dirname(dir); + } + if (!dir) return null; + + const base = path.basename(dir); + const lib = normalizeRelPath(path.join(dir, 'lib.rs')); + const main = normalizeRelPath(path.join(dir, 'main.rs')); + const modRs = normalizeRelPath(path.join(dir, 'mod.rs')); + if (context.fileExists(lib)) return lib; + if (context.fileExists(main)) return main; + if (context.fileExists(modRs)) return modRs; + const asRs = normalizeRelPath(path.join(dir, base + '.rs')); + if (context.fileExists(asRs)) return asRs; + return null; +} + +/** Directory under which the current file's module declares its submodules. */ +export function rustSelfModuleDirAbs(fromFileAbs: string): string { + const base = path.basename(fromFileAbs); + const dir = path.dirname(fromFileAbs); + if (base === 'mod.rs' || base === 'lib.rs' || base === 'main.rs') return dir; + return path.join(dir, base.replace(/\.rs$/, '')); +}