From 2e1a691cebe7166525adbb876a3b96dc89208e3c Mon Sep 17 00:00:00 2001 From: Jeeho Lee Date: Fri, 9 Oct 2026 07:49:13 +0000 Subject: [PATCH] fix(headers): preserve mixed-case tracking tokens Merge every case variant of the two tracking fields without losing custom tokens. Let MCP use the shared merge instead of collapsing aliases first. --- .../adk/tools/mcp_tool/mcp_session_manager.py | 11 +--- .../adk/utils/_google_client_headers.py | 25 ++++---- .../mcp_tool/test_mcp_session_manager.py | 30 ++++++++++ .../utils/test_google_client_headers.py | 59 +++++++++++++++++++ 4 files changed, 103 insertions(+), 22 deletions(-) diff --git a/src/google/adk/tools/mcp_tool/mcp_session_manager.py b/src/google/adk/tools/mcp_tool/mcp_session_manager.py index c83fbcda12..66c13b21d4 100644 --- a/src/google/adk/tools/mcp_tool/mcp_session_manager.py +++ b/src/google/adk/tools/mcp_tool/mcp_session_manager.py @@ -108,12 +108,6 @@ class AsyncAuthorizedSession: # pylint: disable=g-bad-classes # the manager, so credentials granted while the process runs are picked up. _MTLS_PROBE_RETRY_INTERVAL_SECONDS = 300.0 -# The headers `merge_tracking_headers` writes, spelled the way it spells them. -# HTTP header names are case-insensitive and a caller may have used any casing, -# but a dict is not: leaving their spelling alongside ours would put the header -# on the wire twice, so theirs is folded onto ours before merging. -_TRACKING_HEADER_NAMES = frozenset(('user-agent', 'x-goog-api-client')) - def create_mcp_http_client( headers: dict[str, str] | None = None, @@ -998,10 +992,7 @@ def _merge_headers( if additional_headers: base_headers.update(additional_headers) - return merge_tracking_headers({ - key.lower() if key.lower() in _TRACKING_HEADER_NAMES else key: value - for key, value in base_headers.items() - }) + return merge_tracking_headers(base_headers) def _is_session_disconnected(self, session: ClientSession) -> bool: """Checks if a session is disconnected or closed. diff --git a/src/google/adk/utils/_google_client_headers.py b/src/google/adk/utils/_google_client_headers.py index 9f03fafc0a..980f35a1fc 100644 --- a/src/google/adk/utils/_google_client_headers.py +++ b/src/google/adk/utils/_google_client_headers.py @@ -62,19 +62,20 @@ def merge_tracking_headers( Returns: A dictionary of HTTP headers with tracking headers merged. """ - new_headers = (headers or {}).copy() - for key, tracking_header_value in get_tracking_headers( - framework_label=framework_label - ).items(): - custom_value = new_headers.get(key, None) - if not custom_value: - new_headers[key] = tracking_header_value - continue - + tracking_headers = get_tracking_headers(framework_label=framework_label) + new_headers = { + key: value + for key, value in (headers or {}).items() + if key.lower() not in tracking_headers + } + for key, tracking_header_value in tracking_headers.items(): # Merge tracking headers with existing headers and avoid duplicates. value_parts = tracking_header_value.split(" ") - for custom_value_part in custom_value.split(" "): - if custom_value_part not in value_parts: - value_parts.append(custom_value_part) + for header_key, custom_value in (headers or {}).items(): + if header_key.lower() != key or not custom_value: + continue + for custom_value_part in custom_value.split(" "): + if custom_value_part not in value_parts: + value_parts.append(custom_value_part) new_headers[key] = " ".join(value_parts) return new_headers diff --git a/tests/unittests/tools/mcp_tool/test_mcp_session_manager.py b/tests/unittests/tools/mcp_tool/test_mcp_session_manager.py index f747f07171..dc76befd04 100644 --- a/tests/unittests/tools/mcp_tool/test_mcp_session_manager.py +++ b/tests/unittests/tools/mcp_tool/test_mcp_session_manager.py @@ -446,6 +446,36 @@ def test_merge_headers_keeps_custom_user_agent(self): assert merged["user-agent"].startswith("google-adk/") assert merged["user-agent"].endswith(" my-app/1.0") + def test_merge_headers_preserves_tracking_case_aliases(self): + """MCP keeps base and additional custom tokens when their casing differs.""" + base_headers = { + "User-Agent": "base-client/1 shared/1", + "X-Goog-Api-Client": "base-sdk/1 shared/1", + "X-Custom": "base", + } + additional = { + "user-agent": "extra-client/1 shared/1", + "x-goog-api-client": "extra-sdk/1 shared/1", + "X-Custom": "extra", + } + manager = MCPSessionManager( + SseConnectionParams(url="https://example.com/mcp", headers=base_headers) + ) + + merged = manager._merge_headers(additional) + request = httpx.Request("GET", "https://example.com/mcp", headers=merged) + + for key, expected_custom in ( + ("user-agent", "base-client/1 shared/1 extra-client/1"), + ("x-goog-api-client", "base-sdk/1 shared/1 extra-sdk/1"), + ): + values = request.headers.get_list(key) + assert len(values) == 1 + assert values[0].endswith(" " + expected_custom) + assert merged["X-Custom"] == "extra" + assert base_headers["User-Agent"] == "base-client/1 shared/1" + assert additional["user-agent"] == "extra-client/1 shared/1" + def test_is_session_disconnected(self): """Test session disconnection detection.""" manager = MCPSessionManager(self.mock_stdio_connection_params) diff --git a/tests/unittests/utils/test_google_client_headers.py b/tests/unittests/utils/test_google_client_headers.py index 67670d17c5..b4e08c7e6d 100644 --- a/tests/unittests/utils/test_google_client_headers.py +++ b/tests/unittests/utils/test_google_client_headers.py @@ -16,6 +16,7 @@ from google.adk import version from google.adk.utils import _google_client_headers +import httpx import pytest _EXPECTED_BASE_HEADER = ( @@ -79,6 +80,64 @@ def test_merge_tracking_headers(input_headers, expected_headers): assert headers == expected_headers +@pytest.mark.parametrize( + "user_agent_key, api_client_key", + [ + ("User-Agent", "X-Goog-Api-Client"), + ("USER-AGENT", "X-GOOG-API-CLIENT"), + ], +) +def test_merge_tracking_headers_merges_case_insensitively( + user_agent_key, api_client_key +): + """Each tracking header reaches the transport once with custom tokens kept.""" + input_headers = { + user_agent_key: "custom-client/1", + api_client_key: "custom-sdk/1", + "X-Custom": "value", + } + + headers = _google_client_headers.merge_tracking_headers(input_headers) + request = httpx.Request("GET", "https://example.test", headers=headers) + + assert request.headers.get_list("user-agent") == [ + f"{_EXPECTED_BASE_HEADER} custom-client/1" + ] + assert request.headers.get_list("x-goog-api-client") == [ + f"{_EXPECTED_BASE_HEADER} custom-sdk/1" + ] + assert headers["X-Custom"] == "value" + assert input_headers == { + user_agent_key: "custom-client/1", + api_client_key: "custom-sdk/1", + "X-Custom": "value", + } + + +def test_merge_tracking_headers_preserves_all_case_aliases(): + """Case aliases merge into one transport field without losing custom tokens.""" + input_headers = { + "User-Agent": f"first-client/1 shared/1 {_EXPECTED_BASE_HEADER}", + "user-agent": "second-client/1 shared/1", + "X-Goog-Api-Client": "first-sdk/1 shared/1", + "x-goog-api-client": "second-sdk/1 shared/1", + "X-Custom": "value", + } + original_headers = input_headers.copy() + + headers = _google_client_headers.merge_tracking_headers(input_headers) + request = httpx.Request("GET", "https://example.test", headers=headers) + + assert request.headers.get_list("user-agent") == [ + f"{_EXPECTED_BASE_HEADER} first-client/1 shared/1 second-client/1" + ] + assert request.headers.get_list("x-goog-api-client") == [ + f"{_EXPECTED_BASE_HEADER} first-sdk/1 shared/1 second-sdk/1" + ] + assert headers["X-Custom"] == "value" + assert input_headers == original_headers + + def test_get_tracking_http_options(): """get_tracking_http_options returns HttpOptions carrying tracking headers.""" http_options = _google_client_headers.get_tracking_http_options()