From 6a7aff76023c4fe636ee9347143c3151a663f785 Mon Sep 17 00:00:00 2001 From: Donach <39565367+Donach@users.noreply.github.com> Date: Sat, 1 Aug 2026 07:07:31 +0000 Subject: [PATCH] Fix N+1 query in injectPerUserOAuthTokens hook Refactored the `injectPerUserOAuthTokens` hook to extract all `mcp_server_id`s from the results upfront and issue a single batched database lookup using the new `getTokensForServers` method in `UserMCPOAuthTokenRepository`. This avoids N+1 queries when resolving OAuth tokens over multiple servers. --- .jules/bolt.md | 3 + apps/agor-daemon/src/register-hooks.ts | 26 ++++- .../db/repositories/user-mcp-oauth-tokens.ts | 38 ++++++- plan.md | 14 +++ test-repo.sh | 105 ++++++++++++++++++ test-repo2.sh | 50 +++++++++ update-repo.sh | 49 ++++++++ 7 files changed, 281 insertions(+), 4 deletions(-) create mode 100644 plan.md create mode 100644 test-repo.sh create mode 100644 test-repo2.sh create mode 100644 update-repo.sh diff --git a/.jules/bolt.md b/.jules/bolt.md index 8c4087435e..79dac7f60f 100644 --- a/.jules/bolt.md +++ b/.jules/bolt.md @@ -4,3 +4,6 @@ ## 2026-05-16 - [Batch FeathersJS user fetches with $in operator to fix N+1 query] **Learning:** FeathersJS allows passing `$in` clauses through the query parameter (e.g. `user_id: { $in: ownerIds }`). When writing custom Feathers service logic, you can easily parse this array and pass it to Drizzle's `inArray()` to perform a batched query, instead of looping over `service.get(id)` causing N+1 database roundtrips. **Action:** When implementing or updating custom Feathers `find()` methods, extract and parse the `$in` parameters to support batched Drizzle `inArray()` lookups, and always replace `Promise.all(ids.map(id => service.get(id)))` with a single batched `find()` call. +## 2025-06-25 - [Batch FeathersJS user fetches with `getTokensForServers` to fix N+1 query] +**Learning:** By utilizing `inArray()` in Drizzle, we can fetch all related OAuth tokens (per-user and shared modes) for an array of servers in one go, dramatically reducing database round-trips compared to calling `getToken(server.mcp_server_id)` inside a `Promise.all()` loop. +**Action:** When working on FeathersJS hooks that enrich multiple records (e.g. iterating over `context.result`), always extract the array of entity IDs and run a single batched database lookup query via `inArray()`, then resolve records against this pre-fetched Map synchronously. diff --git a/apps/agor-daemon/src/register-hooks.ts b/apps/agor-daemon/src/register-hooks.ts index 05321d3669..c3825ef675 100755 --- a/apps/agor-daemon/src/register-hooks.ts +++ b/apps/agor-daemon/src/register-hooks.ts @@ -798,6 +798,28 @@ export function registerHooks(ctx: RegisterHooksContext): void { return context; } + // Handle both single result and array/paginated results + const servers = Array.isArray(context.result) + ? context.result + : context.result?.data && Array.isArray(context.result.data) + ? context.result.data + : context.result?.mcp_server_id + ? [context.result] + : []; + + const serverIds = servers + .filter((s: MCPServer) => s.auth?.type === 'oauth') + .map((s: MCPServer) => s.mcp_server_id); + + const userTokenRepo = new UserMCPOAuthTokenRepository(db); + const tokenMap = + serverIds.length > 0 + ? await userTokenRepo.getTokensForServers( + userId as import('@agor/core/types').UserID, + serverIds + ) + : new Map(); + const injectToken = async (server: MCPServer) => { if (server.auth?.type !== 'oauth') { return server; @@ -811,8 +833,7 @@ export function registerHooks(ctx: RegisterHooksContext): void { mode === 'per_user' ? (userId as import('@agor/core/types').UserID) : null; try { - const userTokenRepo = new UserMCPOAuthTokenRepository(db); - const row = await userTokenRepo.getToken(tokenUserId, server.mcp_server_id); + const row = tokenMap.get(`${server.mcp_server_id}:${tokenUserId || ''}`); if (!row) { console.log( @@ -875,7 +896,6 @@ export function registerHooks(ctx: RegisterHooksContext): void { return server; }; - // Handle both single result and array/paginated results if (Array.isArray(context.result)) { context.result = await Promise.all(context.result.map(injectToken)); } else if (context.result?.data && Array.isArray(context.result.data)) { diff --git a/packages/core/src/db/repositories/user-mcp-oauth-tokens.ts b/packages/core/src/db/repositories/user-mcp-oauth-tokens.ts index c5ea01f893..a1ad623e05 100644 --- a/packages/core/src/db/repositories/user-mcp-oauth-tokens.ts +++ b/packages/core/src/db/repositories/user-mcp-oauth-tokens.ts @@ -11,7 +11,7 @@ */ import type { MCPServerID, UserID } from '@agor/core/types'; -import { and, eq, isNull } from 'drizzle-orm'; +import { and, eq, inArray, isNull, or } from 'drizzle-orm'; import type { Database } from '../client'; import { deleteFrom, insert, select, update } from '../database-wrapper'; import { @@ -240,6 +240,42 @@ export class UserMCPOAuthTokenRepository { } } + async getTokensForServers( + userId: UserID | null, + serverIds: MCPServerID[] + ): Promise> { + if (serverIds.length === 0) { + return new Map(); + } + try { + const conditions = []; + if (userId === null) { + conditions.push(isNull(userMcpOauthTokens.user_id)); + } else { + conditions.push( + or(eq(userMcpOauthTokens.user_id, userId), isNull(userMcpOauthTokens.user_id)) + ); + } + + const rows = await select(this.db) + .from(userMcpOauthTokens) + .where(and(inArray(userMcpOauthTokens.mcp_server_id, serverIds), ...conditions)) + .all(); + + const result = new Map(); + for (const row of rows) { + const token = rowToToken(row); + result.set(`${token.mcp_server_id}:${token.user_id || ''}`, token); + } + return result; + } catch (error) { + throw new RepositoryError( + `Failed to get tokens for servers: ${error instanceof Error ? error.message : String(error)}`, + error + ); + } + } + async listForUser(userId: UserID): Promise { try { const rows = await select(this.db) diff --git a/plan.md b/plan.md new file mode 100644 index 0000000000..464490a3de --- /dev/null +++ b/plan.md @@ -0,0 +1,14 @@ +1. **Add `getTokensForServers` batched method to `UserMCPOAuthTokenRepository`** + - Update `packages/core/src/db/repositories/user-mcp-oauth-tokens.ts`. + - The method uses `inArray` to query `mcp_server_id` for multiple servers in one go, filtering by `user_id` appropriately (either the specified `userId` or `null` for shared modes). + - We will combine both queries by requesting either the given `userId` or `null`. + - Store the fetched tokens in a Map using a compound key `:` (where `user_id` is an empty string if null). + +2. **Modify `injectPerUserOAuthTokens` hook to utilize the batched method** + - Update `apps/agor-daemon/src/register-hooks.ts`. + - Before applying the `injectToken` map function, extract all `mcp_server_id`s from the result context. + - Call `getTokensForServers` once to retrieve a map of tokens. + - Use the pre-fetched map inside `injectToken` to lookup `oauth_access_token` synchronously, instead of `userTokenRepo.getToken` hitting the database on every server. + +3. **Complete pre-commit steps to ensure proper testing, verification, review, and reflection are done.** + - Run `pnpm lint:fix` and `pnpm test` to ensure stability. diff --git a/test-repo.sh b/test-repo.sh new file mode 100644 index 0000000000..cc9d0377d4 --- /dev/null +++ b/test-repo.sh @@ -0,0 +1,105 @@ +<<<<<<< SEARCH + // Handle both single result and array/paginated results + if (Array.isArray(context.result)) { + context.result = await Promise.all(context.result.map(injectToken)); + } else if (context.result?.data && Array.isArray(context.result.data)) { + context.result.data = await Promise.all(context.result.data.map(injectToken)); + } else if (context.result?.mcp_server_id) { + context.result = await injectToken(context.result); + } +======= + // Handle both single result and array/paginated results + const servers = Array.isArray(context.result) + ? context.result + : context.result?.data && Array.isArray(context.result.data) + ? context.result.data + : context.result?.mcp_server_id + ? [context.result] + : []; + + if (servers.length > 0) { + const serverIds = servers + .filter((s: MCPServer) => s.auth?.type === 'oauth') + .map((s: MCPServer) => s.mcp_server_id); + + const userTokenRepo = new UserMCPOAuthTokenRepository(db); + const tokenMap = serverIds.length > 0 ? await userTokenRepo.getTokensForServers(userId as import('@agor/core/types').UserID, serverIds) : new Map(); + + const injectToken = async (server: MCPServer) => { + if (server.auth?.type !== 'oauth') { + return server; + } + + const mode = server.auth.oauth_mode ?? 'per_user'; + const tokenUserId = mode === 'per_user' ? userId : null; + const row = tokenMap.get(`${server.mcp_server_id}:${tokenUserId || ''}`); + + if (!row) { + console.log( + `[MCP OAuth] No token row for user=${tokenUserId ?? ''} server=${server.name}` + ); + return server; + } + + try { + // JIT refresh — see `refreshAndPersistToken` for mutexing + invalid_grant cleanup. + let accessToken = row.oauth_access_token; + let expiresAt = row.oauth_token_expires_at; + const { needsRefresh, refreshAndPersistToken, InvalidGrantError } = await import( + '@agor/core/tools/mcp/oauth-refresh' + ); + if (needsRefresh(row.oauth_token_expires_at) && row.oauth_refresh_token) { + console.log(`[MCP OAuth] Token near/past expiry for ${server.name} — refreshing`); + try { + accessToken = await refreshAndPersistToken({ + db, + userId: tokenUserId as import('@agor/core/types').UserID | null, + mcpServerId: server.mcp_server_id, + }); + // Re-read to pick up the rotated expiry for the UI. + const fresh = await userTokenRepo.getToken(tokenUserId as import('@agor/core/types').UserID | null, server.mcp_server_id); + if (fresh) expiresAt = fresh.oauth_token_expires_at; + } catch (refreshErr) { + if (refreshErr instanceof InvalidGrantError) { + console.warn( + `[MCP OAuth] invalid_grant refreshing ${server.name} — user must re-auth` + ); + return server; + } + // Transient error: fall through with the stale access_token. The + // MCP call may still succeed or fail cleanly at the transport. + console.warn( + `[MCP OAuth] Refresh failed for ${server.name} (using stale token):`, + refreshErr instanceof Error ? refreshErr.message : refreshErr + ); + } + } + + return { + ...server, + auth: { + ...server.auth, + oauth_access_token: accessToken, + oauth_token_expires_at: + expiresAt instanceof Date ? expiresAt.getTime() : (expiresAt ?? undefined), + }, + }; + } catch (error) { + console.warn( + `[MCP OAuth] Failed to resolve OAuth token for ${server.name}:`, + error instanceof Error ? error.message : error + ); + } + + return server; + }; + + if (Array.isArray(context.result)) { + context.result = await Promise.all(context.result.map(injectToken)); + } else if (context.result?.data && Array.isArray(context.result.data)) { + context.result.data = await Promise.all(context.result.data.map(injectToken)); + } else if (context.result?.mcp_server_id) { + context.result = await injectToken(context.result); + } + } +>>>>>>> REPLACE diff --git a/test-repo2.sh b/test-repo2.sh new file mode 100644 index 0000000000..dbf74cce13 --- /dev/null +++ b/test-repo2.sh @@ -0,0 +1,50 @@ +<<<<<<< SEARCH + const injectToken = async (server: MCPServer) => { + if (server.auth?.type !== 'oauth') { + return server; + } + + // Tokens for both modes live in user_mcp_oauth_tokens: + // - per_user → row keyed by (userId, serverId) + // - shared → row keyed by (NULL, serverId) + const mode = server.auth.oauth_mode ?? 'per_user'; + const tokenUserId: import('@agor/core/types').UserID | null = + mode === 'per_user' ? (userId as import('@agor/core/types').UserID) : null; + + try { + const userTokenRepo = new UserMCPOAuthTokenRepository(db); + const row = await userTokenRepo.getToken(tokenUserId, server.mcp_server_id); +======= + // Handle both single result and array/paginated results + const servers = Array.isArray(context.result) + ? context.result + : context.result?.data && Array.isArray(context.result.data) + ? context.result.data + : context.result?.mcp_server_id + ? [context.result] + : []; + + const serverIds = servers + .filter((s: MCPServer) => s.auth?.type === 'oauth') + .map((s: MCPServer) => s.mcp_server_id); + + const userTokenRepo = new UserMCPOAuthTokenRepository(db); + const tokenMap = serverIds.length > 0 + ? await userTokenRepo.getTokensForServers(userId as import('@agor/core/types').UserID, serverIds) + : new Map(); + + const injectToken = async (server: MCPServer) => { + if (server.auth?.type !== 'oauth') { + return server; + } + + // Tokens for both modes live in user_mcp_oauth_tokens: + // - per_user → row keyed by (userId, serverId) + // - shared → row keyed by (NULL, serverId) + const mode = server.auth.oauth_mode ?? 'per_user'; + const tokenUserId: import('@agor/core/types').UserID | null = + mode === 'per_user' ? (userId as import('@agor/core/types').UserID) : null; + + try { + const row = tokenMap.get(`${server.mcp_server_id}:${tokenUserId || ''}`); +>>>>>>> REPLACE diff --git a/update-repo.sh b/update-repo.sh new file mode 100644 index 0000000000..c707cf4355 --- /dev/null +++ b/update-repo.sh @@ -0,0 +1,49 @@ +<<<<<<< SEARCH +import { and, eq, isNull } from 'drizzle-orm'; +======= +import { and, eq, inArray, isNull, or } from 'drizzle-orm'; +>>>>>>> REPLACE +<<<<<<< SEARCH + async listForUser(userId: UserID): Promise { +======= + async getTokensForServers( + userId: UserID | null, + serverIds: MCPServerID[] + ): Promise> { + if (serverIds.length === 0) { + return new Map(); + } + try { + const conditions = []; + if (userId === null) { + conditions.push(isNull(userMcpOauthTokens.user_id)); + } else { + conditions.push(or(eq(userMcpOauthTokens.user_id, userId), isNull(userMcpOauthTokens.user_id))); + } + + const rows = await select(this.db) + .from(userMcpOauthTokens) + .where( + and( + inArray(userMcpOauthTokens.mcp_server_id, serverIds), + ...conditions + ) + ) + .all(); + + const result = new Map(); + for (const row of rows) { + const token = rowToToken(row); + result.set(`${token.mcp_server_id}:${token.user_id || ''}`, token); + } + return result; + } catch (error) { + throw new RepositoryError( + `Failed to get tokens for servers: ${error instanceof Error ? error.message : String(error)}`, + error + ); + } + } + + async listForUser(userId: UserID): Promise { +>>>>>>> REPLACE