268 lines
9.4 KiBLFS
Python
268 lines
9.4 KiBLFS
Python
"""ACP-to-A2A bridge used by the AgentBeats BenchFlow worker."""
|
|
|
|
from __future__ import annotations
|
|
|
|
import base64
|
|
import json
|
|
import os
|
|
import sys
|
|
import urllib.error
|
|
import urllib.request
|
|
from pathlib import Path, PurePosixPath
|
|
from typing import Any
|
|
from uuid import uuid4
|
|
|
|
AGENT_NAME = "agentbeats-a2a"
|
|
BRIDGE_PATH = "/opt/skillsbench-agentbeats/bin/agentbeats-a2a"
|
|
LAUNCH_COMMAND = f"python3 {BRIDGE_PATH}"
|
|
ENDPOINT_ENV = "SKILLSBENCH_A2A_ENDPOINT_URL"
|
|
TIMEOUT_ENV = "SKILLSBENCH_A2A_TIMEOUT_SEC"
|
|
MAX_FILE_BYTES = 1_000_000
|
|
|
|
|
|
def register_agentbeats_a2a_agent() -> None:
|
|
"""Register the stdio ACP bridge as a BenchFlow agent at runtime."""
|
|
|
|
from benchflow.agents.registry import register_agent
|
|
|
|
register_agent(
|
|
name=AGENT_NAME,
|
|
install_cmd=install_command(),
|
|
launch_cmd=LAUNCH_COMMAND,
|
|
protocol="acp",
|
|
requires_env=[],
|
|
description="SkillsBench AgentBeats ACP bridge to an A2A participant endpoint.",
|
|
install_timeout=120,
|
|
supports_acp_set_model=False,
|
|
)
|
|
|
|
|
|
def install_command() -> str:
|
|
"""Return a shell command that installs this module as a sandbox executable."""
|
|
|
|
source = Path(__file__).read_text(encoding="utf-8")
|
|
encoded = base64.b64encode(source.encode("utf-8")).decode("ascii")
|
|
parent = str(Path(BRIDGE_PATH).parent)
|
|
return (
|
|
"( command -v python3 >/dev/null 2>&1 || "
|
|
"(apt-get update -qq && apt-get install -y -qq python3 >/dev/null 2>&1) ) && "
|
|
f"mkdir -p {parent} && "
|
|
f"printf '%s' {encoded!r} | base64 -d > {BRIDGE_PATH} && "
|
|
f"chmod +x {BRIDGE_PATH}"
|
|
)
|
|
|
|
|
|
class AgentBeatsAcpBridge:
|
|
def __init__(self, *, endpoint_url: str | None = None, timeout_sec: float | None = None) -> None:
|
|
self.endpoint_url = (endpoint_url or os.environ.get(ENDPOINT_ENV, "")).rstrip("/")
|
|
self.timeout_sec = timeout_sec if timeout_sec is not None else _env_float(TIMEOUT_ENV, 900.0)
|
|
self.sessions: set[str] = set()
|
|
|
|
def handle(self, message: dict[str, Any]) -> tuple[dict[str, Any] | None, list[dict[str, Any]]]:
|
|
method = message.get("method")
|
|
request_id = message.get("id")
|
|
if method == "initialize":
|
|
protocol_version = _protocol_version(message.get("params"))
|
|
return self._response(
|
|
request_id,
|
|
{
|
|
"protocolVersion": protocol_version,
|
|
"agentCapabilities": {
|
|
"loadSession": False,
|
|
"promptCapabilities": {"image": False, "audio": False, "embeddedContext": False},
|
|
},
|
|
"agentInfo": {"name": AGENT_NAME, "version": "0.1.0"},
|
|
"authMethods": [],
|
|
},
|
|
), []
|
|
if method == "session/new":
|
|
session_id = uuid4().hex
|
|
self.sessions.add(session_id)
|
|
return self._response(request_id, {"sessionId": session_id}), []
|
|
if method == "session/prompt":
|
|
try:
|
|
result = self._handle_prompt(message.get("params", {}))
|
|
except Exception as exc:
|
|
return self._error(request_id, -32000, str(exc)), []
|
|
text = _extract_text(result) or "A2A participant completed without text output."
|
|
files = _materialize_files(result, Path.cwd())
|
|
if files:
|
|
text = f"{text}\nMaterialized {len(files)} file(s): {', '.join(files)}"
|
|
notification = {
|
|
"jsonrpc": "2.0",
|
|
"method": "session/update",
|
|
"params": {"update": {"sessionUpdate": "agent_message_chunk", "content": {"type": "text", "text": text}}},
|
|
}
|
|
return self._response(request_id, {"stopReason": "end_turn"}), [notification]
|
|
if method in {"session/set_model", "session/set_config_option"}:
|
|
return self._response(request_id, {}), []
|
|
if method == "session/cancel":
|
|
return None, []
|
|
return self._error(request_id, -32601, f"Method not found: {method}"), []
|
|
|
|
def _handle_prompt(self, params: dict[str, Any]) -> dict[str, Any]:
|
|
if not self.endpoint_url:
|
|
raise RuntimeError(f"{ENDPOINT_ENV} is required")
|
|
prompt = _prompt_text(params)
|
|
payload = {
|
|
"jsonrpc": "2.0",
|
|
"id": uuid4().hex,
|
|
"method": "message/send",
|
|
"params": {
|
|
"message": {
|
|
"kind": "message",
|
|
"role": "user",
|
|
"messageId": uuid4().hex,
|
|
"parts": [{"kind": "text", "text": prompt}],
|
|
},
|
|
"configuration": {"blocking": True},
|
|
},
|
|
}
|
|
request = urllib.request.Request(
|
|
self.endpoint_url + "/",
|
|
data=json.dumps(payload).encode("utf-8"),
|
|
headers={"content-type": "application/json"},
|
|
method="POST",
|
|
)
|
|
try:
|
|
with urllib.request.urlopen(request, timeout=self.timeout_sec) as response:
|
|
body = response.read().decode("utf-8", errors="replace")
|
|
except urllib.error.URLError as exc:
|
|
raise RuntimeError(f"A2A request failed: {exc}") from exc
|
|
response_payload = json.loads(body)
|
|
if not isinstance(response_payload, dict):
|
|
raise RuntimeError("A2A response was not a JSON object")
|
|
if response_payload.get("error"):
|
|
raise RuntimeError(f"A2A error response: {response_payload['error']}")
|
|
result = response_payload.get("result")
|
|
if not isinstance(result, dict):
|
|
raise RuntimeError("A2A response missing object result")
|
|
return result
|
|
|
|
@staticmethod
|
|
def _response(request_id: Any, result: dict[str, Any]) -> dict[str, Any]:
|
|
return {"jsonrpc": "2.0", "id": request_id, "result": result}
|
|
|
|
@staticmethod
|
|
def _error(request_id: Any, code: int, message: str) -> dict[str, Any]:
|
|
return {"jsonrpc": "2.0", "id": request_id, "error": {"code": code, "message": message}}
|
|
|
|
|
|
def run_stdio() -> int:
|
|
bridge = AgentBeatsAcpBridge()
|
|
for raw_line in sys.stdin:
|
|
line = raw_line.strip()
|
|
if not line:
|
|
continue
|
|
try:
|
|
message = json.loads(line)
|
|
response, notifications = bridge.handle(message)
|
|
except Exception as exc:
|
|
response = {"jsonrpc": "2.0", "id": None, "error": {"code": -32700, "message": str(exc)}}
|
|
notifications = []
|
|
for notification in notifications:
|
|
_write_rpc(notification)
|
|
if response is not None:
|
|
_write_rpc(response)
|
|
return 0
|
|
|
|
|
|
def _write_rpc(message: dict[str, Any]) -> None:
|
|
sys.stdout.write(json.dumps(message, separators=(",", ":")) + "\n")
|
|
sys.stdout.flush()
|
|
|
|
|
|
def _protocol_version(params: Any) -> int:
|
|
if isinstance(params, dict) and isinstance(params.get("protocolVersion"), int):
|
|
return int(params["protocolVersion"])
|
|
return 1
|
|
|
|
|
|
def _prompt_text(params: Any) -> str:
|
|
if not isinstance(params, dict):
|
|
return ""
|
|
prompt = params.get("prompt")
|
|
if not isinstance(prompt, list):
|
|
return ""
|
|
chunks: list[str] = []
|
|
for part in prompt:
|
|
if isinstance(part, dict) and part.get("type") == "text" and isinstance(part.get("text"), str):
|
|
chunks.append(part["text"])
|
|
return "\n".join(chunks)
|
|
|
|
|
|
def _extract_text(value: Any) -> str:
|
|
chunks: list[str] = []
|
|
|
|
def visit(item: Any) -> None:
|
|
if isinstance(item, dict):
|
|
if item.get("kind") == "text" and isinstance(item.get("text"), str):
|
|
chunks.append(item["text"])
|
|
for nested in item.values():
|
|
visit(nested)
|
|
elif isinstance(item, list):
|
|
for nested in item:
|
|
visit(nested)
|
|
|
|
visit(value)
|
|
return "\n".join(chunk for chunk in chunks if chunk)
|
|
|
|
|
|
def _materialize_files(value: Any, cwd: Path) -> list[str]:
|
|
materialized: list[str] = []
|
|
for file_payload in _iter_file_payloads(value):
|
|
path_value = file_payload.get("path")
|
|
content = file_payload.get("content")
|
|
if not isinstance(path_value, str) or not isinstance(content, str):
|
|
continue
|
|
rel = _safe_relative_path(path_value)
|
|
if rel is None:
|
|
continue
|
|
data = content.encode("utf-8")
|
|
if len(data) > MAX_FILE_BYTES:
|
|
continue
|
|
target = cwd / rel
|
|
target.parent.mkdir(parents=True, exist_ok=True)
|
|
target.write_bytes(data)
|
|
materialized.append(rel.as_posix())
|
|
return materialized
|
|
|
|
|
|
def _iter_file_payloads(value: Any) -> list[dict[str, Any]]:
|
|
found: list[dict[str, Any]] = []
|
|
|
|
def visit(item: Any) -> None:
|
|
if isinstance(item, dict):
|
|
files = item.get("files")
|
|
if isinstance(files, list):
|
|
found.extend(file_item for file_item in files if isinstance(file_item, dict))
|
|
for nested in item.values():
|
|
visit(nested)
|
|
elif isinstance(item, list):
|
|
for nested in item:
|
|
visit(nested)
|
|
|
|
visit(value)
|
|
return found
|
|
|
|
|
|
def _safe_relative_path(value: str) -> Path | None:
|
|
pure = PurePosixPath(value)
|
|
if pure.is_absolute() or ".." in pure.parts or not pure.parts:
|
|
return None
|
|
return Path(pure.as_posix())
|
|
|
|
|
|
def _env_float(name: str, default: float) -> float:
|
|
value = os.environ.get(name)
|
|
if not value:
|
|
return default
|
|
try:
|
|
return float(value)
|
|
except ValueError:
|
|
return default
|
|
|
|
|
|
if __name__ == "__main__":
|
|
raise SystemExit(run_stdio())
|