diff --git a/src/google/adk/integrations/agent_registry/__init__.py b/src/google/adk/integrations/agent_registry/__init__.py index 3c3bd9b2f57..0705d14463e 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 2a9989f2508..7b87c899724 100644 --- a/src/google/adk/integrations/agent_registry/agent_registry.py +++ b/src/google/adk/integrations/agent_registry/agent_registry.py @@ -33,6 +33,8 @@ from google.adk.auth.auth_credential import AuthCredential from google.adk.auth.auth_schemes import AuthScheme 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 @@ -172,6 +174,38 @@ def _is_google_api(url: str) -> bool: ) +_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) + + +_PublishedSkillsAccessor = PublishedSkills + + class AgentRegistry: """Client for interacting with the Google Cloud Agent Registry service. @@ -189,6 +223,7 @@ def __init__( header_provider: ( Callable[[ReadonlyContext], Dict[str, str]] | None ) = None, + project: str | None = None, ): """Initializes the AgentRegistry client. @@ -196,13 +231,16 @@ 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 + 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 try: @@ -232,6 +270,14 @@ def __init__( else AGENT_REGISTRY_BASE_URL ) + @property + def project(self) -> str | None: + return self.project_id + + @property + def published_skills(self) -> PublishedSkills: + return self._published_skills + def _get_auth_headers(self) -> Dict[str, str]: """Refreshes credentials and returns authorization headers.""" try: @@ -275,9 +321,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 @@ -689,6 +738,98 @@ def get_remote_a2a_agent( ) + 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: + 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._base_url}/{path_or_url}" + else: + url = f"{self._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, + ) + if 300 <= response.status_code < 400 and ( + "location" in response.headers or "Location" in response.headers + ): + redirect_url = response.headers.get("location") or response.headers.get( + "Location" + ) + response = self._session.get(redirect_url, 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) or not _SKILL_RESOURCE_NAME_PATTERN.match( + name + ): + raise ValueError( + f"Invalid skill resource name '{name}'. Expected format: " + "'projects/{project}/locations/{location}/skills/{skill_id}'." + ) + + skill_data = self._make_request(name) + default_revision = skill_data.get("defaultRevision") or skill_data.get( + "default_revision" + ) + if not default_revision: + raise ValueError(f"Skill '{name}' does not contain default revision.") + + if default_revision.startswith("http://") or default_revision.startswith( + "https://" + ): + revision_url = default_revision + elif default_revision.startswith("projects/"): + revision_url = f"{self._base_url}/{default_revision}" + else: + clean_revision = default_revision.lstrip("/") + revision_url = f"{self._base_url}/{name}/{clean_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.""" try: diff --git a/tests/unittests/integrations/agent_registry/test_agent_registry.py b/tests/unittests/integrations/agent_registry/test_agent_registry.py index b178b1b8b79..8e82dc9513f 100644 --- a/tests/unittests/integrations/agent_registry/test_agent_registry.py +++ b/tests/unittests/integrations/agent_registry/test_agent_registry.py @@ -13,10 +13,12 @@ # limitations under the License. +import io import os 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 @@ -26,10 +28,13 @@ 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.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 httpx from mcp import ClientSession @@ -94,6 +99,21 @@ def _agent_with_binding_side_effect(path, *_args, **_kwargs): return {} +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 @@ -234,6 +254,40 @@ 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(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="alias-project", + location="global", + ) + assert registry.project_id == "primary-project" + assert registry.project == "primary-project" + def test_get_connection_uri_mcp_interfaces_top_level(self, registry): resource_details = { "interfaces": [ @@ -954,6 +1008,283 @@ 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"{registry._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"{registry._base_url}/{skill_resource_name}" + ) + media_call = registry._session.get.call_args_list[1] + assert ( + media_call.args[0] == f"{registry._base_url}/{revision_resource_name}" + ) + assert media_call.kwargs.get("params") == {"alt": "media"} + + 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" + + 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" + + def test_published_skills_get_with_redirect(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" + ) + redirect_url = "https://storage.googleapis.com/download/bundle.zip" + 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_redirect_response = MagicMock() + mock_redirect_response.status_code = 302 + mock_redirect_response.headers = {"Location": redirect_url} + + mock_final_media_response = MagicMock() + mock_final_media_response.status_code = 200 + mock_final_media_response.headers = {} + mock_final_media_response.content = fake_zip + mock_final_media_response.raise_for_status = MagicMock() + + def mock_get(url, *args, **kwargs): + if url == redirect_url: + return mock_final_media_response + if kwargs.get("params") and kwargs.get("params").get("alt") == "media": + return mock_redirect_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 registry._session.get.call_count == 3 + assert registry._session.get.call_args_list[2].args[0] == redirect_url + + 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/", + "", + 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) + + def test_published_skills_get_missing_default_revision_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, + } + 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: