From 89816e398b9bf830a419baca93837b2a2cafecb3 Mon Sep 17 00:00:00 2001 From: Savannah Ostrowski Date: Tue, 25 Aug 2026 10:15:45 -0700 Subject: [PATCH 1/2] =?UTF-8?q?=F0=9F=90=9B=20Resolve=20imported=20custom?= =?UTF-8?q?=20APIRouter=20subclasses?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- src/core/analyzer.ts | 11 ++- src/core/extractors.ts | 16 ++-- src/core/internal.ts | 8 +- src/core/routerResolver.ts | 89 ++++++++++++++----- src/test/core/analyzer.test.ts | 23 +++++ src/test/core/routerResolver.test.ts | 31 +++++++ .../custom-router/auth/testing_router.py | 5 ++ src/test/fixtures/custom-router/main.py | 5 ++ .../services/test/test_service.py | 8 ++ .../fixtures/inference-boundaries/main.py | 6 ++ .../fixtures/inference-boundaries/subapp.py | 8 ++ src/test/testUtils.ts | 8 ++ 12 files changed, 176 insertions(+), 42 deletions(-) create mode 100644 src/test/fixtures/custom-router/auth/testing_router.py create mode 100644 src/test/fixtures/custom-router/main.py create mode 100644 src/test/fixtures/custom-router/services/test/test_service.py create mode 100644 src/test/fixtures/inference-boundaries/main.py create mode 100644 src/test/fixtures/inference-boundaries/subapp.py diff --git a/src/core/analyzer.ts b/src/core/analyzer.ts index 736834e..3acdd38 100644 --- a/src/core/analyzer.ts +++ b/src/core/analyzer.ts @@ -5,10 +5,10 @@ import type { Tree } from "web-tree-sitter" import { logError } from "../utils/logger" import { + callAssignmentExtractor, collectRecognizedNames, collectStringVariables, decoratorExtractor, - factoryCallExtractor, getNodesByType, importExtractor, includeRouterExtractor, @@ -44,12 +44,11 @@ export function analyzeTree(tree: Tree, filePath: string): FileAnalysis { const decoratedDefs = nodesByType.get("decorated_definition") ?? [] const routes = decoratedDefs.flatMap(decoratorExtractor) - // Get all router assignments + // Get module-level call assignments and recognized router assignments const assignments = nodesByType.get("assignment") ?? [] const { fastAPINames, apiRouterNames } = collectRecognizedNames(nodesByType) - const knownConstructors = new Set([...fastAPINames, ...apiRouterNames]) - const factoryCalls = assignments - .map((node) => factoryCallExtractor(node, knownConstructors)) + const callAssignments = assignments + .map(callAssignmentExtractor) .filter(notNull) const routers = assignments .map((node) => routerExtractor(node, apiRouterNames, fastAPINames)) @@ -89,7 +88,7 @@ export function analyzeTree(tree: Tree, filePath: string): FileAnalysis { includeRouters, mounts, imports, - factoryCalls, + callAssignments, } } diff --git a/src/core/extractors.ts b/src/core/extractors.ts index e464043..f244ee8 100644 --- a/src/core/extractors.ts +++ b/src/core/extractors.ts @@ -4,7 +4,7 @@ import type { Node } from "web-tree-sitter" import type { - FactoryCallInfo, + CallAssignmentInfo, ImportedName, ImportInfo, IncludeRouterInfo, @@ -601,10 +601,7 @@ export function mountExtractor(node: Node): MountInfo | null { } } -export function factoryCallExtractor( - node: Node, - knownConstructors: Set, -): FactoryCallInfo | null { +export function callAssignmentExtractor(node: Node): CallAssignmentInfo | null { if (node.type !== "assignment") { return null } @@ -620,11 +617,6 @@ export function factoryCallExtractor( return null } - const functionName = functionNode.text - if (knownConstructors.has(functionName)) { - return null - } - // Skip function and class-local variables to avoid false positives if ( hasAncestor(node, "function_definition") || @@ -635,7 +627,9 @@ export function factoryCallExtractor( return { variableName: variableNameNode.text, - functionName: functionName, + callee: functionNode.text, + line: node.startPosition.row + 1, + column: node.startPosition.column, } } diff --git a/src/core/internal.ts b/src/core/internal.ts index a282b62..a92b4ee 100644 --- a/src/core/internal.ts +++ b/src/core/internal.ts @@ -78,9 +78,11 @@ export interface MountInfo { app: string } -export interface FactoryCallInfo { +export interface CallAssignmentInfo { variableName: string - functionName: string + callee: string + line: number + column: number } export interface FileAnalysis { @@ -90,7 +92,7 @@ export interface FileAnalysis { includeRouters: IncludeRouterInfo[] mounts: MountInfo[] imports: ImportInfo[] - factoryCalls: FactoryCallInfo[] + callAssignments: CallAssignmentInfo[] } export interface RouterNode { diff --git a/src/core/routerResolver.ts b/src/core/routerResolver.ts index c5f2d98..46b1b9f 100644 --- a/src/core/routerResolver.ts +++ b/src/core/routerResolver.ts @@ -19,6 +19,17 @@ interface ResolutionContext { visited: Set } +interface ResolveReferenceOptions { + reference: string + analysis: FileAnalysis + currentFileUri: string + ctx: ResolutionContext +} + +interface InternalResolveReferenceOptions extends ResolveReferenceOptions { + kind: "includedRouter" | "mountedApp" +} + /** * Finds the main FastAPI app or APIRouter in the list of routers. * If targetVariable is specified, only returns the router with that variable name. @@ -37,6 +48,27 @@ function findAppRouter( ) } +function inferIncludedRouter( + analysis: FileAnalysis, + variableName: string, +): RouterInfo | undefined { + const assignment = analysis.callAssignments.find( + (candidate) => candidate.variableName === variableName, + ) + if (!assignment) { + return undefined + } + + return { + variableName: assignment.variableName, + type: "APIRouter", + prefix: "", + tags: [], + line: assignment.line, + column: assignment.column, + } +} + function createRouterNode( router: RouterInfo, routes: RouteInfo[], @@ -55,6 +87,14 @@ function createRouterNode( } } +function resolveIncludedRouter(options: ResolveReferenceOptions) { + return resolveReference({ ...options, kind: "includedRouter" }) +} + +function resolveMountedApp(options: ResolveReferenceOptions) { + return resolveReference({ ...options, kind: "mountedApp" }) +} + async function processIncludeRouters( analysis: FileAnalysis, ownerRouter: RouterNode, @@ -68,12 +108,12 @@ async function processIncludeRouters( log( `Resolving include_router: ${include.router} (prefix: ${include.prefix || "none"})`, ) - const childRouter = await resolveRouterReference( - include.router, + const childRouter = await resolveIncludedRouter({ + reference: include.router, analysis, currentFileUri, ctx, - ) + }) if (childRouter) { // Merge tags from include_router call with the router's own tags if (include.tags.length > 0) { @@ -198,18 +238,18 @@ async function buildRouterGraphInternal( // `app = FastAPI()` and `app.include_router(...)` inside `create_app` are visible // when analyzing the factory file directly. if (!appRouter && targetVariable) { - const factoryCall = analysis.factoryCalls.find( - (fc) => fc.variableName === targetVariable, + const callAssignment = analysis.callAssignments.find( + (assignment) => assignment.variableName === targetVariable, ) - if (factoryCall) { + if (callAssignment) { const matchingImport = analysis.imports.find((imp) => - imp.names.includes(factoryCall.functionName), + imp.names.includes(callAssignment.callee), ) if (matchingImport) { const namedImport = matchingImport.namedImports.find( - (ni) => (ni.alias ?? ni.name) === factoryCall.functionName, + (ni) => (ni.alias ?? ni.name) === callAssignment.callee, ) - const originalName = namedImport?.name ?? factoryCall.functionName + const originalName = namedImport?.name ?? callAssignment.callee const factoryFileUri = await resolveNamedImport( { modulePath: matchingImport.modulePath, @@ -252,12 +292,12 @@ async function buildRouterGraphInternal( // Process mount() calls for subapps for (const mount of analysis.mounts) { - const childRouter = await resolveRouterReference( - mount.app, + const childRouter = await resolveMountedApp({ + reference: mount.app, analysis, - resolvedEntryUri, + currentFileUri: resolvedEntryUri, ctx, - ) + }) if (childRouter) { rootRouter.children.push({ router: childRouter, @@ -277,12 +317,13 @@ async function buildRouterGraphInternal( * Handles both simple references (e.g., "router") and dotted references * (e.g., "api_routes.router" where api_routes is an imported module). */ -async function resolveRouterReference( - reference: string, - analysis: FileAnalysis, - currentFileUri: string, - ctx: ResolutionContext, -): Promise { +async function resolveReference({ + reference, + analysis, + currentFileUri, + ctx, + kind, +}: InternalResolveReferenceOptions): Promise { const { projectRootUri, parser, fs, visited } = ctx const parts = reference.split(".") const moduleName = parts[0] @@ -358,9 +399,13 @@ async function resolveRouterReference( } // Find the router with the matching variable name - const targetRouter = importedAnalysis.routers.find( - (r) => r.variableName === attributeName, - ) + const targetRouter = + importedAnalysis.routers.find( + (router) => router.variableName === attributeName, + ) ?? + (kind === "includedRouter" + ? inferIncludedRouter(importedAnalysis, attributeName) + : undefined) if (targetRouter) { const visitedKey = `${importedFileUri}#${attributeName}` diff --git a/src/test/core/analyzer.test.ts b/src/test/core/analyzer.test.ts index df572be..4952135 100644 --- a/src/test/core/analyzer.test.ts +++ b/src/test/core/analyzer.test.ts @@ -63,6 +63,29 @@ router = APIRouter(prefix="/api") assert.strictEqual(result.routers[1].prefix, "/api") }) + test("records module-level direct call assignments as neutral facts", () => { + const code = ` +from auth.testing_router import ProtectedRouter + +def build_router(): + local_router = ProtectedRouter() + +router = ProtectedRouter() +` + const tree = parse(code) + const result = analyzeTree(tree, "/test/file.py") + + assert.strictEqual(result.routers.length, 0) + assert.deepStrictEqual(result.callAssignments, [ + { + variableName: "router", + callee: "ProtectedRouter", + line: 7, + column: 0, + }, + ]) + }) + test("extracts include_router calls", () => { const code = ` app.include_router(users.router, prefix="/users") diff --git a/src/test/core/routerResolver.test.ts b/src/test/core/routerResolver.test.ts index 7ba7a9d..13a4d8a 100644 --- a/src/test/core/routerResolver.test.ts +++ b/src/test/core/routerResolver.test.ts @@ -441,6 +441,21 @@ suite("routerResolver", () => { assert.strictEqual(mountChild.router.type, "FastAPI") }) + test("does not infer an APIRouter for a mounted call assignment", async () => { + const result = await buildRouterGraph( + fixtures.inferenceBoundaries.mainPy, + parser, + fixtures.inferenceBoundaries.root, + nodeFileSystem, + ) + + assert.ok(result) + const mountedChild = result.children.find( + (child) => child.prefix === "/mounted", + ) + assert.ok(mountedChild?.router.type !== "APIRouter") + }) + test("merges tags from include_router call with router tags", async () => { const result = await buildRouterGraph( fixtures.standard.mainPy, @@ -676,6 +691,22 @@ suite("routerResolver", () => { assert.ok(methods.includes("post")) }) + test("resolves a router instantiated from an imported APIRouter subclass", async () => { + const result = await buildRouterGraph( + fixtures.customRouter.mainPy, + parser, + fixtures.customRouter.root, + nodeFileSystem, + ) + + assert.ok(result) + const router = result.children.find( + (child) => child.router.variableName === "router", + )?.router + assert.ok(router) + assert.ok(router.routes.some((route) => route.path === "/items")) + }) + test("resolves aliased FastAPI and APIRouter class imports", async () => { const result = await buildRouterGraph( fixtures.aliasedClass.mainPy, diff --git a/src/test/fixtures/custom-router/auth/testing_router.py b/src/test/fixtures/custom-router/auth/testing_router.py new file mode 100644 index 0000000..c8593e0 --- /dev/null +++ b/src/test/fixtures/custom-router/auth/testing_router.py @@ -0,0 +1,5 @@ +from fastapi import APIRouter + + +class ProtectedRouter(APIRouter): + pass diff --git a/src/test/fixtures/custom-router/main.py b/src/test/fixtures/custom-router/main.py new file mode 100644 index 0000000..5a2e566 --- /dev/null +++ b/src/test/fixtures/custom-router/main.py @@ -0,0 +1,5 @@ +from fastapi import FastAPI +from services.test import test_service + +app = FastAPI() +app.include_router(test_service.router) diff --git a/src/test/fixtures/custom-router/services/test/test_service.py b/src/test/fixtures/custom-router/services/test/test_service.py new file mode 100644 index 0000000..bcb89b2 --- /dev/null +++ b/src/test/fixtures/custom-router/services/test/test_service.py @@ -0,0 +1,8 @@ +from auth.testing_router import ProtectedRouter + +router = ProtectedRouter() + + +@router.get("/items") +def list_items(): + return [] diff --git a/src/test/fixtures/inference-boundaries/main.py b/src/test/fixtures/inference-boundaries/main.py new file mode 100644 index 0000000..17e406d --- /dev/null +++ b/src/test/fixtures/inference-boundaries/main.py @@ -0,0 +1,6 @@ +from fastapi import FastAPI + +import subapp + +app = FastAPI() +app.mount("/mounted", subapp.app) diff --git a/src/test/fixtures/inference-boundaries/subapp.py b/src/test/fixtures/inference-boundaries/subapp.py new file mode 100644 index 0000000..b30c749 --- /dev/null +++ b/src/test/fixtures/inference-boundaries/subapp.py @@ -0,0 +1,8 @@ +from fastapi import FastAPI + + +def create_app(): + return FastAPI() + + +app = create_app() diff --git a/src/test/testUtils.ts b/src/test/testUtils.ts index d53a1f3..cd69e8d 100644 --- a/src/test/testUtils.ts +++ b/src/test/testUtils.ts @@ -43,6 +43,10 @@ export const fixtures = { root: uri(join(fixturesPath, "custom-subclass")), mainPy: uri(join(fixturesPath, "custom-subclass", "main.py")), }, + customRouter: { + root: uri(join(fixturesPath, "custom-router")), + mainPy: uri(join(fixturesPath, "custom-router", "main.py")), + }, errorCases: { root: uri(join(fixturesPath, "error-cases")), mainPy: uri(join(fixturesPath, "error-cases", "main.py")), @@ -56,6 +60,10 @@ export const fixtures = { root: uri(join(fixturesPath, "flat")), mainPy: uri(join(fixturesPath, "flat", "main.py")), }, + inferenceBoundaries: { + root: uri(join(fixturesPath, "inference-boundaries")), + mainPy: uri(join(fixturesPath, "inference-boundaries", "main.py")), + }, monorepo: { workspaceRoot: uri(join(fixturesPath, "monorepo")), projectRoot: uri(join(fixturesPath, "monorepo", "service")), From a7ccd6627eae10fcf61cdb91631c7969035db115 Mon Sep 17 00:00:00 2001 From: Savannah Ostrowski Date: Tue, 25 Aug 2026 11:13:07 -0700 Subject: [PATCH 2/2] =?UTF-8?q?=F0=9F=90=9B=20Preserve=20metadata=20for=20?= =?UTF-8?q?inferred=20routers?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- src/core/analyzer.ts | 14 ++---- src/core/extractors.ts | 45 ++++++++++++------- src/core/internal.ts | 2 + src/core/routerResolver.ts | 4 +- src/test/core/analyzer.test.ts | 4 +- src/test/core/routerResolver.test.ts | 19 ++++++++ .../services/test/test_service.py | 2 +- .../fixtures/inference-boundaries/main.py | 1 + .../fixtures/inference-boundaries/subapp.py | 1 + 9 files changed, 62 insertions(+), 30 deletions(-) diff --git a/src/core/analyzer.ts b/src/core/analyzer.ts index 3acdd38..1ca1d42 100644 --- a/src/core/analyzer.ts +++ b/src/core/analyzer.ts @@ -68,17 +68,11 @@ export function analyzeTree(tree: Tree, filePath: string): FileAnalysis { const stringVariables = collectStringVariables(nodesByType) - for (const route of routes) { - route.path = resolveVariables(route.path, stringVariables) + for (const item of [...routes, ...mounts]) { + item.path = resolveVariables(item.path, stringVariables) } - for (const router of routers) { - router.prefix = resolveVariables(router.prefix, stringVariables) - } - for (const ir of includeRouters) { - ir.prefix = resolveVariables(ir.prefix, stringVariables) - } - for (const mount of mounts) { - mount.path = resolveVariables(mount.path, stringVariables) + for (const item of [...routers, ...callAssignments, ...includeRouters]) { + item.prefix = resolveVariables(item.prefix, stringVariables) } return { diff --git a/src/core/extractors.ts b/src/core/extractors.ts index f244ee8..6979fe1 100644 --- a/src/core/extractors.ts +++ b/src/core/extractors.ts @@ -265,6 +265,30 @@ function extractTags(listNode: Node): string[] { .filter((v): v is string => v !== null) } +function extractRouterCallMetadata(callNode: Node): { + prefix: string + tags: string[] +} { + let prefix = "" + let tags: string[] = [] + const argumentsNode = callNode.childForFieldName("arguments") + for (const child of argumentsNode?.namedChildren ?? []) { + if (child.type !== "keyword_argument") { + continue + } + const argName = child.childForFieldName("name")?.text + const argValue = child.childForFieldName("value") + + if (argName === "prefix" && argValue) { + prefix = extractPathFromNode(argValue) + } else if (argName === "tags" && argValue?.type === "list") { + tags = extractTags(argValue) + } + } + + return { prefix, tags } +} + export function routerExtractor( node: Node, apiRouterNames?: Set, @@ -298,22 +322,7 @@ export function routerExtractor( return null } - let prefix = "" - let tags: string[] = [] - const argumentsNode = valueNode.childForFieldName("arguments") - for (const child of argumentsNode?.namedChildren ?? []) { - if (child.type !== "keyword_argument") { - continue - } - const argName = child.childForFieldName("name")?.text - const argValue = child.childForFieldName("value") - - if (argName === "prefix" && argValue) { - prefix = extractPathFromNode(argValue) - } else if (argName === "tags" && argValue?.type === "list") { - tags = extractTags(argValue) - } - } + const { prefix, tags } = extractRouterCallMetadata(valueNode) return { variableName: variableNameNode.text, @@ -625,9 +634,13 @@ export function callAssignmentExtractor(node: Node): CallAssignmentInfo | null { return null } + const { prefix, tags } = extractRouterCallMetadata(valueNode) + return { variableName: variableNameNode.text, callee: functionNode.text, + prefix, + tags, line: node.startPosition.row + 1, column: node.startPosition.column, } diff --git a/src/core/internal.ts b/src/core/internal.ts index a92b4ee..2af74a5 100644 --- a/src/core/internal.ts +++ b/src/core/internal.ts @@ -81,6 +81,8 @@ export interface MountInfo { export interface CallAssignmentInfo { variableName: string callee: string + prefix: string + tags: string[] line: number column: number } diff --git a/src/core/routerResolver.ts b/src/core/routerResolver.ts index 46b1b9f..127b85e 100644 --- a/src/core/routerResolver.ts +++ b/src/core/routerResolver.ts @@ -62,8 +62,8 @@ function inferIncludedRouter( return { variableName: assignment.variableName, type: "APIRouter", - prefix: "", - tags: [], + prefix: assignment.prefix, + tags: assignment.tags, line: assignment.line, column: assignment.column, } diff --git a/src/test/core/analyzer.test.ts b/src/test/core/analyzer.test.ts index 4952135..5ed7fde 100644 --- a/src/test/core/analyzer.test.ts +++ b/src/test/core/analyzer.test.ts @@ -70,7 +70,7 @@ from auth.testing_router import ProtectedRouter def build_router(): local_router = ProtectedRouter() -router = ProtectedRouter() +router = ProtectedRouter(prefix="/protected", tags=["protected"]) ` const tree = parse(code) const result = analyzeTree(tree, "/test/file.py") @@ -80,6 +80,8 @@ router = ProtectedRouter() { variableName: "router", callee: "ProtectedRouter", + prefix: "/protected", + tags: ["protected"], line: 7, column: 0, }, diff --git a/src/test/core/routerResolver.test.ts b/src/test/core/routerResolver.test.ts index 13a4d8a..d4d9cf7 100644 --- a/src/test/core/routerResolver.test.ts +++ b/src/test/core/routerResolver.test.ts @@ -2,6 +2,7 @@ import * as assert from "node:assert" import { Parser } from "../../core/parser" import { findProjectRoot } from "../../core/pathUtils" import { buildRouterGraph } from "../../core/routerResolver" +import { routerNodeToAppDefinition } from "../../core/transformer" import { fixtures, fixturesPath, @@ -456,6 +457,22 @@ suite("routerResolver", () => { assert.ok(mountedChild?.router.type !== "APIRouter") }) + test("does not display an empty non-router include assignment", async () => { + const result = await buildRouterGraph( + fixtures.inferenceBoundaries.mainPy, + parser, + fixtures.inferenceBoundaries.root, + nodeFileSystem, + ) + + assert.ok(result) + const app = routerNodeToAppDefinition( + result, + fixtures.inferenceBoundaries.root, + ) + assert.ok(app.routers.every((router) => router.name !== "not_router")) + }) + test("merges tags from include_router call with router tags", async () => { const result = await buildRouterGraph( fixtures.standard.mainPy, @@ -704,6 +721,8 @@ suite("routerResolver", () => { (child) => child.router.variableName === "router", )?.router assert.ok(router) + assert.strictEqual(router.prefix, "/service") + assert.deepStrictEqual(router.tags, ["service"]) assert.ok(router.routes.some((route) => route.path === "/items")) }) diff --git a/src/test/fixtures/custom-router/services/test/test_service.py b/src/test/fixtures/custom-router/services/test/test_service.py index bcb89b2..aba7bde 100644 --- a/src/test/fixtures/custom-router/services/test/test_service.py +++ b/src/test/fixtures/custom-router/services/test/test_service.py @@ -1,6 +1,6 @@ from auth.testing_router import ProtectedRouter -router = ProtectedRouter() +router = ProtectedRouter(prefix="/service", tags=["service"]) @router.get("/items") diff --git a/src/test/fixtures/inference-boundaries/main.py b/src/test/fixtures/inference-boundaries/main.py index 17e406d..14ec176 100644 --- a/src/test/fixtures/inference-boundaries/main.py +++ b/src/test/fixtures/inference-boundaries/main.py @@ -3,4 +3,5 @@ import subapp app = FastAPI() +app.include_router(subapp.not_router) app.mount("/mounted", subapp.app) diff --git a/src/test/fixtures/inference-boundaries/subapp.py b/src/test/fixtures/inference-boundaries/subapp.py index b30c749..290a0d8 100644 --- a/src/test/fixtures/inference-boundaries/subapp.py +++ b/src/test/fixtures/inference-boundaries/subapp.py @@ -6,3 +6,4 @@ def create_app(): app = create_app() +not_router = object()