From 68336d1cfbff26d15d98c1fe4ba93a674bc77b7d Mon Sep 17 00:00:00 2001 From: chuhanku Date: Wed, 19 Aug 2026 13:25:42 +0800 Subject: [PATCH 1/2] feat(skill sandbox): forward inbound_auth to skills --- .../builtin_tools/test_run_sandbox_agent.py | 233 ++++++++++++++++-- veadk/tools/builtin_tools/execute_skills.py | 134 +++++++++- 2 files changed, 351 insertions(+), 16 deletions(-) diff --git a/tests/tools/builtin_tools/test_run_sandbox_agent.py b/tests/tools/builtin_tools/test_run_sandbox_agent.py index 09d440684..9c9d9f59c 100644 --- a/tests/tools/builtin_tools/test_run_sandbox_agent.py +++ b/tests/tools/builtin_tools/test_run_sandbox_agent.py @@ -13,6 +13,7 @@ # limitations under the License. import importlib.util +import hashlib import json import sys import types @@ -80,8 +81,7 @@ 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, + logger=None, ): module_path = ( Path(__file__).resolve().parents[3] @@ -95,6 +95,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,30 +114,42 @@ 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") + + def fake_build_auth_config(**kwargs): + credential = ( + types.SimpleNamespace(api_key=kwargs["token"]) + if kwargs.get("token") + else None + ) + return types.SimpleNamespace( + **kwargs, + exchanged_auth_credential=credential, + ) + + fake_auth.build_auth_config = fake_build_auth_config fake_logger = types.ModuleType("veadk.utils.logger") - fake_logger.get_logger = lambda _name: types.SimpleNamespace( + fake_logger.get_logger = lambda _name: logger or types.SimpleNamespace( debug=lambda *_args, **_kwargs: None, warning=lambda *_args, **_kwargs: None, error=lambda *_args, **_kwargs: None, ) - 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,11 +161,13 @@ 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 +def _headers_lower(request_obj): + return {key.lower(): value for key, value in request_obj.headers.items()} + + class TestMergeExecutionEnvVars(unittest.TestCase): @classmethod def setUpClass(cls): @@ -202,14 +225,41 @@ def test_runner_code_overrides_the_sandbox_process_environment(self): class TestExecuteSkillsSkillApi(unittest.TestCase): - def _tool_context(self): + def _tool_context(self, *, inbound_credential=None, credentials_by_key=None): + credentials_by_key = credentials_by_key or {} + + class FakeCredentialService: + def __init__(self): + self.stored_credentials = {} + + async def load_credential(self, *, auth_config, callback_context): + self.auth_config = auth_config + self.callback_context = callback_context + if auth_config.credential_key in credentials_by_key: + return credentials_by_key[auth_config.credential_key] + return inbound_credential + + async def set_credential( + self, *, app_name, user_id, credential_key, credential + ): + self.stored_credentials[(app_name, user_id, credential_key)] = ( + credential + ) + + credential_service = ( + FakeCredentialService() + if inbound_credential or credentials_by_key + else None + ) invocation_context = types.SimpleNamespace( session=types.SimpleNamespace(id="session-1"), agent=types.SimpleNamespace(name="agent"), + app_name="assistant", user_id="user", + credential_service=credential_service, ) return types.SimpleNamespace( - state={"TIP_TOKEN_KEY": "tip-from-state"}, + state={}, _invocation_context=invocation_context, ) @@ -265,12 +315,143 @@ 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()) + headers = _headers_lower(request_obj) + self.assertEqual({"content-type"}, set(headers)) 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_forwards_inbound_auth_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", + ) + inbound_credential = types.SimpleNamespace( + auth_type="HTTP", + http=types.SimpleNamespace( + credentials=types.SimpleNamespace(token="inbound-user-jwt") + ), + api_key=None, + ) + tool_context = self._tool_context( + credentials_by_key={"inbound_auth": inbound_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("inbound_auth", credential_service.auth_config.credential_key) + self.assertEqual("header", credential_service.auth_config.auth_method) + self.assertEqual("bearer", credential_service.auth_config.header_scheme) + request_obj, _timeout = captured_requests[0] + headers = _headers_lower(request_obj) + self.assertEqual({"content-type", "inbound_auth"}, set(headers)) + self.assertEqual("inbound-user-jwt", headers["inbound_auth"]) + + def test_logs_inbound_auth_summaries_without_secret_value(self): + captured_logs = [] + logger = types.SimpleNamespace( + debug=lambda message, *args: captured_logs.append(message % args), + warning=lambda *_args, **_kwargs: None, + error=lambda *_args, **_kwargs: None, + ) + + 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() + + module = _load_execute_skills_module( + ensure_agentkit_session_endpoint=lambda **_kwargs: "https://sandbox.test", + logger=logger, + ) + token = "inbound-secret-token" + inbound_credential = types.SimpleNamespace( + auth_type="HTTP", + http=types.SimpleNamespace(credentials=types.SimpleNamespace(token=token)), + api_key=None, + ) + tool_context = self._tool_context( + credentials_by_key={"inbound_auth": inbound_credential} + ) + + with patch.object( + module.request, "urlopen", lambda *_args, **_kwargs: FakeResponse() + ): + result = module.execute_skills("do work", tool_context=tool_context) + + self.assertEqual("a2a result", result) + logs = "\n".join(captured_logs) + self.assertNotIn(token, logs) + self.assertEqual(2, len(captured_logs)) + + received_prefix = "execute_skills inbound_auth received before sandbox send: " + send_prefix = "execute_skills inbound_auth header before sandbox request: " + self.assertTrue(captured_logs[0].startswith(received_prefix)) + self.assertTrue(captured_logs[1].startswith(send_prefix)) + + expected = { + "present": True, + "len": len(token), + "sha256_8": hashlib.sha256(token.encode("utf-8")).hexdigest()[:8], + } + self.assertEqual( + expected, json.loads(captured_logs[0].removeprefix(received_prefix)) + ) + self.assertEqual( + expected, json.loads(captured_logs[1].removeprefix(send_prefix)) + ) + def test_a2a_retries_502_until_upstream_is_ready(self): attempts = [] @@ -393,12 +574,22 @@ def fake_urlopen(request, timeout=None): module = _load_execute_skills_module( ensure_agentkit_session_endpoint=lambda **_kwargs: "https://sandbox.test", ) + inbound_credential = types.SimpleNamespace( + auth_type="HTTP", + http=types.SimpleNamespace( + credentials=types.SimpleNamespace(token="inbound-user-jwt") + ), + api_key=None, + ) + tool_context = self._tool_context( + credentials_by_key={"inbound_auth": inbound_credential} + ) with ( patch.object(module.request, "urlopen", fake_urlopen), patch.object(module.time, "sleep") as sleep, ): - result = module.execute_skills("do work", tool_context=self._tool_context()) + result = module.execute_skills("do work", tool_context=tool_context) self.assertEqual(result, "poll result") self.assertEqual(2, len(captured_requests)) @@ -406,6 +597,18 @@ 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( + "inbound-user-jwt", _headers_lower(send_request)["inbound_auth"] + ) + self.assertEqual( + "inbound-user-jwt", _headers_lower(get_request)["inbound_auth"] + ) + self.assertEqual( + {"content-type", "inbound_auth"}, set(_headers_lower(send_request)) + ) + self.assertEqual( + {"content-type", "inbound_auth"}, set(_headers_lower(get_request)) + ) self.assertEqual("message/send", send_payload["method"]) self.assertEqual("tasks/get", get_payload["method"]) self.assertEqual("task-1", get_payload["params"]["id"]) @@ -517,10 +720,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 7c56e544a..d788e5b97 100644 --- a/veadk/tools/builtin_tools/execute_skills.py +++ b/veadk/tools/builtin_tools/execute_skills.py @@ -14,19 +14,25 @@ from __future__ import annotations +import asyncio +import hashlib import json +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 build_auth_config +from veadk.utils.logger import get_logger _SKILL_API_TIMEOUT = 1800 _A2A_POLL_INTERVAL = 2.0 @@ -44,6 +50,8 @@ "auth-required", } ) +logger = get_logger(__name__) +_INBOUND_AUTH_CREDENTIAL_KEY = "inbound_auth" def _validate_timeout(timeout: int) -> None: @@ -61,6 +69,115 @@ 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 _inbound_auth_debug_summary(value: object | None) -> dict[str, object]: + if value is None: + return {"present": False, "len": 0, "sha256_8": ""} + text = str(value) + return { + "present": bool(text), + "len": len(text), + "sha256_8": hashlib.sha256(text.encode("utf-8")).hexdigest()[:8] + if text + else "", + } + + +def _load_credential_from_service( + *, + tool_context: ToolContext, + auth_config: object, +) -> object | None: + invocation_context = getattr(tool_context, "_invocation_context", None) + credential_service = getattr(invocation_context, "credential_service", None) + if not credential_service: + return None + credential = credential_service.load_credential( + auth_config=auth_config, + callback_context=CallbackContext(invocation_context), + ) + if hasattr(credential, "__await__"): + credential = _await_sync(credential) + return credential + + +def _credential_token_value(credential: object | None) -> str | None: + if credential is None: + return None + api_key = getattr(credential, "api_key", None) + if api_key: + return str(api_key) + http = getattr(credential, "http", None) + http_credentials = getattr(http, "credentials", None) if http is not None else None + http_token = ( + getattr(http_credentials, "token", None) + if http_credentials is not None + else None + ) + if http_token: + return str(http_token) + oauth2 = getattr(credential, "oauth2", None) + oauth2_access_token = ( + getattr(oauth2, "access_token", None) if oauth2 is not None else None + ) + if oauth2_access_token: + return str(oauth2_access_token) + return None + + +def _inbound_auth_token_from_credential_service( + tool_context: ToolContext, +) -> str | None: + invocation_context = getattr(tool_context, "_invocation_context", None) + credential_service = getattr(invocation_context, "credential_service", None) + inbound_auth_token = None + if credential_service: + auth_config = build_auth_config( + credential_key=_INBOUND_AUTH_CREDENTIAL_KEY, + auth_method="header", + header_scheme="bearer", + ) + credential = _load_credential_from_service( + tool_context=tool_context, + auth_config=auth_config, + ) + inbound_auth_token = _credential_token_value(credential) + logger.debug( + "execute_skills inbound_auth received before sandbox send: %s", + json.dumps( + _inbound_auth_debug_summary(inbound_auth_token), + ensure_ascii=False, + sort_keys=True, + ), + ) + return inbound_auth_token + + +def _inbound_auth_token(tool_context: ToolContext) -> str | None: + return _inbound_auth_token_from_credential_service(tool_context) + + def _skill_api_url(endpoint: str, path: str) -> str: if not endpoint: raise RuntimeError("AgentKit session endpoint is empty") @@ -210,6 +327,7 @@ def _post_a2a_jsonrpc( payload: dict[str, object], timeout: int, retry_until: float | None = None, + inbound_auth: str | None = None, ) -> dict: url = _a2a_jsonrpc_url(endpoint) while True: @@ -221,7 +339,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", + **({"inbound_auth": inbound_auth} if inbound_auth else {}), + }, method="POST", ) try: @@ -277,6 +398,15 @@ def _execute_skills_via_a2a( "user_id": invocation_context.user_id, "session_id": invocation_context.session.id, } + inbound_auth = _inbound_auth_token(tool_context) + logger.debug( + "execute_skills inbound_auth header before sandbox request: %s", + json.dumps( + _inbound_auth_debug_summary(inbound_auth), + ensure_ascii=False, + sort_keys=True, + ), + ) task = _a2a_result_task( "A2ASendMessage", _post_a2a_jsonrpc( @@ -296,6 +426,7 @@ def _execute_skills_via_a2a( }, timeout=_a2a_request_timeout(deadline), retry_until=deadline, + inbound_auth=inbound_auth, ), ) task_id = _a2a_task_id(task) @@ -321,6 +452,7 @@ def _execute_skills_via_a2a( }, timeout=_a2a_request_timeout(deadline), retry_until=deadline, + inbound_auth=inbound_auth, ), ) poll_interval = min(poll_interval * 2, _A2A_MAX_POLL_INTERVAL) From 67a9d3dbaf4d3685a5bceda9f9aefedfc9ac2bb7 Mon Sep 17 00:00:00 2001 From: chuhanku Date: Wed, 19 Aug 2026 17:01:37 +0800 Subject: [PATCH 2/2] fix: support sdk version --- pyproject.toml | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/pyproject.toml b/pyproject.toml index 4a0dc278b..9e736d4e6 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", # 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