diff --git a/tests/tools/builtin_tools/test_agentkit.py b/tests/tools/builtin_tools/test_agentkit.py index 3233974c..c5a5aa69 100644 --- a/tests/tools/builtin_tools/test_agentkit.py +++ b/tests/tools/builtin_tools/test_agentkit.py @@ -203,5 +203,280 @@ def test_builds_exec_bash_invoke_tool_request(self): ) +class TestEnsureAgentkitSessionEndpoint(unittest.TestCase): + @classmethod + def setUpClass(cls): + cls.agentkit_module = _load_agentkit_module() + + def test_creates_session_and_prefers_public_endpoint(self): + captured = {} + + class FakeCreateSessionRequest: + def __init__(self, **kwargs): + captured["create_request"] = kwargs + + class FakeGetSessionRequest: + def __init__(self, **kwargs): + captured["get_request"] = kwargs + + class FakeClient: + def __init__(self, **kwargs): + captured["client"] = kwargs + + def create_session(self, _request): + return types.SimpleNamespace(session_id="session-1") + + def get_session(self, _request): + return types.SimpleNamespace( + endpoint="https://public.example", + internal_endpoint="http://internal.example", + status="Ready", + ) + + fake_tools_types = types.ModuleType("agentkit.sdk.tools.types") + fake_tools_types.CreateSessionRequest = FakeCreateSessionRequest + fake_tools_types.GetSessionRequest = FakeGetSessionRequest + fake_tools_client = types.ModuleType("agentkit.sdk.tools.client") + fake_tools_client.AgentkitToolsClient = FakeClient + fake_tools_package = types.ModuleType("agentkit.sdk.tools") + fake_tools_package.types = fake_tools_types + fake_sdk_package = types.ModuleType("agentkit.sdk") + fake_agentkit_package = types.ModuleType("agentkit") + + with patch.dict( + sys.modules, + { + "agentkit": fake_agentkit_package, + "agentkit.sdk": fake_sdk_package, + "agentkit.sdk.tools": fake_tools_package, + "agentkit.sdk.tools.types": fake_tools_types, + "agentkit.sdk.tools.client": fake_tools_client, + }, + ): + with ( + patch.object( + self.agentkit_module, + "get_agentkit_endpoint_config", + return_value=("agentkit", "cn-beijing", "host", "https"), + ), + patch.object( + self.agentkit_module, + "get_agentkit_credentials", + return_value=("ak", "sk", {"X-Security-Token": "token"}), + ), + ): + endpoint = self.agentkit_module.ensure_agentkit_session_endpoint( + tool_id="tool-1", + tool_user_session_id="user-session-1", + tool_state={"state": "value"}, + ttl=900, + ) + + self.assertEqual(endpoint, "https://public.example") + self.assertEqual( + captured["client"], + { + "access_key": "ak", + "secret_key": "sk", + "region": "cn-beijing", + "session_token": "token", + }, + ) + self.assertEqual( + captured["create_request"], + { + "ToolId": "tool-1", + "UserSessionId": "user-session-1", + "Ttl": 900, + }, + ) + self.assertEqual( + captured["get_request"], + { + "ToolId": "tool-1", + "SessionId": "session-1", + }, + ) + + def test_uses_create_session_endpoint_without_waiting_by_default(self): + captured = {"get_calls": 0} + + class FakeCreateSessionRequest: + def __init__(self, **_kwargs): + pass + + class FakeGetSessionRequest: + def __init__(self, **_kwargs): + pass + + class FakeClient: + def __init__(self, **_kwargs): + pass + + def create_session(self, _request): + return types.SimpleNamespace( + session_id="session-1", + endpoint="https://public.example", + internal_endpoint="http://internal.example", + ) + + def get_session(self, _request): + captured["get_calls"] += 1 + raise AssertionError( + "get_session should not be called when waiting is disabled" + ) + + fake_tools_types = types.ModuleType("agentkit.sdk.tools.types") + fake_tools_types.CreateSessionRequest = FakeCreateSessionRequest + fake_tools_types.GetSessionRequest = FakeGetSessionRequest + fake_tools_client = types.ModuleType("agentkit.sdk.tools.client") + fake_tools_client.AgentkitToolsClient = FakeClient + fake_tools_package = types.ModuleType("agentkit.sdk.tools") + fake_tools_package.types = fake_tools_types + + with patch.dict( + sys.modules, + { + "agentkit": types.ModuleType("agentkit"), + "agentkit.sdk": types.ModuleType("agentkit.sdk"), + "agentkit.sdk.tools": fake_tools_package, + "agentkit.sdk.tools.types": fake_tools_types, + "agentkit.sdk.tools.client": fake_tools_client, + }, + ): + with ( + patch.object( + self.agentkit_module, + "get_agentkit_endpoint_config", + return_value=("agentkit", "cn-beijing", "host", "https"), + ), + patch.object( + self.agentkit_module, + "get_agentkit_credentials", + return_value=("ak", "sk", {}), + ), + ): + endpoint = self.agentkit_module.ensure_agentkit_session_endpoint( + tool_id="tool-1", + tool_user_session_id="user-session-1", + ) + + self.assertEqual(endpoint, "https://public.example") + self.assertEqual(captured["get_calls"], 0) + + def test_polls_until_session_is_ready(self): + statuses = iter(["Starting", "Ready"]) + + class FakeRequest: + def __init__(self, **_kwargs): + pass + + class FakeClient: + def __init__(self, **_kwargs): + pass + + def create_session(self, _request): + return types.SimpleNamespace(session_id="session-1") + + def get_session(self, _request): + return types.SimpleNamespace( + status=next(statuses), + endpoint="https://public.example", + internal_endpoint=None, + ) + + fake_tools_types = types.ModuleType("agentkit.sdk.tools.types") + fake_tools_types.CreateSessionRequest = FakeRequest + fake_tools_types.GetSessionRequest = FakeRequest + fake_tools_client = types.ModuleType("agentkit.sdk.tools.client") + fake_tools_client.AgentkitToolsClient = FakeClient + fake_tools_package = types.ModuleType("agentkit.sdk.tools") + fake_tools_package.types = fake_tools_types + + with patch.dict( + sys.modules, + { + "agentkit": types.ModuleType("agentkit"), + "agentkit.sdk": types.ModuleType("agentkit.sdk"), + "agentkit.sdk.tools": fake_tools_package, + "agentkit.sdk.tools.types": fake_tools_types, + "agentkit.sdk.tools.client": fake_tools_client, + }, + ): + with ( + patch.object( + self.agentkit_module, + "get_agentkit_endpoint_config", + return_value=("agentkit", "cn-beijing", "host", "https"), + ), + patch.object( + self.agentkit_module, + "get_agentkit_credentials", + return_value=("ak", "sk", {}), + ), + patch.object(self.agentkit_module.time, "sleep") as sleep, + ): + endpoint = self.agentkit_module.ensure_agentkit_session_endpoint( + tool_id="tool-1", + tool_user_session_id="user-session-1", + wait_until_ready=True, + ) + + self.assertEqual(endpoint, "https://public.example") + sleep.assert_called_once_with(1.0) + + def test_raises_when_session_enters_failed_status(self): + class FakeRequest: + def __init__(self, **_kwargs): + pass + + class FakeClient: + def __init__(self, **_kwargs): + pass + + def create_session(self, _request): + return types.SimpleNamespace(session_id="session-1") + + def get_session(self, _request): + return types.SimpleNamespace(status="Failed") + + fake_tools_types = types.ModuleType("agentkit.sdk.tools.types") + fake_tools_types.CreateSessionRequest = FakeRequest + fake_tools_types.GetSessionRequest = FakeRequest + fake_tools_client = types.ModuleType("agentkit.sdk.tools.client") + fake_tools_client.AgentkitToolsClient = FakeClient + fake_tools_package = types.ModuleType("agentkit.sdk.tools") + fake_tools_package.types = fake_tools_types + + with patch.dict( + sys.modules, + { + "agentkit": types.ModuleType("agentkit"), + "agentkit.sdk": types.ModuleType("agentkit.sdk"), + "agentkit.sdk.tools": fake_tools_package, + "agentkit.sdk.tools.types": fake_tools_types, + "agentkit.sdk.tools.client": fake_tools_client, + }, + ): + with ( + patch.object( + self.agentkit_module, + "get_agentkit_endpoint_config", + return_value=("agentkit", "cn-beijing", "host", "https"), + ), + patch.object( + self.agentkit_module, + "get_agentkit_credentials", + return_value=("ak", "sk", {}), + ), + ): + with self.assertRaisesRegex(RuntimeError, "terminal status Failed"): + self.agentkit_module.ensure_agentkit_session_endpoint( + tool_id="tool-1", + tool_user_session_id="user-session-1", + wait_until_ready=True, + ) + + if __name__ == "__main__": unittest.main() diff --git a/tests/tools/builtin_tools/test_run_sandbox_agent.py b/tests/tools/builtin_tools/test_run_sandbox_agent.py index bd7f9d8a..59391a02 100644 --- a/tests/tools/builtin_tools/test_run_sandbox_agent.py +++ b/tests/tools/builtin_tools/test_run_sandbox_agent.py @@ -76,7 +76,12 @@ def _load_run_sandbox_agent_module(): return module -def _load_execute_skills_module(run_sandbox_agent): +def _load_execute_skills_module( + *, + ensure_agentkit_session_endpoint=lambda **_kwargs: "", + run_sandbox_agent=lambda **_kwargs: "", + wait_for_skill_api_health=lambda **_kwargs: None, +): module_path = ( Path(__file__).resolve().parents[3] / "veadk" @@ -101,12 +106,17 @@ def _load_execute_skills_module(run_sandbox_agent): fake_agentkit = types.ModuleType("veadk.tools.builtin_tools._agentkit") fake_agentkit.get_agentkit_account_id = lambda _state: "test-account" fake_agentkit.resolve_agentkit_tool_id = lambda _name: "test-tool" + fake_agentkit.ensure_agentkit_session_endpoint = ensure_agentkit_session_endpoint fake_runner = types.ModuleType("veadk.tools.builtin_tools.run_sandbox_agent") fake_runner.run_sandbox_agent = run_sandbox_agent fake_utils = types.ModuleType("veadk.utils") fake_utils.__path__ = [] # type: ignore[attr-defined] fake_logger = types.ModuleType("veadk.utils.logger") - fake_logger.get_logger = lambda _name: object() + fake_logger.get_logger = lambda _name: types.SimpleNamespace( + debug=lambda *_args, **_kwargs: None, + warning=lambda *_args, **_kwargs: None, + error=lambda *_args, **_kwargs: None, + ) stub_modules = { "google": fake_google, @@ -129,6 +139,8 @@ def _load_execute_skills_module(run_sandbox_agent): assert spec.loader is not None module = importlib.util.module_from_spec(spec) spec.loader.exec_module(module) + if wait_for_skill_api_health is not None: + module._wait_for_skill_api_health = wait_for_skill_api_health return module @@ -188,29 +200,351 @@ def test_runner_code_overrides_the_sandbox_process_environment(self): self.assertIn('srv_pythonpath = env.get("SRV_PYTHONPATH")', code) -class TestExecuteSkillsEnvVars(unittest.TestCase): - def test_passes_custom_env_vars_to_each_sandbox_execution(self): +class TestExecuteSkillsSkillApi(unittest.TestCase): + def _tool_context(self): + invocation_context = types.SimpleNamespace( + session=types.SimpleNamespace(id="session-1"), + agent=types.SimpleNamespace(name="agent"), + user_id="user", + ) + return types.SimpleNamespace( + state={"TIP_TOKEN_KEY": "tip-from-state"}, + _invocation_context=invocation_context, + ) + + def test_prefers_new_skill_execute_api_when_endpoint_is_available(self): + captured_requests = [] + health_endpoints = [] + session_kwargs = [] + + class FakeResponse: + def __enter__(self): + return self + + def __exit__(self, *_args): + return None + + def read(self): + return b'{"content": "api result"}' + + def fake_urlopen(request, timeout=None): + captured_requests.append((request, timeout)) + return FakeResponse() + + module = _load_execute_skills_module( + ensure_agentkit_session_endpoint=lambda **kwargs: ( + session_kwargs.append(kwargs) or "https://sandbox.test" + ), + wait_for_skill_api_health=lambda **kwargs: health_endpoints.append( + kwargs["endpoint"] + ), + ) + + with patch.object(module.request, "urlopen", fake_urlopen): + result = module.execute_skills("do work", tool_context=self._tool_context()) + + self.assertEqual(result, "api result") + self.assertTrue(session_kwargs[0]["wait_until_ready"]) + self.assertEqual(["https://sandbox.test"], health_endpoints) + self.assertEqual(1, len(captured_requests)) + request_obj, timeout = captured_requests[0] + self.assertEqual("https://sandbox.test/v1/skills/execute", request_obj.full_url) + self.assertEqual(900, timeout) + self.assertEqual("POST", request_obj.get_method()) + self.assertEqual("tip-from-state", request_obj.headers["X-tip-token-key"]) + self.assertIn(b'"prompt": "do work"', request_obj.data) + + def test_health_check_retries_502_until_upstream_is_ready(self): + attempts = [] + + class HealthyResponse: + def __enter__(self): + return self + + def __exit__(self, *_args): + return None + + class ErrorResponse: + def read(self): + return b"bad gateway" + + def close(self): + return None + + module = _load_execute_skills_module(wait_for_skill_api_health=None) + + def fake_urlopen(req, **_kwargs): + attempts.append((req.full_url, req.get_method())) + if len(attempts) == 1: + raise module.error.HTTPError( + url=req.full_url, + code=502, + msg="Bad Gateway", + hdrs={}, + fp=ErrorResponse(), + ) + return HealthyResponse() + + with ( + patch.object(module.request, "urlopen", fake_urlopen), + patch.object(module.time, "sleep") as sleep, + ): + module._wait_for_skill_api_health(endpoint="https://sandbox.test") + + self.assertEqual( + [ + ("https://sandbox.test/v1/skills/healthz", "GET"), + ("https://sandbox.test/v1/skills/healthz", "GET"), + ], + attempts, + ) + sleep.assert_called_once_with(1.0) + + def test_health_check_allows_images_without_health_endpoint(self): + class NotFoundResponse: + def read(self): + return b"not found" + + def close(self): + return None + + module = _load_execute_skills_module(wait_for_skill_api_health=None) + + def fake_urlopen(req, **_kwargs): + raise module.error.HTTPError( + url=req.full_url, + code=404, + msg="Not Found", + hdrs={}, + fp=NotFoundResponse(), + ) + + with patch.object(module.request, "urlopen", fake_urlopen): + module._wait_for_skill_api_health(endpoint="https://sandbox.test") + + def test_env_vars_use_legacy_runcode_execution(self): captured_kwargs = {} def fake_run_sandbox_agent(**kwargs): captured_kwargs.update(kwargs) - return "done" + return "legacy result" - module = _load_execute_skills_module(fake_run_sandbox_agent) - tool_context = types.SimpleNamespace(state={}) + module = _load_execute_skills_module( + ensure_agentkit_session_endpoint=lambda **_kwargs: self.fail( + "Skill API must not be used when env_vars are provided" + ), + run_sandbox_agent=fake_run_sandbox_agent, + ) result = module.execute_skills( "do work", - tool_context=tool_context, + tool_context=self._tool_context(), env_vars={"CUSTOM_VALUE": "custom", "TOS_SKILLS_DIR": ""}, ) - self.assertEqual(result, "done") + self.assertEqual(result, "legacy result") self.assertEqual( - captured_kwargs["extra_env_vars"], {"CUSTOM_VALUE": "custom", "TOS_SKILLS_DIR": ""}, + captured_kwargs["extra_env_vars"], ) + def test_skill_api_url_preserves_agentkit_endpoint_query_auth(self): + module = _load_execute_skills_module( + ensure_agentkit_session_endpoint=lambda **_kwargs: "", + ) + + self.assertEqual( + "https://sandbox.test/v1/skills/execute?faasInstanceName=inst&Authorization=key", + module._skill_api_url( + "https://sandbox.test/?faasInstanceName=inst&Authorization=key", + "/v1/skills/execute", + ), + ) + + def test_requires_sandbox_upgrade_when_skill_api_returns_404(self): + class NotFoundResponse: + def read(self): + return b"not found" + + def close(self): + return None + + def fake_urlopen(_request, **_kwargs): + raise module.error.HTTPError( + url="https://sandbox.test/v1/skills/execute", + code=404, + msg="Not Found", + hdrs={}, + fp=NotFoundResponse(), + ) + + module = _load_execute_skills_module( + ensure_agentkit_session_endpoint=lambda **_kwargs: "https://sandbox.test", + ) + + with patch.object(module.request, "urlopen", fake_urlopen): + with self.assertRaisesRegex( + RuntimeError, + r"HTTP 404.*(?:升级|upgrade).*Skill", + ): + module.execute_skills("do work", tool_context=self._tool_context()) + + def test_requires_sandbox_upgrade_when_skill_api_returns_405(self): + class MethodNotAllowedResponse: + def read(self): + return b"method not allowed" + + def close(self): + return None + + def fake_urlopen(_request, **_kwargs): + raise module.error.HTTPError( + url="https://sandbox.test/v1/skills/execute", + code=405, + msg="Method Not Allowed", + hdrs={}, + fp=MethodNotAllowedResponse(), + ) + + module = _load_execute_skills_module( + ensure_agentkit_session_endpoint=lambda **_kwargs: "https://sandbox.test", + ) + + with patch.object(module.request, "urlopen", fake_urlopen): + with self.assertRaisesRegex( + RuntimeError, + r"HTTP 405.*(?:升级|upgrade).*Skill", + ): + module.execute_skills("do work", tool_context=self._tool_context()) + + def test_raises_runtime_error_when_session_endpoint_is_unavailable(self): + def raise_endpoint_error(**_kwargs): + raise RuntimeError("session unsupported") + + module = _load_execute_skills_module( + ensure_agentkit_session_endpoint=raise_endpoint_error, + ) + + with self.assertRaisesRegex( + RuntimeError, r"AgentKit session endpoint is not available" + ): + module.execute_skills("do work", tool_context=self._tool_context()) + + def test_missing_tool_context_raises_value_error(self): + module = _load_execute_skills_module( + ensure_agentkit_session_endpoint=lambda **_kwargs: self.fail( + "Skill API requires a tool_context" + ), + ) + + with self.assertRaisesRegex(ValueError, r"tool_context is required"): + module.execute_skills("do work", tool_context=None) + + def test_non_compatibility_skill_api_http_error_is_not_swallowed(self): + class ServerErrorResponse: + def read(self): + return b"internal error" + + def close(self): + return None + + def fake_urlopen(_request, **_kwargs): + raise module.error.HTTPError( + url="https://sandbox.test/v1/skills/execute", + code=500, + msg="Internal Server Error", + hdrs={}, + fp=ServerErrorResponse(), + ) + + module = _load_execute_skills_module( + ensure_agentkit_session_endpoint=lambda **_kwargs: "https://sandbox.test", + ) + + with patch.object(module.request, "urlopen", fake_urlopen): + with self.assertRaisesRegex(RuntimeError, "HTTP 500: internal error"): + module.execute_skills("do work", tool_context=self._tool_context()) + + def test_stream_mode_aggregates_text_chunks_from_skill_api_sse(self): + sse_body = ( + "event: chunk\n" + 'data: {"request_id":"req_1","type":"progress","content":"started","metadata":{}}\n\n' + "event: chunk\n" + 'data: {"request_id":"req_1","type":"text","content":"hello ","metadata":{}}\n\n' + "event: chunk\n" + 'data: {"request_id":"req_1","type":"text","content":"world","metadata":{}}\n\n' + "event: done\n" + 'data: {"request_id":"req_1","type":"progress","content":"done","metadata":{}}\n\n' + ).encode() + + class FakeResponse: + def __enter__(self): + return self + + def __exit__(self, *_args): + return None + + def read(self): + raise AssertionError("stream response must not be buffered with read()") + + def __iter__(self): + return iter(sse_body.splitlines(keepends=True)) + + captured_urls = [] + + def fake_urlopen(request, **_kwargs): + captured_urls.append(request.full_url) + return FakeResponse() + + module = _load_execute_skills_module( + ensure_agentkit_session_endpoint=lambda **_kwargs: "https://sandbox.test/", + ) + + with patch.object(module.request, "urlopen", fake_urlopen): + result = module.execute_skills( + "do work", + tool_context=self._tool_context(), + prefer_stream=True, + ) + + self.assertEqual(result, "hello world") + self.assertEqual(["https://sandbox.test/v1/skills/stream"], captured_urls) + + def test_stream_mode_raises_skill_api_error_event(self): + sse_body = ( + "event: error\n" + 'data: {"request_id":"req_1","type":"text","content":"skill failed","metadata":{}}\n\n' + ).encode() + + class FakeResponse: + def __enter__(self): + return self + + def __exit__(self, *_args): + return None + + def read(self): + raise AssertionError("stream response must not be buffered with read()") + + def __iter__(self): + return iter(sse_body.splitlines(keepends=True)) + + module = _load_execute_skills_module( + ensure_agentkit_session_endpoint=lambda **_kwargs: "https://sandbox.test", + ) + + with patch.object( + module.request, + "urlopen", + lambda *_args, **_kwargs: FakeResponse(), + ): + with self.assertRaisesRegex(RuntimeError, "skill failed"): + module.execute_skills( + "do work", + tool_context=self._tool_context(), + prefer_stream=True, + ) + if __name__ == "__main__": unittest.main() diff --git a/veadk/tools/builtin_tools/_agentkit.py b/veadk/tools/builtin_tools/_agentkit.py index c51621f2..10abdeec 100644 --- a/veadk/tools/builtin_tools/_agentkit.py +++ b/veadk/tools/builtin_tools/_agentkit.py @@ -14,6 +14,7 @@ import json import os +import time from typing import Any, Optional from veadk.auth.veauth.utils import get_credential_from_vefaas_iam @@ -24,6 +25,11 @@ logger = get_logger(__name__) +_SESSION_READY_TIMEOUT = 120.0 +_SESSION_POLL_INTERVAL = 1.0 +_SESSION_TERMINAL_STATUSES = frozenset({"failed", "terminating", "terminated"}) + + def resolve_agentkit_tool_id(*preferred_env_names: str) -> str: """Resolve the first configured AgentKit tool id with AGENTKIT_TOOL_ID fallback.""" for env_name in [*preferred_env_names, "AGENTKIT_TOOL_ID"]: @@ -206,3 +212,112 @@ def invoke_agentkit_exec_bash( header=header, scheme=scheme, ) + + +def ensure_agentkit_session_endpoint( + *, + tool_id: str, + tool_user_session_id: str, + tool_state: Optional[dict[str, Any]] = None, + ttl: int = 1800, + prefer_internal_endpoint: bool = False, + wait_until_ready: bool = False, + ready_timeout: float = _SESSION_READY_TIMEOUT, + poll_interval: float = _SESSION_POLL_INTERVAL, +) -> str: + """Create or reuse an AgentKit tool session and return its endpoint.""" + from agentkit.sdk.tools import types as tools_types + from agentkit.sdk.tools.client import AgentkitToolsClient + + if wait_until_ready: + if ready_timeout < 0: + raise ValueError("ready_timeout must be greater than or equal to 0") + if poll_interval <= 0: + raise ValueError("poll_interval must be greater than 0") + + _, region, _, _ = get_agentkit_endpoint_config() + ak, sk, header = get_agentkit_credentials(tool_state) + session_token = header.get("X-Security-Token", "") + client = AgentkitToolsClient( + access_key=ak, + secret_key=sk, + region=region, + session_token=session_token, + ) + session = client.create_session( + tools_types.CreateSessionRequest( + ToolId=tool_id, + UserSessionId=tool_user_session_id, + Ttl=ttl, + ) + ) + if not wait_until_ready: + public_endpoint = getattr(session, "endpoint", None) + internal_endpoint = getattr(session, "internal_endpoint", None) + endpoint = ( + internal_endpoint or public_endpoint + if prefer_internal_endpoint + else public_endpoint or internal_endpoint + ) + if endpoint: + return endpoint + + session_id = session.session_id + if not session_id: + return "" + current_session = client.get_session( + tools_types.GetSessionRequest( + ToolId=tool_id, + SessionId=session_id, + ) + ) + if prefer_internal_endpoint: + return current_session.internal_endpoint or current_session.endpoint or "" + return current_session.endpoint or current_session.internal_endpoint or "" + + session_id = session.session_id + if not session_id: + raise RuntimeError("AgentKit CreateSession response is missing SessionId") + + deadline = time.monotonic() + ready_timeout + last_status = "Unknown" + while True: + current_session = client.get_session( + tools_types.GetSessionRequest( + ToolId=tool_id, + SessionId=session_id, + ) + ) + status = (getattr(current_session, "status", None) or "").strip() + last_status = status or "Unknown" + logger.debug(f"AgentKit session {session_id} status: {last_status}") + normalized_status = status.lower() + if normalized_status == "ready": + public_endpoint = getattr(current_session, "endpoint", None) or getattr( + session, "endpoint", None + ) + internal_endpoint = getattr( + current_session, "internal_endpoint", None + ) or getattr(session, "internal_endpoint", None) + endpoint = ( + internal_endpoint or public_endpoint + if prefer_internal_endpoint + else public_endpoint or internal_endpoint + ) + if endpoint: + return endpoint + raise RuntimeError( + f"AgentKit session {session_id} is Ready but has no endpoint" + ) + if normalized_status in _SESSION_TERMINAL_STATUSES: + raise RuntimeError( + f"AgentKit session {session_id} entered terminal status {last_status}" + ) + + remaining = deadline - time.monotonic() + if remaining <= 0: + raise TimeoutError( + f"Timed out waiting for AgentKit session {session_id} to become " + f"Ready; last status: {last_status}" + ) + time.sleep(min(poll_interval, remaining)) diff --git a/veadk/tools/builtin_tools/execute_skills.py b/veadk/tools/builtin_tools/execute_skills.py index 029e7416..75059ab5 100644 --- a/veadk/tools/builtin_tools/execute_skills.py +++ b/veadk/tools/builtin_tools/execute_skills.py @@ -12,24 +12,260 @@ # See the License for the specific language governing permissions and # limitations under the License. +from __future__ import annotations + +import json +import os +import time +from collections.abc import Iterable from typing import Optional +from urllib import error, request +from urllib.parse import urlsplit, urlunsplit from google.adk.tools import ToolContext from veadk.tools.builtin_tools._agentkit import ( + ensure_agentkit_session_endpoint, get_agentkit_account_id, resolve_agentkit_tool_id, ) from veadk.tools.builtin_tools.run_sandbox_agent import run_sandbox_agent -from veadk.utils.logger import get_logger -logger = get_logger(__name__) + +_SKILL_API_UPGRADE_STATUS_CODES = frozenset({404, 405}) +_SKILL_API_TRANSIENT_STATUS_CODES = frozenset({502, 503, 504}) +_SKILL_API_TIMEOUT = 900 +_SKILL_API_HEALTH_TIMEOUT = 30.0 +_SKILL_API_HEALTH_POLL_INTERVAL = 1.0 +_SKILL_API_HEALTH_REQUEST_TIMEOUT = 5.0 + + +def _skill_api_upgrade_hint(path: str) -> str: + api_path = ( + "/v1/skills/stream" + if path.rstrip("/").endswith("/stream") + else "/v1/skills/execute" + ) + return ( + f"提示:当前 Skill 沙箱镜像未实现 {api_path} 接口,可能是旧版沙箱镜像。" + "请升级 Skill 沙箱镜像或切换到支持 Skill HTTP API 的新版沙箱。" + ) + + +def _tool_user_session_id(tool_context: ToolContext) -> str: + invocation_context = tool_context._invocation_context + session_id = invocation_context.session.id + agent_name = invocation_context.agent.name + user_id = invocation_context.user_id + return agent_name + "_" + user_id + "_" + session_id + + +def _tip_token_key(tool_context: ToolContext) -> str | None: + state = tool_context.state or {} + return ( + state.get("TIP_TOKEN_KEY") + or state.get("tip_token_key") + or os.getenv("TIP_TOKEN_KEY") + or None + ) + + +def _skill_api_url(endpoint: str, path: str) -> str: + if not endpoint: + raise RuntimeError("AgentKit session endpoint is empty") + parts = urlsplit(endpoint) + endpoint_path = parts.path.rstrip("/") + skill_path = path.lstrip("/") + joined_path = f"{endpoint_path}/{skill_path}" if endpoint_path else f"/{skill_path}" + return urlunsplit( + (parts.scheme, parts.netloc, joined_path, parts.query, parts.fragment) + ) + + +def _post_skill_api_json( + *, + endpoint: str, + path: str, + payload: dict[str, object], + tip_token_key: str | None, + timeout: int, + stream: bool, +) -> str: + headers = { + "Content-Type": "application/json", + "Accept": "application/json, text/event-stream", + } + if tip_token_key: + headers["X-Tip-Token-Key"] = tip_token_key + + req = request.Request( + _skill_api_url(endpoint, path), + data=json.dumps(payload).encode("utf-8"), + headers=headers, + method="POST", + ) + try: + with request.urlopen(req, timeout=timeout) as response: + if stream: + return _parse_skill_stream_response(response) + return _parse_skill_execute_response(response.read()) + except error.HTTPError as exc: + if exc.code in _SKILL_API_UPGRADE_STATUS_CODES: + raise RuntimeError( + f"Skill HTTP API returned HTTP {exc.code}. " + f"{_skill_api_upgrade_hint(path)}" + ) from exc + detail = exc.read().decode("utf-8", errors="replace") + raise RuntimeError( + f"Skill HTTP API request failed with HTTP {exc.code}: {detail}" + ) from exc + except error.URLError as exc: + raise RuntimeError( + f"Skill HTTP API endpoint is not reachable: {exc.reason}" + ) from exc + + +def _wait_for_skill_api_health( + *, + endpoint: str, + timeout: float = _SKILL_API_HEALTH_TIMEOUT, + poll_interval: float = _SKILL_API_HEALTH_POLL_INTERVAL, +) -> None: + """Wait until the Skill API upstream is reachable through the session endpoint.""" + deadline = time.monotonic() + timeout + last_error = "unknown error" + while True: + req = request.Request( + _skill_api_url(endpoint, "/v1/skills/healthz"), + headers={"Accept": "application/json"}, + method="GET", + ) + try: + remaining = max(0.001, deadline - time.monotonic()) + with request.urlopen( + req, + timeout=min(_SKILL_API_HEALTH_REQUEST_TIMEOUT, remaining), + ): + return + except error.HTTPError as exc: + if exc.code in _SKILL_API_UPGRADE_STATUS_CODES: + # Some compatible images predate the dedicated health endpoint. + return + if exc.code not in _SKILL_API_TRANSIENT_STATUS_CODES: + detail = exc.read().decode("utf-8", errors="replace") + raise RuntimeError( + f"Skill HTTP API health check failed with HTTP {exc.code}: {detail}" + ) from exc + last_error = f"HTTP {exc.code}" + except error.URLError as exc: + last_error = str(exc.reason) + + remaining = deadline - time.monotonic() + if remaining <= 0: + raise RuntimeError( + f"Timed out waiting for Skill HTTP API health check: {last_error}" + ) + time.sleep(min(poll_interval, remaining)) + + +def _parse_skill_execute_response(raw: bytes) -> str: + try: + payload = json.loads(raw.decode("utf-8")) + except json.JSONDecodeError: + return raw.decode("utf-8", errors="replace") + + if isinstance(payload, dict): + if isinstance(payload.get("content"), str): + return payload["content"] + data = payload.get("data") + if isinstance(data, dict) and isinstance(data.get("content"), str): + return data["content"] + return json.dumps(payload, ensure_ascii=False) + + +def _parse_skill_stream_response(raw: bytes | Iterable[bytes]) -> str: + chunks: list[str] = [] + event_name = "message" + data_lines: list[str] = [] + + def flush_event() -> None: + nonlocal event_name, data_lines + if not data_lines: + event_name = "message" + return + data = "\n".join(data_lines) + try: + payload = json.loads(data) + except json.JSONDecodeError: + payload = {} + + if event_name == "error": + content = payload.get("content") if isinstance(payload, dict) else None + if isinstance(content, str): + raise RuntimeError(content) + raise RuntimeError(data) + if isinstance(payload, dict) and payload.get("type") == "text": + content = payload.get("content") + if isinstance(content, str): + chunks.append(content) + + event_name = "message" + data_lines = [] + + raw_lines = raw.splitlines() if isinstance(raw, bytes) else raw + for raw_line in raw_lines: + line = raw_line.decode("utf-8", errors="replace").rstrip("\r\n") + if not line: + flush_event() + continue + if line.startswith(":"): + continue + if line.startswith("event:"): + event_name = line[len("event:") :].strip() + elif line.startswith("data:"): + data_lines.append(line[len("data:") :].strip()) + + flush_event() + return "".join(chunks) + + +def _execute_skills_via_skill_api( + *, + workflow_prompt: str, + tool_id: str, + tool_context: ToolContext, + prefer_stream: bool, + timeout: int, +) -> str: + try: + endpoint = ensure_agentkit_session_endpoint( + tool_id=tool_id, + tool_user_session_id=_tool_user_session_id(tool_context), + tool_state=tool_context.state, + ttl=max(timeout, 1800), + wait_until_ready=True, + ) + except Exception as exc: + raise RuntimeError( + f"AgentKit session endpoint is not available: {exc}" + ) from exc + _wait_for_skill_api_health(endpoint=endpoint) + path = "/v1/skills/stream" if prefer_stream else "/v1/skills/execute" + return _post_skill_api_json( + endpoint=endpoint, + path=path, + payload={"prompt": workflow_prompt}, + tip_token_key=_tip_token_key(tool_context), + timeout=timeout, + stream=prefer_stream, + ) def execute_skills( workflow_prompt: str, tool_context: ToolContext = None, env_vars: Optional[dict[str, str]] = None, + prefer_stream: bool = False, ) -> str: """Execute skills in a sandbox and return the output. @@ -38,27 +274,36 @@ def execute_skills( Args: workflow_prompt (str): instruction of workflow env_vars (Optional[dict[str, str]]): Environment variables passed to the - skill agent process for this execution only. + skill agent process for this execution only. Requests with custom + environment variables use the legacy RunCode execution path. Returns: str: The output of the code execution. """ - timeout = 900 - tool_id = resolve_agentkit_tool_id("AGENTKIT_TOOL_ID_SKILLS") - account_id = get_agentkit_account_id(tool_context.state if tool_context else None) + if tool_context is None: + raise ValueError("tool_context is required for execute_skills") - extra_env_vars = {} - if account_id: - extra_env_vars["TOS_SKILLS_DIR"] = ( - f"tos://agentkit-platform-{account_id}/skills/" - ) + tool_id = resolve_agentkit_tool_id("AGENTKIT_TOOL_ID_SKILLS") if env_vars: - extra_env_vars.update(env_vars) + account_id = get_agentkit_account_id(tool_context.state) + extra_env_vars = dict(env_vars) + if account_id: + extra_env_vars.setdefault( + "TOS_SKILLS_DIR", + f"tos://agentkit-platform-{account_id}/skills/", + ) + return run_sandbox_agent( + workflow_prompt=workflow_prompt, + tool_id=tool_id, + tool_context=tool_context, + timeout=_SKILL_API_TIMEOUT, + extra_env_vars=extra_env_vars, + ) - return run_sandbox_agent( + return _execute_skills_via_skill_api( workflow_prompt=workflow_prompt, tool_id=tool_id, tool_context=tool_context, - timeout=timeout, - extra_env_vars=extra_env_vars, + prefer_stream=prefer_stream, + timeout=_SKILL_API_TIMEOUT, )