refactor(types): type the integrations, skills and MCP tool packages for strict mypy
Co-authored-by: George Weale <gweale@google.com> PiperOrigin-RevId: 961125207
This commit is contained in:
committed by
Copybara-Service
parent
03f44c8e10
commit
cd36dbc338
@@ -0,0 +1,99 @@
|
||||
# Copyright 2026 Google LLC
|
||||
#
|
||||
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||
# you may not use this file except in compliance with the License.
|
||||
# You may obtain a copy of the License at
|
||||
#
|
||||
# http://www.apache.org/licenses/LICENSE-2.0
|
||||
#
|
||||
# Unless required by applicable law or agreed to in writing, software
|
||||
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
# See the License for the specific language governing permissions and
|
||||
# limitations under the License.
|
||||
|
||||
"""Typed construction boundary for unannotated Google SDK classes."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from collections.abc import Mapping
|
||||
import json
|
||||
from typing import cast
|
||||
from typing import Protocol
|
||||
|
||||
from google.api_core.client_info import ClientInfo
|
||||
from google.api_core.gapic_v1.client_info import ClientInfo as GapicClientInfo
|
||||
from google.auth.credentials import Credentials
|
||||
from google.oauth2 import credentials as user_credentials
|
||||
from google.oauth2 import service_account
|
||||
|
||||
|
||||
class _ApiRepresentable(Protocol):
|
||||
|
||||
def to_api_repr(self) -> dict[str, object]:
|
||||
...
|
||||
|
||||
|
||||
class _ClientInfoFactory(Protocol):
|
||||
|
||||
def __call__(self, *, user_agent: str) -> ClientInfo:
|
||||
...
|
||||
|
||||
|
||||
class _GapicClientInfoFactory(Protocol):
|
||||
|
||||
def __call__(self, *, user_agent: str) -> GapicClientInfo:
|
||||
...
|
||||
|
||||
|
||||
class _ServiceAccountCredentialsFactory(Protocol):
|
||||
|
||||
def __call__(self, info: Mapping[str, object]) -> Credentials:
|
||||
...
|
||||
|
||||
|
||||
class _UserCredentialsFactory(Protocol):
|
||||
|
||||
def __call__(self, *, token: str) -> user_credentials.Credentials:
|
||||
...
|
||||
|
||||
|
||||
def read_api_repr(obj: object) -> dict[str, object]:
|
||||
"""Read the API representation of an unannotated SDK object."""
|
||||
return cast(_ApiRepresentable, obj).to_api_repr()
|
||||
|
||||
|
||||
def create_client_info(*, user_agent: str) -> ClientInfo:
|
||||
"""Create client metadata through the SDK's unannotated constructor."""
|
||||
factory = cast(_ClientInfoFactory, ClientInfo)
|
||||
return factory(user_agent=user_agent)
|
||||
|
||||
|
||||
def create_gapic_client_info(*, user_agent: str) -> GapicClientInfo:
|
||||
"""Create GAPIC client metadata through its unannotated constructor."""
|
||||
factory = cast(_GapicClientInfoFactory, GapicClientInfo)
|
||||
return factory(user_agent=user_agent)
|
||||
|
||||
|
||||
def load_service_account_credentials(raw_json: str) -> Credentials:
|
||||
"""Parse service-account JSON and construct typed credentials."""
|
||||
try:
|
||||
info: object = json.loads(raw_json)
|
||||
except json.JSONDecodeError as e:
|
||||
raise ValueError(f"Invalid service account JSON: {e}") from e
|
||||
if not isinstance(info, dict) or not all(
|
||||
isinstance(key, str) for key in info
|
||||
):
|
||||
raise ValueError("Service account JSON must contain an object.")
|
||||
|
||||
factory = cast(
|
||||
_ServiceAccountCredentialsFactory,
|
||||
service_account.Credentials.from_service_account_info,
|
||||
)
|
||||
return factory(cast(dict[str, object], info))
|
||||
|
||||
|
||||
def create_user_credentials(*, token: str) -> user_credentials.Credentials:
|
||||
"""Create OAuth user credentials through the unannotated constructor."""
|
||||
factory = cast(_UserCredentialsFactory, user_credentials.Credentials)
|
||||
return factory(token=token)
|
||||
@@ -273,6 +273,8 @@ class AgentRegistry:
|
||||
data: Dict[str, Any] = response.json()
|
||||
return data
|
||||
except requests.exceptions.HTTPError as e:
|
||||
if e.response is None:
|
||||
raise RuntimeError(f"API request failed: {e}") from e
|
||||
raise RuntimeError(
|
||||
f"API request failed with status {e.response.status_code}:"
|
||||
f" {e.response.text}"
|
||||
|
||||
@@ -22,9 +22,9 @@ from __future__ import annotations
|
||||
import typing
|
||||
|
||||
if typing.TYPE_CHECKING:
|
||||
from .bigquery_credentials import BigQueryCredentialsConfig
|
||||
from .bigquery_skill import get_bigquery_skill
|
||||
from .bigquery_toolset import BigQueryToolset
|
||||
from .bigquery_credentials import BigQueryCredentialsConfig as BigQueryCredentialsConfig
|
||||
from .bigquery_skill import get_bigquery_skill as get_bigquery_skill
|
||||
from .bigquery_toolset import BigQueryToolset as BigQueryToolset
|
||||
|
||||
# Map attribute names to relative module paths
|
||||
_lazy_imports = {
|
||||
|
||||
@@ -33,10 +33,10 @@ class BigQueryCredentialsConfig(BaseGoogleCredentialsConfig):
|
||||
|
||||
def __post_init__(self) -> BigQueryCredentialsConfig:
|
||||
"""Populate default scope if scopes is None."""
|
||||
super().__post_init__()
|
||||
super().__post_init__() # type: ignore[misc]
|
||||
|
||||
if not self.scopes:
|
||||
self.scopes = BIGQUERY_SCOPES
|
||||
self.scopes = BIGQUERY_SCOPES.copy()
|
||||
# Set the token cache key
|
||||
self._token_cache_key = BIGQUERY_TOKEN_CACHE_KEY
|
||||
|
||||
|
||||
@@ -18,7 +18,6 @@ from typing import List
|
||||
from typing import Optional
|
||||
from typing import Union
|
||||
|
||||
import google.api_core.client_info
|
||||
from google.api_core.gapic_v1 import client_info as gapic_client_info
|
||||
from google.auth.credentials import Credentials
|
||||
from google.cloud import bigquery
|
||||
@@ -26,6 +25,8 @@ from google.cloud import dataplex_v1
|
||||
|
||||
from ... import version
|
||||
from ...utils._telemetry_context import _is_visual_builder
|
||||
from .._google_sdk import create_client_info as _create_client_info
|
||||
from .._google_sdk import create_gapic_client_info as _create_gapic_client_info
|
||||
|
||||
USER_AGENT_BASE = f"google-adk/{version.__version__}"
|
||||
BQ_USER_AGENT = f"adk-bigquery-tool {USER_AGENT_BASE}"
|
||||
@@ -66,9 +67,7 @@ def get_bigquery_client(
|
||||
else:
|
||||
user_agents.extend([ua for ua in user_agent if ua])
|
||||
|
||||
client_info = google.api_core.client_info.ClientInfo(
|
||||
user_agent=" ".join(user_agents)
|
||||
)
|
||||
client_info = _create_client_info(user_agent=" ".join(user_agents))
|
||||
|
||||
bigquery_client = bigquery.Client(
|
||||
project=project,
|
||||
@@ -106,7 +105,9 @@ def get_dataplex_catalog_client(
|
||||
else:
|
||||
user_agents.extend([ua for ua in user_agent if ua])
|
||||
|
||||
client_info = gapic_client_info.ClientInfo(user_agent=" ".join(user_agents))
|
||||
client_info: gapic_client_info.ClientInfo = _create_gapic_client_info(
|
||||
user_agent=" ".join(user_agents)
|
||||
)
|
||||
|
||||
return dataplex_v1.CatalogServiceClient(
|
||||
credentials=credentials,
|
||||
|
||||
@@ -25,7 +25,7 @@ from .config import BigQueryToolConfig
|
||||
|
||||
def list_dataset_ids(
|
||||
project_id: str, credentials: Credentials, settings: BigQueryToolConfig
|
||||
) -> list[str]:
|
||||
) -> list[str] | dict[str, str]:
|
||||
"""List BigQuery dataset ids in a Google Cloud project.
|
||||
|
||||
Args:
|
||||
@@ -143,7 +143,7 @@ def list_table_ids(
|
||||
dataset_id: str,
|
||||
credentials: Credentials,
|
||||
settings: BigQueryToolConfig,
|
||||
) -> list[str]:
|
||||
) -> list[str] | dict[str, str]:
|
||||
"""List table ids in a BigQuery dataset.
|
||||
|
||||
Args:
|
||||
|
||||
@@ -27,12 +27,24 @@ from google.cloud import bigquery
|
||||
|
||||
from . import client
|
||||
from ...tools.tool_context import ToolContext
|
||||
from .._google_sdk import read_api_repr as _read_api_repr
|
||||
from .config import BigQueryToolConfig
|
||||
from .config import WriteMode
|
||||
|
||||
BIGQUERY_SESSION_INFO_KEY = "bigquery_session_info"
|
||||
|
||||
|
||||
def _parse_session_info(value: object) -> tuple[str, str] | None:
|
||||
"""Validate persisted BigQuery session state."""
|
||||
if not isinstance(value, (list, tuple)) or len(value) != 2:
|
||||
return None
|
||||
session_id: object = value[0]
|
||||
dataset_id: object = value[1]
|
||||
if not isinstance(session_id, str) or not isinstance(dataset_id, str):
|
||||
return None
|
||||
return session_id, dataset_id
|
||||
|
||||
|
||||
def _execute_sql(
|
||||
project_id: str,
|
||||
query: str,
|
||||
@@ -96,8 +108,11 @@ def _execute_sql(
|
||||
# allowed. This artifact must have been created in a BigQuery session. In
|
||||
# such a scenario, the session info (session id and the anonymous dataset
|
||||
# containing the artifact) is persisted in the tool context.
|
||||
bq_session_info = tool_context.state.get(BIGQUERY_SESSION_INFO_KEY, None)
|
||||
if bq_session_info:
|
||||
stored_session_info: object = tool_context.state.get(
|
||||
BIGQUERY_SESSION_INFO_KEY
|
||||
)
|
||||
bq_session_info = _parse_session_info(stored_session_info)
|
||||
if bq_session_info is not None:
|
||||
bq_session_id, bq_session_dataset_id = bq_session_info
|
||||
else:
|
||||
session_creator_job = bq_client.query(
|
||||
@@ -107,8 +122,18 @@ def _execute_sql(
|
||||
dry_run=True, create_session=True, labels=bq_job_labels
|
||||
),
|
||||
)
|
||||
bq_session_id = session_creator_job.session_info.session_id
|
||||
bq_session_dataset_id = session_creator_job.destination.dataset_id
|
||||
session_info = session_creator_job.session_info
|
||||
destination = session_creator_job.destination
|
||||
session_id = (
|
||||
session_info.session_id if session_info is not None else None
|
||||
)
|
||||
if session_id is None or destination is None:
|
||||
raise RuntimeError(
|
||||
"BigQuery did not return session metadata for the protected"
|
||||
" query."
|
||||
)
|
||||
bq_session_id = session_id
|
||||
bq_session_dataset_id = destination.dataset_id
|
||||
|
||||
# Remember the BigQuery session info for subsequent queries
|
||||
tool_context.state[BIGQUERY_SESSION_INFO_KEY] = (
|
||||
@@ -155,7 +180,8 @@ def _execute_sql(
|
||||
labels=bq_job_labels,
|
||||
),
|
||||
)
|
||||
return {"status": "SUCCESS", "dry_run_info": dry_run_job.to_api_repr()}
|
||||
dry_run_info = _read_api_repr(dry_run_job)
|
||||
return {"status": "SUCCESS", "dry_run_info": dry_run_info}
|
||||
|
||||
# Finally execute the query, fetch the result, and return it
|
||||
job_config = bigquery.QueryJobConfig(
|
||||
@@ -792,7 +818,7 @@ def forecast(
|
||||
timestamp_col: str,
|
||||
data_col: str,
|
||||
horizon: int = 10,
|
||||
id_cols: Optional[list[str]] = None,
|
||||
id_cols: list[str] | None = None,
|
||||
*,
|
||||
credentials: Credentials,
|
||||
settings: BigQueryToolConfig,
|
||||
@@ -1165,10 +1191,10 @@ def detect_anomalies(
|
||||
history_data: str,
|
||||
times_series_timestamp_col: str,
|
||||
times_series_data_col: str,
|
||||
horizon: Optional[int] = 1000,
|
||||
target_data: Optional[str] = None,
|
||||
times_series_id_cols: Optional[list[str]] = None,
|
||||
anomaly_prob_threshold: Optional[float] = 0.95,
|
||||
horizon: int | None = 1000,
|
||||
target_data: str | None = None,
|
||||
times_series_id_cols: list[str] | None = None,
|
||||
anomaly_prob_threshold: float | None = 0.95,
|
||||
*,
|
||||
credentials: Credentials,
|
||||
settings: BigQueryToolConfig,
|
||||
|
||||
@@ -69,7 +69,7 @@ class CloudRunSandboxCodeExecutor(BaseCodeExecutor):
|
||||
# Overrides the BaseCodeExecutor attribute: this executor cannot optimize_data_file.
|
||||
optimize_data_file: bool = Field(default=False, frozen=True, exclude=True)
|
||||
|
||||
def __init__(self, **data):
|
||||
def __init__(self, **data: object) -> None:
|
||||
if 'stateful' in data and data['stateful']:
|
||||
raise ValueError(
|
||||
'Cannot set `stateful=True` in CloudRunSandboxCodeExecutor.'
|
||||
|
||||
@@ -215,6 +215,7 @@ class DaytonaEnvironment(BaseEnvironment):
|
||||
if self._timeout > 0 and auto_stop_interval_mins == 0:
|
||||
auto_stop_interval_mins = 1
|
||||
|
||||
params: CreateSandboxFromImageParams | CreateSandboxFromSnapshotParams
|
||||
if self._image:
|
||||
params = CreateSandboxFromImageParams(
|
||||
image=self._image,
|
||||
|
||||
@@ -14,18 +14,19 @@
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import google.api_core.client_info
|
||||
from google.api_core.client_info import ClientInfo
|
||||
from google.auth.credentials import Credentials
|
||||
from google.cloud import storage
|
||||
|
||||
from ... import version
|
||||
from .._google_sdk import create_client_info as _create_client_info
|
||||
|
||||
USER_AGENT = f"adk-gcs-tool google-adk/{version.__version__}"
|
||||
|
||||
|
||||
def _get_client_info() -> google.api_core.client_info.ClientInfo:
|
||||
def _get_client_info() -> ClientInfo:
|
||||
"""Get client info."""
|
||||
return google.api_core.client_info.ClientInfo(user_agent=USER_AGENT)
|
||||
return _create_client_info(user_agent=USER_AGENT)
|
||||
|
||||
|
||||
_client_cache: dict[tuple[int, str | None], storage.Client] = {}
|
||||
|
||||
@@ -30,10 +30,10 @@ class GCSCredentialsConfig(BaseGoogleCredentialsConfig):
|
||||
|
||||
def __post_init__(self) -> GCSCredentialsConfig:
|
||||
"""Populate default scope if scopes is None."""
|
||||
super().__post_init__()
|
||||
super().__post_init__() # type: ignore[misc]
|
||||
|
||||
if not self.scopes:
|
||||
self.scopes = GCS_DEFAULT_SCOPE
|
||||
self.scopes = GCS_DEFAULT_SCOPE.copy()
|
||||
|
||||
# Set the token cache key
|
||||
self._token_cache_key = GCS_TOKEN_CACHE_KEY
|
||||
|
||||
@@ -90,6 +90,8 @@ class LangchainTool(FunctionTool):
|
||||
type(tool),
|
||||
)
|
||||
|
||||
if func is None:
|
||||
raise ValueError('Langchain tool must define a sync or async callable.')
|
||||
super().__init__(func)
|
||||
# run_manager is a special parameter for langchain tool
|
||||
self._ignore_params.append('run_manager')
|
||||
@@ -157,7 +159,7 @@ class LangchainTool(FunctionTool):
|
||||
False,
|
||||
self.name,
|
||||
self.description,
|
||||
tool_wrapper.func,
|
||||
self.func,
|
||||
tool_wrapper.args,
|
||||
)
|
||||
|
||||
|
||||
@@ -14,17 +14,16 @@
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
from typing import Optional
|
||||
|
||||
from google.api_core.gapic_v1 import client_info
|
||||
from google.auth import default as default_service_credential
|
||||
from google.cloud import parametermanager_v1
|
||||
from google.oauth2 import credentials as user_credentials
|
||||
from google.oauth2 import service_account
|
||||
|
||||
from ... import version
|
||||
from ...utils._mtls_utils import get_api_endpoint
|
||||
from .._google_sdk import create_gapic_client_info as _create_gapic_client_info
|
||||
from .._google_sdk import create_user_credentials as _create_user_credentials
|
||||
from .._google_sdk import load_service_account_credentials as _load_service_account_credentials
|
||||
|
||||
USER_AGENT = f"google-adk/{version.__version__}"
|
||||
|
||||
@@ -80,14 +79,9 @@ class ParameterManagerClient:
|
||||
)
|
||||
|
||||
if service_account_json:
|
||||
try:
|
||||
credentials = service_account.Credentials.from_service_account_info(
|
||||
json.loads(service_account_json)
|
||||
)
|
||||
except json.JSONDecodeError as e:
|
||||
raise ValueError(f"Invalid service account JSON: {e}") from e
|
||||
credentials = _load_service_account_credentials(service_account_json)
|
||||
elif auth_token:
|
||||
credentials = user_credentials.Credentials(token=auth_token)
|
||||
credentials = _create_user_credentials(token=auth_token)
|
||||
else:
|
||||
try:
|
||||
credentials, _ = default_service_credential(
|
||||
@@ -121,7 +115,7 @@ class ParameterManagerClient:
|
||||
self._client = parametermanager_v1.ParameterManagerClient(
|
||||
credentials=self._credentials,
|
||||
client_options=client_options,
|
||||
client_info=client_info.ClientInfo(user_agent=USER_AGENT),
|
||||
client_info=_create_gapic_client_info(user_agent=USER_AGENT),
|
||||
)
|
||||
|
||||
def get_parameter(self, resource_name: str) -> str:
|
||||
|
||||
@@ -14,17 +14,16 @@
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
from typing import Optional
|
||||
|
||||
from google.api_core.gapic_v1 import client_info
|
||||
from google.auth import default as default_service_credential
|
||||
from google.cloud import secretmanager
|
||||
from google.oauth2 import credentials as user_credentials
|
||||
from google.oauth2 import service_account
|
||||
|
||||
from ... import version
|
||||
from ...utils import _mtls_utils
|
||||
from .._google_sdk import create_gapic_client_info as _create_gapic_client_info
|
||||
from .._google_sdk import create_user_credentials as _create_user_credentials
|
||||
from .._google_sdk import load_service_account_credentials as _load_service_account_credentials
|
||||
|
||||
USER_AGENT = f"google-adk/{version.__version__}"
|
||||
|
||||
@@ -83,14 +82,9 @@ class SecretManagerClient:
|
||||
)
|
||||
|
||||
if service_account_json:
|
||||
try:
|
||||
credentials = service_account.Credentials.from_service_account_info(
|
||||
json.loads(service_account_json)
|
||||
)
|
||||
except json.JSONDecodeError as e:
|
||||
raise ValueError(f"Invalid service account JSON: {e}") from e
|
||||
credentials = _load_service_account_credentials(service_account_json)
|
||||
elif auth_token:
|
||||
credentials = user_credentials.Credentials(token=auth_token)
|
||||
credentials = _create_user_credentials(token=auth_token)
|
||||
else:
|
||||
try:
|
||||
credentials, _ = default_service_credential(
|
||||
@@ -123,7 +117,7 @@ class SecretManagerClient:
|
||||
self._client = secretmanager.SecretManagerServiceClient(
|
||||
credentials=self._credentials,
|
||||
client_options=client_options,
|
||||
client_info=client_info.ClientInfo(user_agent=USER_AGENT),
|
||||
client_info=_create_gapic_client_info(user_agent=USER_AGENT),
|
||||
)
|
||||
|
||||
def get_secret(self, resource_name: str) -> str:
|
||||
|
||||
@@ -16,6 +16,8 @@ from __future__ import annotations
|
||||
|
||||
import logging
|
||||
from typing import Any
|
||||
from typing import cast
|
||||
from typing import Protocol
|
||||
|
||||
from google.adk.runners import Runner
|
||||
from google.genai import types
|
||||
@@ -32,6 +34,12 @@ except ImportError as e:
|
||||
logger = logging.getLogger("google_adk." + __name__)
|
||||
|
||||
|
||||
class _SocketModeHandler(Protocol):
|
||||
|
||||
async def start_async(self) -> None:
|
||||
...
|
||||
|
||||
|
||||
class SlackRunner:
|
||||
"""Runner for ADK agents on Slack."""
|
||||
|
||||
@@ -119,5 +127,8 @@ class SlackRunner:
|
||||
|
||||
async def start(self, app_token: str) -> None:
|
||||
"""Starts the Slack app using Socket Mode."""
|
||||
handler = AsyncSocketModeHandler(self.slack_app, app_token)
|
||||
handler = cast(
|
||||
_SocketModeHandler,
|
||||
AsyncSocketModeHandler(self.slack_app, app_token),
|
||||
)
|
||||
await handler.start_async()
|
||||
|
||||
@@ -23,6 +23,7 @@ from __future__ import annotations
|
||||
import base64
|
||||
import logging
|
||||
from typing import Any
|
||||
from typing import cast
|
||||
from typing import Literal
|
||||
from typing import TYPE_CHECKING
|
||||
|
||||
@@ -132,7 +133,12 @@ class SandboxClient:
|
||||
import json
|
||||
|
||||
if hasattr(response, "body") and response.body:
|
||||
return json.loads(response.body)
|
||||
parsed: object = json.loads(response.body)
|
||||
if not isinstance(parsed, dict) or not all(
|
||||
isinstance(key, str) for key in parsed
|
||||
):
|
||||
raise ValueError("Sandbox response body must be a JSON object.")
|
||||
return parsed
|
||||
return {}
|
||||
|
||||
def update_access_token(self, access_token: str) -> None:
|
||||
@@ -206,7 +212,7 @@ class SandboxClient:
|
||||
request_dict=request_dict,
|
||||
)
|
||||
parsed = self._parse_response(response)
|
||||
return parsed.get("results", [])
|
||||
return cast(list[dict[str, Any]], parsed.get("results", []))
|
||||
except Exception as e:
|
||||
# Batch endpoint not available, fall back to sequential
|
||||
if "404" in str(e) or "not found" in str(e).lower():
|
||||
@@ -215,7 +221,7 @@ class SandboxClient:
|
||||
logger.warning("Batch CDP failed: %s, falling back to sequential", e)
|
||||
|
||||
# Sequential fallback
|
||||
results = []
|
||||
results: list[dict[str, Any]] = []
|
||||
for cmd in commands:
|
||||
try:
|
||||
result = await self.make_cdp_request(
|
||||
@@ -298,9 +304,15 @@ class SandboxClient:
|
||||
if active_tab_id is None:
|
||||
return None
|
||||
|
||||
for tab in parsed.get("all_tabs", []):
|
||||
all_tabs = parsed.get("all_tabs")
|
||||
if not isinstance(all_tabs, list):
|
||||
return None
|
||||
for tab in all_tabs:
|
||||
if not isinstance(tab, dict):
|
||||
continue
|
||||
if tab.get("id") == active_tab_id:
|
||||
return tab.get("url")
|
||||
url = tab.get("url")
|
||||
return url if isinstance(url, str) else None
|
||||
|
||||
return None
|
||||
except Exception as e:
|
||||
|
||||
@@ -24,11 +24,13 @@ import asyncio
|
||||
import logging
|
||||
import time
|
||||
from typing import Any
|
||||
from typing import cast
|
||||
from typing import Literal
|
||||
from typing import TYPE_CHECKING
|
||||
|
||||
from ...features import experimental
|
||||
from ...features import FeatureName
|
||||
from ...sessions.state import State
|
||||
from ...tools.computer_use.base_computer import BaseComputer
|
||||
from ...tools.computer_use.base_computer import ComputerEnvironment
|
||||
from ...tools.computer_use.base_computer import ComputerState
|
||||
@@ -158,7 +160,7 @@ class AgentEngineSandboxComputer(BaseComputer):
|
||||
self._client = vertexai_client
|
||||
|
||||
# Session state for sharing sandbox/tokens across invocations
|
||||
self._session_state: dict[str, Any] | None = None
|
||||
self._session_state: State | None = None
|
||||
|
||||
async def prepare(self, tool_context: "ToolContext") -> None:
|
||||
"""Bind session state for sandbox resource sharing."""
|
||||
@@ -184,8 +186,12 @@ class AgentEngineSandboxComputer(BaseComputer):
|
||||
if self._agent_engine_name:
|
||||
return self._agent_engine_name
|
||||
|
||||
state = cast(State, self._session_state)
|
||||
|
||||
# Check session state
|
||||
agent_engine_name = self._session_state.get(_STATE_KEY_AGENT_ENGINE_NAME)
|
||||
agent_engine_name = cast(
|
||||
"str | None", state.get(_STATE_KEY_AGENT_ENGINE_NAME)
|
||||
)
|
||||
if agent_engine_name:
|
||||
return agent_engine_name
|
||||
|
||||
@@ -194,15 +200,15 @@ class AgentEngineSandboxComputer(BaseComputer):
|
||||
client = self._get_client()
|
||||
|
||||
agent_engine = await asyncio.to_thread(client.agent_engines.create)
|
||||
agent_engine_name = agent_engine.api_resource.name
|
||||
agent_engine_name = cast(str, agent_engine.api_resource.name)
|
||||
|
||||
# Store in session state for sharing
|
||||
self._session_state[_STATE_KEY_AGENT_ENGINE_NAME] = agent_engine_name
|
||||
state[_STATE_KEY_AGENT_ENGINE_NAME] = agent_engine_name
|
||||
logger.info("Created agent engine: %s", agent_engine_name)
|
||||
|
||||
return agent_engine_name
|
||||
|
||||
async def _get_sandbox(self) -> tuple[str, Any]:
|
||||
async def _get_sandbox(self) -> tuple[str, object]:
|
||||
"""Get the sandbox, creating one if needed.
|
||||
|
||||
Returns:
|
||||
@@ -213,13 +219,14 @@ class AgentEngineSandboxComputer(BaseComputer):
|
||||
# Check if provided in constructor (BYOS mode)
|
||||
if self._sandbox_name:
|
||||
# Get sandbox object from name
|
||||
sandbox = await asyncio.to_thread(
|
||||
sandbox: object = await asyncio.to_thread(
|
||||
client.agent_engines.sandboxes.get, name=self._sandbox_name
|
||||
)
|
||||
return self._sandbox_name, sandbox
|
||||
|
||||
# Check session state for existing sandbox
|
||||
sandbox_name = self._session_state.get(_STATE_KEY_SANDBOX_NAME)
|
||||
state = cast(State, self._session_state)
|
||||
sandbox_name = state.get(_STATE_KEY_SANDBOX_NAME)
|
||||
if sandbox_name:
|
||||
sandbox = await asyncio.to_thread(
|
||||
client.agent_engines.sandboxes.get, name=sandbox_name
|
||||
@@ -262,7 +269,7 @@ class AgentEngineSandboxComputer(BaseComputer):
|
||||
sandbox_name = operation.response.name
|
||||
|
||||
# Store in session state for sharing
|
||||
self._session_state[_STATE_KEY_SANDBOX_NAME] = sandbox_name
|
||||
state[_STATE_KEY_SANDBOX_NAME] = sandbox_name
|
||||
logger.info("Created sandbox: %s", sandbox_name)
|
||||
|
||||
return sandbox_name, operation.response
|
||||
@@ -276,9 +283,11 @@ class AgentEngineSandboxComputer(BaseComputer):
|
||||
Returns:
|
||||
The access token.
|
||||
"""
|
||||
state = cast(State, self._session_state)
|
||||
|
||||
# Check session state
|
||||
token = self._session_state.get(_STATE_KEY_ACCESS_TOKEN)
|
||||
expiry = self._session_state.get(_STATE_KEY_TOKEN_EXPIRY, 0)
|
||||
token = cast("str | None", state.get(_STATE_KEY_ACCESS_TOKEN))
|
||||
expiry = cast(float, state.get(_STATE_KEY_TOKEN_EXPIRY, 0))
|
||||
if token and time.time() < expiry - _TOKEN_REFRESH_BUFFER:
|
||||
return token
|
||||
|
||||
@@ -286,17 +295,18 @@ class AgentEngineSandboxComputer(BaseComputer):
|
||||
logger.debug("Generating new access token for sandbox: %s", sandbox_name)
|
||||
client = self._get_client()
|
||||
|
||||
token = await asyncio.to_thread(
|
||||
client.agent_engines.sandboxes.generate_access_token,
|
||||
service_account_email=self._service_account_email,
|
||||
timeout=_DEFAULT_TOKEN_TIMEOUT,
|
||||
token = cast(
|
||||
str,
|
||||
await asyncio.to_thread(
|
||||
client.agent_engines.sandboxes.generate_access_token,
|
||||
service_account_email=self._service_account_email,
|
||||
timeout=_DEFAULT_TOKEN_TIMEOUT,
|
||||
),
|
||||
)
|
||||
|
||||
# Store in session state
|
||||
self._session_state[_STATE_KEY_ACCESS_TOKEN] = token
|
||||
self._session_state[_STATE_KEY_TOKEN_EXPIRY] = (
|
||||
time.time() + _DEFAULT_TOKEN_TIMEOUT
|
||||
)
|
||||
state[_STATE_KEY_ACCESS_TOKEN] = token
|
||||
state[_STATE_KEY_TOKEN_EXPIRY] = time.time() + _DEFAULT_TOKEN_TIMEOUT
|
||||
|
||||
return token
|
||||
|
||||
@@ -313,8 +323,9 @@ class AgentEngineSandboxComputer(BaseComputer):
|
||||
except Exception as e:
|
||||
# Token generation failed - clear cached token and retry
|
||||
logger.warning("Token generation failed, clearing cache: %s", e)
|
||||
self._session_state[_STATE_KEY_ACCESS_TOKEN] = None
|
||||
self._session_state[_STATE_KEY_TOKEN_EXPIRY] = 0
|
||||
state = cast(State, self._session_state)
|
||||
state[_STATE_KEY_ACCESS_TOKEN] = None
|
||||
state[_STATE_KEY_TOKEN_EXPIRY] = 0
|
||||
token = await self._get_access_token(sandbox_name)
|
||||
|
||||
return SandboxClient(
|
||||
|
||||
@@ -59,7 +59,7 @@ def _load_dir(directory: pathlib.Path) -> dict[str, str]:
|
||||
Returns:
|
||||
Dictionary mapping relative file paths to their string content.
|
||||
"""
|
||||
files = {}
|
||||
files: dict[str, str] = {}
|
||||
if directory.exists() and directory.is_dir():
|
||||
for file_path in directory.rglob("*"):
|
||||
if "__pycache__" in file_path.parts:
|
||||
@@ -74,7 +74,9 @@ def _load_dir(directory: pathlib.Path) -> dict[str, str]:
|
||||
return files
|
||||
|
||||
|
||||
def _parse_skill_md_content(content: str) -> tuple[dict, str]:
|
||||
def _parse_skill_md_content(
|
||||
content: str,
|
||||
) -> tuple[dict[str, object], str]:
|
||||
"""Parse SKILL.md from raw content string.
|
||||
|
||||
Args:
|
||||
@@ -104,12 +106,17 @@ def _parse_skill_md_content(content: str) -> tuple[dict, str]:
|
||||
if not isinstance(parsed, dict):
|
||||
raise ValueError("SKILL.md frontmatter must be a YAML mapping")
|
||||
|
||||
return parsed, body
|
||||
frontmatter: dict[str, object] = {}
|
||||
for key, value in parsed.items():
|
||||
if not isinstance(key, str):
|
||||
raise ValueError("SKILL.md frontmatter keys must be strings")
|
||||
frontmatter[key] = value
|
||||
return frontmatter, body
|
||||
|
||||
|
||||
def _parse_skill_md(
|
||||
skill_dir: pathlib.Path,
|
||||
) -> tuple[dict, str, pathlib.Path]:
|
||||
) -> tuple[dict[str, object], str, pathlib.Path]:
|
||||
"""Parse SKILL.md from a skill directory.
|
||||
|
||||
Args:
|
||||
@@ -477,7 +484,7 @@ def _list_skills_in_dir(
|
||||
Dictionary mapping skill IDs to their frontmatter.
|
||||
"""
|
||||
skills_base_path = pathlib.Path(skills_base_path).resolve()
|
||||
skills = {}
|
||||
skills: dict[str, models.Frontmatter] = {}
|
||||
|
||||
if not skills_base_path.is_dir():
|
||||
logging.warning(
|
||||
@@ -546,7 +553,7 @@ def _list_skills_in_gcs_dir(
|
||||
pass
|
||||
logging.info("Found %s skills in GCS.", iterator.prefixes)
|
||||
|
||||
skills = {}
|
||||
skills: dict[str, models.Frontmatter] = {}
|
||||
for skill_prefix in sorted(iterator.prefixes):
|
||||
manifest_blob = bucket.blob(f"{skill_prefix}SKILL.md")
|
||||
|
||||
@@ -628,10 +635,10 @@ def _load_skill_from_gcs_dir(
|
||||
f" name '{skill_name_expected}'."
|
||||
)
|
||||
|
||||
def _load_files_in_dir(subdir: str) -> Dict[str, Union[str, bytes]]:
|
||||
def _load_files_in_dir(subdir: str) -> dict[str, Union[str, bytes]]:
|
||||
prefix = f"{skill_dir_prefix}{subdir}/"
|
||||
blobs = bucket.list_blobs(prefix=prefix)
|
||||
result = {}
|
||||
result: dict[str, str | bytes] = {}
|
||||
|
||||
for blob in blobs:
|
||||
relative_path = blob.name[len(prefix) :]
|
||||
@@ -648,7 +655,7 @@ def _load_skill_from_gcs_dir(
|
||||
assets = _load_files_in_dir("assets")
|
||||
raw_scripts = _load_files_in_dir("scripts")
|
||||
|
||||
scripts = {}
|
||||
scripts: dict[str, models.Script] = {}
|
||||
for name, src in raw_scripts.items():
|
||||
if isinstance(src, bytes):
|
||||
try:
|
||||
|
||||
@@ -16,6 +16,7 @@ from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
from collections import deque
|
||||
import concurrent.futures
|
||||
from contextlib import AbstractAsyncContextManager
|
||||
from contextlib import AsyncExitStack
|
||||
import contextvars
|
||||
@@ -26,14 +27,17 @@ import logging
|
||||
import os
|
||||
import sys
|
||||
import threading
|
||||
from types import TracebackType
|
||||
from typing import Any
|
||||
from typing import AsyncIterator
|
||||
from typing import Callable
|
||||
from typing import cast
|
||||
from typing import Dict
|
||||
from typing import Optional
|
||||
from typing import Protocol
|
||||
from typing import runtime_checkable
|
||||
from typing import TextIO
|
||||
from typing import TYPE_CHECKING
|
||||
import urllib.parse
|
||||
|
||||
import google.auth
|
||||
@@ -41,20 +45,26 @@ import google.auth.credentials
|
||||
from google.auth.transport.requests import Request
|
||||
import httpx
|
||||
|
||||
try:
|
||||
_AIO_SUPPORTED = False
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from google.auth.aio.credentials import Credentials as AsyncCredentials
|
||||
from google.auth.aio.transport import Response as AsyncResponse
|
||||
from google.auth.aio.transport.sessions import AsyncAuthorizedSession
|
||||
else:
|
||||
try:
|
||||
from google.auth.aio.credentials import Credentials as AsyncCredentials
|
||||
from google.auth.aio.transport.sessions import AsyncAuthorizedSession
|
||||
|
||||
_AIO_SUPPORTED = True
|
||||
except ImportError:
|
||||
_AIO_SUPPORTED = True
|
||||
except ImportError:
|
||||
|
||||
class AsyncCredentials: # pylint: disable=g-bad-classes
|
||||
pass
|
||||
class AsyncCredentials: # pylint: disable=g-bad-classes
|
||||
pass
|
||||
|
||||
class AsyncAuthorizedSession: # pylint: disable=g-bad-classes
|
||||
pass
|
||||
class AsyncAuthorizedSession: # pylint: disable=g-bad-classes
|
||||
pass
|
||||
|
||||
_AIO_SUPPORTED = False
|
||||
|
||||
from mcp import ClientSession
|
||||
from mcp import SamplingCapability
|
||||
@@ -63,8 +73,8 @@ from mcp.client.session import ElicitationFnT
|
||||
from mcp.client.session import SamplingFnT
|
||||
from mcp.client.sse import sse_client
|
||||
from mcp.client.stdio import stdio_client
|
||||
from mcp.client.streamable_http import create_mcp_http_client as _create_mcp_http_client
|
||||
from mcp.client.streamable_http import McpHttpClientFactory
|
||||
from mcp.client.streamable_http import create_mcp_http_client as _create_mcp_http_client # type: ignore[attr-defined]
|
||||
from mcp.client.streamable_http import McpHttpClientFactory # type: ignore[attr-defined]
|
||||
from mcp.client.streamable_http import streamable_http_client
|
||||
from pydantic import BaseModel
|
||||
from pydantic import ConfigDict
|
||||
@@ -122,7 +132,7 @@ class _StreamableHttpClientWrapper:
|
||||
url: str,
|
||||
http_client: httpx.AsyncClient,
|
||||
terminate_on_close: bool = True,
|
||||
):
|
||||
) -> None:
|
||||
self.url = url
|
||||
self.http_client = http_client
|
||||
self.terminate_on_close = terminate_on_close
|
||||
@@ -148,7 +158,12 @@ class _StreamableHttpClientWrapper:
|
||||
await self.http_client.__aexit__(type(e), e, e.__traceback__)
|
||||
raise
|
||||
|
||||
async def __aexit__(self, exc_type, exc_val, exc_tb) -> None:
|
||||
async def __aexit__(
|
||||
self,
|
||||
exc_type: type[BaseException] | None,
|
||||
exc_val: BaseException | None,
|
||||
exc_tb: TracebackType | None,
|
||||
) -> None:
|
||||
try:
|
||||
await self.ctx_mgr.__aexit__(exc_type, exc_val, exc_tb)
|
||||
finally:
|
||||
@@ -231,7 +246,7 @@ class _DebugHttpxClientFactory:
|
||||
self,
|
||||
base_factory: CheckableMcpHttpClientFactory,
|
||||
session_manager: MCPSessionManager | None = None,
|
||||
):
|
||||
) -> None:
|
||||
self._base_factory = base_factory
|
||||
self._session_manager = session_manager
|
||||
|
||||
@@ -255,7 +270,7 @@ class _DebugHttpxClientFactory:
|
||||
or query_params.get('session_id', [None])[0]
|
||||
)
|
||||
|
||||
async def _response_hook(self, response: httpx.Response):
|
||||
async def _response_hook(self, response: httpx.Response) -> None:
|
||||
debug_list = None
|
||||
if self._session_manager is not None:
|
||||
session_id = self._extract_session_id(response)
|
||||
@@ -377,14 +392,18 @@ def retry_on_errors(func):
|
||||
return wrapper
|
||||
|
||||
|
||||
class _RefreshableAsyncCredentials(AsyncCredentials):
|
||||
# `google.auth.*` is resolved with `follow_imports = "skip"`, so the base class
|
||||
# is `Any` here and strict mode rejects subclassing it. The alternative is to
|
||||
# swap in a fake base class under `TYPE_CHECKING`, which makes the checker read
|
||||
# a class hierarchy that does not exist at runtime.
|
||||
class _RefreshableAsyncCredentials(AsyncCredentials): # type: ignore[misc]
|
||||
"""Adapter to refresh sync credentials asynchronously."""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
creds: google.auth.credentials.Credentials,
|
||||
target_host: str | None = None,
|
||||
):
|
||||
) -> None:
|
||||
super().__init__()
|
||||
self._creds = creds
|
||||
self._target_host = target_host
|
||||
@@ -422,11 +441,11 @@ class _RefreshableAsyncCredentials(AsyncCredentials):
|
||||
class _GoogleAuthAsyncByteStream(httpx.AsyncByteStream):
|
||||
"""Adapter to bridge google-auth Response.content with httpx.AsyncByteStream."""
|
||||
|
||||
def __init__(self, auth_response: Any):
|
||||
def __init__(self, auth_response: AsyncResponse) -> None:
|
||||
self._auth_response = auth_response
|
||||
|
||||
async def __aiter__(self) -> AsyncIterator[bytes]:
|
||||
async for chunk in self._auth_response.content():
|
||||
async for chunk in self._auth_response.content(1024):
|
||||
yield chunk
|
||||
|
||||
async def aclose(self) -> None:
|
||||
@@ -436,7 +455,7 @@ class _GoogleAuthAsyncByteStream(httpx.AsyncByteStream):
|
||||
class _GoogleAuthAsyncTransport(httpx.AsyncBaseTransport):
|
||||
"""Adapter to bridge google-auth AsyncAuthorizedSession with httpx.AsyncBaseTransport."""
|
||||
|
||||
def __init__(self, auth_session: Any):
|
||||
def __init__(self, auth_session: AsyncAuthorizedSession) -> None:
|
||||
self._auth_session = auth_session
|
||||
|
||||
async def handle_async_request(
|
||||
@@ -457,7 +476,7 @@ class _GoogleAuthAsyncTransport(httpx.AsyncBaseTransport):
|
||||
# prevent aiohttp from forcibly closing the stream after sse_read_timeout.
|
||||
timeout_val = 0.0
|
||||
|
||||
auth_response: Any = await self._auth_session.request(
|
||||
auth_response = await self._auth_session.request(
|
||||
method=request.method,
|
||||
url=str(request.url),
|
||||
data=content if content else None,
|
||||
@@ -489,7 +508,7 @@ class _GoogleAuthAsyncTransport(httpx.AsyncBaseTransport):
|
||||
class _SharedAsyncTransport(httpx.AsyncBaseTransport):
|
||||
"""Wrapper transport that prevents the wrapped transport from being closed."""
|
||||
|
||||
def __init__(self, transport: httpx.AsyncBaseTransport):
|
||||
def __init__(self, transport: httpx.AsyncBaseTransport) -> None:
|
||||
self._transport = transport
|
||||
|
||||
async def handle_async_request(
|
||||
@@ -507,7 +526,7 @@ def _create_mtls_client_factory(
|
||||
"""Returns a factory that creates httpx.AsyncClient using the mtls_transport."""
|
||||
|
||||
def factory(
|
||||
headers: dict[str, Any] | None = None,
|
||||
headers: dict[str, str] | None = None,
|
||||
timeout: httpx.Timeout | None = None,
|
||||
auth: httpx.Auth | None = None,
|
||||
) -> httpx.AsyncClient:
|
||||
@@ -543,7 +562,7 @@ class MCPSessionManager:
|
||||
sampling_callback: SamplingFnT | None = None,
|
||||
sampling_capabilities: SamplingCapability | None = None,
|
||||
elicitation_callback: ElicitationFnT | None = None,
|
||||
):
|
||||
) -> None:
|
||||
"""Initializes the MCP session manager.
|
||||
|
||||
Args:
|
||||
@@ -562,6 +581,11 @@ class MCPSessionManager:
|
||||
self._sampling_callback = sampling_callback
|
||||
self._sampling_capabilities = sampling_capabilities
|
||||
self._elicitation_callback = elicitation_callback
|
||||
self._connection_params: (
|
||||
StdioConnectionParams
|
||||
| SseConnectionParams
|
||||
| StreamableHTTPConnectionParams
|
||||
)
|
||||
|
||||
if isinstance(connection_params, StdioServerParameters):
|
||||
# So far timeout is not configurable. Given MCP is still evolving, we
|
||||
@@ -604,7 +628,8 @@ class MCPSessionManager:
|
||||
] = {}
|
||||
|
||||
def _make_on_session_created(self, session_key: str) -> Callable[[str], None]:
|
||||
def on_session_created(session_id: str):
|
||||
|
||||
def on_session_created(session_id: str) -> None:
|
||||
logger.debug('Session created: %s -> %s', session_id, session_key)
|
||||
self._session_id_to_key[session_id] = session_key
|
||||
|
||||
@@ -612,7 +637,7 @@ class MCPSessionManager:
|
||||
|
||||
def _set_active_debug_list(
|
||||
self, session_key: str, debug_list: list[dict[str, Any]]
|
||||
):
|
||||
) -> None:
|
||||
self._active_debug_lists[session_key] = debug_list
|
||||
|
||||
def _get_active_debug_list_by_session_id(
|
||||
@@ -720,18 +745,18 @@ class MCPSessionManager:
|
||||
Returns:
|
||||
Merged headers dictionary, or None if no headers are provided.
|
||||
"""
|
||||
if isinstance(self._connection_params, StdioConnectionParams) or isinstance(
|
||||
self._connection_params, StdioServerParameters
|
||||
):
|
||||
if isinstance(self._connection_params, StdioConnectionParams):
|
||||
# Stdio connections don't support headers
|
||||
return None
|
||||
|
||||
base_headers = {}
|
||||
base_headers: Dict[str, str] = {}
|
||||
if (
|
||||
hasattr(self._connection_params, 'headers')
|
||||
and self._connection_params.headers
|
||||
):
|
||||
base_headers = self._connection_params.headers.copy()
|
||||
base_headers = cast(
|
||||
'Dict[str, str]', self._connection_params.headers
|
||||
).copy()
|
||||
|
||||
if additional_headers:
|
||||
base_headers.update(additional_headers)
|
||||
@@ -774,7 +799,7 @@ class MCPSessionManager:
|
||||
session_key: str,
|
||||
exit_stack: AsyncExitStack,
|
||||
stored_loop: asyncio.AbstractEventLoop,
|
||||
):
|
||||
) -> None:
|
||||
"""Cleans up a session, handling different event loops safely.
|
||||
|
||||
Args:
|
||||
@@ -803,7 +828,7 @@ class MCPSessionManager:
|
||||
)
|
||||
|
||||
# Attach a callback so errors don't go unnoticed
|
||||
def cleanup_done(f: asyncio.Future):
|
||||
def cleanup_done(f: concurrent.futures.Future[None]) -> None:
|
||||
try:
|
||||
if f.exception():
|
||||
logger.warning(
|
||||
@@ -844,18 +869,19 @@ class MCPSessionManager:
|
||||
) -> AbstractAsyncContextManager[Any]:
|
||||
"""Creates an MCP client based on the connection parameters.
|
||||
|
||||
Args:
|
||||
session_key: Optional session key for this client.
|
||||
merged_headers: Optional headers to include in the connection. Only
|
||||
applicable for SSE and StreamableHTTP connections.
|
||||
mtls_transport: Optional mTLS transport for the HTTP client.
|
||||
Args:
|
||||
session_key: Optional session key for this client.
|
||||
merged_headers: Optional headers to include in the connection. Only
|
||||
applicable for SSE and StreamableHTTP connections.
|
||||
mtls_transport: Optional mTLS transport for the HTTP client.
|
||||
|
||||
Returns:
|
||||
The appropriate MCP client instance.
|
||||
Returns:
|
||||
The appropriate MCP client instance.
|
||||
|
||||
Raises:
|
||||
ValueError: If the connection parameters are not supported.
|
||||
ValueError: If the connection parameters are not supported.
|
||||
"""
|
||||
client: AbstractAsyncContextManager[Any]
|
||||
if isinstance(self._connection_params, StdioConnectionParams):
|
||||
client = stdio_client(
|
||||
server=self._connection_params.server_params,
|
||||
@@ -974,15 +1000,10 @@ class MCPSessionManager:
|
||||
|
||||
# Create a new session (either first time or replacing disconnected one)
|
||||
exit_stack = AsyncExitStack()
|
||||
timeout_in_seconds = (
|
||||
self._connection_params.timeout
|
||||
if hasattr(self._connection_params, 'timeout')
|
||||
else None
|
||||
)
|
||||
sse_read_timeout_in_seconds = (
|
||||
self._connection_params.sse_read_timeout
|
||||
if hasattr(self._connection_params, 'sse_read_timeout')
|
||||
else None
|
||||
# Connection params are extensible, so neither timeout is guaranteed.
|
||||
timeout_in_seconds = getattr(self._connection_params, 'timeout', None)
|
||||
sse_read_timeout_in_seconds = getattr(
|
||||
self._connection_params, 'sse_read_timeout', None
|
||||
)
|
||||
|
||||
try:
|
||||
@@ -1038,7 +1059,7 @@ class MCPSessionManager:
|
||||
)
|
||||
raise ConnectionError(f'Failed to create MCP session: {e}') from e
|
||||
|
||||
def __getstate__(self):
|
||||
def __getstate__(self) -> dict[str, Any]:
|
||||
"""Custom pickling to exclude non-picklable runtime objects."""
|
||||
state = self.__dict__.copy()
|
||||
# Remove unpicklable entries or those that shouldn't persist across pickle
|
||||
@@ -1055,7 +1076,7 @@ class MCPSessionManager:
|
||||
|
||||
return state
|
||||
|
||||
def __setstate__(self, state):
|
||||
def __setstate__(self, state: dict[str, Any]) -> None:
|
||||
"""Custom unpickling to restore state."""
|
||||
self.__dict__.update(state)
|
||||
# Re-initialize members that were not pickled
|
||||
@@ -1070,7 +1091,7 @@ class MCPSessionManager:
|
||||
if not hasattr(self, '_errlog') or self._errlog is None:
|
||||
self._errlog = sys.stderr
|
||||
|
||||
async def close(self):
|
||||
async def close(self) -> None:
|
||||
"""Closes all sessions and cleans up resources."""
|
||||
async with self._session_lock:
|
||||
for session_key in list(self._sessions.keys()):
|
||||
|
||||
@@ -14,7 +14,6 @@
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import base64
|
||||
from collections.abc import Awaitable
|
||||
import inspect
|
||||
@@ -24,8 +23,10 @@ from typing import Callable
|
||||
from typing import cast
|
||||
from typing import Protocol
|
||||
from typing import runtime_checkable
|
||||
from typing import TypeGuard
|
||||
import warnings
|
||||
|
||||
from fastapi.openapi.models import APIKey
|
||||
from fastapi.openapi.models import APIKeyIn
|
||||
from google.genai.types import FunctionDeclaration
|
||||
from mcp.shared.exceptions import McpError
|
||||
@@ -59,6 +60,8 @@ from .session_context import SessionContext
|
||||
|
||||
logger = logging.getLogger("google_adk." + __name__)
|
||||
|
||||
_ConfirmationPredicate = Callable[..., bool | Awaitable[bool]]
|
||||
|
||||
|
||||
@runtime_checkable
|
||||
class ProgressCallbackFactory(Protocol):
|
||||
@@ -122,6 +125,23 @@ class ProgressCallbackFactory(Protocol):
|
||||
...
|
||||
|
||||
|
||||
def _is_async_callable(value: object) -> bool:
|
||||
return callable(value) and (
|
||||
inspect.iscoroutinefunction(value)
|
||||
or inspect.iscoroutinefunction(getattr(value, "__call__", None))
|
||||
)
|
||||
|
||||
|
||||
def _is_progress_callback(value: object) -> TypeGuard[ProgressFnT]:
|
||||
return _is_async_callable(value)
|
||||
|
||||
|
||||
def _is_progress_callback_factory(
|
||||
value: object,
|
||||
) -> TypeGuard[ProgressCallbackFactory]:
|
||||
return callable(value) and not _is_async_callable(value)
|
||||
|
||||
|
||||
class McpTool(BaseAuthenticatedTool):
|
||||
"""Turns an MCP Tool into an ADK Tool.
|
||||
|
||||
@@ -148,7 +168,7 @@ class McpTool(BaseAuthenticatedTool):
|
||||
| None
|
||||
) = None,
|
||||
progress_callback: ProgressFnT | ProgressCallbackFactory | None = None,
|
||||
):
|
||||
) -> None:
|
||||
"""Initializes an McpTool.
|
||||
|
||||
This tool wraps an MCP Tool interface and uses a session manager to
|
||||
@@ -234,7 +254,9 @@ class McpTool(BaseAuthenticatedTool):
|
||||
# Format: meta.ui.visibility
|
||||
ui = meta.get("ui", {})
|
||||
if isinstance(ui, dict):
|
||||
return ui.get("visibility", [])
|
||||
visibility = ui.get("visibility", [])
|
||||
if isinstance(visibility, list):
|
||||
return [item for item in visibility if isinstance(item, str)]
|
||||
return []
|
||||
|
||||
@property
|
||||
@@ -267,8 +289,10 @@ class McpTool(BaseAuthenticatedTool):
|
||||
return None
|
||||
|
||||
async def _invoke_callable(
|
||||
self, target: Callable[..., Any], args_to_call: dict[str, Any]
|
||||
) -> Any:
|
||||
self,
|
||||
target: _ConfirmationPredicate,
|
||||
args_to_call: dict[str, Any],
|
||||
) -> bool:
|
||||
"""Invokes a callable, handling both sync and async cases."""
|
||||
|
||||
# Functions are callable objects, but not all callable objects are functions
|
||||
@@ -279,9 +303,10 @@ class McpTool(BaseAuthenticatedTool):
|
||||
and inspect.iscoroutinefunction(target.__call__)
|
||||
)
|
||||
if is_async:
|
||||
return await target(**args_to_call)
|
||||
awaitable_result = cast(Awaitable[bool], target(**args_to_call))
|
||||
return await awaitable_result
|
||||
else:
|
||||
return target(**args_to_call)
|
||||
return cast(bool, target(**args_to_call))
|
||||
|
||||
def _prepare_callable_args(
|
||||
self,
|
||||
@@ -325,9 +350,8 @@ class McpTool(BaseAuthenticatedTool):
|
||||
args_to_call = self._prepare_callable_args(
|
||||
self._require_confirmation, args, tool_context
|
||||
)
|
||||
return cast(
|
||||
bool,
|
||||
await self._invoke_callable(self._require_confirmation, args_to_call),
|
||||
return await self._invoke_callable(
|
||||
self._require_confirmation, args_to_call
|
||||
)
|
||||
return bool(self._require_confirmation)
|
||||
|
||||
@@ -395,7 +419,11 @@ class McpTool(BaseAuthenticatedTool):
|
||||
@retry_on_errors
|
||||
@override
|
||||
async def _run_async_impl(
|
||||
self, *, args, tool_context: ToolContext, credential: AuthCredential
|
||||
self,
|
||||
*,
|
||||
args: dict[str, Any],
|
||||
tool_context: ToolContext,
|
||||
credential: AuthCredential,
|
||||
) -> dict[str, Any]:
|
||||
"""Runs the tool asynchronously.
|
||||
|
||||
@@ -408,13 +436,16 @@ class McpTool(BaseAuthenticatedTool):
|
||||
"""
|
||||
# Extract headers from credential for session pooling
|
||||
auth_headers = await self._get_headers(tool_context, credential)
|
||||
dynamic_headers = None
|
||||
dynamic_headers: dict[str, str] | None = None
|
||||
if self._header_provider:
|
||||
dynamic_headers = self._header_provider(
|
||||
provided_headers = self._header_provider(
|
||||
ReadonlyContext(tool_context._invocation_context) # pylint: disable=protected-access
|
||||
)
|
||||
if inspect.isawaitable(dynamic_headers):
|
||||
dynamic_headers = await dynamic_headers
|
||||
dynamic_headers = (
|
||||
await provided_headers
|
||||
if inspect.isawaitable(provided_headers)
|
||||
else provided_headers
|
||||
)
|
||||
|
||||
headers: dict[str, str] = {}
|
||||
if auth_headers:
|
||||
@@ -513,22 +544,20 @@ class McpTool(BaseAuthenticatedTool):
|
||||
):
|
||||
return None
|
||||
|
||||
# Determine if callback is a factory by checking if it's a coroutine
|
||||
# function. ProgressFnT is an async function, while ProgressCallbackFactory
|
||||
# is a sync function that returns an async function.
|
||||
if asyncio.iscoroutinefunction(self._progress_callback):
|
||||
return self._progress_callback
|
||||
progress_callback = self._progress_callback
|
||||
|
||||
# If it's a regular callable (not async), treat it as a factory
|
||||
if callable(self._progress_callback) and not inspect.iscoroutinefunction(
|
||||
self._progress_callback
|
||||
):
|
||||
return self._progress_callback(self.name, callback_context=tool_context)
|
||||
# ProgressFnT is asynchronous, while ProgressCallbackFactory is a
|
||||
# synchronous function that returns an asynchronous callback.
|
||||
if _is_progress_callback(progress_callback):
|
||||
return progress_callback
|
||||
|
||||
return self._progress_callback
|
||||
if _is_progress_callback_factory(progress_callback):
|
||||
return progress_callback(self.name, callback_context=tool_context)
|
||||
|
||||
raise TypeError("Invalid MCP progress callback")
|
||||
|
||||
async def _get_headers(
|
||||
self, tool_context: ToolContext, credential: AuthCredential
|
||||
self, tool_context: ToolContext, credential: AuthCredential | None
|
||||
) -> dict[str, str] | None:
|
||||
"""Extracts authentication headers from credentials.
|
||||
|
||||
@@ -580,33 +609,33 @@ class McpTool(BaseAuthenticatedTool):
|
||||
headers = headers or {}
|
||||
headers.update(credential.http.additional_headers)
|
||||
elif credential.api_key:
|
||||
if (
|
||||
not self._credentials_manager
|
||||
or not self._credentials_manager._auth_config
|
||||
):
|
||||
credentials_manager = self._credentials_manager
|
||||
auth_config = (
|
||||
credentials_manager._auth_config if credentials_manager else None
|
||||
)
|
||||
if auth_config is None:
|
||||
error_msg = (
|
||||
"Cannot find corresponding auth scheme for API key credential"
|
||||
f" {credential}"
|
||||
)
|
||||
logger.error(error_msg)
|
||||
raise ValueError(error_msg)
|
||||
elif (
|
||||
self._credentials_manager._auth_config.auth_scheme.in_
|
||||
!= APIKeyIn.header
|
||||
):
|
||||
auth_scheme = auth_config.auth_scheme
|
||||
if not isinstance(auth_scheme, APIKey):
|
||||
error_msg = (
|
||||
"McpTool only supports header-based API key authentication."
|
||||
" Configured location:"
|
||||
f" {self._credentials_manager._auth_config.auth_scheme.in_}"
|
||||
"API key credentials require an APIKey authentication scheme,"
|
||||
f" got {type(auth_scheme).__name__}."
|
||||
)
|
||||
logger.error(error_msg)
|
||||
raise ValueError(error_msg)
|
||||
else:
|
||||
headers = {
|
||||
self._credentials_manager._auth_config.auth_scheme.name: (
|
||||
credential.api_key
|
||||
)
|
||||
}
|
||||
if auth_scheme.in_ != APIKeyIn.header:
|
||||
error_msg = (
|
||||
"McpTool only supports header-based API key authentication."
|
||||
f" Configured location: {auth_scheme.in_}"
|
||||
)
|
||||
logger.error(error_msg)
|
||||
raise ValueError(error_msg)
|
||||
headers = {auth_scheme.name: credential.api_key}
|
||||
elif credential.service_account:
|
||||
# Service accounts should be exchanged for access tokens before reaching this point
|
||||
logger.warning(
|
||||
@@ -620,7 +649,7 @@ class McpTool(BaseAuthenticatedTool):
|
||||
class MCPTool(McpTool):
|
||||
"""Deprecated name, use `McpTool` instead."""
|
||||
|
||||
def __init__(self, *args, **kwargs):
|
||||
def __init__(self, *args: Any, **kwargs: Any) -> None:
|
||||
warnings.warn(
|
||||
"MCPTool class is deprecated, use `McpTool` instead.",
|
||||
DeprecationWarning,
|
||||
|
||||
@@ -30,6 +30,8 @@ from typing import TypeVar
|
||||
from typing import Union
|
||||
import warnings
|
||||
|
||||
from fastapi.openapi.models import APIKeyIn
|
||||
from mcp import ClientSession
|
||||
from mcp import SamplingCapability
|
||||
from mcp import StdioServerParameters
|
||||
from mcp.client.session import ElicitationFnT
|
||||
@@ -63,6 +65,12 @@ logger = logging.getLogger("google_adk." + __name__)
|
||||
|
||||
|
||||
T = TypeVar("T")
|
||||
_ConnectionParams = Union[
|
||||
StdioServerParameters,
|
||||
StdioConnectionParams,
|
||||
SseConnectionParams,
|
||||
StreamableHTTPConnectionParams,
|
||||
]
|
||||
|
||||
|
||||
class McpToolset(BaseToolset):
|
||||
@@ -98,12 +106,7 @@ class McpToolset(BaseToolset):
|
||||
def __init__(
|
||||
self,
|
||||
*,
|
||||
connection_params: (
|
||||
StdioServerParameters
|
||||
| StdioConnectionParams
|
||||
| SseConnectionParams
|
||||
| StreamableHTTPConnectionParams
|
||||
),
|
||||
connection_params: _ConnectionParams,
|
||||
tool_filter: ToolPredicate | list[str] | None = None,
|
||||
tool_name_prefix: str | None = None,
|
||||
errlog: TextIO = sys.stderr,
|
||||
@@ -123,7 +126,7 @@ class McpToolset(BaseToolset):
|
||||
sampling_capabilities: SamplingCapability | None = None,
|
||||
elicitation_callback: ElicitationFnT | None = None,
|
||||
credential_key: str | None = None,
|
||||
):
|
||||
) -> None:
|
||||
"""Initializes the McpToolset.
|
||||
|
||||
Args:
|
||||
@@ -222,7 +225,7 @@ class McpToolset(BaseToolset):
|
||||
return None
|
||||
|
||||
credential = None
|
||||
if readonly_context:
|
||||
if readonly_context and self._auth_config.credential_key:
|
||||
credential = readonly_context.get_credential(
|
||||
self._auth_config.credential_key
|
||||
)
|
||||
@@ -274,31 +277,24 @@ class McpToolset(BaseToolset):
|
||||
headers.update(credential.http.additional_headers)
|
||||
elif credential.api_key:
|
||||
# For API key, use the auth scheme to determine header name
|
||||
if self._auth_config.auth_scheme:
|
||||
from fastapi.openapi.models import APIKeyIn
|
||||
|
||||
if hasattr(self._auth_config.auth_scheme, "in_"):
|
||||
if self._auth_config.auth_scheme.in_ == APIKeyIn.header:
|
||||
headers = {self._auth_config.auth_scheme.name: credential.api_key}
|
||||
auth_scheme = self._auth_config.auth_scheme
|
||||
if auth_scheme:
|
||||
if hasattr(auth_scheme, "in_"):
|
||||
if auth_scheme.in_ == APIKeyIn.header:
|
||||
headers = {auth_scheme.name: credential.api_key}
|
||||
else:
|
||||
logger.warning(
|
||||
raise ValueError(
|
||||
"McpToolset only supports header-based API key authentication."
|
||||
" Configured location: %s",
|
||||
self._auth_config.auth_scheme.in_,
|
||||
f" Configured location: {auth_scheme.in_}"
|
||||
)
|
||||
else:
|
||||
# Default to using scheme name as header
|
||||
headers = {self._auth_config.auth_scheme.name: credential.api_key}
|
||||
headers = {auth_scheme.name: credential.api_key}
|
||||
|
||||
return headers
|
||||
|
||||
@property
|
||||
def connection_params(self) -> Union[
|
||||
StdioServerParameters,
|
||||
StdioConnectionParams,
|
||||
SseConnectionParams,
|
||||
StreamableHTTPConnectionParams,
|
||||
]:
|
||||
def connection_params(self) -> _ConnectionParams:
|
||||
return self._connection_params
|
||||
|
||||
@property
|
||||
@@ -329,7 +325,7 @@ class McpToolset(BaseToolset):
|
||||
|
||||
async def _execute_with_session(
|
||||
self,
|
||||
coroutine_func: Callable[[Any], Awaitable[T]],
|
||||
coroutine_func: Callable[[ClientSession], Awaitable[T]],
|
||||
error_message: str,
|
||||
readonly_context: Optional[ReadonlyContext] = None,
|
||||
) -> T:
|
||||
@@ -344,9 +340,12 @@ class McpToolset(BaseToolset):
|
||||
|
||||
# Add headers from header_provider if available
|
||||
if self._header_provider and readonly_context:
|
||||
provider_headers = self._header_provider(readonly_context)
|
||||
if inspect.isawaitable(provider_headers):
|
||||
provider_headers = await provider_headers
|
||||
provided_headers = self._header_provider(readonly_context)
|
||||
provider_headers = (
|
||||
await provided_headers
|
||||
if inspect.isawaitable(provided_headers)
|
||||
else provided_headers
|
||||
)
|
||||
if provider_headers:
|
||||
headers.update(provider_headers)
|
||||
|
||||
@@ -406,7 +405,7 @@ class McpToolset(BaseToolset):
|
||||
)
|
||||
|
||||
# Apply filtering based on context and tool_filter
|
||||
tools = []
|
||||
tools: List[BaseTool] = []
|
||||
for tool in tools_response.tools:
|
||||
mcp_tool = MCPTool(
|
||||
mcp_tool=tool,
|
||||
@@ -515,6 +514,7 @@ class McpToolset(BaseToolset):
|
||||
"""Creates an McpToolset from a configuration object."""
|
||||
mcp_toolset_config = McpToolsetConfig.model_validate(config.model_dump())
|
||||
|
||||
connection_params: _ConnectionParams
|
||||
if mcp_toolset_config.stdio_server_params:
|
||||
connection_params = mcp_toolset_config.stdio_server_params
|
||||
elif mcp_toolset_config.stdio_connection_params:
|
||||
@@ -536,14 +536,14 @@ class McpToolset(BaseToolset):
|
||||
use_mcp_resources=mcp_toolset_config.use_mcp_resources,
|
||||
)
|
||||
|
||||
def __getstate__(self):
|
||||
def __getstate__(self) -> dict[str, Any]:
|
||||
"""Custom pickling to exclude non-picklable runtime objects."""
|
||||
state = self.__dict__.copy()
|
||||
# Remove unpicklable file-like objects
|
||||
state.pop("_errlog", None)
|
||||
return state
|
||||
|
||||
def __setstate__(self, state):
|
||||
def __setstate__(self, state: dict[str, Any]) -> None:
|
||||
"""Custom unpickling to restore state."""
|
||||
self.__dict__.update(state)
|
||||
# Default to sys.stderr if _errlog was removed during pickling
|
||||
@@ -554,7 +554,7 @@ class McpToolset(BaseToolset):
|
||||
class MCPToolset(McpToolset):
|
||||
"""Deprecated name, use `McpToolset` instead."""
|
||||
|
||||
def __init__(self, *args, **kwargs):
|
||||
def __init__(self, *args: Any, **kwargs: Any) -> None:
|
||||
warnings.warn(
|
||||
"MCPToolset class is deprecated, use `McpToolset` instead.",
|
||||
DeprecationWarning,
|
||||
@@ -589,7 +589,7 @@ class McpToolsetConfig(BaseToolConfig):
|
||||
use_mcp_resources: bool = False
|
||||
|
||||
@model_validator(mode="after")
|
||||
def _check_only_one_params_field(self):
|
||||
def _check_only_one_params_field(self) -> McpToolsetConfig:
|
||||
param_fields = [
|
||||
self.stdio_server_params,
|
||||
self.stdio_connection_params,
|
||||
|
||||
@@ -342,10 +342,12 @@ class SessionContext:
|
||||
# to the read/write MemoryObjectStreams needed to build the
|
||||
# ClientSession. We limit to the first two values to be compatible
|
||||
# with all clients.
|
||||
read_stream, write_stream = transports[:2]
|
||||
if self._is_stdio:
|
||||
session = await exit_stack.enter_async_context(
|
||||
ClientSession(
|
||||
*transports[:2],
|
||||
read_stream,
|
||||
write_stream,
|
||||
read_timeout_seconds=timedelta(seconds=self._timeout)
|
||||
if self._timeout is not None
|
||||
else None,
|
||||
@@ -359,7 +361,8 @@ class SessionContext:
|
||||
# instead of the connection timeout as the read_timeout for the session.
|
||||
session = await exit_stack.enter_async_context(
|
||||
ClientSession(
|
||||
*transports[:2],
|
||||
read_stream,
|
||||
write_stream,
|
||||
read_timeout_seconds=timedelta(seconds=self._sse_read_timeout)
|
||||
if self._sse_read_timeout is not None
|
||||
else None,
|
||||
|
||||
@@ -690,6 +690,13 @@ class TestAgentRegistry:
|
||||
):
|
||||
registry._make_request("test-path")
|
||||
|
||||
def test_make_request_handles_http_error_without_response(self, registry):
|
||||
error = requests.exceptions.HTTPError("Connection closed")
|
||||
registry._session.get.side_effect = error
|
||||
|
||||
with pytest.raises(RuntimeError, match="API request failed:"):
|
||||
registry._make_request("test-path")
|
||||
|
||||
def test_make_request_raises_request_error(self, registry):
|
||||
error = requests.exceptions.RequestException(
|
||||
"Connection failed", request=MagicMock()
|
||||
|
||||
@@ -674,6 +674,36 @@ def test_execute_sql_select_stmt(write_mode):
|
||||
assert result == {"status": "SUCCESS", "rows": query_result}
|
||||
|
||||
|
||||
def test_execute_sql_protected_requires_session_metadata():
|
||||
"""Test that protected mode rejects an incomplete session response."""
|
||||
credentials = mock.create_autospec(Credentials, instance=True)
|
||||
tool_settings = BigQueryToolConfig(write_mode=WriteMode.PROTECTED)
|
||||
tool_context = mock.create_autospec(ToolContext, instance=True)
|
||||
tool_context.state.get.return_value = None
|
||||
|
||||
with mock.patch.object(bigquery, "Client", autospec=True) as Client:
|
||||
bq_client = Client.return_value
|
||||
session_creator_job = mock.create_autospec(bigquery.QueryJob)
|
||||
session_creator_job.session_info = None
|
||||
bq_client.query.return_value = session_creator_job
|
||||
|
||||
result = query_tool.execute_sql(
|
||||
"my_project",
|
||||
"SELECT 1",
|
||||
credentials,
|
||||
tool_settings,
|
||||
tool_context,
|
||||
)
|
||||
|
||||
assert result == {
|
||||
"status": "ERROR",
|
||||
"error_details": (
|
||||
"BigQuery did not return session metadata for the protected query."
|
||||
),
|
||||
}
|
||||
bq_client.query_and_wait.assert_not_called()
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("query", "statement_type"),
|
||||
[
|
||||
|
||||
@@ -23,7 +23,7 @@ from unittest.mock import patch
|
||||
from google.adk.integrations.vmaas.sandbox_client import SandboxClient
|
||||
|
||||
|
||||
def _make_response(data: dict) -> MagicMock:
|
||||
def _make_response(data: object) -> MagicMock:
|
||||
"""Create a mock HttpResponse with a JSON body."""
|
||||
response = MagicMock()
|
||||
response.body = json.dumps(data)
|
||||
@@ -56,6 +56,11 @@ class TestSandboxClient(unittest.IsolatedAsyncioTestCase):
|
||||
self.client.update_access_token(new_token)
|
||||
self.assertEqual(self.client._access_token, new_token)
|
||||
|
||||
def test_parse_response_rejects_non_object_json(self):
|
||||
"""Test that malformed sandbox response shapes fail explicitly."""
|
||||
with self.assertRaisesRegex(ValueError, "must be a JSON object"):
|
||||
self.client._parse_response(_make_response(["unexpected"]))
|
||||
|
||||
@patch("asyncio.to_thread")
|
||||
async def test_make_cdp_request(self, mock_to_thread):
|
||||
"""Test making a single CDP request."""
|
||||
|
||||
@@ -14,6 +14,7 @@
|
||||
|
||||
"""Unit tests for the AgentEngineSandboxComputer class."""
|
||||
|
||||
import asyncio
|
||||
import time
|
||||
import unittest
|
||||
from unittest.mock import AsyncMock
|
||||
|
||||
@@ -1489,8 +1489,10 @@ class TestGoogleAuthAsyncByteStream:
|
||||
@pytest.mark.asyncio
|
||||
async def test_iteration_yields_chunks(self):
|
||||
mock_auth_response = AsyncMock()
|
||||
requested_chunk_sizes: list[int] = []
|
||||
|
||||
async def mock_content():
|
||||
async def mock_content(chunk_size: int):
|
||||
requested_chunk_sizes.append(chunk_size)
|
||||
yield b"chunk1"
|
||||
yield b"chunk2"
|
||||
|
||||
@@ -1502,6 +1504,7 @@ class TestGoogleAuthAsyncByteStream:
|
||||
chunks.append(chunk)
|
||||
|
||||
assert chunks == [b"chunk1", b"chunk2"]
|
||||
assert requested_chunk_sizes == [1024]
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_aclose_closes_response(self):
|
||||
|
||||
@@ -244,8 +244,8 @@ class TestMcpToolsetGetAuthHeaders:
|
||||
assert headers is not None
|
||||
assert headers["X-API-Key"] == "test-api-key-12345"
|
||||
|
||||
def test_get_auth_headers_api_key_non_header_logs_warning(self, caplog):
|
||||
"""Test that non-header API key logs a warning."""
|
||||
def test_get_auth_headers_api_key_non_header_fails_closed(self):
|
||||
"""Non-header API keys must not degrade to unauthenticated requests."""
|
||||
# Note: fastapi's APIKey model uses 'in' not 'in_'
|
||||
auth_scheme = APIKeyScheme(**{
|
||||
"in": APIKeyIn.query, # Query param, not header
|
||||
@@ -263,10 +263,10 @@ class TestMcpToolsetGetAuthHeaders:
|
||||
api_key="test-api-key",
|
||||
)
|
||||
|
||||
headers = toolset._get_auth_headers()
|
||||
|
||||
# Should return None for non-header API key
|
||||
assert headers is None
|
||||
with pytest.raises(
|
||||
ValueError, match="only supports header-based API key authentication"
|
||||
):
|
||||
toolset._get_auth_headers()
|
||||
|
||||
def test_get_auth_headers_reads_from_readonly_context(
|
||||
self, toolset_with_oauth2
|
||||
|
||||
Reference in New Issue
Block a user