Refactor our own client to use the fasta2a client, had no idea this existed

This commit is contained in:
Yiorgis Gozadinos 2025-10-15 11:54:58 +03:00
parent eca68cd30c
commit 509954a866
No known key found for this signature in database

View file

@ -7,9 +7,18 @@ from rich.console import Console
from rich.markdown import Markdown
from rich.prompt import Prompt
try:
from fasta2a.client import A2AClient as FastA2AClient
from fasta2a.schema import Message, TextPart
except ImportError as e:
raise ImportError(
"A2A support requires the 'a2a' extra. "
"Install with: uv pip install 'haiku.rag[a2a]'"
) from e
class A2AClient:
"""Simple A2A protocol client."""
"""Interactive A2A protocol client."""
def __init__(self, base_url: str = "http://localhost:8000"):
"""Initialize A2A client.
@ -18,11 +27,12 @@ class A2AClient:
base_url: Base URL of the A2A server
"""
self.base_url = base_url.rstrip("/")
self.client = httpx.AsyncClient(timeout=60.0)
http_client = httpx.AsyncClient(timeout=60.0)
self._client = FastA2AClient(base_url=base_url, http_client=http_client)
async def close(self):
"""Close the HTTP client."""
await self.client.aclose()
await self._client.http_client.aclose()
async def get_agent_card(self) -> dict[str, Any]:
"""Fetch the agent card from the A2A server.
@ -30,7 +40,9 @@ class A2AClient:
Returns:
Agent card dictionary with agent capabilities and metadata
"""
response = await self.client.get(f"{self.base_url}/.well-known/agent-card.json")
response = await self._client.http_client.get(
f"{self.base_url}/.well-known/agent-card.json"
)
response.raise_for_status()
return response.json()
@ -53,46 +65,38 @@ class A2AClient:
if context_id is None:
context_id = str(uuid.uuid4())
message_id = str(uuid.uuid4())
payload: dict[str, Any] = {
"jsonrpc": "2.0",
"method": "message/send",
"params": {
"contextId": context_id,
"message": {
"kind": "message",
"role": "user",
"messageId": message_id,
"parts": [{"kind": "text", "text": text}],
},
},
"id": 1,
}
if skill_id:
payload["params"]["skillId"] = skill_id
response = await self.client.post(
self.base_url,
json=payload,
headers={"Content-Type": "application/json"},
message = Message(
kind="message",
role="user",
message_id=str(uuid.uuid4()),
parts=[TextPart(kind="text", text=text)],
)
response.raise_for_status()
initial_response = response.json()
# Extract task ID from response
result = initial_response.get("result", {})
task_id = result.get("id")
metadata: dict[str, Any] = {"contextId": context_id}
if skill_id:
metadata["skillId"] = skill_id
if not task_id:
return initial_response
response = await self._client.send_message(message, metadata=metadata)
# Poll for task completion
return await self.wait_for_task(task_id)
if "error" in response:
return {"error": response["error"]}
result = response.get("result")
if not result:
return {"result": result}
# Result can be either Task or Message - check if it's a Task with an id
if result.get("kind") == "task":
task_id = result.get("id")
if task_id:
# Poll for task completion
return await self.wait_for_task(task_id)
# Return the message directly
return {"result": result}
async def wait_for_task(
self, task_id: str, max_wait: int = 60, poll_interval: float = 0.5
self, task_id: str, max_wait: int = 120, poll_interval: float = 0.5
) -> dict[str, Any]:
"""Poll for task completion.
@ -109,27 +113,19 @@ class A2AClient:
start_time = time.time()
while time.time() - start_time < max_wait:
payload = {
"jsonrpc": "2.0",
"method": "tasks/get",
"params": {"id": task_id},
"id": 2,
}
task_response = await self._client.get_task(task_id)
response = await self.client.post(
self.base_url,
json=payload,
headers={"Content-Type": "application/json"},
)
response.raise_for_status()
task = response.json()
if "error" in task_response:
return {"error": task_response["error"]}
result = task.get("result", {})
status = result.get("status", {})
state = status.get("state")
task = task_response.get("result")
if not task:
raise Exception("No task in response")
state = task.get("status", {}).get("state")
if state == "completed":
return task
return {"result": task}
elif state == "failed":
raise Exception(f"Task failed: {task}")