diff --git a/contributing/samples/integrations/gcp_skill_registry_agent/agent.py b/contributing/samples/integrations/gcp_skill_registry_agent/agent.py index 558b6ecfffb..19d916626fc 100644 --- a/contributing/samples/integrations/gcp_skill_registry_agent/agent.py +++ b/contributing/samples/integrations/gcp_skill_registry_agent/agent.py @@ -24,7 +24,11 @@ ) # Initialize SkillToolset with registry -skill_toolset = SkillToolset(skills=[], registry=registry) +skill_toolset = SkillToolset( + skills=[], + registry=registry, + registry_skills=["your-pinned-skill-id"], +) root_agent = Agent( model="gemini-2.5-flash", diff --git a/src/google/adk/tools/skill_toolset.py b/src/google/adk/tools/skill_toolset.py index fb9a3149909..d5d979e4f7d 100644 --- a/src/google/adk/tools/skill_toolset.py +++ b/src/google/adk/tools/skill_toolset.py @@ -108,8 +108,8 @@ class SkillDiscoveryMode(Enum): The `list_skills` tool is not offered, and the model can call `load_skill` straight away. Preferable for a small, stable catalog, where the discovery - turn costs more than the names do. Registry skills are unaffected: they are - still reachable only through `search_skills`. + turn costs more than the names do. Unpinned registry skills are unaffected: + they are still reachable only through `search_skills`. """ @@ -325,7 +325,10 @@ async def run_async( results = await self._toolset._registry.search_skills(query=query) formatted_results = [] for r in results: - if r.name in self._toolset._skills: + if ( + r.name in self._toolset._skills + or r.name in self._toolset._registry_skill_aliases + ): logger.warning( "Skill naming conflict: skill '%s' already exists locally." " Registry skill is filtered.", @@ -1362,6 +1365,7 @@ def __init__( skills: list[models.Skill] | None = None, *, registry: SkillRegistry | None = None, + registry_skills: list[str] | None = None, code_executor: BaseCodeExecutor | None = None, environment: BaseEnvironment | None = None, skills_folder: Path | str | None = None, @@ -1376,6 +1380,12 @@ def __init__( Args: skills: List of skills to register. registry: Optional skill registry for dynamic loading. + registry_skills: Optional list of skill names in the registry to pin. + Pinned skills are fetched once per process on first use, then appear in + `list_skills` / the EAGER catalog and are served locally without further + registry calls; unpinned registry skills are still reachable via + `search_skills`. The skill is stored under the name in the archive's + frontmatter, which may differ from the registry resource name. code_executor: Optional code executor for script execution. environment: Optional environment for executing scripts. skills_folder: Optional absolute path where skills are stored in the @@ -1405,6 +1415,15 @@ def __init__( self._skills = {skill.name: skill for skill in skills} self._registry = registry + if registry_skills and registry is None: + raise ValueError("Cannot specify registry_skills without a registry") + self._registry_skills: list[str] | None = ( + list(dict.fromkeys(registry_skills)) if registry_skills else None + ) + self._registry_skills_loaded = False + self._fetched_registry_skills: set[str] = set() + self._registry_skill_aliases: dict[str, str] = {} + self._registry_skills_lock: asyncio.Lock | None = None self._code_executor = code_executor self._env = environment if code_executor and environment: @@ -1467,6 +1486,68 @@ def skills_folder(self) -> Path | None: return self._env.working_dir / "skills" return None + @property + def registry_skills(self) -> list[str]: + """The list of pinned registry skills.""" + return list(self._registry_skills) if self._registry_skills else [] + + def _get_registry_skills_lock(self) -> asyncio.Lock: + if self._registry_skills_lock is None: + self._registry_skills_lock = asyncio.Lock() + return self._registry_skills_lock + + async def prefetch(self) -> None: + """Fetches pinned registry skills into the local skill catalog.""" + if ( + self._registry_skills_loaded + or not self._registry_skills + or not self._registry + ): + return + + async with self._get_registry_skills_lock(): + if self._registry_skills_loaded: + return + missing = [ + n + for n in self._registry_skills + if n not in self._fetched_registry_skills + ] + if not missing: + self._registry_skills_loaded = True + return + + results = await asyncio.gather( + *(self._registry.get_skill(name=n) for n in missing), + return_exceptions=True, + ) + for name, result in zip(missing, results): + if isinstance(result, asyncio.CancelledError): + raise result + if isinstance(result, BaseException): + logger.warning( + "Failed to fetch registry skill '%s': %s", + name, + result, + exc_info=result, + ) + continue + + skill = result + if skill.name in self._skills: + logger.warning( + "Skill naming conflict: skill '%s' already exists locally." + " Registry skill is filtered.", + skill.name, + ) + else: + self._skills[skill.name] = skill + self._registry_skill_aliases[name] = skill.name + self._fetched_registry_skills.add(name) + + if len(self._fetched_registry_skills) == len(self._registry_skills): + self._registry_skills_loaded = True + def _has_script_execution(self, context: ReadonlyContext | None) -> bool: """Whether scripts can be run; an unknown agent counts as yes.""" if self._env is not None or self._code_executor is not None: @@ -1482,6 +1563,7 @@ async def get_tools( self, readonly_context: ReadonlyContext | None = None ) -> list[BaseTool]: """Returns the list of tools in this toolset.""" + await self.prefetch() dynamic_tools = await self._resolve_additional_tools_from_state( readonly_context ) @@ -1569,10 +1651,15 @@ async def _get_or_fetch_skill( self, skill_name: str, invocation_id: str | None = None ) -> models.Skill | None: """Retrieves a skill by name, falling back to the registry if configured.""" + await self.prefetch() skill = self._get_skill(skill_name) if skill: return skill + if aliased_name := self._registry_skill_aliases.get(skill_name): + if skill := self._get_skill(aliased_name): + return skill + if not self._registry: return None @@ -1700,6 +1787,7 @@ def clone_with_updated_skills( return SkillToolset( skills=skills, registry=self._registry, + registry_skills=self._registry_skills, code_executor=self._code_executor, environment=self._env, skills_folder=self._skills_folder, @@ -1736,6 +1824,7 @@ async def process_llm_request( self, *, tool_context: ToolContext, llm_request: LlmRequest ) -> None: """Processes the outgoing LLM request to include available skills.""" + await self.prefetch() if self._env is not None and not self._env.is_initialized: await self._env.initialize() selected_core_tools = { diff --git a/tests/unittests/tools/test_skill_toolset.py b/tests/unittests/tools/test_skill_toolset.py index 207ce5c02d6..7bb5f353bac 100644 --- a/tests/unittests/tools/test_skill_toolset.py +++ b/tests/unittests/tools/test_skill_toolset.py @@ -3669,3 +3669,327 @@ def test_unload_skill_api_leaves_other_skills_active( assert toolset.unload_skill(stateful_context, "skill1") is True assert toolset.list_active_skills(stateful_context) == ["skill2"] + + +def test_skill_toolset_init_registry_skills_without_registry_raises(): + with pytest.raises( + ValueError, + match="Cannot specify registry_skills without a registry", + ): + skill_toolset.SkillToolset(registry_skills=["skill1"]) + + +@pytest.mark.asyncio +async def test_skill_toolset_prefetch_pins_registry_skills( + mock_registry, mock_skill1, mock_skill2 +): + mock_registry.get_skill.side_effect = ( + lambda name: mock_skill1 if name == "skill1" else mock_skill2 + ) + toolset = skill_toolset.SkillToolset( + registry=mock_registry, registry_skills=["skill1", "skill2"] + ) + assert toolset.registry_skills == ["skill1", "skill2"] + assert toolset.skills == [] + + await toolset.prefetch() + + assert mock_registry.get_skill.call_count == 2 + assert toolset.skills == [mock_skill1, mock_skill2] + + # Subsequent prefetch calls should be a no-op + await toolset.prefetch() + assert mock_registry.get_skill.call_count == 2 + + +@pytest.mark.asyncio +async def test_skill_toolset_prefetch_called_by_get_tools( + mock_registry, mock_skill1 +): + mock_registry.get_skill.return_value = mock_skill1 + toolset = skill_toolset.SkillToolset( + registry=mock_registry, registry_skills=["skill1"] + ) + + tools = await toolset.get_tools() + mock_registry.get_skill.assert_called_once_with(name="skill1") + assert mock_skill1 in toolset.skills + assert len(tools) == 5 + + +@pytest.mark.asyncio +async def test_skill_toolset_prefetch_called_by_process_llm_request_in_eager_mode( + mock_registry, mock_skill1, tool_context_instance +): + mock_registry.get_skill.return_value = mock_skill1 + toolset = skill_toolset.SkillToolset( + registry=mock_registry, + registry_skills=["skill1"], + discovery_mode=skill_toolset.SkillDiscoveryMode.EAGER, + ) + llm_req = mock.create_autospec(llm_request_model.LlmRequest, instance=True) + + await toolset.process_llm_request( + tool_context=tool_context_instance, llm_request=llm_req + ) + + mock_registry.get_skill.assert_called_once_with(name="skill1") + llm_req.append_instructions.assert_called_once() + args, _ = llm_req.append_instructions.call_args + instructions = args[0] + assert "" in instructions[1] + assert "skill1" in instructions[1] + + +@pytest.mark.asyncio +async def test_skill_toolset_prefetch_collision_with_local_skill( + mock_registry, mock_skill1 +): + mock_skill_reg = mock.create_autospec(models.Skill, instance=True) + mock_skill_reg.name = "skill1" + mock_registry.get_skill.return_value = mock_skill_reg + + toolset = skill_toolset.SkillToolset( + skills=[mock_skill1], + registry=mock_registry, + registry_skills=["skill1"], + ) + + await toolset.prefetch() + + mock_registry.get_skill.assert_called_once_with(name="skill1") + assert toolset._skills["skill1"] is mock_skill1 + + +@pytest.mark.asyncio +async def test_skill_toolset_prefetch_failure_is_logged_and_retried( + mock_registry, mock_skill1, mock_skill2 +): + call_counts = {"fail_skill": 0, "ok_skill": 0} + + async def get_skill_impl(*, name): + if name == "fail_skill": + call_counts["fail_skill"] += 1 + if call_counts["fail_skill"] == 1: + raise RuntimeError("Transient network error") + return mock_skill1 + call_counts["ok_skill"] += 1 + return mock_skill2 + + mock_registry.get_skill.side_effect = get_skill_impl + + toolset = skill_toolset.SkillToolset( + registry=mock_registry, + registry_skills=["fail_skill", "ok_skill"], + ) + + # First prefetch attempt: fail_skill raises exception, ok_skill succeeds + await toolset.prefetch() + + assert toolset.skills == [mock_skill2] + assert toolset._registry_skills_loaded is False + assert call_counts["fail_skill"] == 1 + assert call_counts["ok_skill"] == 1 + + # Second prefetch attempt: only fail_skill is retried and succeeds + await toolset.prefetch() + + assert set(s.name for s in toolset.skills) == {"skill1", "skill2"} + assert toolset._registry_skills_loaded is True + assert call_counts["fail_skill"] == 2 + assert call_counts["ok_skill"] == 1 + + +@pytest.mark.asyncio +async def test_skill_toolset_search_skills_filters_pinned_registry_skill( + mock_registry, mock_skill1, tool_context_instance +): + mock_frontmatter1 = mock.create_autospec(models.Frontmatter, instance=True) + mock_frontmatter1.name = "skill1" + mock_frontmatter1.model_dump.return_value = {"name": "skill1"} + + mock_frontmatter2 = mock.create_autospec(models.Frontmatter, instance=True) + mock_frontmatter2.name = "skill2" + mock_frontmatter2.model_dump.return_value = {"name": "skill2"} + + mock_registry.get_skill.return_value = mock_skill1 + mock_registry.search_skills.return_value = [ + mock_frontmatter1, + mock_frontmatter2, + ] + + toolset = skill_toolset.SkillToolset( + registry=mock_registry, registry_skills=["skill1"] + ) + tools = await toolset.get_tools() + tool = next(t for t in tools if isinstance(t, skill_toolset.SearchSkillsTool)) + + result = await tool.run_async( + args={"query": "test"}, tool_context=tool_context_instance + ) + + # skill1 should be prefetched and filtered out of search results + assert result == [{"name": "skill2"}] + + +def test_skill_toolset_clone_with_updated_skills_forwards_registry_skills( + mock_registry, mock_skill1 +): + toolset = skill_toolset.SkillToolset( + registry=mock_registry, registry_skills=["skill1"] + ) + cloned = toolset.clone_with_updated_skills([mock_skill1]) + + assert cloned.registry_skills == ["skill1"] + assert cloned._registry == mock_registry + + +@pytest.mark.asyncio +async def test_skill_toolset_prefetch_concurrent_calls( + mock_registry, mock_skill1 +): + mock_registry.get_skill.return_value = mock_skill1 + toolset = skill_toolset.SkillToolset( + registry=mock_registry, registry_skills=["skill1"] + ) + + await asyncio.gather( + toolset.prefetch(), + toolset.prefetch(), + toolset.prefetch(), + ) + + mock_registry.get_skill.assert_called_once_with(name="skill1") + + +@pytest.mark.asyncio +async def test_skill_toolset_list_skills_tool_includes_pinned_registry_skill( + mock_registry, mock_skill1, tool_context_instance +): + """LAZY mode: the list_skills tool output must contain pinned skills.""" + mock_registry.get_skill.return_value = mock_skill1 + toolset = skill_toolset.SkillToolset( + registry=mock_registry, registry_skills=["skill1"] + ) + tools = await toolset.get_tools() + tool = next(t for t in tools if isinstance(t, skill_toolset.ListSkillsTool)) + + result = await tool.run_async(args={}, tool_context=tool_context_instance) + + assert "" in result + assert "skill1" in result + mock_registry.get_skill.assert_called_once_with(name="skill1") + + +@pytest.mark.asyncio +async def test_skill_toolset_load_skill_does_not_refetch_pinned_registry_skill( + mock_registry, mock_skill1, tool_context_instance +): + """After prefetch, load_skill for a pinned name is served locally.""" + mock_registry.get_skill.return_value = mock_skill1 + toolset = skill_toolset.SkillToolset( + registry=mock_registry, registry_skills=["skill1"] + ) + tool_context_instance.state.get.return_value = None + + await toolset.prefetch() + assert mock_registry.get_skill.call_count == 1 + + tool = skill_toolset.LoadSkillTool(toolset) + # Two invocations: the per-invocation cache must not be what serves this. + for invocation_id in ("inv-1", "inv-2"): + tool_context_instance.invocation_id = invocation_id + result = await tool.run_async( + args={"skill_name": "skill1"}, tool_context=tool_context_instance + ) + assert result["skill_name"] == "skill1" + + assert mock_registry.get_skill.call_count == 1 + + +@pytest.mark.asyncio +async def test_skill_toolset_prefetch_stores_skill_under_frontmatter_name( + mock_registry, tool_context_instance +): + """A pinned name is marked fetched even when the archive's frontmatter name + differs; the skill is stored under the frontmatter name, and can be resolved + by either name without refetching.""" + fetched = mock.create_autospec(models.Skill, instance=True) + fetched.name = "bar" + fetched.instructions = "Instructions for bar" + fetched.frontmatter = mock.create_autospec(models.Frontmatter, instance=True) + fetched.frontmatter.metadata = {} + mock_registry.get_skill.return_value = fetched + toolset = skill_toolset.SkillToolset( + registry=mock_registry, registry_skills=["foo"] + ) + + await toolset.prefetch() + + mock_registry.get_skill.assert_called_once_with(name="foo") + assert toolset._get_skill("bar") is fetched + assert toolset._get_skill("foo") is None + assert toolset._registry_skills_loaded is True + + # Resolving via _get_or_fetch_skill with the requested registry id "foo" + # uses the alias map and does not call registry again. + resolved = await toolset._get_or_fetch_skill("foo") + assert resolved is fetched + assert mock_registry.get_skill.call_count == 1 + + # Also via LoadSkillTool using the registry id "foo" + tool_context_instance.state.get.return_value = None + load_tool = skill_toolset.LoadSkillTool(toolset) + res = await load_tool.run_async( + args={"skill_name": "foo"}, tool_context=tool_context_instance + ) + assert res["skill_name"] == "foo" + assert res["instructions"] == "Instructions for bar" + assert mock_registry.get_skill.call_count == 1 + + # No refetch on the next prefetch. + await toolset.prefetch() + assert mock_registry.get_skill.call_count == 1 + + +@pytest.mark.asyncio +async def test_skill_toolset_search_skills_filters_aliased_registry_skill( + mock_registry, tool_context_instance +): + """SearchSkillsTool filters out pinned skills by both registry ID and frontmatter name.""" + fetched = mock.create_autospec(models.Skill, instance=True) + fetched.name = "bar" + mock_registry.get_skill.return_value = fetched + + mock_frontmatter_reg = mock.create_autospec(models.Frontmatter, instance=True) + mock_frontmatter_reg.name = "foo" + mock_frontmatter_reg.model_dump.return_value = {"name": "foo"} + + mock_frontmatter_fm = mock.create_autospec(models.Frontmatter, instance=True) + mock_frontmatter_fm.name = "bar" + mock_frontmatter_fm.model_dump.return_value = {"name": "bar"} + + mock_frontmatter_other = mock.create_autospec( + models.Frontmatter, instance=True + ) + mock_frontmatter_other.name = "other" + mock_frontmatter_other.model_dump.return_value = {"name": "other"} + + mock_registry.search_skills.return_value = [ + mock_frontmatter_reg, + mock_frontmatter_fm, + mock_frontmatter_other, + ] + + toolset = skill_toolset.SkillToolset( + registry=mock_registry, registry_skills=["foo"] + ) + tools = await toolset.get_tools() + tool = next(t for t in tools if isinstance(t, skill_toolset.SearchSkillsTool)) + + result = await tool.run_async( + args={"query": "test"}, tool_context=tool_context_instance + ) + + # Both "foo" (aliased registry id) and "bar" (stored frontmatter name) filtered + assert result == [{"name": "other"}]