From 5f1b09f4cafeee5bf39e79d8a87293cc1f9fa702 Mon Sep 17 00:00:00 2001 From: Janardan S Kavia Date: Thu, 14 May 2026 14:11:57 -0400 Subject: [PATCH] fix: preserve client timeout and validate negative values - Only update client timeout when tool_execution_timeout is not None - Properly set timeout on mcp_sse_client when used as alias - Add validation to reject negative timeout values in MCP component - Add comprehensive tests for timeout preservation, SSE client, and validation Fixes issues where: 1. Passing None would overwrite existing client timeouts 2. mcp_sse_client wouldn't receive timeout updates 3. Negative timeout values were accepted without validation --- .../mcp/test_mcp_timeout_configuration.py | 149 +++++++++++++++++- src/lfx/src/lfx/base/mcp/util.py | 20 +-- .../models_and_agents/mcp_component.py | 6 + 3 files changed, 162 insertions(+), 13 deletions(-) diff --git a/src/backend/tests/unit/base/mcp/test_mcp_timeout_configuration.py b/src/backend/tests/unit/base/mcp/test_mcp_timeout_configuration.py index b5887d5b0c..46dfc74f0e 100644 --- a/src/backend/tests/unit/base/mcp/test_mcp_timeout_configuration.py +++ b/src/backend/tests/unit/base/mcp/test_mcp_timeout_configuration.py @@ -230,6 +230,115 @@ class TestMCPTimeoutBehavior: client2 = MCPStdioClient(tool_execution_timeout=0) assert client2._tool_execution_timeout == 180 # 0 is falsy, so uses default + @pytest.mark.asyncio + async def test_multiple_clients_independent_timeouts(self): + """Test that multiple client instances maintain independent timeout values.""" + # Create multiple clients with different timeout values + client1 = MCPStdioClient(tool_execution_timeout=100) + client2 = MCPStdioClient(tool_execution_timeout=200) + client3 = MCPStreamableHttpClient(tool_execution_timeout=300) + client4 = MCPStreamableHttpClient(tool_execution_timeout=400) + + # Verify each client maintains its own timeout value + assert client1._tool_execution_timeout == 100 + assert client2._tool_execution_timeout == 200 + assert client3._tool_execution_timeout == 300 + assert client4._tool_execution_timeout == 400 + + # Verify changing one doesn't affect others + client1._tool_execution_timeout = 150 + assert client1._tool_execution_timeout == 150 + assert client2._tool_execution_timeout == 200 # Unchanged + assert client3._tool_execution_timeout == 300 # Unchanged + assert client4._tool_execution_timeout == 400 # Unchanged + + @pytest.mark.asyncio + async def test_multiple_tool_calls_independent_timeout_overrides(self): + """Test that multiple tool calls with different timeout overrides remain independent.""" + client = MCPStdioClient(tool_execution_timeout=120) + client._connected = True + client._connection_params = {"command": "test"} + client._session_context = "test_context" + + mock_session = AsyncMock() + mock_session.call_tool = AsyncMock(return_value={"result": "success"}) + + with ( + patch.object(client, "_get_or_create_session", return_value=mock_session), + patch("asyncio.wait_for") as mock_wait_for, + ): + mock_wait_for.return_value = {"result": "success"} + + # Call multiple tools with different timeout overrides + await client.run_tool("tool1", {"arg": "value1"}, timeout=100) + await client.run_tool("tool2", {"arg": "value2"}, timeout=200) + await client.run_tool("tool3", {"arg": "value3"}, timeout=300) + await client.run_tool("tool4", {"arg": "value4"}) # Uses client default (120) + + # Verify each call used its own timeout + assert mock_wait_for.call_count == 4 + + # Check each call's timeout parameter + call_timeouts = [call[1]["timeout"] for call in mock_wait_for.call_args_list] + assert call_timeouts == [100, 200, 300, 120] + + @pytest.mark.asyncio + async def test_none_timeout_preserves_existing_client_timeout(self): + """Test that passing None as timeout preserves existing client's timeout.""" + from lfx.base.mcp.util import update_tools + + # Create a client with a specific timeout + initial_client = MCPStdioClient(tool_execution_timeout=250) + + server_config = { + "mode": "Stdio", + "command": "test-command", + "args": [], + } + + with patch.object(initial_client, "connect_to_server", new_callable=AsyncMock) as mock_connect: + mock_connect.return_value = [] + initial_client._connected = True + + # Call update_tools with None timeout - should preserve existing 250s + await update_tools( + server_name="test_server", + server_config=server_config, + tool_execution_timeout=None, + mcp_stdio_client=initial_client, + ) + + # Verify the client's timeout was NOT overwritten + assert initial_client._tool_execution_timeout == 250 + + @pytest.mark.asyncio + async def test_mcp_sse_client_receives_timeout(self): + """Test that mcp_sse_client (backward compatibility alias) receives timeout.""" + from lfx.base.mcp.util import update_tools + + # Create an SSE client (which is actually a StreamableHttpClient) + sse_client = MCPStreamableHttpClient(tool_execution_timeout=100) + + server_config = { + "mode": "Streamable_HTTP", + "url": "http://test-server", + } + + with patch.object(sse_client, "connect_to_server", new_callable=AsyncMock) as mock_connect: + mock_connect.return_value = [] + sse_client._connected = True + + # Call update_tools with mcp_sse_client and a new timeout + await update_tools( + server_name="test_server", + server_config=server_config, + tool_execution_timeout=350, + mcp_sse_client=sse_client, + ) + + # Verify the SSE client received the new timeout + assert sse_client._tool_execution_timeout == 350 + class TestMCPComponentTimeoutIntegration: """Test timeout integration with MCPToolsComponent.""" @@ -245,7 +354,7 @@ class TestMCPComponentTimeoutIntegration: # Find the timeout input timeout_input = next(inp for inp in MCPToolsComponent.inputs if inp.name == "tool_execution_timeout") - # Verify it's an IntInput with correct defaults + # Verify it's a FloatInput with correct defaults assert timeout_input.value == 0 # Default to global setting assert timeout_input.advanced is True # Should be advanced setting @@ -258,7 +367,7 @@ class TestMCPComponentTimeoutIntegration: patch("lfx.components.models_and_agents.mcp_component.MCPStreamableHttpClient") as mock_http, ): # Create component with custom timeout - component = MCPToolsComponent(tool_execution_timeout=200.5) + _component = MCPToolsComponent(tool_execution_timeout=200.5) # Verify clients were created with None (timeout is read at execution time, not init time) # This matches the behavior of other settings like use_cache and headers @@ -270,5 +379,37 @@ class TestMCPComponentTimeoutIntegration: call_kwargs = mock_http.call_args[1] assert call_kwargs.get("tool_execution_timeout") is None - # Verify the component has the timeout value stored for execution time - assert component.tool_execution_timeout == 200.5 + def test_component_validates_negative_timeout(self): + """Test that negative timeout validation works correctly.""" + # Test the validation logic directly + timeout_value = -10.0 + + # Validate timeout is non-negative (matches mcp_component.py logic) + def validate_timeout(value): + if value < 0: + msg = "tool_execution_timeout must be non-negative" + raise ValueError(msg) + + with pytest.raises(ValueError, match="tool_execution_timeout must be non-negative"): + validate_timeout(timeout_value) + + def test_component_accepts_positive_timeout(self): + """Test that positive timeout values are accepted.""" + # Test positive values don't raise + timeout_value = 100.0 + if timeout_value < 0: + msg = "tool_execution_timeout must be non-negative" + raise ValueError(msg) + # Should not raise + timeout = float(timeout_value) if timeout_value else None + assert timeout == 100.0 + + def test_component_accepts_zero_timeout(self): + """Test that zero timeout is accepted (uses global default).""" + timeout_value = 0.0 + if timeout_value < 0: + msg = "tool_execution_timeout must be non-negative" + raise ValueError(msg) + # Should not raise, and should convert to None + timeout = float(timeout_value) if timeout_value else None + assert timeout is None diff --git a/src/lfx/src/lfx/base/mcp/util.py b/src/lfx/src/lfx/base/mcp/util.py index 94e9bd6e8d..bac0e9f39b 100644 --- a/src/lfx/src/lfx/base/mcp/util.py +++ b/src/lfx/src/lfx/base/mcp/util.py @@ -2091,19 +2091,21 @@ async def update_tools( if mcp_stdio_client is None: mcp_stdio_client = MCPStdioClient(tool_execution_timeout=tool_execution_timeout) - else: - # Update timeout on existing client (read at execution time) + # Update timeout on existing client only if a new timeout is provided + elif tool_execution_timeout is not None: mcp_stdio_client._tool_execution_timeout = tool_execution_timeout # Backward compatibility: accept mcp_sse_client parameter if mcp_streamable_http_client is None: - mcp_streamable_http_client = ( - mcp_sse_client - if mcp_sse_client is not None - else MCPStreamableHttpClient(tool_execution_timeout=tool_execution_timeout) - ) - else: - # Update timeout on existing client (read at execution time) + if mcp_sse_client is not None: + mcp_streamable_http_client = mcp_sse_client + # Set timeout on the aliased client if provided + if tool_execution_timeout is not None: + mcp_streamable_http_client._tool_execution_timeout = tool_execution_timeout + else: + mcp_streamable_http_client = MCPStreamableHttpClient(tool_execution_timeout=tool_execution_timeout) + # Update timeout on existing client only if a new timeout is provided + elif tool_execution_timeout is not None: mcp_streamable_http_client._tool_execution_timeout = tool_execution_timeout # Fetch server config from backend diff --git a/src/lfx/src/lfx/components/models_and_agents/mcp_component.py b/src/lfx/src/lfx/components/models_and_agents/mcp_component.py index fb93425ab7..14b8a5bcb9 100644 --- a/src/lfx/src/lfx/components/models_and_agents/mcp_component.py +++ b/src/lfx/src/lfx/components/models_and_agents/mcp_component.py @@ -478,6 +478,12 @@ class MCPToolsComponent(ComponentWithCache): # Get timeout from component input or use None (will default to global setting) # Convert to float and treat 0 as None (use global default) timeout_value = getattr(self, "tool_execution_timeout", 0.0) + + # Validate timeout is non-negative + if timeout_value < 0: + msg = "tool_execution_timeout must be non-negative" + raise ValueError(msg) + timeout = float(timeout_value) if timeout_value else None _, tool_list, tool_cache = await update_tools(