diff --git a/src/google/adk/integrations/agent_registry/__init__.py b/src/google/adk/integrations/agent_registry/__init__.py index 3c3bd9b2f5..0705d14463 100644 --- a/src/google/adk/integrations/agent_registry/__init__.py +++ b/src/google/adk/integrations/agent_registry/__init__.py @@ -13,7 +13,9 @@ # limitations under the License. from .agent_registry import AgentRegistry +from .agent_registry import PublishedSkills __all__ = [ 'AgentRegistry', + 'PublishedSkills', ] diff --git a/src/google/adk/integrations/agent_registry/agent_registry.py b/src/google/adk/integrations/agent_registry/agent_registry.py index 35f2c994e9..daaf0bbed4 100644 --- a/src/google/adk/integrations/agent_registry/agent_registry.py +++ b/src/google/adk/integrations/agent_registry/agent_registry.py @@ -35,6 +35,8 @@ from google.adk.auth.auth_schemes import AuthScheme from google.adk.auth.auth_tool import AuthConfig from google.adk.integrations.agent_identity.gcp_auth_provider_scheme import GcpAuthProviderScheme +from google.adk.skills import _utils +from google.adk.skills.models import Skill from google.adk.telemetry.tracing import GCP_MCP_SERVER_DESTINATION_ID from google.adk.tools.base_tool import BaseTool from google.adk.tools.mcp_tool.mcp_session_manager import SseConnectionParams @@ -68,6 +70,18 @@ AGENT_REGISTRY_BASE_URL = "https://agentregistry.googleapis.com/v1" AGENT_REGISTRY_MTLS_BASE_URL = "https://agentregistry.mtls.googleapis.com/v1" +AGENT_REGISTRY_V1ALPHA_BASE_URL = "https://agentregistry.googleapis.com/v1alpha" +AGENT_REGISTRY_V1ALPHA_MTLS_BASE_URL = ( + "https://agentregistry.mtls.googleapis.com/v1alpha" +) + +_SAFE_REGISTRY_ID_PATTERN = re.compile(r"^[a-z0-9]+(?:[._-][a-z0-9]+)*$") + + +def _is_safe_registry_id(name: str) -> bool: + """True if `name` is safe to use as a single skill-registry path segment.""" + return len(name) <= 256 and bool(_SAFE_REGISTRY_ID_PATTERN.fullmatch(name)) + _TRANSPORT_MAPPING = { "HTTP_JSON": _compat.TP_HTTP_JSON, @@ -317,6 +331,35 @@ def _select_auth_provider( return chosen +_SKILL_RESOURCE_NAME_PATTERN = re.compile( + r"^projects/([^/]+)/locations/([^/]+)/skills/([^/]+)$" +) + + +class PublishedSkills: + """Accessor for interacting with published skills in Agent Registry.""" + + def __init__(self, registry: AgentRegistry): + self._registry = registry + + def get(self, name: str) -> Skill: + """Retrieves and loads a published skill by full resource name. + + Args: + name: Full resource name of the skill, in the format + ``projects/{project}/locations/{location}/skills/{skill_id}``. + + Returns: + A loaded `Skill` ready to pass to `SkillToolset(skills=[...])`. + + Raises: + ValueError: If the skill name does not match the expected resource name + format, or the skill does not contain a default revision. + RuntimeError: If an API request to fetch metadata or media fails. + """ + return self._registry._fetch_published_skill_sync(name) + + class AgentRegistry: """Client for interacting with the Google Cloud Agent Registry service. @@ -329,6 +372,7 @@ class AgentRegistry: # The class default also covers registries pickled by older versions. _deployed_agent_id: str | None = None + _skills_base_url_override: str | None = None def __init__( self, @@ -337,6 +381,7 @@ def __init__( header_provider: ( Callable[[ReadonlyContext], Dict[str, str]] | None ) = None, + project: str | None = None, ): """Initializes the AgentRegistry client. @@ -344,13 +389,21 @@ def __init__( project_id: The Google Cloud project ID. location: The Google Cloud location (region). header_provider: Optional provider for custom headers. + project: Optional alias for project_id. """ - self.project_id = project_id + if project_id and project and project_id != project: + raise ValueError( + "Cannot specify both 'project_id' and 'project' with different" + f" values: {project_id!r} vs {project!r}." + ) + self.project_id = project_id or project self.location = location if not self.project_id or not self.location: raise ValueError("project_id and location must be provided") + self._published_skills = PublishedSkills(self) + self._base_path = f"projects/{self.project_id}/locations/{self.location}" self._header_provider = header_provider self._connect_lock = threading.Lock() @@ -407,6 +460,33 @@ def __setstate__(self, state: Dict[str, Any]) -> None: self.__dict__.update(state) self._connect_lock = threading.Lock() + @property + def project(self) -> str | None: + return self.project_id + + @property + def published_skills(self) -> PublishedSkills: + return self._published_skills + + @property + def _skills_base_url(self) -> str: + """Returns the base URL for skill endpoints, defaulting to v1alpha.""" + if self._skills_base_url_override: + return self._skills_base_url_override + if "AGENT_REGISTRY_ENDPOINT" in os.environ: + return os.environ["AGENT_REGISTRY_ENDPOINT"] + if getattr(self, "_use_mtls", False): + return AGENT_REGISTRY_V1ALPHA_MTLS_BASE_URL + return _mtls_utils.get_api_endpoint( + location="", + default_template=AGENT_REGISTRY_V1ALPHA_BASE_URL, + mtls_template=AGENT_REGISTRY_V1ALPHA_MTLS_BASE_URL, + ) + + @_skills_base_url.setter + def _skills_base_url(self, value: str | None) -> None: + self._skills_base_url_override = value + def _get_auth_headers(self) -> Dict[str, str]: """Refreshes credentials and returns authorization headers.""" self._ensure_connected() @@ -430,13 +510,15 @@ def _make_request( params: Dict[str, Any] | None = None, json_data: Dict[str, Any] | None = None, timeout: float | None = None, + base_url: str | None = None, ) -> Dict[str, Any]: """Helper function to make requests to the Agent Registry API.""" self._ensure_connected() + root_url = base_url or self._base_url if path.startswith("projects/"): - url = f"{self._base_url}/{path}" + url = f"{root_url}/{path}" else: - url = f"{self._base_url}/{self._base_path}/{path}" + url = f"{root_url}/{self._base_path}/{path}" quota_project_id = ( getattr(self._credentials, "quota_project_id", None) or self.project_id ) @@ -457,9 +539,12 @@ def _make_request( data: Dict[str, Any] = response.json() return data except requests.exceptions.HTTPError as e: + status_code = ( + e.response.status_code if e.response is not None else "unknown" + ) + error_text = e.response.text if e.response is not None else str(e) raise RuntimeError( - f"API request failed with status {e.response.status_code}:" - f" {e.response.text}" + f"API request failed with status {status_code}: {error_text}" ) from e except requests.exceptions.RequestException as e: raise RuntimeError(f"API request failed (network error): {e}") from e @@ -984,6 +1069,91 @@ def get_remote_a2a_agent( auth_credential=auth_credential, ) + def get_published_skill(self, name: str) -> Skill: + """Retrieves and loads a published skill by full resource name. + + Args: + name: Full resource name of the skill, in the format + ``projects/{project}/locations/{location}/skills/{skill_id}``. + + Returns: + A loaded `Skill` ready to pass to `SkillToolset(skills=[...])`. + """ + return self.published_skills.get(name) + + def _download_media( + self, + path_or_url: str, + params: Dict[str, Any] | None = None, + ) -> bytes: + self._ensure_connected() + if path_or_url.startswith("http://") or path_or_url.startswith("https://"): + url = path_or_url + elif path_or_url.startswith("projects/"): + url = f"{self._skills_base_url}/{path_or_url}" + else: + url = f"{self._skills_base_url}/{self._base_path}/{path_or_url}" + + quota_project_id = ( + getattr(self._credentials, "quota_project_id", None) or self.project_id + ) + headers = merge_tracking_headers( + {"x-goog-user-project": quota_project_id} if quota_project_id else {} + ) + try: + response = self._session.get( + url, + headers=headers, + params=params, + allow_redirects=True, + ) + response.raise_for_status() + return bytes(response.content) + except requests.exceptions.HTTPError as e: + status_code = ( + e.response.status_code if e.response is not None else "unknown" + ) + error_text = e.response.text if e.response is not None else str(e) + raise RuntimeError( + f"API request failed with status {status_code}: {error_text}" + ) from e + except requests.exceptions.RequestException as e: + raise RuntimeError(f"API request failed (network error): {e}") from e + except Exception as e: + raise RuntimeError(f"API request failed: {e}") from e + + def _fetch_published_skill_sync(self, name: str) -> Skill: + if not isinstance(name, str): + raise ValueError( + f"Invalid skill resource name {name!r}. Expected format: " + "'projects/{project}/locations/{location}/skills/{skill_id}'." + ) + match = _SKILL_RESOURCE_NAME_PATTERN.match(name) + if not match or not all( + _is_safe_registry_id(seg) for seg in match.groups() + ): + raise ValueError( + f"Invalid skill resource name '{name}'. Expected format: " + "'projects/{project}/locations/{location}/skills/{skill_id}'." + ) + + skill_data = self._make_request(name, base_url=self._skills_base_url) + default_revision = skill_data.get("defaultRevision") or skill_data.get( + "default_revision" + ) + if ( + not default_revision + or not isinstance(default_revision, str) + or not default_revision.startswith("projects/") + ): + raise ValueError(f"Skill '{name}' does not contain default revision.") + + revision_url = f"{self._skills_base_url}/{default_revision}" + zip_bytes = self._download_media(revision_url, params={"alt": "media"}) + skill = _utils._load_skill_from_zip_bytes(zip_bytes) + skill._uri = revision_url + return skill + def _use_client_cert_effective() -> bool: """Returns whether client certificate should be used for mTLS.""" diff --git a/tests/unittests/integrations/agent_registry/test_agent_registry.py b/tests/unittests/integrations/agent_registry/test_agent_registry.py index f1a49f6b59..d5b12c7d8c 100644 --- a/tests/unittests/integrations/agent_registry/test_agent_registry.py +++ b/tests/unittests/integrations/agent_registry/test_agent_registry.py @@ -13,11 +13,13 @@ # limitations under the License. +import io import os import pickle from unittest.mock import AsyncMock from unittest.mock import MagicMock from unittest.mock import patch +import zipfile from fastapi.openapi.models import OAuth2 from google.adk.a2a import _compat @@ -27,11 +29,16 @@ from google.adk.auth.auth_credential import OAuth2Auth from google.adk.integrations.agent_identity.gcp_auth_provider_scheme import GcpAuthProviderScheme from google.adk.integrations.agent_registry import AgentRegistry +from google.adk.integrations.agent_registry import PublishedSkills from google.adk.integrations.agent_registry.agent_registry import _ProtocolType from google.adk.integrations.agent_registry.agent_registry import _should_use_mtls_endpoint +from google.adk.integrations.agent_registry.agent_registry import AGENT_REGISTRY_V1ALPHA_BASE_URL +from google.adk.integrations.agent_registry.agent_registry import AGENT_REGISTRY_V1ALPHA_MTLS_BASE_URL from google.adk.integrations.agent_registry.agent_registry import AgentRegistrySingleMcpToolset +from google.adk.skills.models import Skill from google.adk.telemetry.tracing import GCP_MCP_SERVER_DESTINATION_ID from google.adk.tools.mcp_tool.mcp_toolset import McpToolset +from google.adk.tools.skill_toolset import SkillToolset from google.adk.utils._google_client_headers import merge_tracking_headers import google.oauth2.credentials import httpx @@ -190,6 +197,21 @@ def _enter_agent_engine(monkeypatch): monkeypatch.setenv("GOOGLE_CLOUD_AGENT_ENGINE_LOCATION", "l") +def _create_fake_zip_bytes( + name: str = "my-skill", + description: str = "test", + instructions: str = "# My Skill", +) -> bytes: + """Creates a fake zip file in memory and returns its bytes.""" + zip_buffer = io.BytesIO() + with zipfile.ZipFile(zip_buffer, "w") as z: + z.writestr( + "SKILL.md", + f"---\nname: {name}\ndescription: {description}\n---\n{instructions}\n", + ) + return zip_buffer.getvalue() + + class TestAgentRegistry: @pytest.fixture @@ -330,6 +352,63 @@ def test_init_raises_value_error_if_params_missing(self): ): AgentRegistry(project_id=None, location=None) + def test_init_with_project_alias(self): + mock_creds = MagicMock() + mock_creds.quota_project_id = None + with ( + patch("google.auth.default", return_value=(mock_creds, "project-id")), + patch( + "google.auth.transport.requests.AuthorizedSession", + autospec=True, + ), + ): + registry = AgentRegistry(project="my-project", location="global") + assert registry.project_id == "my-project" + assert registry.project == "my-project" + assert registry.location == "global" + assert isinstance(registry.published_skills, PublishedSkills) + + def test_init_with_project_and_project_id_matching(self): + mock_creds = MagicMock() + mock_creds.quota_project_id = None + with ( + patch("google.auth.default", return_value=(mock_creds, "project-id")), + patch( + "google.auth.transport.requests.AuthorizedSession", + autospec=True, + ), + ): + registry = AgentRegistry( + project_id="primary-project", + project="primary-project", + location="global", + ) + assert registry.project_id == "primary-project" + assert registry.project == "primary-project" + + def test_init_with_project_and_project_id_differing_raises(self): + mock_creds = MagicMock() + mock_creds.quota_project_id = None + with ( + patch("google.auth.default", return_value=(mock_creds, "project-id")), + patch( + "google.auth.transport.requests.AuthorizedSession", + autospec=True, + ), + ): + with pytest.raises( + ValueError, + match=( + "Cannot specify both 'project_id' and 'project' with different" + " values" + ), + ): + AgentRegistry( + project_id="primary-project", + project="alias-project", + location="global", + ) + def test_get_connection_uri_mcp_interfaces_top_level(self, registry): resource_details = { "interfaces": [ @@ -1393,6 +1472,357 @@ def side_effect(path, *args, **kwargs): assert agent._auth_config is None + def test_published_skills_get_success(self, registry): + skill_resource_name = ( + "projects/test-project/locations/global/skills/my-skill" + ) + revision_resource_name = ( + "projects/test-project/locations/global/skills/my-skill/revisions/rev-1" + ) + fake_zip = _create_fake_zip_bytes( + name="my-skill", + description="A test skill", + instructions="# Instructions", + ) + + mock_metadata_response = MagicMock() + mock_metadata_response.json.return_value = { + "name": skill_resource_name, + "defaultRevision": revision_resource_name, + "description": "A test skill", + } + mock_metadata_response.raise_for_status = MagicMock() + + mock_media_response = MagicMock() + mock_media_response.status_code = 200 + mock_media_response.headers = {} + mock_media_response.content = fake_zip + mock_media_response.raise_for_status = MagicMock() + + def mock_get(url, *args, **kwargs): + if kwargs.get("params") and kwargs.get("params").get("alt") == "media": + return mock_media_response + return mock_metadata_response + + registry._session.get.side_effect = mock_get + + skill = registry.published_skills.get(name=skill_resource_name) + + assert isinstance(skill, Skill) + assert skill.name == "my-skill" + assert skill.description == "A test skill" + assert skill.instructions == "# Instructions" + assert ( + skill._uri + == f"{AGENT_REGISTRY_V1ALPHA_BASE_URL}/{revision_resource_name}" + ) + + assert registry._session.get.call_count == 2 + metadata_call = registry._session.get.call_args_list[0] + assert ( + metadata_call.args[0] + == f"{AGENT_REGISTRY_V1ALPHA_BASE_URL}/{skill_resource_name}" + ) + media_call = registry._session.get.call_args_list[1] + assert ( + media_call.args[0] + == f"{AGENT_REGISTRY_V1ALPHA_BASE_URL}/{revision_resource_name}" + ) + assert media_call.kwargs.get("params") == {"alt": "media"} + assert media_call.kwargs.get("allow_redirects") is True + + def test_published_skills_get_positional_arg(self, registry): + skill_resource_name = ( + "projects/test-project/locations/global/skills/my-skill" + ) + revision_resource_name = ( + "projects/test-project/locations/global/skills/my-skill/revisions/rev-1" + ) + fake_zip = _create_fake_zip_bytes() + + mock_metadata_response = MagicMock() + mock_metadata_response.json.return_value = { + "name": skill_resource_name, + "defaultRevision": revision_resource_name, + } + mock_metadata_response.raise_for_status = MagicMock() + + mock_media_response = MagicMock() + mock_media_response.status_code = 200 + mock_media_response.headers = {} + mock_media_response.content = fake_zip + mock_media_response.raise_for_status = MagicMock() + + def mock_get(url, *args, **kwargs): + if kwargs.get("params") and kwargs.get("params").get("alt") == "media": + return mock_media_response + return mock_metadata_response + + registry._session.get.side_effect = mock_get + + skill = registry.published_skills.get(skill_resource_name) + assert isinstance(skill, Skill) + assert skill.name == "my-skill" + assert ( + skill._uri + == f"{AGENT_REGISTRY_V1ALPHA_BASE_URL}/{revision_resource_name}" + ) + + def test_get_published_skill_convenience_method(self, registry): + skill_resource_name = ( + "projects/test-project/locations/global/skills/my-skill" + ) + revision_resource_name = ( + "projects/test-project/locations/global/skills/my-skill/revisions/rev-1" + ) + fake_zip = _create_fake_zip_bytes() + + mock_metadata_response = MagicMock() + mock_metadata_response.json.return_value = { + "name": skill_resource_name, + "defaultRevision": revision_resource_name, + } + mock_metadata_response.raise_for_status = MagicMock() + + mock_media_response = MagicMock() + mock_media_response.status_code = 200 + mock_media_response.headers = {} + mock_media_response.content = fake_zip + mock_media_response.raise_for_status = MagicMock() + + def mock_get(url, *args, **kwargs): + if kwargs.get("params") and kwargs.get("params").get("alt") == "media": + return mock_media_response + return mock_metadata_response + + registry._session.get.side_effect = mock_get + + skill = registry.get_published_skill(skill_resource_name) + assert isinstance(skill, Skill) + assert skill.name == "my-skill" + assert ( + skill._uri + == f"{AGENT_REGISTRY_V1ALPHA_BASE_URL}/{revision_resource_name}" + ) + + def test_published_skills_get_with_endpoint_override(self, registry): + skill_resource_name = ( + "projects/test-project/locations/global/skills/my-skill" + ) + revision_resource_name = ( + "projects/test-project/locations/global/skills/my-skill/revisions/rev-1" + ) + fake_zip = _create_fake_zip_bytes() + + mock_metadata_response = MagicMock() + mock_metadata_response.json.return_value = { + "name": skill_resource_name, + "defaultRevision": revision_resource_name, + } + mock_metadata_response.raise_for_status = MagicMock() + + mock_media_response = MagicMock() + mock_media_response.status_code = 200 + mock_media_response.headers = {} + mock_media_response.content = fake_zip + mock_media_response.raise_for_status = MagicMock() + + def mock_get(url, *args, **kwargs): + if kwargs.get("params") and kwargs.get("params").get("alt") == "media": + return mock_media_response + return mock_metadata_response + + registry._session.get.side_effect = mock_get + + with patch.dict( + os.environ, {"AGENT_REGISTRY_ENDPOINT": "https://custom.endpoint.com"} + ): + skill = registry.published_skills.get(name=skill_resource_name) + + assert skill._uri == f"https://custom.endpoint.com/{revision_resource_name}" + assert registry._session.get.call_count == 2 + metadata_call = registry._session.get.call_args_list[0] + assert ( + metadata_call.args[0] + == f"https://custom.endpoint.com/{skill_resource_name}" + ) + media_call = registry._session.get.call_args_list[1] + assert ( + media_call.args[0] + == f"https://custom.endpoint.com/{revision_resource_name}" + ) + + def test_published_skills_get_dotted_id_success(self, registry): + skill_resource_name = ( + "projects/test-project/locations/global/skills/" + "cloud.google.com-google-cloud-networking-observability" + ) + revision_resource_name = f"{skill_resource_name}/revisions/rev-1" + fake_zip = _create_fake_zip_bytes(name="networking-observability") + + mock_metadata_response = MagicMock() + mock_metadata_response.json.return_value = { + "name": skill_resource_name, + "defaultRevision": revision_resource_name, + } + mock_metadata_response.raise_for_status = MagicMock() + + mock_media_response = MagicMock() + mock_media_response.status_code = 200 + mock_media_response.headers = {} + mock_media_response.content = fake_zip + mock_media_response.raise_for_status = MagicMock() + + def mock_get(url, *args, **kwargs): + if kwargs.get("params") and kwargs.get("params").get("alt") == "media": + return mock_media_response + return mock_metadata_response + + registry._session.get.side_effect = mock_get + + skill = registry.published_skills.get(name=skill_resource_name) + assert skill.name == "networking-observability" + + def test_published_skills_passed_to_skill_toolset(self, registry): + skill_resource_name = ( + "projects/test-project/locations/global/skills/my-skill" + ) + revision_resource_name = ( + "projects/test-project/locations/global/skills/my-skill/revisions/rev-1" + ) + fake_zip = _create_fake_zip_bytes() + + mock_metadata_response = MagicMock() + mock_metadata_response.json.return_value = { + "name": skill_resource_name, + "defaultRevision": revision_resource_name, + } + mock_metadata_response.raise_for_status = MagicMock() + + mock_media_response = MagicMock() + mock_media_response.status_code = 200 + mock_media_response.headers = {} + mock_media_response.content = fake_zip + mock_media_response.raise_for_status = MagicMock() + + def mock_get(url, *args, **kwargs): + if kwargs.get("params") and kwargs.get("params").get("alt") == "media": + return mock_media_response + return mock_metadata_response + + registry._session.get.side_effect = mock_get + + skill = registry.published_skills.get(skill_resource_name) + toolset = SkillToolset(skills=[skill]) + assert toolset is not None + assert "my-skill" in toolset._skills + + @pytest.mark.parametrize( + "invalid_name", + [ + "my-skill", + "skills/my-skill", + "projects/test-project/skills/my-skill", + "projects/test-project/locations/global/agents/my-agent", + "projects/test-project/locations/global/skills/", + "projects/../locations/global/skills/my-skill", + "projects/test-project/locations/../skills/my-skill", + "projects/test-project/locations/global/skills/..", + "projects/test-project/locations/global/skills/.", + "projects/test-project/locations/global/skills/x?alt=media", + "projects/test-project/locations/global/skills/x%2f", + "projects/test-project/locations/global/skills/MySkill", + f"projects/test-project/locations/global/skills/{'a' * 257}", + f"projects/{'a' * 257}/locations/global/skills/my-skill", + f"projects/test-project/locations/{'a' * 257}/skills/my-skill", + "", + 12345, + None, + ], + ) + def test_published_skills_get_invalid_name_raises( + self, registry, invalid_name + ): + with pytest.raises(ValueError, match="Invalid skill resource name"): + registry.published_skills.get(invalid_name) + registry._session.get.assert_not_called() + + @pytest.mark.parametrize( + "bad_revision", + [ + None, + "", + "https://storage.googleapis.com/evil/archive.zip", + "http://evil.com/archive.zip", + "revisions/rev-1", + ], + ) + def test_published_skills_get_missing_or_invalid_default_revision_raises( + self, registry, bad_revision + ): + skill_resource_name = ( + "projects/test-project/locations/global/skills/my-skill" + ) + mock_metadata_response = MagicMock() + mock_metadata_response.json.return_value = { + "name": skill_resource_name, + "defaultRevision": bad_revision, + } + mock_metadata_response.raise_for_status = MagicMock() + registry._session.get.return_value = mock_metadata_response + + with pytest.raises(ValueError, match="does not contain default revision"): + registry.published_skills.get(name=skill_resource_name) + + def test_published_skills_get_metadata_http_error_raises(self, registry): + skill_resource_name = ( + "projects/test-project/locations/global/skills/my-skill" + ) + mock_response = MagicMock() + mock_response.status_code = 404 + mock_response.text = "Not Found" + error = requests.exceptions.HTTPError( + "404 Client Error", request=MagicMock(), response=mock_response + ) + registry._session.get.side_effect = error + + with pytest.raises( + RuntimeError, match="API request failed with status 404" + ): + registry.published_skills.get(name=skill_resource_name) + + def test_published_skills_get_media_http_error_raises(self, registry): + skill_resource_name = ( + "projects/test-project/locations/global/skills/my-skill" + ) + mock_metadata_response = MagicMock() + mock_metadata_response.json.return_value = { + "name": skill_resource_name, + "defaultRevision": ( + "projects/test-project/locations/global/skills/my-skill/revisions/r1" + ), + } + mock_metadata_response.raise_for_status = MagicMock() + + mock_error_response = MagicMock() + mock_error_response.status_code = 500 + mock_error_response.text = "Internal Server Error" + error = requests.exceptions.HTTPError( + "500 Server Error", request=MagicMock(), response=mock_error_response + ) + + def mock_get(url, *args, **kwargs): + if kwargs.get("params") and kwargs.get("params").get("alt") == "media": + raise error + return mock_metadata_response + + registry._session.get.side_effect = mock_get + + with pytest.raises( + RuntimeError, match="API request failed with status 500" + ): + registry.published_skills.get(name=skill_resource_name) + class TestAgentRegistryMtls: @@ -1464,6 +1894,69 @@ def test_make_request_configures_mtls(self, mock_cert_source, registry): args, kwargs = mock_session.get.call_args assert "agentregistry.mtls.googleapis.com" in args[0] + @patch( + "google.auth.transport.mtls.has_default_client_cert_source", + return_value=True, + ) + @patch("google.auth.transport.mtls.default_client_cert_source") + @patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "true"}) + def test_published_skills_get_configures_mtls( + self, mock_cert_source, registry + ): + mock_cert_source.return_value = lambda: (b"cert", b"key") + with ( + patch( + "google.auth.default", return_value=(MagicMock(), "test-project") + ), + patch( + "google.adk.integrations.agent_registry.agent_registry._use_client_cert_effective", + return_value=True, + ), + patch("google.auth.transport.requests.AuthorizedSession"), + ): + registry = AgentRegistry(project_id="test-project", location="global") + skill_resource_name = ( + "projects/test-project/locations/global/skills/my-skill" + ) + revision_resource_name = "projects/test-project/locations/global/skills/my-skill/revisions/rev-1" + fake_zip = _create_fake_zip_bytes() + + mock_metadata_response = MagicMock() + mock_metadata_response.json.return_value = { + "name": skill_resource_name, + "defaultRevision": revision_resource_name, + } + mock_metadata_response.raise_for_status = MagicMock() + + mock_media_response = MagicMock() + mock_media_response.status_code = 200 + mock_media_response.headers = {} + mock_media_response.content = fake_zip + mock_media_response.raise_for_status = MagicMock() + + def mock_get(url, *args, **kwargs): + if kwargs.get("params") and kwargs.get("params").get("alt") == "media": + return mock_media_response + return mock_metadata_response + + registry._session.get.side_effect = mock_get + + skill = registry.published_skills.get(name=skill_resource_name) + assert ( + skill._uri + == f"{AGENT_REGISTRY_V1ALPHA_MTLS_BASE_URL}/{revision_resource_name}" + ) + metadata_call = registry._session.get.call_args_list[0] + assert ( + metadata_call.args[0] + == f"{AGENT_REGISTRY_V1ALPHA_MTLS_BASE_URL}/{skill_resource_name}" + ) + media_call = registry._session.get.call_args_list[1] + assert ( + media_call.args[0] + == f"{AGENT_REGISTRY_V1ALPHA_MTLS_BASE_URL}/{revision_resource_name}" + ) + @pytest.mark.parametrize( "env_val, has_cert, expected", [