如何在Apollo Server首请求前统计数据源调用次数并拦截超限请求?
针对你需要在第一个数据源请求发送前统计总调用次数、超过阈值则阻止所有请求的需求,结合Apollo Server和RESTDataSource的特性,提供以下几种可行方案:
方案1:静态分析查询AST + 解析器元数据标记(适用于固定调用次数的解析器)
如果你的解析器中数据源调用次数是固定的(不依赖运行时参数或请求结果),可以通过给解析器添加元数据,在请求开始时解析查询AST统计总调用次数。
步骤1:给解析器添加调用次数元数据
修改resolvers.ts,为每个查询字段的解析器添加元数据,标记该字段会触发的API调用次数:
export const resolvers: Resolvers = { Query: { getCompanies: { resolve: (_, __, { dataSources }) => dataSources.companyDatasource.getCompanies(), meta: { apiCalls: 1 } // 固定1次调用 }, getCompany: { resolve: (_, { name }, { dataSources }) => dataSources.companyDatasource.getCompanyByName(name), meta: { apiCalls: 1 } }, getCompanyCEOs: { resolve: async (_, { name }, { dataSources }) => { const company = await dataSources.companyDatasource.getCompanyByName(name); return dataSources.companyDatasource.getCEOs(company.id); }, meta: { apiCalls: 2 } // 固定2次调用 } } };
步骤2:在Apollo插件中解析查询并检查阈值
在main.ts中添加Apollo插件,利用requestDidStart生命周期钩子解析查询AST,累加元数据中的调用次数,与阈值对比:
import { visit } from 'graphql'; import { ApolloError } from 'apollo-server'; // 自定义方法:从数据库获取API调用阈值 const getThresholdFromDB = async () => { // 替换为你的实际数据库查询逻辑 return 1; }; const server = new ApolloServer({ typeDefs: schema, schema, resolvers, dataSources, cache: 'bounded', plugins: [ { async requestDidStart(requestContext) { const { document } = requestContext.request; let totalApiCalls = 0; // 遍历查询AST,收集所有要执行的字段 visit(document, { Field(node) { const fieldResolver = resolvers.Query?.[node.name.value]; if (fieldResolver?.meta?.apiCalls) { totalApiCalls += fieldResolver.meta.apiCalls; } } }); // 获取阈值并检查 const threshold = await getThresholdFromDB(); if (totalApiCalls > threshold) { throw new ApolloError( `当前请求将触发${totalApiCalls}次API调用,超过允许的阈值${threshold}`, 'EXCEEDED_API_CALL_LIMIT' ); } } } ] }); await server.start();
局限:无法处理动态调用次数的场景(比如解析器根据运行时结果决定调用N次数据源)。
方案2:自定义数据源包装器(支持动态调用场景)
如果存在动态调用次数的解析器,可通过包装数据源方法,先收集所有待执行的请求,统一检查阈值后再执行。
步骤1:创建追踪数据源的包装类
import CompanyDatasource from './company.datasource'; class TrackedCompanyDatasource { private originalDS: CompanyDatasource; private pendingCalls: Array<() => Promise<any>> = []; private threshold: number; private hasChecked = false; constructor(originalDS: CompanyDatasource, threshold: number) { this.originalDS = originalDS; this.threshold = threshold; } // 包装所有数据源方法 getCompanies() { this.pendingCalls.push(() => this.originalDS.getCompanies()); return this.executeWhenChecked(this.pendingCalls.length - 1); } getCompanyByName(name: string) { this.pendingCalls.push(() => this.originalDS.getCompanyByName(name)); return this.executeWhenChecked(this.pendingCalls.length - 1); } getCEOs(companyId: string) { this.pendingCalls.push(() => this.originalDS.getCEOs(companyId)); return this.executeWhenChecked(this.pendingCalls.length - 1); } private async executeWhenChecked(index: number) { // 第一次调用时,等待当前事件循环结束,确保所有同步解析逻辑完成 if (!this.hasChecked) { await new Promise(resolve => setImmediate(resolve)); if (this.pendingCalls.length > this.threshold) { throw new Error(`待执行API调用次数${this.pendingCalls.length}超过阈值${this.threshold}`); } this.hasChecked = true; } // 执行实际请求 return this.pendingCalls[index](); } }
步骤2:在dataSources中使用包装后的实例
const dataSources = async () => { const threshold = await getThresholdFromDB(); return { companyDatasource: new TrackedCompanyDatasource(new CompanyDatasource(), threshold) }; };
说明:该方案利用setImmediate等待解析器的同步逻辑完成,收集所有待执行的请求。但对于完全动态的依赖(比如循环调用数据源,次数由前一个请求结果决定),仍无法提前统计,此时需要在循环前先检查剩余阈值。
方案3:修改RESTDataSource底层逻辑(对解析器透明)
直接修改CompanyDatasource的请求发送逻辑,将所有请求加入队列,统一检查阈值后执行:
export default class CompanyDatasource extends RESTDataSource { private pendingRequests: Array<{ request: Request; resolve: (value: any) => void; reject: (reason?: any) => void; }> = []; private hasChecked = false; private threshold: number; constructor(threshold: number) { super(); this.baseURL = 'your_api_base_url'; this.threshold = threshold; } override async willSendRequest(request: Request) { // 拦截请求,加入队列 return new Promise((resolve, reject) => { this.pendingRequests.push({ request, resolve, reject }); if (!this.hasChecked) { this.checkAndProcessRequests(); } }); } private async checkAndProcessRequests() { this.hasChecked = true; // 等待当前事件循环结束,收集所有待请求 await new Promise(resolve => setImmediate(resolve)); if (this.pendingRequests.length > this.threshold) { // 拒绝所有请求 this.pendingRequests.forEach(({ reject }) => reject(new Error(`请求次数${this.pendingRequests.length}超过阈值${this.threshold}`)) ); this.pendingRequests = []; return; } // 批量执行所有请求 for (const { request, resolve, reject } of this.pendingRequests) { try { const response = await super.sendRequest(request); resolve(response); } catch (err) { reject(err); } } this.pendingRequests = []; } // 原有数据源方法不变 async getCompanies() { return this.get(`/companies`); } async getCompanyByName(name: string) { return this.get(`/companies?name=${encodeURIComponent(name)}`); } async getCEOs(companyId: string) { return this.get(`/companies/${companyId}/ceos`); } }
然后在main.ts的dataSources中传入阈值:
const dataSources = async () => { const threshold = await getThresholdFromDB(); return { companyDatasource: new CompanyDatasource(threshold) }; };
注意:对于依赖前一个请求结果的动态调用(如getCompanyCEOs),由于第二个请求在第一个请求完成后才会被加入队列,该方案无法提前统计这类动态请求的次数。此时建议结合方案1的元数据标记,预先预留足够的额度,或调整团队约定为“每次调用前检查剩余额度,超过则阻止后续请求”。
内容的提问来源于stack exchange,提问作者Kit A.

