Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension


Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
2 changes: 1 addition & 1 deletion pyproject.toml
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down
87 changes: 76 additions & 11 deletions tests/tools/builtin_tools/test_run_sandbox_agent.py
Original file line number Diff line number Diff line change
Expand Up @@ -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]
Expand All @@ -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

Expand All @@ -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,
Expand All @@ -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,
}

Expand All @@ -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


Expand Down Expand Up @@ -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"},
Expand Down Expand Up @@ -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 = []

Expand Down Expand Up @@ -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"])
Expand Down Expand Up @@ -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",
),
)

Expand Down
102 changes: 73 additions & 29 deletions veadk/tools/builtin_tools/execute_skills.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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")
Expand Down Expand Up @@ -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 ""
Expand Down Expand Up @@ -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:
Expand All @@ -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:
Expand Down Expand Up @@ -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(
Expand All @@ -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)
Expand All @@ -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)
Expand Down