diff --git a/tests/skills/test_skill_space_policy.py b/tests/skills/test_skill_space_policy.py new file mode 100644 index 000000000..0f984cec7 --- /dev/null +++ b/tests/skills/test_skill_space_policy.py @@ -0,0 +1,170 @@ +# Copyright (c) 2025 Beijing Volcano Engine Technology Co., Ltd. and/or its affiliates. +# +# 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. + +import asyncio +import json + +import pytest + +from veadk.skills import utils +from veadk.skills import registry as registry_module +from veadk.skills.policy import ( + MAX_SKILL_SPACE_POLICY_BYTES, + SkillSpacePolicyError, + parse_skill_space_policy, +) +from veadk.skills.skill import Skill +from veadk.skills.registry import VeSkillRegistry + + +def _skill(skill_id: str | None, name: str) -> Skill: + return Skill( + id=skill_id, + name=name, + description=f"{name} description", + path=f"skills/{name}.zip", + skill_space_id="ss-test", + ) + + +@pytest.mark.parametrize( + ("raw_value", "message"), + [ + ('{"mode":"other","ids":[]}', "mode"), + ('{"mode":"allow","ids":"skill-1"}', "ids must be a list"), + ('{"mode":"allow","ids":[""]}', "non-empty strings"), + ('{"mode":"allow","ids":[],"v":1}', "exactly 'mode' and 'ids'"), + ('{"ref":"skill-policy-1"}', "exactly 'mode' and 'ids'"), + ], +) +def test_parse_skill_space_policy_rejects_unsupported_shapes(raw_value, message): + with pytest.raises(SkillSpacePolicyError, match=message): + parse_skill_space_policy(raw_value) + + +def test_parse_skill_space_policy_rejects_values_over_create_session_limit(): + raw_value = "x" * (MAX_SKILL_SPACE_POLICY_BYTES + 1) + + with pytest.raises(SkillSpacePolicyError, match="must not exceed 8192 bytes"): + parse_skill_space_policy(raw_value) + + +def test_policy_json_is_compact_and_deduplicated(): + policy = parse_skill_space_policy( + '{"mode": "deny", "ids": ["skill-2", "skill-1", "skill-2"]}' + ) + + assert policy.to_json() == '{"mode":"deny","ids":["skill-1","skill-2"]}' + + +@pytest.mark.parametrize( + ("mode", "ids", "expected"), + [ + ("deny", ["skill-2"], ["skill-1"]), + ("allow", ["skill-2"], ["skill-2"]), + ("deny", [], ["skill-1", "skill-2"]), + ("allow", [], []), + ], +) +def test_load_skills_from_cloud_applies_policy( + monkeypatch: pytest.MonkeyPatch, + mode: str, + ids: list[str], + expected: list[str], +): + monkeypatch.setenv( + "SKILL_SPACE_POLICY", + json.dumps({"mode": mode, "ids": ids}), + ) + monkeypatch.setattr( + utils, + "_load_skills_from_space_id", + lambda _space_id, *, raise_on_error=False: [ + _skill("skill-1", "one"), + _skill("skill-2", "two"), + ], + ) + + skills = utils.load_skills_from_cloud("ss-test") + + assert [skill.id for skill in skills] == expected + + +def test_missing_policy_keeps_all_remote_skills(monkeypatch: pytest.MonkeyPatch): + monkeypatch.delenv("SKILL_SPACE_POLICY", raising=False) + monkeypatch.setattr( + utils, + "_load_skills_from_space_id", + lambda _space_id, *, raise_on_error=False: [ + _skill("skill-1", "one"), + _skill("skill-2", "two"), + ], + ) + + skills = utils.load_skills_from_cloud("ss-test") + + assert [skill.id for skill in skills] == ["skill-1", "skill-2"] + + +def test_policy_excludes_remote_skills_without_stable_ids( + monkeypatch: pytest.MonkeyPatch, +): + monkeypatch.setenv( + "SKILL_SPACE_POLICY", + '{"mode":"deny","ids":[]}', + ) + monkeypatch.setattr( + utils, + "_load_skills_from_space_id", + lambda _space_id, *, raise_on_error=False: [_skill(None, "missing-id")], + ) + + assert utils.load_skills_from_cloud("ss-test") == [] + + +def test_invalid_policy_disables_remote_skills_without_listing_space( + monkeypatch: pytest.MonkeyPatch, +): + monkeypatch.setenv("SKILL_SPACE_POLICY", '{"mode":"allow","ids":[],"v":1}') + monkeypatch.setattr( + utils, + "_load_skills_from_space_id", + lambda *_args, **_kwargs: pytest.fail("invalid policy must fail closed"), + ) + + assert utils.load_skills_from_cloud("ss-test") == [] + + +def test_registry_get_skill_cannot_bypass_deny_policy( + monkeypatch: pytest.MonkeyPatch, +): + monkeypatch.setenv( + "SKILL_SPACE_POLICY", + '{"mode":"deny","ids":["skill-1"]}', + ) + monkeypatch.setattr( + utils, + "_load_skills_from_space_id", + lambda _space_id, *, raise_on_error=False: [_skill("skill-1", "one")], + ) + monkeypatch.setattr( + registry_module, + "materialize_remote_skill", + lambda *_args, **_kwargs: pytest.fail("denied Skill must not be downloaded"), + ) + + registry = VeSkillRegistry(skill_source_id="ss-test") + + with pytest.raises(ValueError, match="not found"): + asyncio.run(registry.get_skill(name="one")) diff --git a/tests/tools/builtin_tools/test_agentkit.py b/tests/tools/builtin_tools/test_agentkit.py index 76f26b837..270c5bc10 100644 --- a/tests/tools/builtin_tools/test_agentkit.py +++ b/tests/tools/builtin_tools/test_agentkit.py @@ -374,6 +374,57 @@ def get_session(self, _request): }, ) + def test_create_session_injects_compact_skill_space_policy(self): + captured = {} + + class FakeClient: + def list_sessions(self, _request): + return types.SimpleNamespace(session_infos=[]) + + def create_session(self, request): + captured["request"] = request + return types.SimpleNamespace(session_id="session-1") + + with patch.dict( + os.environ, + { + "SKILL_SPACE_POLICY": ( + '{"mode": "deny", "ids": ["skill-2", "skill-1", "skill-2"]}' + ) + }, + ): + self.agentkit_module._get_or_create_agentkit_session( + client=FakeClient(), + tool_id="tool-1", + tool_user_session_id="user-session-1", + ttl=900, + ) + + request = captured["request"] + assert len(request.envs) == 1 + assert request.envs[0].key == "SKILL_SPACE_POLICY" + assert request.envs[0].value == '{"mode":"deny","ids":["skill-1","skill-2"]}' + + def test_create_session_rejects_unsupported_skill_space_policy(self): + class FakeClient: + def list_sessions(self, _request): + return types.SimpleNamespace(session_infos=[]) + + def create_session(self, _request): + raise AssertionError("invalid policy must not reach CreateSession") + + with patch.dict( + os.environ, + {"SKILL_SPACE_POLICY": '{"mode":"allow","ids":[],"ref":"x"}'}, + ): + with self.assertRaisesRegex(ValueError, "exactly 'mode' and 'ids'"): + self.agentkit_module._get_or_create_agentkit_session( + client=FakeClient(), + tool_id="tool-1", + tool_user_session_id="user-session-1", + ttl=900, + ) + def test_uses_create_session_endpoint_without_waiting_by_default(self): captured = {"get_calls": 0} diff --git a/tests/tools/builtin_tools/test_run_sandbox_agent.py b/tests/tools/builtin_tools/test_run_sandbox_agent.py index 9c9d9f59c..87a026cf5 100644 --- a/tests/tools/builtin_tools/test_run_sandbox_agent.py +++ b/tests/tools/builtin_tools/test_run_sandbox_agent.py @@ -15,6 +15,7 @@ import importlib.util import hashlib import json +import os import sys import types import unittest @@ -223,6 +224,74 @@ def test_runner_code_overrides_the_sandbox_process_environment(self): self.assertNotIn("if key not in env", code) self.assertIn('srv_pythonpath = env.get("SRV_PYTHONPATH")', code) + def test_run_sandbox_agent_forwards_compact_skill_space_policy(self): + invocation_context = types.SimpleNamespace( + session=types.SimpleNamespace(id="session-1"), + agent=types.SimpleNamespace(name="agent"), + user_id="user", + ) + tool_context = types.SimpleNamespace( + _invocation_context=invocation_context, + state={}, + ) + response = {"Result": {"Result": "done"}} + + with ( + patch.dict( + os.environ, + { + "SKILL_SPACE_ID": "ss-test", + "SKILL_SPACE_POLICY": ( + '{"mode": "deny", "ids": ["skill-2", "skill-1", "skill-2"]}' + ), + }, + ), + patch.object( + self.module, + "invoke_agentkit_run_code", + return_value=response, + ) as invoke, + ): + self.module.run_sandbox_agent( + "do work", + "tool-1", + tool_context=tool_context, + ) + + runner_code = invoke.call_args.kwargs["code"] + self.assertIn("SKILL_SPACE_POLICY", runner_code) + self.assertIn('{"mode":"deny","ids":["skill-1","skill-2"]}', runner_code) + + def test_run_sandbox_agent_rejects_unsupported_skill_space_policy(self): + invocation_context = types.SimpleNamespace( + session=types.SimpleNamespace(id="session-1"), + agent=types.SimpleNamespace(name="agent"), + user_id="user", + ) + tool_context = types.SimpleNamespace( + _invocation_context=invocation_context, + state={}, + ) + + with ( + patch.dict( + os.environ, + {"SKILL_SPACE_POLICY": '{"mode":"allow","ids":[],"ref":"x"}'}, + ), + patch.object( + self.module, + "invoke_agentkit_run_code", + ) as invoke, + ): + with self.assertRaisesRegex(ValueError, "exactly 'mode' and 'ids'"): + self.module.run_sandbox_agent( + "do work", + "tool-1", + tool_context=tool_context, + ) + + invoke.assert_not_called() + class TestExecuteSkillsSkillApi(unittest.TestCase): def _tool_context(self, *, inbound_credential=None, credentials_by_key=None): diff --git a/veadk/skills/policy.py b/veadk/skills/policy.py new file mode 100644 index 000000000..301810029 --- /dev/null +++ b/veadk/skills/policy.py @@ -0,0 +1,95 @@ +# Copyright (c) 2025 Beijing Volcano Engine Technology Co., Ltd. and/or its affiliates. +# +# 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. + +"""Session-scoped filtering for remote Skill Space skills.""" + +from __future__ import annotations + +import json +from dataclasses import dataclass +from typing import Literal + + +SKILL_SPACE_POLICY_ENV = "SKILL_SPACE_POLICY" +MAX_SKILL_SPACE_POLICY_BYTES = 8192 + + +class SkillSpacePolicyError(ValueError): + """Raised when ``SKILL_SPACE_POLICY`` is malformed.""" + + +@dataclass(frozen=True) +class SkillSpacePolicy: + mode: Literal["allow", "deny"] + ids: frozenset[str] + + def allows(self, skill_id: str | None) -> bool: + if not skill_id: + return False + selected = skill_id in self.ids + return selected if self.mode == "allow" else not selected + + def to_json(self) -> str: + return json.dumps( + {"mode": self.mode, "ids": sorted(self.ids)}, + ensure_ascii=False, + separators=(",", ":"), + ) + + +def parse_skill_space_policy(raw_value: str) -> SkillSpacePolicy: + """Parse the only supported inline policy shape: ``mode`` plus ``ids``.""" + if len(raw_value.encode("utf-8")) > MAX_SKILL_SPACE_POLICY_BYTES: + raise SkillSpacePolicyError( + f"{SKILL_SPACE_POLICY_ENV} must not exceed " + f"{MAX_SKILL_SPACE_POLICY_BYTES} bytes" + ) + + try: + payload = json.loads(raw_value) + except json.JSONDecodeError as exc: + raise SkillSpacePolicyError( + f"{SKILL_SPACE_POLICY_ENV} must be valid JSON" + ) from exc + + if not isinstance(payload, dict): + raise SkillSpacePolicyError(f"{SKILL_SPACE_POLICY_ENV} must be a JSON object") + if set(payload) != {"mode", "ids"}: + raise SkillSpacePolicyError( + f"{SKILL_SPACE_POLICY_ENV} supports exactly 'mode' and 'ids'" + ) + + mode = payload["mode"] + if mode not in {"allow", "deny"}: + raise SkillSpacePolicyError( + f"{SKILL_SPACE_POLICY_ENV}.mode must be 'allow' or 'deny'" + ) + + raw_ids = payload["ids"] + if not isinstance(raw_ids, list): + raise SkillSpacePolicyError(f"{SKILL_SPACE_POLICY_ENV}.ids must be a list") + + ids: set[str] = set() + for skill_id in raw_ids: + if not isinstance(skill_id, str) or not skill_id.strip(): + raise SkillSpacePolicyError( + f"{SKILL_SPACE_POLICY_ENV}.ids must contain non-empty strings" + ) + if skill_id != skill_id.strip(): + raise SkillSpacePolicyError( + f"{SKILL_SPACE_POLICY_ENV}.ids must not contain surrounding whitespace" + ) + ids.add(skill_id) + + return SkillSpacePolicy(mode=mode, ids=frozenset(ids)) diff --git a/veadk/skills/utils.py b/veadk/skills/utils.py index 72fdfc58e..aec5131f7 100644 --- a/veadk/skills/utils.py +++ b/veadk/skills/utils.py @@ -22,6 +22,11 @@ from typing import Any, Dict, Optional, Callable from veadk.skills.skill import Skill +from veadk.skills.policy import ( + SKILL_SPACE_POLICY_ENV, + SkillSpacePolicyError, + parse_skill_space_policy, +) from veadk.utils.logger import get_logger from veadk.utils.volcengine_sign import ve_request, volcengine_signed_request @@ -140,6 +145,18 @@ def load_skills_from_cloud( skill_space_ids_list = [x.strip() for x in skill_space_ids.split(",") if x.strip()] logger.info(f"Load skills from cloud skill sources: {skill_space_ids_list}") + raw_policy = os.getenv(SKILL_SPACE_POLICY_ENV) + policy = None + if raw_policy is not None: + try: + policy = parse_skill_space_policy(raw_policy) + except SkillSpacePolicyError as exc: + logger.error( + f"Invalid {SKILL_SPACE_POLICY_ENV}; remote Skill Space skills " + f"are disabled for this session: {exc}" + ) + return [] + skills = [] for skill_space_id in skill_space_ids_list: @@ -160,7 +177,15 @@ def load_skills_from_cloud( ) ) - return skills + if policy is None: + return skills + + filtered_skills = [skill for skill in skills if policy.allows(skill.id)] + logger.info( + f"Applied {SKILL_SPACE_POLICY_ENV} mode={policy.mode}: " + f"kept {len(filtered_skills)} of {len(skills)} remote skills" + ) + return filtered_skills def _get_cloud_credentials() -> tuple[str, str, str]: diff --git a/veadk/tools/builtin_tools/_agentkit.py b/veadk/tools/builtin_tools/_agentkit.py index b779f0dba..814aa7a8b 100644 --- a/veadk/tools/builtin_tools/_agentkit.py +++ b/veadk/tools/builtin_tools/_agentkit.py @@ -282,13 +282,27 @@ def _get_or_create_agentkit_session( ) return chosen - return client.create_session( - tools_types.CreateSessionRequest( - ToolId=tool_id, - UserSessionId=tool_user_session_id, - Ttl=ttl, + request_kwargs: dict[str, Any] = { + "ToolId": tool_id, + "UserSessionId": tool_user_session_id, + "Ttl": ttl, + } + raw_skill_space_policy = os.getenv("SKILL_SPACE_POLICY") + if raw_skill_space_policy is not None: + from veadk.skills.policy import ( + SKILL_SPACE_POLICY_ENV, + parse_skill_space_policy, ) - ) + + policy = parse_skill_space_policy(raw_skill_space_policy) + request_kwargs["Envs"] = [ + tools_types.EnvsItemForCreateSession( + Key=SKILL_SPACE_POLICY_ENV, + Value=policy.to_json(), + ) + ] + + return client.create_session(tools_types.CreateSessionRequest(**request_kwargs)) def ensure_agentkit_session_endpoint( diff --git a/veadk/tools/builtin_tools/run_sandbox_agent.py b/veadk/tools/builtin_tools/run_sandbox_agent.py index bde11cadf..520ac9894 100644 --- a/veadk/tools/builtin_tools/run_sandbox_agent.py +++ b/veadk/tools/builtin_tools/run_sandbox_agent.py @@ -19,6 +19,7 @@ from google.adk.tools import ToolContext +from veadk.skills.policy import SKILL_SPACE_POLICY_ENV, parse_skill_space_policy from veadk.tools.builtin_tools._agentkit import invoke_agentkit_run_code from veadk.utils.logger import get_logger @@ -202,6 +203,11 @@ def run_sandbox_agent( skill_space_id = os.getenv("SKILL_SPACE_ID", "") if skill_space_id: base_env_vars["SKILL_SPACE_ID"] = skill_space_id + raw_skill_space_policy = os.getenv(SKILL_SPACE_POLICY_ENV) + if raw_skill_space_policy is not None: + base_env_vars[SKILL_SPACE_POLICY_ENV] = parse_skill_space_policy( + raw_skill_space_policy + ).to_json() env_vars = _merge_execution_env_vars(base_env_vars, extra_env_vars) logger.debug(