A2A worker for simple ask()
This commit is contained in:
parent
6f893b1626
commit
469352673e
2 changed files with 221 additions and 8 deletions
167
src/haiku/rag/a2a.py
Normal file
167
src/haiku/rag/a2a.py
Normal file
|
|
@ -0,0 +1,167 @@
|
||||||
|
import uuid
|
||||||
|
from contextlib import asynccontextmanager
|
||||||
|
from pathlib import Path
|
||||||
|
|
||||||
|
import logfire
|
||||||
|
|
||||||
|
from haiku.rag.client import HaikuRAG
|
||||||
|
from haiku.rag.config import Config
|
||||||
|
|
||||||
|
try:
|
||||||
|
from fasta2a import FastA2A, Worker # type: ignore
|
||||||
|
from fasta2a.broker import InMemoryBroker # type: ignore
|
||||||
|
from fasta2a.schema import ( # type: ignore
|
||||||
|
Artifact,
|
||||||
|
Message,
|
||||||
|
TaskIdParams,
|
||||||
|
TaskSendParams,
|
||||||
|
TextPart,
|
||||||
|
)
|
||||||
|
from fasta2a.storage import InMemoryStorage # type: ignore
|
||||||
|
except ImportError as e:
|
||||||
|
raise ImportError(
|
||||||
|
"A2A support requires the 'a2a' extra. "
|
||||||
|
"Install with: uv pip install 'haiku.rag[a2a]'"
|
||||||
|
) from e
|
||||||
|
|
||||||
|
logfire.configure(send_to_logfire="if-token-present", service_name="a2a")
|
||||||
|
logfire.instrument_pydantic_ai()
|
||||||
|
|
||||||
|
|
||||||
|
def create_qa_a2a_app(
|
||||||
|
db_path: Path,
|
||||||
|
deep: bool = False,
|
||||||
|
):
|
||||||
|
"""Create an A2A app for the QA agent.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
db_path: Path to the LanceDB database
|
||||||
|
deep: Use deep multi-agent QA for complex questions
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
A FastA2A ASGI application
|
||||||
|
"""
|
||||||
|
if deep:
|
||||||
|
raise NotImplementedError("Deep QA agent not yet implemented for A2A")
|
||||||
|
|
||||||
|
from haiku.rag.qa.agent import Dependencies, QuestionAnswerAgent
|
||||||
|
|
||||||
|
# Create the agent (client will be provided per-task in custom worker)
|
||||||
|
temp_client = HaikuRAG(db_path)
|
||||||
|
qa_agent = QuestionAnswerAgent(
|
||||||
|
client=temp_client,
|
||||||
|
provider=Config.QA_PROVIDER,
|
||||||
|
model=Config.QA_MODEL,
|
||||||
|
)
|
||||||
|
|
||||||
|
# Create custom worker using base Worker class
|
||||||
|
storage = InMemoryStorage()
|
||||||
|
broker = InMemoryBroker()
|
||||||
|
|
||||||
|
class QAWorker(Worker[list[Message]]):
|
||||||
|
async def run_task(self, params: TaskSendParams) -> None:
|
||||||
|
task = await self.storage.load_task(params["id"])
|
||||||
|
if task is None:
|
||||||
|
raise ValueError(f"Task {params['id']} not found")
|
||||||
|
|
||||||
|
if task["status"]["state"] != "submitted":
|
||||||
|
raise ValueError(
|
||||||
|
f"Task {params['id']} already processed: {task['status']['state']}"
|
||||||
|
)
|
||||||
|
|
||||||
|
await self.storage.update_task(task["id"], state="working")
|
||||||
|
|
||||||
|
# Load context and build simple message for agent
|
||||||
|
context = await self.storage.load_context(task["context_id"]) or []
|
||||||
|
context.extend(task.get("history", []))
|
||||||
|
|
||||||
|
# Extract the user's question from the latest message
|
||||||
|
user_messages = [
|
||||||
|
msg for msg in task.get("history", []) if msg["role"] == "user"
|
||||||
|
]
|
||||||
|
if not user_messages:
|
||||||
|
await self.storage.update_task(task["id"], state="failed")
|
||||||
|
return
|
||||||
|
|
||||||
|
last_user_msg = user_messages[-1]
|
||||||
|
question = ""
|
||||||
|
for part in last_user_msg.get("parts", []):
|
||||||
|
if part.get("kind") == "text":
|
||||||
|
question = part.get("text", "")
|
||||||
|
break
|
||||||
|
|
||||||
|
try:
|
||||||
|
# Create fresh client for this task and run QA agent
|
||||||
|
async with HaikuRAG(db_path) as client:
|
||||||
|
deps = Dependencies(client=client)
|
||||||
|
result = await qa_agent._agent.run(question, deps=deps)
|
||||||
|
|
||||||
|
# Build response message
|
||||||
|
response_message = Message(
|
||||||
|
role="agent",
|
||||||
|
parts=[TextPart(kind="text", text=str(result.output))],
|
||||||
|
kind="message",
|
||||||
|
message_id=str(uuid.uuid4()),
|
||||||
|
)
|
||||||
|
|
||||||
|
# Update context with new message
|
||||||
|
context.append(response_message)
|
||||||
|
await self.storage.update_context(task["context_id"], context)
|
||||||
|
|
||||||
|
# Build artifacts (optional)
|
||||||
|
artifacts = self.build_artifacts(result.output)
|
||||||
|
|
||||||
|
await self.storage.update_task(
|
||||||
|
task["id"],
|
||||||
|
state="completed",
|
||||||
|
new_messages=[response_message],
|
||||||
|
new_artifacts=artifacts,
|
||||||
|
)
|
||||||
|
except Exception:
|
||||||
|
await self.storage.update_task(task["id"], state="failed")
|
||||||
|
raise
|
||||||
|
|
||||||
|
async def cancel_task(self, params: TaskIdParams) -> None:
|
||||||
|
pass
|
||||||
|
|
||||||
|
def build_message_history(self, history: list[Message]) -> list[Message]:
|
||||||
|
return history
|
||||||
|
|
||||||
|
def build_artifacts(self, result: str) -> list[Artifact]:
|
||||||
|
# Simple artifact with the result text
|
||||||
|
return [
|
||||||
|
Artifact(
|
||||||
|
artifact_id=str(uuid.uuid4()),
|
||||||
|
name="result",
|
||||||
|
parts=[TextPart(kind="text", text=result)],
|
||||||
|
)
|
||||||
|
]
|
||||||
|
|
||||||
|
worker = QAWorker(storage=storage, broker=broker)
|
||||||
|
|
||||||
|
# Create FastA2A app with custom worker lifecycle
|
||||||
|
@asynccontextmanager
|
||||||
|
async def lifespan(app):
|
||||||
|
async with app.task_manager:
|
||||||
|
async with worker.run():
|
||||||
|
yield
|
||||||
|
|
||||||
|
return FastA2A(
|
||||||
|
storage=storage,
|
||||||
|
broker=broker,
|
||||||
|
name="haiku-rag-qa",
|
||||||
|
description="Question answering agent powered by haiku.rag RAG system",
|
||||||
|
lifespan=lifespan,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def create_research_a2a_app(db_path: Path):
|
||||||
|
"""Create an A2A app for the research agent.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
db_path: Path to the LanceDB database
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
A FastA2A ASGI application
|
||||||
|
"""
|
||||||
|
raise NotImplementedError("Research agent not yet implemented for A2A")
|
||||||
|
|
@ -366,7 +366,7 @@ def download_models_cmd():
|
||||||
|
|
||||||
|
|
||||||
@cli.command(
|
@cli.command(
|
||||||
"serve", help="Start the haiku.rag MCP server (by default in streamable HTTP mode)"
|
"serve", help="Start the haiku.rag server (MCP by default, or A2A with --a2a)"
|
||||||
)
|
)
|
||||||
def serve(
|
def serve(
|
||||||
db: Path = typer.Option(
|
db: Path = typer.Option(
|
||||||
|
|
@ -379,17 +379,63 @@ def serve(
|
||||||
"--stdio",
|
"--stdio",
|
||||||
help="Run MCP server on stdio Transport",
|
help="Run MCP server on stdio Transport",
|
||||||
),
|
),
|
||||||
|
a2a: bool = typer.Option(
|
||||||
|
False,
|
||||||
|
"--a2a",
|
||||||
|
help="Run A2A (Agent-to-Agent) server instead of MCP",
|
||||||
|
),
|
||||||
|
a2a_agent: str = typer.Option(
|
||||||
|
"qa",
|
||||||
|
"--a2a-agent",
|
||||||
|
help="Which agent to serve via A2A: 'qa', 'qa-deep', or 'research'",
|
||||||
|
),
|
||||||
|
a2a_host: str = typer.Option(
|
||||||
|
"127.0.0.1",
|
||||||
|
"--a2a-host",
|
||||||
|
help="Host to bind A2A server to",
|
||||||
|
),
|
||||||
|
a2a_port: int = typer.Option(
|
||||||
|
8000,
|
||||||
|
"--a2a-port",
|
||||||
|
help="Port to bind A2A server to",
|
||||||
|
),
|
||||||
) -> None:
|
) -> None:
|
||||||
"""Start the MCP server."""
|
"""Start the MCP or A2A server."""
|
||||||
from haiku.rag.app import HaikuRAGApp
|
if a2a:
|
||||||
|
try:
|
||||||
|
from haiku.rag.a2a import create_qa_a2a_app, create_research_a2a_app
|
||||||
|
except ImportError as e:
|
||||||
|
typer.echo(f"Error: {e}")
|
||||||
|
raise typer.Exit(1)
|
||||||
|
|
||||||
app = HaikuRAGApp(db_path=db)
|
import uvicorn
|
||||||
|
|
||||||
transport = None
|
if a2a_agent == "qa":
|
||||||
if stdio:
|
typer.echo(f"Starting QA agent A2A server on {a2a_host}:{a2a_port}")
|
||||||
transport = "stdio"
|
app = create_qa_a2a_app(db_path=db, deep=False)
|
||||||
|
elif a2a_agent == "qa-deep":
|
||||||
|
typer.echo(f"Starting deep QA agent A2A server on {a2a_host}:{a2a_port}")
|
||||||
|
app = create_qa_a2a_app(db_path=db, deep=True)
|
||||||
|
elif a2a_agent == "research":
|
||||||
|
typer.echo(f"Starting research agent A2A server on {a2a_host}:{a2a_port}")
|
||||||
|
app = create_research_a2a_app(db_path=db)
|
||||||
|
else:
|
||||||
|
typer.echo(
|
||||||
|
f"Error: Unknown agent type '{a2a_agent}'. Use 'qa', 'qa-deep', or 'research'"
|
||||||
|
)
|
||||||
|
raise typer.Exit(1)
|
||||||
|
|
||||||
asyncio.run(app.serve(transport=transport))
|
uvicorn.run(app, host=a2a_host, port=a2a_port)
|
||||||
|
else:
|
||||||
|
from haiku.rag.app import HaikuRAGApp
|
||||||
|
|
||||||
|
app = HaikuRAGApp(db_path=db)
|
||||||
|
|
||||||
|
transport = None
|
||||||
|
if stdio:
|
||||||
|
transport = "stdio"
|
||||||
|
|
||||||
|
asyncio.run(app.serve(transport=transport))
|
||||||
|
|
||||||
|
|
||||||
@cli.command("migrate", help="Migrate an SQLite database to LanceDB")
|
@cli.command("migrate", help="Migrate an SQLite database to LanceDB")
|
||||||
|
|
|
||||||
Loading…
Reference in a new issue