Skip to content

Commit 59d0e43

Browse files
committed
feat(web): add recovery for connector tool-load failures
1 parent fde95a5 commit 59d0e43

25 files changed

Lines changed: 561 additions & 100 deletions

packages/web/src/app/api/(client)/client.ts

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -320,7 +320,7 @@ export const getOffers = async (): Promise<OffersResponse | ServiceError> => {
320320
return result as OffersResponse | ServiceError;
321321
}
322322

323-
export const connectMcpToAsk = async (body: { serverId: string; returnTo?: string }): Promise<ConnectMcpResponse | ServiceError> => {
323+
export const connectMcpToAsk = async (body: { serverId: string; returnTo?: string; forceAuthorization?: boolean }): Promise<ConnectMcpResponse | ServiceError> => {
324324
const result = await fetch('/api/ee/askmcp/connect', {
325325
method: 'POST',
326326
headers: {

packages/web/src/app/api/(server)/ee/askmcp/connect/route.test.ts

Lines changed: 55 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -48,7 +48,7 @@ vi.mock('@ai-sdk/mcp', () => ({
4848
const { POST } = await import('./route');
4949
const { getMcpOAuthReturnToFromState } = await import('@/ee/features/chat/mcp/mcpOAuthReturnTo');
5050

51-
function createRequest(body: { serverId: string; returnTo?: string } = { serverId: 'server-1' }) {
51+
function createRequest(body: { serverId: string; returnTo?: string; forceAuthorization?: boolean } = { serverId: 'server-1' }) {
5252
return new NextRequest('https://sourcebot.example.com/api/ee/askmcp/connect', {
5353
method: 'POST',
5454
headers: { 'content-type': 'application/json' },
@@ -229,6 +229,60 @@ describe('POST /api/ee/askmcp/connect', () => {
229229
});
230230
});
231231

232+
test('forces an interactive OAuth redirect for reconnect recovery', async () => {
233+
const prisma = createPrismaMock();
234+
const tx = createTransactionMock();
235+
tx.userMcpServer.findUnique.mockResolvedValue({
236+
tokens: 'encrypted:{"access_token":"stale-token"}',
237+
codeVerifier: null,
238+
state: null,
239+
});
240+
mocks.authContext = {
241+
org: { id: 1 },
242+
user: { id: 'user-1' },
243+
prisma,
244+
};
245+
mocks.unsafePrisma.$transaction.mockImplementation(async (callback, _options) => callback(tx));
246+
mocks.mcpAuth.mockImplementation(async (provider) => {
247+
await expect(provider.tokens()).resolves.toBeUndefined();
248+
provider.authorizationUrl = 'https://oauth.example.com/authorize';
249+
return 'REDIRECT';
250+
});
251+
252+
const response = await POST(createRequest({
253+
serverId: 'server-1',
254+
returnTo: '/chat/abc123',
255+
forceAuthorization: true,
256+
}));
257+
258+
expect(await response.json()).toEqual({
259+
authorizationUrl: 'https://oauth.example.com/authorize',
260+
});
261+
});
262+
263+
test('does not report a forced reconnect as successful without an OAuth redirect', async () => {
264+
const prisma = createPrismaMock();
265+
const tx = createTransactionMock();
266+
mocks.authContext = {
267+
org: { id: 1 },
268+
user: { id: 'user-1' },
269+
prisma,
270+
};
271+
mocks.unsafePrisma.$transaction.mockImplementation(async (callback, _options) => callback(tx));
272+
mocks.mcpAuth.mockResolvedValue('AUTHORIZED');
273+
274+
const response = await POST(createRequest({
275+
serverId: 'server-1',
276+
returnTo: '/chat/abc123',
277+
forceAuthorization: true,
278+
}));
279+
280+
expect(response.status).toBe(502);
281+
expect(await response.json()).toMatchObject({
282+
message: 'Could not start connector reauthorization.',
283+
});
284+
});
285+
232286
test('ignores unsafe return paths', async () => {
233287
const prisma = createPrismaMock();
234288
const tx = createTransactionMock();

packages/web/src/app/api/(server)/ee/askmcp/connect/route.ts

Lines changed: 15 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -24,6 +24,7 @@ import { getEnabledMcpOAuthScopeNames } from '@/ee/features/chat/mcp/oauthScopeU
2424
const bodySchema = z.object({
2525
serverId: z.string(),
2626
returnTo: z.string().optional(),
27+
forceAuthorization: z.boolean().optional().default(false),
2728
});
2829
const logger = createLogger('mcp-connect');
2930
const MCP_AUTH_FETCH_TIMEOUT_MS = Math.min(env.SOURCEBOT_MCP_TOOL_CALL_TIMEOUT_MS, 30000);
@@ -146,6 +147,7 @@ export const POST = apiHandler(async (request: NextRequest) => {
146147
callbackReturnTo,
147148
allowClientRegistration: true,
148149
requestedOAuthScopes: getEnabledMcpOAuthScopeNames(mcpServer.oauthScopes),
150+
forceAuthorization: parsed.data.forceAuthorization,
149151
});
150152

151153
let authResult: Awaited<ReturnType<typeof mcpAuth>>;
@@ -200,7 +202,7 @@ export const POST = apiHandler(async (request: NextRequest) => {
200202
throw error;
201203
}
202204

203-
if (connectResult.authResult === 'AUTHORIZED') {
205+
if (connectResult.authResult === 'AUTHORIZED' && !parsed.data.forceAuthorization) {
204206
// Already has valid tokens (e.g., refreshed)
205207
void captureEvent('ask_mcp_connector_connection_completed', {
206208
...eventProperties,
@@ -209,6 +211,18 @@ export const POST = apiHandler(async (request: NextRequest) => {
209211
return { authorizationUrl: null } satisfies ConnectMcpResponse;
210212
}
211213

214+
if (connectResult.authResult === 'AUTHORIZED') {
215+
void captureEvent('ask_mcp_connector_connection_failed', {
216+
...eventProperties,
217+
failureReason: 'missing_authorization_url',
218+
});
219+
throw new ServiceErrorException({
220+
statusCode: StatusCodes.BAD_GATEWAY,
221+
errorCode: ErrorCode.UNEXPECTED_ERROR,
222+
message: 'Could not start connector reauthorization.',
223+
});
224+
}
225+
212226
if (!connectResult.authorizationUrl) {
213227
void captureEvent('ask_mcp_connector_connection_failed', {
214228
...eventProperties,

packages/web/src/ee/features/chat/agent.test.ts

Lines changed: 44 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -321,6 +321,50 @@ describe('createMessageStream approval continuation', () => {
321321
});
322322
});
323323

324+
test('streams the connector ID when its tools fail to load', async () => {
325+
const { getConnectedMcpClients } = await import('@/ee/features/chat/mcp/mcpClientFactory');
326+
const { getMcpTools } = await import('@/ee/features/chat/mcp/mcpToolSets');
327+
vi.mocked(getConnectedMcpClients).mockResolvedValueOnce([
328+
{ serverId: 'server-linear', serverName: 'Linear' },
329+
] as never);
330+
vi.mocked(getMcpTools).mockResolvedValueOnce({
331+
tools: {},
332+
failedServers: [{ serverId: 'server-linear', serverName: 'Linear' }],
333+
serverFaviconUrls: {},
334+
toolDisplayNames: {},
335+
cleanup: vi.fn(),
336+
});
337+
mockAi.streamText.mockReturnValue(createFakeStreamResult());
338+
339+
await createMessageStream({
340+
chatId: 'chat-id',
341+
messages: [createUserMessage()],
342+
selectedRepos: [],
343+
disabledMcpServerIds: [],
344+
prisma: {},
345+
model: {},
346+
modelName: 'test-model',
347+
promptCacheStrategy: noopStrategy,
348+
onFinish: vi.fn(),
349+
onError: () => 'error',
350+
userId: 'user-id',
351+
orgId: 1,
352+
} as unknown as Parameters<typeof createMessageStream>[0]);
353+
354+
const execute = mockAi.latestCreateUIMessageStreamOptions?.execute;
355+
if (!execute) {
356+
throw new Error('Expected createUIMessageStream to capture execute callback.');
357+
}
358+
359+
const write = vi.fn();
360+
await execute({ writer: { merge: vi.fn(), write } });
361+
362+
expect(write).toHaveBeenCalledWith({
363+
type: 'data-mcp-failed-server',
364+
data: { serverId: 'server-linear', serverName: 'Linear' },
365+
});
366+
});
367+
324368
test.each([
325369
['dynamic', dynamicApprovalRespondedPart],
326370
['static', staticApprovalRespondedPart],

packages/web/src/ee/features/chat/agent.ts

Lines changed: 5 additions & 5 deletions
Original file line numberDiff line numberDiff line change
@@ -393,10 +393,10 @@ export const createMessageStream = async ({
393393
data: { modelToolName, rawToolName },
394394
});
395395
},
396-
onMcpServerFailed: (serverName) => {
396+
onMcpServerFailed: (server) => {
397397
writer.write({
398398
type: 'data-mcp-failed-server',
399-
data: { serverName },
399+
data: server,
400400
});
401401
},
402402
onMcpAuthRequired: (failure) => {
@@ -536,7 +536,7 @@ interface AgentOptions {
536536
onWriteSource: (source: Source) => void;
537537
onMcpServerDiscovered: (sanitizedName: string, faviconUrl: string) => void;
538538
onMcpToolDiscovered: (modelToolName: string, rawToolName: string) => void;
539-
onMcpServerFailed: (serverName: string) => void;
539+
onMcpServerFailed: (server: { serverId: string; serverName: string }) => void;
540540
// Fired at most once per connector per response when a tool call fails
541541
// with a reconnect-required authentication failure.
542542
onMcpAuthRequired: (failure: McpToolAuthFailure) => void;
@@ -646,8 +646,8 @@ const createAgentStream = async ({
646646
}
647647
}
648648

649-
for (const serverName of mcpToolSetsObj.failedServers) {
650-
onMcpServerFailed(serverName);
649+
for (const server of mcpToolSetsObj.failedServers) {
650+
onMcpServerFailed(server);
651651
}
652652

653653
const mcpRegistry = buildMcpToolRegistry(mcpToolSetsObj.tools);

packages/web/src/ee/features/chat/components/chatThread/chatThread.tsx

Lines changed: 63 additions & 23 deletions
Original file line numberDiff line numberDiff line change
@@ -30,7 +30,7 @@ import { isServiceError } from '@/lib/utils';
3030
import { NotConfiguredErrorBanner } from '@/features/chat/components/notConfiguredErrorBanner';
3131
import { McpServerIconContext, McpServerIconMap, McpToolNameContext, McpToolNameMap } from '../../mcpDisplayMetadataContext';
3232
import { McpReconnectContext } from '../../mcpReconnectContext';
33-
import { McpAuthRequiredData, useMcpReconnectController } from './useMcpReconnectController';
33+
import { McpAuthRequiredData, McpServerLoadFailureData, useMcpReconnectController } from './useMcpReconnectController';
3434
import { McpReconnectBanner } from './mcpReconnectBanner';
3535
import { ToolApprovalProvider } from '../../toolApprovalContext';
3636
import useCaptureEvent from '@/hooks/useCaptureEvent';
@@ -122,19 +122,7 @@ export const ChatThread = ({
122122
return map;
123123
});
124124

125-
const [failedMcpServers, setFailedMcpServers] = useState<string[]>(() => {
126-
const names: string[] = [];
127-
initialMessages?.forEach((message) => {
128-
message.parts
129-
.filter((part) => part.type === 'data-mcp-failed-server')
130-
.forEach((part) => {
131-
if (!names.includes(part.data.serverName)) {
132-
names.push(part.data.serverName);
133-
}
134-
});
135-
});
136-
return names;
137-
});
125+
const [failedMcpServers, setFailedMcpServers] = useState<McpServerLoadFailureData[]>([]);
138126
const [isFailedMcpBannerVisible, setIsFailedMcpBannerVisible] = useState(false);
139127

140128
const { selectedLanguageModel } = useSelectedLanguageModel();
@@ -150,6 +138,37 @@ export const ChatThread = ({
150138
useEffect(() => { modelRef.current = selectedLanguageModel; }, [selectedLanguageModel]);
151139
useEffect(() => { disabledMcpRef.current = disabledMcpServerIds; }, [disabledMcpServerIds]);
152140

141+
const forceDisableMcpServer = useCallback((serverId: string) => {
142+
if (disabledMcpRef.current.includes(serverId)) {
143+
return;
144+
}
145+
146+
const nextDisabledServerIds = [...disabledMcpRef.current, serverId];
147+
disabledMcpRef.current = nextDisabledServerIds;
148+
onDisabledMcpServerIdsChange(nextDisabledServerIds);
149+
}, [onDisabledMcpServerIdsChange]);
150+
151+
const reenableMcpServer = useCallback((serverId: string) => {
152+
if (!disabledMcpRef.current.includes(serverId)) {
153+
return;
154+
}
155+
156+
const nextDisabledServerIds = disabledMcpRef.current.filter((id) => id !== serverId);
157+
disabledMcpRef.current = nextDisabledServerIds;
158+
onDisabledMcpServerIdsChange(nextDisabledServerIds);
159+
}, [onDisabledMcpServerIdsChange]);
160+
161+
const registerFailedMcpServer = useCallback((server: McpServerLoadFailureData) => {
162+
setFailedMcpServers((prev) => {
163+
if (prev.some((candidate) => candidate.serverId === server.serverId)) {
164+
return prev;
165+
}
166+
return [...prev, server];
167+
});
168+
setIsFailedMcpBannerVisible(true);
169+
forceDisableMcpServer(server.serverId);
170+
}, [forceDisableMcpServer]);
171+
153172
const getTransportBody = useCallback(() => ({
154173
selectedSearchScopes: searchScopesRef.current,
155174
languageModel: modelRef.current,
@@ -160,6 +179,7 @@ export const ChatThread = ({
160179
// messages and status), so transient auth-required events received in
161180
// onData are forwarded to it through a ref.
162181
const onMcpAuthRequiredRef = useRef<((data: McpAuthRequiredData) => void) | null>(null);
182+
const onMcpServerLoadFailedRef = useRef<((data: McpServerLoadFailureData) => void) | null>(null);
163183

164184
// Transport with dynamic body, resolved on every request, including auto-resends
165185
// triggered by sendAutomaticallyWhen after tool approval.
@@ -203,13 +223,8 @@ export const ChatThread = ({
203223
}));
204224
}
205225
if (dataPart.type === 'data-mcp-failed-server') {
206-
setFailedMcpServers((prev) => {
207-
if (prev.includes(dataPart.data.serverName)) {
208-
return prev;
209-
}
210-
return [...prev, dataPart.data.serverName];
211-
});
212-
setIsFailedMcpBannerVisible(true);
226+
registerFailedMcpServer(dataPart.data);
227+
onMcpServerLoadFailedRef.current?.(dataPart.data);
213228
}
214229
if (dataPart.type === 'data-mcp-auth-required') {
215230
onMcpAuthRequiredRef.current?.(dataPart.data);
@@ -290,6 +305,7 @@ export const ChatThread = ({
290305
const {
291306
contextValue: mcpReconnectContextValue,
292307
onAuthRequired: onMcpAuthRequired,
308+
onServerLoadFailed: onMcpServerLoadFailed,
293309
} = useMcpReconnectController({
294310
status,
295311
messages,
@@ -302,7 +318,30 @@ export const ChatThread = ({
302318

303319
useEffect(() => {
304320
onMcpAuthRequiredRef.current = onMcpAuthRequired;
305-
}, [onMcpAuthRequired]);
321+
onMcpServerLoadFailedRef.current = onMcpServerLoadFailed;
322+
}, [onMcpAuthRequired, onMcpServerLoadFailed]);
323+
324+
const handledLoadReconnectsRef = useRef(new Set<string>());
325+
useEffect(() => {
326+
for (const state of Object.values(mcpReconnectContextValue.reconnectStates)) {
327+
if (state.source !== 'tool-load') {
328+
continue;
329+
}
330+
331+
if (state.status === 'reconnected') {
332+
if (handledLoadReconnectsRef.current.has(state.serverId)) {
333+
continue;
334+
}
335+
handledLoadReconnectsRef.current.add(state.serverId);
336+
setFailedMcpServers((prev) => prev.filter((server) => server.serverId !== state.serverId));
337+
reenableMcpServer(state.serverId);
338+
continue;
339+
}
340+
341+
handledLoadReconnectsRef.current.delete(state.serverId);
342+
registerFailedMcpServer({ serverId: state.serverId, serverName: state.serverName });
343+
}
344+
}, [mcpReconnectContextValue.reconnectStates, reenableMcpServer, registerFailedMcpServer]);
306345

307346
// When the chat is finished, refresh the page to update the chat history.
308347
const prevStatus = usePrevious(status);
@@ -439,7 +478,7 @@ export const ChatThread = ({
439478
/>
440479
)}
441480
<McpFailedServersBanner
442-
serverNames={failedMcpServers}
481+
servers={failedMcpServers}
443482
isVisible={isFailedMcpBannerVisible}
444483
onClose={() => setIsFailedMcpBannerVisible(false)}
445484
/>
@@ -546,6 +585,7 @@ export const ChatThread = ({
546585
isContextSelectorOpen={isContextSelectorOpen}
547586
onContextSelectorOpenChanged={setIsContextSelectorOpen}
548587
disabledMcpServerIds={disabledMcpServerIds}
588+
unavailableMcpServerIds={failedMcpServers.map((server) => server.serverId)}
549589
onDisabledMcpServerIdsChange={onDisabledMcpServerIdsChange}
550590
isAuthenticated={isAuthenticated}
551591
/>

0 commit comments

Comments
 (0)