Skip to content
Merged
275 changes: 275 additions & 0 deletions tests/tools/builtin_tools/test_agentkit.py
Original file line number Diff line number Diff line change
Expand Up @@ -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()
Loading
Loading