feat(cpp): parse CUDA source extensions

This commit is contained in:
azizur100389 2026-06-15 16:27:47 +01:00
parent bd1d446baa
commit acbdac6d3b
8 changed files with 82 additions and 12 deletions

View file

@ -32,7 +32,17 @@ const EXTENSION_MAP: Record<SupportedLanguages, readonly string[]> = {
[SupportedLanguages.Python]: ['.py'],
[SupportedLanguages.Java]: ['.java'],
[SupportedLanguages.C]: ['.c'],
[SupportedLanguages.CPlusPlus]: ['.cpp', '.cc', '.cxx', '.h', '.hpp', '.hxx', '.hh'],
[SupportedLanguages.CPlusPlus]: [
'.cpp',
'.cc',
'.cxx',
'.h',
'.hpp',
'.hxx',
'.hh',
'.cu',
'.cuh',
],
[SupportedLanguages.CSharp]: ['.cs'],
[SupportedLanguages.Go]: ['.go'],
[SupportedLanguages.Ruby]: ['.rb', '.rake', '.gemspec'],

View file

@ -45,13 +45,20 @@ import { logger } from '../../logger.js';
// ---------- constants ----------
const HEADER_EXTENSIONS = new Set(['.h', '.hpp', '.hxx', '.hh']);
const HEADER_EXTENSIONS = new Set(['.h', '.hpp', '.hxx', '.hh', '.cuh']);
// Source = headers (provider-eligible) ∪ implementation files (.c/.cpp/.cc/.cxx).
// Source = headers (provider-eligible) ∪ implementation files (.c/.cpp/.cc/.cxx/.cu).
// Spread keeps the subset relationship explicit so a future contributor adding
// a new header extension to HEADER_EXTENSIONS does not have to remember to
// also add it here.
const SOURCE_EXTENSIONS = new Set<string>([...HEADER_EXTENSIONS, '.c', '.cpp', '.cc', '.cxx']);
const SOURCE_EXTENSIONS = new Set<string>([
...HEADER_EXTENSIONS,
'.c',
'.cpp',
'.cc',
'.cxx',
'.cu',
]);
const INCLUDE_QUERY_SRC = '(preproc_include path: (_) @import.source) @import';
@ -275,6 +282,8 @@ function getLanguageForFile(filePath: string): unknown | null {
case '.hpp':
case '.hxx':
case '.hh':
case '.cu':
case '.cuh':
return Cpp;
default:
return null;
@ -298,7 +307,7 @@ function getLanguageForFile(filePath: string): unknown | null {
function isLocalInclude(cleaned: string, suffixIndex: SuffixIndex): boolean {
const candidates = [cleaned];
if (!/\.[a-zA-Z0-9]+$/.test(cleaned)) {
for (const ext of ['.h', '.hpp', '.hxx', '.hh']) candidates.push(cleaned + ext);
for (const ext of ['.h', '.hpp', '.hxx', '.hh', '.cuh']) candidates.push(cleaned + ext);
}
for (const c of candidates) {
if (suffixIndex.get(c) || suffixIndex.getInsensitive(c)) return true;
@ -428,7 +437,7 @@ export class IncludeExtractor implements ContractExtractor {
try {
const rows = await db(
`MATCH (f:File)
WHERE f.filePath =~ '.*\\\\.(h|hpp|hxx|hh)$'
WHERE f.filePath =~ '.*\\\\.(h|hpp|hxx|hh|cuh)$'
RETURN f.filePath AS filePath, f.id AS fileId`,
);
// gitnexus analyze stores absolute paths in the File.filePath column.

View file

@ -36,6 +36,8 @@ export const EXTENSIONS = [
'.cxx',
'.hxx',
'.hh',
'.cu',
'.cuh',
// C#
'.cs',
// Go

View file

@ -425,7 +425,7 @@ export const cProvider = defineLanguage({
export const cppProvider = defineLanguage({
id: SupportedLanguages.CPlusPlus,
extensions: ['.cpp', '.cc', '.cxx', '.h', '.hpp', '.hxx', '.hh'],
extensions: ['.cpp', '.cc', '.cxx', '.h', '.hpp', '.hxx', '.hh', '.cu', '.cuh'],
entryPointPatterns: [
/^main$/,
/^init_/,

View file

@ -2,14 +2,14 @@ import { readdirSync, type Dirent } from 'fs';
import { join, relative } from 'path';
/** C++ header extensions to scan for in the workspace. */
const HEADER_EXTENSIONS = new Set(['.h', '.hpp', '.hxx', '.hh']);
const HEADER_EXTENSIONS = new Set(['.h', '.hpp', '.hxx', '.hh', '.cuh']);
/**
* Walk `repoPath` recursively and return relative paths of all C++ header files.
* Used by `loadResolutionConfig` so the C++ resolver can resolve `#include`
* targets that live in header files.
*
* Scans for: .h, .hpp, .hxx, .hh
* Scans for: .h, .hpp, .hxx, .hh, .cuh
*/
export function scanCppHeaderFiles(repoPath: string): ReadonlySet<string> {
const headers = new Set<string>();

View file

@ -288,6 +288,21 @@ describe('Tree-sitter multi-language parsing', () => {
expect(names).toContain('helper');
});
it('treats CUDA .cu and .cuh files as C++ for definition extraction', async () => {
expect(getLanguageFromFilename('src/kernels/force.cu')).toBe(SupportedLanguages.CPlusPlus);
expect(getLanguageFromFilename('src/force/nep.cuh')).toBe(SupportedLanguages.CPlusPlus);
await loadLanguage(SupportedLanguages.CPlusPlus, 'src/kernels/force.cu');
const code = `class Force { public: void apply(); };\nvoid launchKernel() {}`;
const provider = getProvider(SupportedLanguages.CPlusPlus);
const { matches } = parseAndQuery(parser, code, provider.treeSitterQueries);
const defs = extractDefinitions(matches);
const names = defs.map((d) => d.name);
expect(defs.some((d) => d.type === 'definition.class' && d.name === 'Force')).toBe(true);
expect(names).toContain('launchKernel');
});
it('captures C++ typedef anonymous structs, enums, and enumerators', async () => {
await loadLanguage(SupportedLanguages.CPlusPlus);
const code = `

View file

@ -67,6 +67,16 @@ describe('IncludeExtractor', () => {
expect(providers[0].contractId).toBe('include::utils/helper.hpp');
});
it('registers .cuh CUDA headers as providers', async () => {
writeFile('src/force/nep.cuh', '#pragma once\nclass NEP {};');
const contracts = await extractor.extract(null, tmpDir, makeRepo(tmpDir));
const providers = contracts.filter((c) => c.role === 'provider');
expect(providers).toHaveLength(1);
expect(providers[0].contractId).toBe('include::src/force/nep.cuh');
});
it('does not register .cpp files as providers', async () => {
writeFile('src/main.cpp', 'int main() { return 0; }');
writeFile('src/utils.h', '#pragma once');
@ -237,6 +247,22 @@ int main() { return 0; }`,
expect(consumers).toHaveLength(0);
});
it('scans .cu files for includes and resolves local .cuh headers', async () => {
writeFile('include/kernel.cuh', '#pragma once\nvoid launchKernel();');
writeFile(
'src/main.cu',
`#include "include/kernel.cuh"
#include "external/gpu_runtime.cuh"
void launch() { launchKernel(); }`,
);
const contracts = await extractor.extract(null, tmpDir, makeRepo(tmpDir));
const consumers = contracts.filter((c) => c.role === 'consumer');
expect(consumers).toHaveLength(1);
expect(consumers[0].contractId).toBe('include::external/gpu_runtime.cuh');
});
it('resolves locally when include omits extension and a matching .h exists', async () => {
writeFile('foo/bar.h', '#pragma once');
writeFile('src/main.cpp', '#include "foo/bar"\nint main(){return 0;}');

View file

@ -68,9 +68,12 @@ describe('getLanguageFromFilename', () => {
});
describe('C++', () => {
it.each(['.cpp', '.cc', '.cxx', '.h', '.hpp', '.hxx', '.hh'])('detects %s files', (ext) => {
expect(getLanguageFromFilename(`file${ext}`)).toBe(SupportedLanguages.CPlusPlus);
});
it.each(['.cpp', '.cc', '.cxx', '.h', '.hpp', '.hxx', '.hh', '.cu', '.cuh'])(
'detects %s files',
(ext) => {
expect(getLanguageFromFilename(`file${ext}`)).toBe(SupportedLanguages.CPlusPlus);
},
);
});
describe('C#', () => {
@ -172,6 +175,11 @@ describe('getProviderForFile', () => {
SupportedLanguages.PHP,
);
});
it('routes CUDA C++ source and header files to the C++ provider', () => {
expect(getProviderForFile('src/kernels/integrate.cu')?.id).toBe(SupportedLanguages.CPlusPlus);
expect(getProviderForFile('src/force/nep.cuh')?.id).toBe(SupportedLanguages.CPlusPlus);
});
});
describe('isBuiltInOrNoise', () => {