Skip to content
Draft
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
11 changes: 1 addition & 10 deletions src/google/adk/tools/mcp_tool/mcp_session_manager.py
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand Down Expand Up @@ -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.
Expand Down
25 changes: 13 additions & 12 deletions src/google/adk/utils/_google_client_headers.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
30 changes: 30 additions & 0 deletions tests/unittests/tools/mcp_tool/test_mcp_session_manager.py
Original file line number Diff line number Diff line change
Expand Up @@ -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)
Expand Down
59 changes: 59 additions & 0 deletions tests/unittests/utils/test_google_client_headers.py
Original file line number Diff line number Diff line change
Expand Up @@ -16,6 +16,7 @@

from google.adk import version
from google.adk.utils import _google_client_headers
import httpx
import pytest

_EXPECTED_BASE_HEADER = (
Expand Down Expand Up @@ -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()
Expand Down
Loading