fix(group): match weak thrift java consumers

This commit is contained in:
liyipeng06 2026-04-28 16:03:59 +08:00
parent d9a1ac3f41
commit d39404b2d7
5 changed files with 232 additions and 3 deletions

View file

@ -10,6 +10,15 @@ import type { ThriftDetection, ThriftLanguagePlugin } from './types.js';
const GENERATED_MEMBER_TYPES = new Set(['Iface', 'Client']);
const SERVICE_TYPE_RE = /^[A-Z][A-Za-z0-9]*(?:Service|Management)$/;
interface VariableBinding {
name: string;
serviceName: string;
scopeStart: number;
scopeEnd: number;
declarationEnd: number;
scopeSize: number;
}
const VARIABLE_PATTERNS = compilePatterns({
name: 'java-thrift-variables',
language: Java,
@ -139,12 +148,71 @@ function methodNamesInClassBody(body: Parser.SyntaxNode): string[] {
return names;
}
function nearestAncestor(node: Parser.SyntaxNode, types: Set<string>): Parser.SyntaxNode | null {
let current: Parser.SyntaxNode | null = node;
while (current) {
if (types.has(current.type)) return current;
current = current.parent;
}
return null;
}
function bindingScope(varNode: Parser.SyntaxNode): {
scope: Parser.SyntaxNode;
declarationEnd: number;
} | null {
const declaration = nearestAncestor(
varNode,
new Set(['field_declaration', 'local_variable_declaration', 'formal_parameter']),
);
if (!declaration) return null;
if (declaration.type === 'field_declaration') {
const classBody = nearestAncestor(declaration, new Set(['class_body']));
if (!classBody) return null;
return { scope: classBody, declarationEnd: 0 };
}
if (declaration.type === 'formal_parameter') {
const callable = nearestAncestor(
declaration,
new Set(['method_declaration', 'constructor_declaration']),
);
if (!callable) return null;
return { scope: callable, declarationEnd: 0 };
}
const block = nearestAncestor(declaration, new Set(['block']));
if (!block) return null;
return { scope: block, declarationEnd: declaration.endIndex };
}
function resolveServiceForReceiver(
bindings: VariableBinding[],
receiver: string,
callNode: Parser.SyntaxNode,
): string | null {
const callStart = callNode.startIndex;
const candidates = bindings.filter(
(binding) =>
binding.name === receiver &&
binding.scopeStart <= callStart &&
callStart <= binding.scopeEnd &&
binding.declarationEnd <= callStart,
);
candidates.sort((a, b) => {
if (a.scopeSize !== b.scopeSize) return a.scopeSize - b.scopeSize;
return b.declarationEnd - a.declarationEnd;
});
return candidates[0]?.serviceName ?? null;
}
export const JAVA_THRIFT_PLUGIN: ThriftLanguagePlugin = {
name: 'java-thrift',
language: Java,
scan(tree) {
const out: ThriftDetection[] = [];
const variables = new Map<string, string>();
const bindings: VariableBinding[] = [];
for (const match of runCompiledPatterns(VARIABLE_PATTERNS, tree)) {
const serviceNode = match.captures.service;
@ -153,14 +221,25 @@ export const JAVA_THRIFT_PLUGIN: ThriftLanguagePlugin = {
const memberNode = match.meta.scoped ? match.captures.member : undefined;
const serviceName = serviceFromType(serviceNode.text, memberNode?.text);
if (!serviceName) continue;
variables.set(varNode.text, serviceName);
const scope = bindingScope(varNode);
if (!scope) continue;
bindings.push({
name: varNode.text,
serviceName,
scopeStart: scope.scope.startIndex,
scopeEnd: scope.scope.endIndex,
declarationEnd: scope.declarationEnd,
scopeSize: scope.scope.endIndex - scope.scope.startIndex,
});
}
for (const match of runCompiledPatterns(CALL_PATTERNS, tree)) {
const receiver = match.captures.receiver?.text;
const methodName = match.captures.method?.text;
const callNode = match.captures.receiver?.parent;
if (!receiver || !methodName) continue;
const serviceName = variables.get(receiver);
if (!callNode) continue;
const serviceName = resolveServiceForReceiver(bindings, receiver, callNode);
if (!serviceName) continue;
out.push({
role: 'consumer',

View file

@ -90,6 +90,31 @@ function findMatchingKeys(contractId: string, index: Map<string, StoredContract[
return matches;
}
if (normalized.startsWith('thrift::')) {
const rest = normalized.substring('thrift::'.length);
const slashIdx = rest.indexOf('/');
if (slashIdx > 0) {
const service = rest.substring(0, slashIdx);
const method = rest.substring(slashIdx + 1);
if (!service.includes('.') && method && method !== '*') {
const matches: string[] = [];
for (const key of index.keys()) {
if (!key.startsWith('thrift::') || key.endsWith('/*')) continue;
const providerRest = key.substring('thrift::'.length);
const providerSlashIdx = providerRest.indexOf('/');
if (providerSlashIdx < 0) continue;
const providerService = providerRest.substring(0, providerSlashIdx);
const providerMethod = providerRest.substring(providerSlashIdx + 1);
if (providerMethod !== method) continue;
if (providerService === service || providerService.endsWith('.' + service)) {
matches.push(key);
}
}
return matches;
}
}
}
return [];
}

View file

@ -482,6 +482,49 @@ describe('runWildcardMatch', () => {
expect(matched[0].contractId).toBe('thrift::OrderService/*');
});
it('matches bare thrift service method to a package-qualified thrift provider method', () => {
const consumer = makeThriftContract(
'thrift::OrderService/PlaceOrder',
'consumer',
'frontend',
);
const provider = makeThriftContract(
'thrift::billing.v1.OrderService/PlaceOrder',
'provider',
'backend',
);
const providerIndex = buildProviderIndex([provider]);
const { matched, unmatched } = runExactMatch([consumer, provider], providerIndex);
expect(matched).toHaveLength(1);
expect(matched[0].type).toBe('thrift');
expect(matched[0].matchType).toBe('exact');
expect(matched[0].contractId).toBe('thrift::OrderService/PlaceOrder');
expect(matched[0].from.repo).toBe('frontend');
expect(matched[0].to.repo).toBe('backend');
expect(unmatched).toHaveLength(0);
});
it('does not match bare thrift service method to a different provider method', () => {
const consumer = makeThriftContract(
'thrift::OrderService/PlaceOrder',
'consumer',
'frontend',
);
const provider = makeThriftContract(
'thrift::billing.v1.OrderService/GetOrderStatus',
'provider',
'backend',
);
const providerIndex = buildProviderIndex([provider]);
const { matched, unmatched } = runExactMatch([consumer, provider], providerIndex);
expect(matched).toHaveLength(0);
expect(unmatched).toEqual([consumer, provider]);
});
it('does not match a thrift wildcard to a gRPC provider', () => {
const consumer = makeThriftContract('thrift::OrderService/*', 'consumer', 'frontend');
const provider = makeGrpcContract(

View file

@ -301,6 +301,44 @@ describe('syncGroup', () => {
expect(result.unmatched).toEqual([provider]);
});
it('matches weak thrift method consumers to namespace-qualified providers during sync', async () => {
const config = makeConfig({ 'app/provider': 'provider-repo', 'app/consumer': 'consumer-repo' });
const provider: StoredContract = {
contractId: 'thrift::billing.v1.OrderService/PlaceOrder',
type: 'thrift',
role: 'provider',
symbolUid: 'uid-provider-place-order',
symbolRef: { filePath: 'idl/order.thrift', name: 'OrderService.PlaceOrder' },
symbolName: 'OrderService.PlaceOrder',
confidence: 0.85,
meta: {},
repo: 'app/provider',
};
const consumer: StoredContract = {
contractId: 'thrift::OrderService/PlaceOrder',
type: 'thrift',
role: 'consumer',
symbolUid: 'uid-consumer-place-order',
symbolRef: { filePath: 'src/BillingWorkflow.java', name: 'orderService.PlaceOrder' },
symbolName: 'orderService.PlaceOrder',
confidence: 0.45,
meta: {},
repo: 'app/consumer',
};
const result = await syncGroup(config, {
extractorOverride: async () => [provider, consumer],
skipWrite: true,
});
expect(result.crossLinks).toHaveLength(1);
expect(result.crossLinks[0].matchType).toBe('exact');
expect(result.crossLinks[0].contractId).toBe('thrift::OrderService/PlaceOrder');
expect(result.crossLinks[0].from.repo).toBe('app/consumer');
expect(result.crossLinks[0].to.repo).toBe('app/provider');
expect(result.unmatched).toHaveLength(0);
});
it('dedupes duplicate wildcard cross-links during sync', async () => {
const config = makeConfig({ 'app/provider': 'provider-repo', 'app/consumer': 'consumer-repo' });
const provider: StoredContract = {

View file

@ -278,6 +278,50 @@ class BillingWorker {
);
});
it('test_extract_java_thrift_consumers_resolve_receiver_by_nearest_scope', async () => {
writeFile(
'idl/order.thrift',
`namespace java billing.v1
service OrderService {
PlaceOrderResponse PlaceOrder(1: PlaceOrderRequest request)
}
service InvoiceService {
Invoice CreateInvoice(1: string orderId)
}`,
);
writeFile(
'src/main/java/example/BillingWorker.java',
`package example;
class BillingWorker {
void submitOrder(OrderService.Iface client, PlaceOrderRequest request) throws Exception {
client.PlaceOrder(request);
}
void submitInvoice() throws Exception {
InvoiceService.Client client = new InvoiceService.Client(null);
client.CreateInvoice("order-1");
}
}`,
);
const contracts = await extractor.extract(null, tmpDir, makeRepo(tmpDir));
const consumers = contracts
.filter((c) => c.role === 'consumer')
.sort((a, b) => a.contractId.localeCompare(b.contractId));
expect(consumers.map((c) => c.contractId)).toEqual([
'thrift::billing.v1.InvoiceService/CreateInvoice',
'thrift::billing.v1.OrderService/PlaceOrder',
]);
expect(consumers.map((c) => c.symbolName).sort()).toEqual([
'client.CreateInvoice',
'client.PlaceOrder',
]);
});
it('test_extract_java_thrift_providers_from_iface_and_service_implements', async () => {
writeFile(
'idl/order.thrift',