209 lines
6.9 KiBLFS
Python
209 lines
6.9 KiBLFS
Python
from __future__ import annotations
|
|
|
|
import json
|
|
import threading
|
|
from collections.abc import Iterator
|
|
from contextlib import contextmanager
|
|
from http.server import BaseHTTPRequestHandler, ThreadingHTTPServer
|
|
from pathlib import Path
|
|
from typing import Any
|
|
|
|
import pytest
|
|
|
|
from skillsbench_agentbeats import acp_bridge
|
|
from skillsbench_agentbeats.acp_bridge import (
|
|
AGENT_NAME,
|
|
BRIDGE_PATH,
|
|
LAUNCH_COMMAND,
|
|
MAX_FILE_BYTES,
|
|
AgentBeatsAcpBridge,
|
|
_materialize_files,
|
|
)
|
|
|
|
|
|
class _A2AServer(ThreadingHTTPServer):
|
|
result: dict[str, Any]
|
|
received: list[dict[str, Any]]
|
|
|
|
|
|
class _A2AHandler(BaseHTTPRequestHandler):
|
|
server: _A2AServer
|
|
|
|
def do_POST(self) -> None:
|
|
length = int(self.headers.get("content-length", "0"))
|
|
payload = json.loads(self.rfile.read(length).decode("utf-8"))
|
|
self.server.received.append(
|
|
{
|
|
"path": self.path,
|
|
"headers": dict(self.headers),
|
|
"payload": payload,
|
|
}
|
|
)
|
|
body = json.dumps(
|
|
{
|
|
"jsonrpc": "2.0",
|
|
"id": payload.get("id"),
|
|
"result": self.server.result,
|
|
}
|
|
).encode("utf-8")
|
|
self.send_response(200)
|
|
self.send_header("content-type", "application/json")
|
|
self.send_header("content-length", str(len(body)))
|
|
self.end_headers()
|
|
self.wfile.write(body)
|
|
|
|
def log_message(self, format: str, *args: Any) -> None:
|
|
del format, args
|
|
|
|
|
|
def test_register_agentbeats_a2a_agent_updates_benchflow_registry() -> None:
|
|
from benchflow.agents import registry
|
|
|
|
agents = registry.AGENTS.copy()
|
|
installers = registry.AGENT_INSTALLERS.copy()
|
|
launch = registry.AGENT_LAUNCH.copy()
|
|
try:
|
|
acp_bridge.register_agentbeats_a2a_agent()
|
|
|
|
config = registry.AGENTS[AGENT_NAME]
|
|
assert config.name == AGENT_NAME
|
|
assert config.protocol == "acp"
|
|
assert config.requires_env == []
|
|
assert config.launch_cmd == LAUNCH_COMMAND
|
|
assert registry.AGENT_LAUNCH[AGENT_NAME] == LAUNCH_COMMAND
|
|
assert BRIDGE_PATH in registry.AGENT_INSTALLERS[AGENT_NAME]
|
|
assert "python3" in registry.AGENT_INSTALLERS[AGENT_NAME]
|
|
assert "base64 -d" in registry.AGENT_INSTALLERS[AGENT_NAME]
|
|
finally:
|
|
registry.AGENTS.clear()
|
|
registry.AGENTS.update(agents)
|
|
registry.AGENT_INSTALLERS.clear()
|
|
registry.AGENT_INSTALLERS.update(installers)
|
|
registry.AGENT_LAUNCH.clear()
|
|
registry.AGENT_LAUNCH.update(launch)
|
|
|
|
|
|
def test_acp_bridge_forwards_prompt_to_a2a_and_materializes_files(
|
|
monkeypatch: pytest.MonkeyPatch,
|
|
tmp_path: Path,
|
|
) -> None:
|
|
result = {
|
|
"kind": "message",
|
|
"role": "agent",
|
|
"messageId": "reply-1",
|
|
"parts": [
|
|
{"kind": "text", "text": "done"},
|
|
{"kind": "data", "data": {"files": [{"path": "answer.json", "content": '{"ok": true}\n'}]}},
|
|
],
|
|
}
|
|
with _running_a2a_server(result) as (url, received):
|
|
bridge = AgentBeatsAcpBridge(endpoint_url=url, timeout_sec=5)
|
|
response, notifications = bridge.handle(
|
|
{
|
|
"jsonrpc": "2.0",
|
|
"id": 1,
|
|
"method": "initialize",
|
|
"params": {"protocolVersion": 1},
|
|
}
|
|
)
|
|
assert notifications == []
|
|
assert response is not None
|
|
assert response["result"]["agentInfo"]["name"] == AGENT_NAME
|
|
|
|
response, notifications = bridge.handle({"jsonrpc": "2.0", "id": 2, "method": "session/new", "params": {}})
|
|
assert notifications == []
|
|
assert response is not None
|
|
session_id = response["result"]["sessionId"]
|
|
|
|
monkeypatch.chdir(tmp_path)
|
|
response, notifications = bridge.handle(
|
|
{
|
|
"jsonrpc": "2.0",
|
|
"id": 3,
|
|
"method": "session/prompt",
|
|
"params": {
|
|
"sessionId": session_id,
|
|
"prompt": [{"type": "text", "text": "write the answer file"}],
|
|
},
|
|
}
|
|
)
|
|
|
|
assert response == {"jsonrpc": "2.0", "id": 3, "result": {"stopReason": "end_turn"}}
|
|
assert notifications == [
|
|
{
|
|
"jsonrpc": "2.0",
|
|
"method": "session/update",
|
|
"params": {
|
|
"update": {
|
|
"sessionUpdate": "agent_message_chunk",
|
|
"content": {"type": "text", "text": "done\nMaterialized 1 file(s): answer.json"},
|
|
}
|
|
},
|
|
}
|
|
]
|
|
assert json.loads((tmp_path / "answer.json").read_text()) == {"ok": True}
|
|
assert len(received) == 1
|
|
assert received[0]["path"] == "/"
|
|
assert received[0]["headers"]["Content-Type"] == "application/json"
|
|
payload = received[0]["payload"]
|
|
assert payload["method"] == "message/send"
|
|
assert payload["params"]["configuration"] == {"blocking": True}
|
|
message = payload["params"]["message"]
|
|
assert message["kind"] == "message"
|
|
assert message["role"] == "user"
|
|
assert message["parts"] == [{"kind": "text", "text": "write the answer file"}]
|
|
|
|
|
|
def test_acp_bridge_returns_error_when_endpoint_is_missing() -> None:
|
|
bridge = AgentBeatsAcpBridge(endpoint_url="", timeout_sec=5)
|
|
|
|
response, notifications = bridge.handle(
|
|
{
|
|
"jsonrpc": "2.0",
|
|
"id": 4,
|
|
"method": "session/prompt",
|
|
"params": {"prompt": [{"type": "text", "text": "hello"}]},
|
|
}
|
|
)
|
|
|
|
assert notifications == []
|
|
assert response is not None
|
|
assert response["id"] == 4
|
|
assert response["error"]["code"] == -32000
|
|
assert "SKILLSBENCH_A2A_ENDPOINT_URL is required" in response["error"]["message"]
|
|
|
|
|
|
def test_materialize_files_rejects_unsafe_paths_and_large_content(tmp_path: Path) -> None:
|
|
materialized = _materialize_files(
|
|
{
|
|
"files": [
|
|
{"path": "nested/answer.txt", "content": "ok"},
|
|
{"path": "../escape.txt", "content": "bad"},
|
|
{"path": "/abs.txt", "content": "bad"},
|
|
{"path": "too-large.txt", "content": "x" * (MAX_FILE_BYTES + 1)},
|
|
]
|
|
},
|
|
tmp_path,
|
|
)
|
|
|
|
assert materialized == ["nested/answer.txt"]
|
|
assert (tmp_path / "nested" / "answer.txt").read_text() == "ok"
|
|
assert not (tmp_path / "escape.txt").exists()
|
|
assert not (tmp_path / "abs.txt").exists()
|
|
assert not (tmp_path / "too-large.txt").exists()
|
|
|
|
|
|
@contextmanager
|
|
def _running_a2a_server(result: dict[str, Any]) -> Iterator[tuple[str, list[dict[str, Any]]]]:
|
|
server = _A2AServer(("127.0.0.1", 0), _A2AHandler)
|
|
server.result = result
|
|
server.received = []
|
|
thread = threading.Thread(target=server.serve_forever, daemon=True)
|
|
thread.start()
|
|
try:
|
|
yield f"http://127.0.0.1:{server.server_port}", server.received
|
|
finally:
|
|
server.shutdown()
|
|
server.server_close()
|
|
thread.join(timeout=5)
|