From c70657528250485c108ce41486068d7bcc5f40d0 Mon Sep 17 00:00:00 2001 From: Praveen Sampath Date: Thu, 30 Jul 2026 09:37:25 +0000 Subject: [PATCH] fix: address critical OAuth MCP foundation issues --- connectors/google/src/connector.rs | 1 + sdk/typescript/src/connector.ts | 28 +++++--- sdk/typescript/src/mcp-adapter.ts | 25 ++++--- sdk/typescript/tests/mcp-adapter.test.ts | 39 +++++++++++ services/ai/streaming/generate.py | 65 ++++++++++++------- services/ai/tests/helpers.py | 23 +++++++ .../integration/test_chat_stream_lifecycle.py | 8 ++- .../src/remote_mcp/gateway.rs | 1 + .../db/repositories/service_credentials.rs | 8 ++- 9 files changed, 153 insertions(+), 45 deletions(-) diff --git a/connectors/google/src/connector.rs b/connectors/google/src/connector.rs index 6c61c7b2..3b6b5003 100644 --- a/connectors/google/src/connector.rs +++ b/connectors/google/src/connector.rs @@ -997,6 +997,7 @@ impl Connector for GoogleConnector { "properties": {}, "required": [] }), + required_scopes: None, source_types: vec![SourceType::GoogleDrive], admin_only: true, hidden: true, diff --git a/sdk/typescript/src/connector.ts b/sdk/typescript/src/connector.ts index b420d142..924a0805 100644 --- a/sdk/typescript/src/connector.ts +++ b/sdk/typescript/src/connector.ts @@ -86,19 +86,24 @@ export abstract class Connector< return this._mcpAdapter as McpAdapter; } + private async discoverMcpCatalog(credentials: TCredentials): Promise { + const adapter = await this.getMcpAdapter(); + if (!adapter) { + return false; + } + const { env, headers } = this.prepareMcpAuth(credentials); + await adapter.discover(env, headers); + return adapter.hasCachedCatalog(); + } + /** * Discover MCP tools/resources/prompts and cache them. Called when * credentials first become available (e.g., during initial sync). */ async bootstrapMcp(credentials: TCredentials): Promise { - const adapter = await this.getMcpAdapter(); - if (!adapter) { - return; - } - const { env, headers } = this.prepareMcpAuth(credentials); logger.info('Bootstrapping MCP: discovering tools'); try { - await adapter.discover(env, headers); + await this.discoverMcpCatalog(credentials); } catch (err) { logger.warn({ err }, 'MCP bootstrap failed'); } @@ -111,12 +116,15 @@ export abstract class Connector< async oauthCredentialReady( request: OAuthCredentialReadyRequest ): Promise { - const adapter = await this.getMcpAdapter(); - if (!adapter) { + logger.info('Refreshing MCP catalog after OAuth credential update'); + try { + return await this.discoverMcpCatalog(request.credentials as TCredentials); + } catch (err) { + const adapter = await this.getMcpAdapter(); + adapter?.clearCachedCatalog(); + logger.warn({ err }, 'OAuth credential-ready MCP refresh failed'); return false; } - await this.bootstrapMcp(request.credentials as TCredentials); - return adapter.hasCachedCatalog(); } prepareMcpAuth(credentials: TCredentials): { diff --git a/sdk/typescript/src/mcp-adapter.ts b/sdk/typescript/src/mcp-adapter.ts index 4324beda..3f14c7f0 100644 --- a/sdk/typescript/src/mcp-adapter.ts +++ b/sdk/typescript/src/mcp-adapter.ts @@ -63,6 +63,12 @@ export class McpAdapter { ); } + clearCachedCatalog(): void { + this.cachedActions = null; + this.cachedResources = null; + this.cachedPrompts = null; + } + private async withSession( env: Record | undefined, headers: Record | undefined, @@ -116,15 +122,18 @@ export class McpAdapter { env?: Record, headers?: Record ): Promise { - await this.withSession(env, headers, async (client) => { - this.cachedActions = await this.fetchActions(client); - this.cachedResources = await this.fetchResources(client); - this.cachedPrompts = await this.fetchPrompts(client); - }); + const catalog = await this.withSession(env, headers, async (client) => ({ + actions: await this.fetchActions(client), + resources: await this.fetchResources(client), + prompts: await this.fetchPrompts(client), + })); + this.cachedActions = catalog.actions; + this.cachedResources = catalog.resources; + this.cachedPrompts = catalog.prompts; logger.info( - `MCP discovery complete: ${this.cachedActions?.length ?? 0} tools, ` + - `${this.cachedResources?.length ?? 0} resources, ` + - `${this.cachedPrompts?.length ?? 0} prompts` + `MCP discovery complete: ${catalog.actions.length} tools, ` + + `${catalog.resources.length} resources, ` + + `${catalog.prompts.length} prompts` ); } diff --git a/sdk/typescript/tests/mcp-adapter.test.ts b/sdk/typescript/tests/mcp-adapter.test.ts index 77eac70b..65a06fb1 100644 --- a/sdk/typescript/tests/mcp-adapter.test.ts +++ b/sdk/typescript/tests/mcp-adapter.test.ts @@ -246,6 +246,45 @@ describe('Connector MCP integration', () => { expect(manifest.prompts).toHaveLength(1); }); + it('does not report a stale catalog as refreshed when OAuth discovery fails', async () => { + class StdioMcpConnector extends Connector { + readonly name = 'mcp-test-stdio'; + readonly version = '0.1.0'; + readonly sourceTypes = ['mcp_test']; + failAuthentication = false; + + get mcpServer(): StdioMcpServer { + return STDIO_SERVER; + } + + prepareMcpEnv(): Record { + if (this.failAuthentication) { + throw new Error('invalid OAuth credential'); + } + return { TEST_MODE: '1' }; + } + + async sync(): Promise {} + } + + const connector = new StdioMcpConnector(); + await connector.bootstrapMcp({}); + connector.failAuthentication = true; + + const refreshed = await connector.oauthCredentialReady({ + source_id: 'source-1', + user_id: 'user-1', + provider: 'example', + flow: 'user_write', + credentials: { access_token: 'invalid' }, + }); + + expect(refreshed).toBe(false); + expect((await connector.getManifest('http://test:8000')).mcp_catalog_loaded).toBe( + false + ); + }); + it('stdio: delegates action execution to MCP tool', async () => { class StdioMcpConnector extends Connector { readonly name = 'mcp-test-stdio'; diff --git a/services/ai/streaming/generate.py b/services/ai/streaming/generate.py index 401ae293..af0c2716 100644 --- a/services/ai/streaming/generate.py +++ b/services/ai/streaming/generate.py @@ -736,7 +736,8 @@ async def stream_generator( event_index = 0 message_stop_received = False - pending_message_start_sse: str | None = None + pending_message_sses: list[str] = [] + message_stream_started = False cancelled = False last_cancel_check_at = 0.0 async for event in stream: @@ -849,18 +850,27 @@ async def stream_generator( event_json = event.to_json(indent=None) event_sse = f"event: message\ndata: {event_json}\n\n" - if event.type == "message_start": - # Hold this until the provider emits actual content. If - # it immediately stops, the retry below stays invisible - # and the persistence wrapper does not create an empty - # assistant row. - pending_message_start_sse = event_sse - elif event.type == "message_stop" and not content_blocks: - pass + has_substantive_content = any( + block["type"] == "tool_use" + or ( + block["type"] == "text" + and str(block.get("text", "")).strip() + ) + for block in content_blocks + ) + if not message_stream_started: + # Buffer the entire provider envelope until it contains + # a tool call or non-whitespace text. Providers commonly + # emit an empty text block before stopping; exposing that + # envelope would create a duplicate assistant stream when + # the empty-response retry below runs. + pending_message_sses.append(event_sse) + if has_substantive_content: + for pending_sse in pending_message_sses: + yield pending_sse + pending_message_sses.clear() + message_stream_started = True else: - if pending_message_start_sse is not None: - yield pending_message_start_sse - pending_message_start_sse = None logger.debug("Yielding event to client: %s", event_json) yield event_sse @@ -880,19 +890,26 @@ async def stream_generator( b["type"] == "text" and str(b.get("text", "")).strip() for b in content_blocks ) - if not tool_calls and not has_text and empty_response_retries < 1: - empty_response_retries += 1 - logger.warning( - "Provider returned an empty response in iteration %s; " - "retrying once with a continuation prompt", - model_iteration, - ) - conversation_messages.append( - MessageParam( - role="user", content=_EMPTY_RESPONSE_RECOVERY_PROMPT + if not tool_calls and not has_text: + if empty_response_retries < 1: + empty_response_retries += 1 + logger.warning( + "Provider returned an empty response in iteration %s; " + "retrying once with a continuation prompt", + model_iteration, ) - ) - continue + conversation_messages.append( + MessageParam( + role="user", content=_EMPTY_RESPONSE_RECOVERY_PROMPT + ) + ) + continue + + # Preserve the provider's final empty response after the + # single recovery attempt has already been exhausted. + for pending_sse in pending_message_sses: + yield pending_sse + pending_message_sses.clear() parse_errors = parse_tool_call_inputs( cast(list[ToolUseBlockParam], tool_calls) ) diff --git a/services/ai/tests/helpers.py b/services/ai/tests/helpers.py index c1a1e52d..ed915fa3 100644 --- a/services/ai/tests/helpers.py +++ b/services/ai/tests/helpers.py @@ -258,6 +258,23 @@ def text_response_events(text: str): yield RawMessageStopEvent(type="message_stop") +def empty_text_response_events(): + """Yield a complete provider envelope containing an empty text block.""" + yield message_start_event() + yield RawContentBlockStartEvent( + type="content_block_start", + index=0, + content_block=TextBlock(type="text", text=""), + ) + yield RawContentBlockStopEvent(type="content_block_stop", index=0) + yield RawMessageDeltaEvent( + type="message_delta", + delta=Delta(stop_reason="end_turn", stop_sequence=None), + usage=MessageDeltaUsage(output_tokens=0), + ) + yield RawMessageStopEvent(type="message_stop") + + def create_mock_llm( tool_call_json: dict[str, Any], response_text: str = "Here are the results.", @@ -437,6 +454,7 @@ class GatedRecordingLLM: Each response entry follows the same convention as ``create_mock_llm_multi``: * ``("empty", None)`` + * ``("empty_text", None)`` * ``("text", "response string")`` * ``("tool_call", {"name": ..., "input": ..., "id": ...})`` @@ -517,6 +535,11 @@ async def stream_response(self, **kwargs): if kind == "empty": yield message_start_event() yield RawMessageStopEvent(type="message_stop") + elif kind == "empty_text": + for event in empty_text_response_events(): + yield event + if self._inter_event_delay: + await asyncio.sleep(self._inter_event_delay) elif kind == "tool_call": for event in tool_call_events( payload["input"], diff --git a/services/ai/tests/integration/test_chat_stream_lifecycle.py b/services/ai/tests/integration/test_chat_stream_lifecycle.py index a7ccabd6..e5643c17 100644 --- a/services/ai/tests/integration/test_chat_stream_lifecycle.py +++ b/services/ai/tests/integration/test_chat_stream_lifecycle.py @@ -2534,10 +2534,10 @@ async def test_context_overflow_triggers_compaction_retry( async def test_empty_provider_response_retries_once( self, seeded_chat, redis_client, redis_keys ): - """Retry a provider turn containing only message_start/message_stop.""" + """Retry an empty provider turn without exposing its text envelope.""" chat_id, _user_id, model_id = seeded_chat llm = GatedRecordingLLM( - [("empty", None), ("text", "Recovered after an empty response.")], + [("empty_text", None), ("text", "Recovered after an empty response.")], model_id, ) @@ -2546,6 +2546,10 @@ async def test_empty_provider_response_retries_once( events = await collect_sse_events(client, chat_id) assert any(et == "end_of_stream" for et, _, _ in events) + message_payloads = [ + json.loads(data) for event_type, data, _ in events if event_type == "message" + ] + assert sum(payload["type"] == "message_start" for payload in message_payloads) == 1 assert len(llm.calls) == 2 retry_messages = llm.calls[1]["messages"] assert retry_messages[-1] == { diff --git a/services/connector-manager/src/remote_mcp/gateway.rs b/services/connector-manager/src/remote_mcp/gateway.rs index 62e07840..f3499ed5 100644 --- a/services/connector-manager/src/remote_mcp/gateway.rs +++ b/services/connector-manager/src/remote_mcp/gateway.rs @@ -1036,6 +1036,7 @@ fn action_from_tool(tool: &JsonValue, write_tools_enabled: bool) -> Option Result<()> { )); } - let client = reqwest::Client::new(); + let client = reqwest::Client::builder() + .timeout(OAUTH_REFRESH_REQUEST_TIMEOUT) + .build() + .context("failed to build OAuth refresh client")?; let mut request = client.post(&token_uri).form(&form); if auth_method == "client_secret_basic" { request = request.basic_auth(