Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
32 changes: 30 additions & 2 deletions apps/agor-daemon/src/register-hooks.ts
Original file line number Diff line number Diff line change
Expand Up @@ -798,6 +798,34 @@ export function registerHooks(ctx: RegisterHooksContext): void {
return context;
}

let servers: MCPServer[] = [];
if (Array.isArray(context.result)) {
servers = context.result;
} else if (context.result?.data && Array.isArray(context.result.data)) {
servers = context.result.data;
} else if (context.result?.mcp_server_id) {
servers = [context.result];
}

const oauthServers = servers.filter((s) => s.auth?.type === 'oauth');
const userTokenRepo = new UserMCPOAuthTokenRepository(db);
const tokenMap = new Map<string, import('@agor/core/db').UserMCPOAuthToken>();

if (oauthServers.length > 0) {
const serverIds = oauthServers.map(
(s) => s.mcp_server_id as import('@agor/core/types').MCPServerID
);
const batchedTokens = await userTokenRepo.getTokensForServers(
serverIds,
userId as import('@agor/core/types').UserID
);

for (const token of batchedTokens) {
const key = `${token.mcp_server_id}:${token.user_id ?? 'shared'}`;
tokenMap.set(key, token);
}
}

const injectToken = async (server: MCPServer) => {
if (server.auth?.type !== 'oauth') {
return server;
Expand All @@ -811,8 +839,8 @@ 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 key = `${server.mcp_server_id}:${tokenUserId ?? 'shared'}`;
const row = tokenMap.get(key);

if (!row) {
console.log(
Expand Down
37 changes: 36 additions & 1 deletion packages/core/src/db/repositories/user-mcp-oauth-tokens.ts
Original file line number Diff line number Diff line change
Expand Up @@ -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 {
Expand Down Expand Up @@ -80,6 +80,41 @@ export class UserMCPOAuthTokenRepository {
* Look up the token row for a (user, server) pair. Pass `null` for userId
* to read the shared-mode row.
*/
/**
* Look up tokens for a list of servers. Fetches both shared-mode and per-user
* tokens for the given userId.
*/
async getTokensForServers(
serverIds: MCPServerID[],
userId: UserID | null
): Promise<UserMCPOAuthToken[]> {
if (serverIds.length === 0) return [];

try {
const conditions: any[] = [inArray(userMcpOauthTokens.mcp_server_id, serverIds)];

if (userId) {
conditions.push(
or(isNull(userMcpOauthTokens.user_id), eq(userMcpOauthTokens.user_id, userId))
);
} else {
conditions.push(isNull(userMcpOauthTokens.user_id));
}

const rows = await select(this.db)
.from(userMcpOauthTokens)
.where(and(...conditions))
.all();

return rows.map(rowToToken);
} catch (error) {
throw new RepositoryError(
`Failed to get OAuth tokens for servers: ${error instanceof Error ? error.message : String(error)}`,
error
);
}
}

async getToken(userId: UserID | null, serverId: MCPServerID): Promise<UserMCPOAuthToken | null> {
try {
const row = await select(this.db)
Expand Down