diff --git a/pyproject.toml b/pyproject.toml index 5992cc52..a1b9edb0 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -17,7 +17,7 @@ authors = [ ] dependencies = [ "pydantic-settings==2.10.1", # Config management - "a2a-sdk==0.3.7", # For Google Agent2Agent protocol + "a2a-sdk>=0.3.7,<0.4.0", #For Google Agent2Agent protocol "deprecated==1.2.18", "google-adk>=1.34.0", # For basic agent architecture # litellm and sqlalchemy are required by code paths veadk always uses diff --git a/tests/tools/builtin_tools/test_run_sandbox_agent.py b/tests/tools/builtin_tools/test_run_sandbox_agent.py index 09d44068..ef5741c7 100644 --- a/tests/tools/builtin_tools/test_run_sandbox_agent.py +++ b/tests/tools/builtin_tools/test_run_sandbox_agent.py @@ -80,8 +80,6 @@ def _load_run_sandbox_agent_module(): 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] @@ -95,6 +93,15 @@ def _load_execute_skills_module( fake_google.__path__ = [] # type: ignore[attr-defined] fake_google_adk = types.ModuleType("google.adk") fake_google_adk.__path__ = [] # type: ignore[attr-defined] + fake_google_adk_agents = types.ModuleType("google.adk.agents") + fake_google_adk_agents.__path__ = [] # type: ignore[attr-defined] + fake_callback_context = types.ModuleType("google.adk.agents.callback_context") + + class FakeCallbackContext: + def __init__(self, invocation_context): + self._invocation_context = invocation_context + + fake_callback_context.CallbackContext = FakeCallbackContext fake_google_adk_tools = types.ModuleType("google.adk.tools") fake_google_adk_tools.ToolContext = object @@ -105,13 +112,14 @@ def _load_execute_skills_module( fake_builtin_tools = types.ModuleType("veadk.tools.builtin_tools") fake_builtin_tools.__path__ = [] # type: ignore[attr-defined] 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_auth = types.ModuleType("veadk.utils.auth") + fake_auth.VE_TIP_TOKEN_CREDENTIAL_KEY = "ve_tip_token" + fake_auth.VE_TIP_TOKEN_HEADER = "X-Ve-TIP-Token" + fake_auth.build_auth_config = lambda **kwargs: types.SimpleNamespace(**kwargs) fake_logger = types.ModuleType("veadk.utils.logger") fake_logger.get_logger = lambda _name: types.SimpleNamespace( debug=lambda *_args, **_kwargs: None, @@ -122,13 +130,15 @@ def _load_execute_skills_module( stub_modules = { "google": fake_google, "google.adk": fake_google_adk, + "google.adk.agents": fake_google_adk_agents, + "google.adk.agents.callback_context": fake_callback_context, "google.adk.tools": fake_google_adk_tools, "veadk": fake_veadk, "veadk.tools": fake_tools, "veadk.tools.builtin_tools": fake_builtin_tools, "veadk.tools.builtin_tools._agentkit": fake_agentkit, - "veadk.tools.builtin_tools.run_sandbox_agent": fake_runner, "veadk.utils": fake_utils, + "veadk.utils.auth": fake_auth, "veadk.utils.logger": fake_logger, } @@ -140,8 +150,6 @@ def _load_execute_skills_module( 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 @@ -202,11 +210,19 @@ def test_runner_code_overrides_the_sandbox_process_environment(self): class TestExecuteSkillsSkillApi(unittest.TestCase): - def _tool_context(self): + def _tool_context(self, *, tip_credential=None): + class FakeCredentialService: + async def load_credential(self, *, auth_config, callback_context): + self.auth_config = auth_config + self.callback_context = callback_context + return tip_credential + + credential_service = FakeCredentialService() if tip_credential else None invocation_context = types.SimpleNamespace( session=types.SimpleNamespace(id="session-1"), agent=types.SimpleNamespace(name="agent"), user_id="user", + credential_service=credential_service, ) return types.SimpleNamespace( state={"TIP_TOKEN_KEY": "tip-from-state"}, @@ -265,12 +281,59 @@ def fake_urlopen(request, timeout=None): self.assertEqual("https://sandbox.test/a2a", request_obj.full_url) self.assertEqual(60, timeout) self.assertEqual("POST", request_obj.get_method()) + self.assertEqual("tip-from-state", request_obj.headers["X-tip-token-key"]) self.assertEqual("message/send", payload["method"]) self.assertEqual("do work", payload["params"]["message"]["parts"][0]["text"]) self.assertFalse(payload["params"]["configuration"]["blocking"]) self.assertEqual("user", payload["params"]["metadata"]["user_id"]) self.assertEqual("session-1", payload["params"]["metadata"]["session_id"]) + def test_a2a_prefers_tip_token_from_credential_service(self): + captured_requests = [] + + class FakeResponse: + def __enter__(self): + return self + + def __exit__(self, *_args): + return None + + def read(self): + return json.dumps( + { + "jsonrpc": "2.0", + "id": "req", + "result": { + "kind": "task", + "id": "task-1", + "status": {"state": "completed"}, + "artifacts": [ + {"parts": [{"kind": "text", "text": "a2a result"}]} + ], + }, + } + ).encode() + + def fake_urlopen(request, timeout=None): + captured_requests.append((request, timeout)) + return FakeResponse() + + module = _load_execute_skills_module( + ensure_agentkit_session_endpoint=lambda **_kwargs: "https://sandbox.test", + ) + tool_context = self._tool_context( + tip_credential=types.SimpleNamespace(api_key="tip-from-credential") + ) + + with patch.object(module.request, "urlopen", fake_urlopen): + result = module.execute_skills("do work", tool_context=tool_context) + + self.assertEqual(result, "a2a result") + credential_service = tool_context._invocation_context.credential_service + self.assertEqual("ve_tip_token", credential_service.auth_config.credential_key) + request_obj, _timeout = captured_requests[0] + self.assertEqual("tip-from-credential", request_obj.headers["X-tip-token-key"]) + def test_a2a_retries_502_until_upstream_is_ready(self): attempts = [] @@ -406,6 +469,8 @@ def fake_urlopen(request, timeout=None): get_request, _get_timeout = captured_requests[1] send_payload = json.loads(send_request.data.decode()) get_payload = json.loads(get_request.data.decode()) + self.assertEqual("tip-from-state", send_request.headers["X-tip-token-key"]) + self.assertEqual("tip-from-state", get_request.headers["X-tip-token-key"]) self.assertEqual("message/send", send_payload["method"]) self.assertEqual("tasks/get", get_payload["method"]) self.assertEqual("task-1", get_payload["params"]["id"]) @@ -517,10 +582,10 @@ def test_skill_api_url_preserves_agentkit_endpoint_query_auth(self): ) self.assertEqual( - "https://sandbox.test/v1/skills/execute?faasInstanceName=inst&Authorization=key", + "https://sandbox.test/a2a?faasInstanceName=inst&Authorization=key", module._skill_api_url( "https://sandbox.test/?faasInstanceName=inst&Authorization=key", - "/v1/skills/execute", + "/a2a", ), ) diff --git a/veadk/tools/builtin_tools/execute_skills.py b/veadk/tools/builtin_tools/execute_skills.py index 7c56e544..2849b948 100644 --- a/veadk/tools/builtin_tools/execute_skills.py +++ b/veadk/tools/builtin_tools/execute_skills.py @@ -14,19 +14,28 @@ from __future__ import annotations +import asyncio import json +import os +import threading import time import uuid from typing import Optional from urllib import error, request from urllib.parse import urlsplit, urlunsplit +from google.adk.agents.callback_context import CallbackContext from google.adk.tools import ToolContext from veadk.tools.builtin_tools._agentkit import ( ensure_agentkit_session_endpoint, resolve_agentkit_tool_id, ) +from veadk.utils.auth import ( + VE_TIP_TOKEN_CREDENTIAL_KEY, + VE_TIP_TOKEN_HEADER, + build_auth_config, +) _SKILL_API_TIMEOUT = 1800 _A2A_POLL_INTERVAL = 2.0 @@ -61,6 +70,62 @@ def _tool_user_session_id(tool_context: ToolContext) -> str: return agent_name + "_" + user_id + "_" + session_id +def _await_sync(awaitable): + try: + asyncio.get_running_loop() + except RuntimeError: + return asyncio.run(awaitable) + + result = {} + + def run_in_thread(): + try: + result["value"] = asyncio.run(awaitable) + except BaseException as exc: + result["error"] = exc + + thread = threading.Thread(target=run_in_thread) + thread.start() + thread.join() + if "error" in result: + raise result["error"] + return result.get("value") + + +def _tip_token_key_from_credential_service(tool_context: ToolContext) -> str | None: + invocation_context = getattr(tool_context, "_invocation_context", None) + credential_service = getattr(invocation_context, "credential_service", None) + if not credential_service: + return None + + auth_config = build_auth_config( + auth_method="apikey", + credential_key=VE_TIP_TOKEN_CREDENTIAL_KEY, + header_name=VE_TIP_TOKEN_HEADER, + ) + credential = credential_service.load_credential( + auth_config=auth_config, + callback_context=CallbackContext(invocation_context), + ) + if hasattr(credential, "__await__"): + credential = _await_sync(credential) + return getattr(credential, "api_key", None) if credential else None + + +def _tip_token_key(tool_context: ToolContext) -> str | None: + tip_token_key = _tip_token_key_from_credential_service(tool_context) + if tip_token_key: + return tip_token_key + + 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") @@ -91,34 +156,6 @@ def _a2a_jsonrpc_url(endpoint: str) -> str: return _skill_api_url(endpoint, "/a2a") -def _post_json( - *, - endpoint: str, - path: str, - payload: dict[str, object], - timeout: int, - accept: str = "application/json", -) -> bytes: - req = request.Request( - _skill_api_url(endpoint, path), - data=json.dumps(payload).encode("utf-8"), - headers={"Content-Type": "application/json", "Accept": accept}, - method="POST", - ) - try: - with request.urlopen(req, timeout=timeout) as response: - return response.read() - except error.HTTPError as exc: - detail = exc.read().decode("utf-8", errors="replace") - raise RuntimeError( - f"AgentKit Skill {path} request failed with HTTP {exc.code}: {detail}" - ) from exc - except error.URLError as exc: - raise RuntimeError( - f"AgentKit Skill {path} endpoint is not reachable: {exc.reason}" - ) from exc - - def _extract_text_from_parts(parts: object) -> str: if not isinstance(parts, list): return "" @@ -210,6 +247,7 @@ def _post_a2a_jsonrpc( payload: dict[str, object], timeout: int, retry_until: float | None = None, + tip_token_key: str | None = None, ) -> dict: url = _a2a_jsonrpc_url(endpoint) while True: @@ -221,7 +259,10 @@ def _post_a2a_jsonrpc( req = request.Request( url, data=json.dumps(payload).encode("utf-8"), - headers={"Content-Type": "application/json"}, + headers={ + "Content-Type": "application/json", + **({"X-Tip-Token-Key": tip_token_key} if tip_token_key else {}), + }, method="POST", ) try: @@ -277,6 +318,7 @@ def _execute_skills_via_a2a( "user_id": invocation_context.user_id, "session_id": invocation_context.session.id, } + tip_token_key = _tip_token_key(tool_context) task = _a2a_result_task( "A2ASendMessage", _post_a2a_jsonrpc( @@ -296,6 +338,7 @@ def _execute_skills_via_a2a( }, timeout=_a2a_request_timeout(deadline), retry_until=deadline, + tip_token_key=tip_token_key, ), ) task_id = _a2a_task_id(task) @@ -321,6 +364,7 @@ def _execute_skills_via_a2a( }, timeout=_a2a_request_timeout(deadline), retry_until=deadline, + tip_token_key=tip_token_key, ), ) poll_interval = min(poll_interval * 2, _A2A_MAX_POLL_INTERVAL)