Merge pull request #95 from ggozad/feat/a2a

Full implementation of A2A protocol for haiku.rag
This commit is contained in:
Yiorgis Gozadinos 2025-10-14 13:26:50 +03:00 committed by GitHub
commit f1d6700bfa
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
28 changed files with 2983 additions and 59 deletions

View file

@ -18,6 +18,7 @@ Retrieval-Augmented Generation (RAG) library built on LanceDB.
- **File monitoring**: Auto-index files when run as server
- **40+ file formats**: PDF, DOCX, HTML, Markdown, code files, URLs
- **MCP server**: Expose as tools for AI assistants
- **A2A agent**: Conversational agent with context and multi-turn dialogue
- **CLI & Python API**: Use from command line or Python
## Quick Start
@ -143,6 +144,24 @@ haiku-rag serve --stdio
Provides tools for document management and search directly in your AI assistant.
## A2A Agent
Run as a conversational agent with the Agent-to-Agent protocol:
```bash
# Start the A2A server
haiku-rag serve --a2a
# Connect with the interactive client (in another terminal)
haiku-rag a2aclient
```
The A2A agent provides:
- Multi-turn dialogue with context
- Intelligent multi-search for complex questions
- Source citations with titles and URIs
- Full document retrieval on request
## Documentation
Full documentation at: https://ggozad.github.io/haiku.rag/

190
docs/a2a.md Normal file
View file

@ -0,0 +1,190 @@
# Agent-to-Agent (A2A) Protocol
The A2A server exposes `haiku.rag` as a conversational agent using the Agent-to-Agent protocol. Unlike the MCP server which provides stateless tools, the A2A agent maintains conversation history and context across multiple turns.
## Features
- **Conversational Context**: Maintains full conversation history including tool calls and results
- **Multi-turn Dialogue**: Supports follow-up questions with pronoun resolution ("he", "it", "that document")
- **Intelligent Search**: Performs single or multiple searches depending on question complexity
- **Source Citations**: Always includes sources with both titles and URIs
- **Full Document Retrieval**: Can fetch complete documents on request
- **Multiple Skills**: Exposes three distinct skills with appropriate artifacts:
- `document-qa`: Conversational question answering (default)
- `document-search`: Semantic search with structured results
- `document-retrieve`: Fetch complete documents by URI
## Starting A2A Server
```bash
haiku-rag serve --a2a
```
Server options:
- `--a2a-host` - Host to bind to (default: 127.0.0.1)
- `--a2a-port` - Port to bind to (default: 8000)
Example:
```bash
haiku-rag serve --a2a --a2a-host 0.0.0.0 --a2a-port 8080
```
## Interactive A2A Client
Test and interact with haiku.rag's A2A server using the built-in interactive client:
```bash
haiku-rag a2aclient
```
Client options:
- `--url` - Base URL of the A2A server (default: http://localhost:8000)
Example:
```bash
# Connect to local server
haiku-rag a2aclient
# Connect to remote server
haiku-rag a2aclient --url https://example.com:8000
```
The interactive client provides:
- Rich markdown rendering of agent responses
- Conversation context across multiple turns
- Agent card discovery and display
- Compact artifact summaries
## Requirements
A2A support requires the `a2a` extra:
```bash
uv pip install 'haiku.rag[a2a]'
```
## Python Usage
```python
from pathlib import Path
from haiku.rag.a2a import create_a2a_app
import uvicorn
# Create A2A app
app = create_a2a_app(Path("database.lancedb"))
# Run with uvicorn
uvicorn.run(app, host="127.0.0.1", port=8000)
```
This installs the `fasta2a` package and its dependencies.
## Architecture
The A2A agent uses:
- **FastA2A**: Python framework implementing the A2A protocol
- **Pydantic AI**: Agent framework with tool support
- **In-Memory Storage**: Context and message history storage (persists during server lifetime)
- **Conversation State**: Full pydantic-ai message history serialized in A2A context
### Message History
The agent stores the complete conversation state including:
- User prompts
- Agent responses
- Tool calls and their arguments
- Tool return values
This enables the agent to:
- Reference previous searches
- Understand pronouns and context
- Maintain coherent multi-turn conversations
### Context Management
Each conversation is identified by a `context_id`. All messages within the same context share conversation history. This allows the agent to:
- Remember what was discussed
- Track which documents were already found
- Provide contextual follow-up answers
### Skills
The agent exposes three skills:
- **document-qa** (default): Conversational question answering including follow-ups and multi-turn dialogue
- **document-search**: Direct semantic search returning formatted results
- **document-retrieve**: Fetch complete document content by URI
### Artifacts
All operations create artifacts for traceability:
- **search_results**: Created for each `search_documents` tool call
- Contains query and array of SearchResult objects (content, score, document_title, document_uri)
- **document**: Created for each `get_full_document` tool call
- Contains complete document text
- **qa_result**: Created for all document-qa operations
- Contains question, answer, and skill identifier
- Always created for Q&A, even when answering from conversation history without tools
### Memory Management
To prevent memory growth, the server uses LRU (Least Recently Used) eviction:
- Maximum 1000 contexts kept in memory (configurable via `A2A_MAX_CONTEXTS`)
- When limit exceeded, least recently used contexts are automatically evicted
Configure via environment variable:
```bash
export A2A_MAX_CONTEXTS=1000
```
## Security
By default, the A2A agent runs without authentication. For production deployments, you should add authentication.
### Adding Authentication
The `create_a2a_app()` function accepts optional security parameters that declare authentication requirements in the agent card:
```python
from haiku.rag.a2a import create_a2a_app
app = create_a2a_app(
db_path,
security_schemes={
"apiKeyAuth": {
"type": "apiKey",
"in": "header",
"name": "X-API-Key",
"description": "API key authentication",
}
},
security=[{"apiKeyAuth": []}],
)
```
This populates the agent card at `/.well-known/agent-card.json` so other agents can discover your authentication requirements.
### Security Examples
Three working examples are provided in `examples/a2a-security/`:
1. **API Key** (`apikey_example.py`) - Simple header-based authentication
2. **OAuth2 GitHub** (`oauth2_github.py`) - GitHub Personal Access Token authentication
3. **OAuth2 Enterprise** (`oauth2_example.py`) - Full OAuth2 with JWT verification
Each example shows:
- How to declare security in the agent card
- How to implement authentication middleware
- How to verify credentials

View file

@ -125,15 +125,49 @@ When `--verbose` is set the CLI also consumes the internal research stream, prin
## Server
Start the MCP server:
Start services (requires at least one flag):
```bash
# HTTP transport (default)
haiku-rag serve
# MCP server only (HTTP transport)
haiku-rag serve --mcp
# stdio transport
haiku-rag serve --stdio
# MCP server (stdio transport)
haiku-rag serve --mcp --stdio
# A2A server only
haiku-rag serve --a2a
# File monitoring only
haiku-rag serve --monitor
# All services
haiku-rag serve --monitor --mcp --a2a
# Custom ports
haiku-rag serve --mcp --mcp-port 9000 --a2a --a2a-port 9001
```
See [Server Mode](server.md) for details on available services.
### A2A Interactive Client
Connect to and chat with haiku.rag's A2A server:
```bash
# Connect to local server
haiku-rag a2aclient
# Connect to remote server
haiku-rag a2aclient --url https://example.com:8000
```
The interactive client provides:
- Rich markdown rendering of agent responses
- Multi-turn conversation with context
- Agent card discovery and display
- Compact artifact summaries
See [A2A documentation](a2a.md) for more details.
## Settings
View current configuration settings:

View file

@ -73,7 +73,7 @@ Configure which LLM provider to use for question answering. Any provider and mod
```bash
QA_PROVIDER="ollama"
QA_MODEL="qwen3"
QA_MODEL="gpt-oss"
OLLAMA_BASE_URL="http://localhost:11434"
```

View file

@ -10,10 +10,11 @@
- **Support for various embedding providers**: Ollama, VoyageAI, OpenAI or add your own
- **Native Hybrid Search**: Vector search combined with full-text search using native LanceDB RRF reranking
- **Reranking**: Optional result reranking with MixedBread AI or Cohere
- **Question Answering**: Built-in QA agents using Ollama, OpenAI, or Anthropic.
- **Question Answering**: Built-in QA agents using Ollama, OpenAI, or Anthropic
- **File monitoring**: Automatically index files when run as a server
- **Extended file format support**: Parse 40+ file formats including PDF, DOCX, HTML, Markdown, code files and more. Or add a URL!
- **MCP server**: Exposes functionality as MCP tools
- **A2A agent**: Conversational agent with context and multi-turn dialogue support
- **CLI commands**: Access all functionality from your terminal
- Add sources from text, files, or URLs, optionally with a humanreadable title
- **Python client**: Call `haiku.rag` from your own python applications
@ -57,6 +58,7 @@ haiku-rag migrate old_database.sqlite # Migrate from SQLite
- [CLI](cli.md) - Command line interface usage
- [Server](server.md) - File monitoring and server mode
- [MCP](mcp.md) - Model Context Protocol integration
- [A2A](a2a.md) - Agent-to-Agent conversational protocol
- [Python](python.md) - Python API reference
- [Agents](agents.md) - QA agent and multi-agent research

View file

@ -1,16 +1,47 @@
# Server Mode
The server provides automatic file monitoring and MCP functionality.
The server provides automatic file monitoring, MCP functionality, and A2A agent support.
## Starting the Server
The `serve` command requires at least one service flag. You can enable file monitoring, MCP server, A2A server, or any combination:
### MCP Server Only
```bash
haiku-rag serve
haiku-rag serve --mcp
```
Transport options:
- Default - Streamable HTTP transport
- Default - Streamable HTTP transport on port 8001
- `--stdio` - Standard input/output transport
- `--mcp-port` - Custom port (default: 8001)
### A2A Server Only
```bash
haiku-rag serve --a2a
```
Options:
- `--a2a-host` - Host to bind to (default: 127.0.0.1)
- `--a2a-port` - Port to bind to (default: 8000)
See [A2A documentation](a2a.md) for details on the conversational agent.
### File Monitoring Only
```bash
haiku-rag serve --monitor
```
### All Services
```bash
haiku-rag serve --monitor --mcp --a2a
```
This will start file monitoring, MCP server on port 8001, and A2A server on port 8000.
## File Monitoring

41
examples/README.md Normal file
View file

@ -0,0 +1,41 @@
# haiku.rag Examples
This directory contains example scripts demonstrating various features of haiku.rag.
## A2A Security Examples
**Directory:** `a2a-security/`
Three examples showing how to add authentication to haiku.rag's A2A server:
### API Key Authentication
**File:** `a2a-security/apikey_example.py`
Simple header-based authentication suitable for internal services and development.
```bash
python examples/a2a-security/apikey_example.py /path/to/database.lancedb
```
### OAuth2 GitHub Authentication
**File:** `a2a-security/oauth2_github.py`
GitHub Personal Access Token authentication for GitHub-integrated services.
```bash
python examples/a2a-security/oauth2_github.py /path/to/database.lancedb
```
### OAuth2 Enterprise Authentication
**File:** `a2a-security/oauth2_example.py`
Full OAuth2 with JWT verification for enterprise environments.
```bash
python examples/a2a-security/oauth2_example.py /path/to/database.lancedb
```
See individual files for detailed setup instructions and usage examples.

View file

@ -0,0 +1,130 @@
"""Example: Adding API Key authentication to haiku.rag A2A agent.
Simple header-based authentication suitable for internal services and development.
Perfect for getting started with A2A authentication.
Setup:
# Run with default key
python apikey_example.py /path/to/database.lancedb
# Or use your own key
export API_KEY='your-secret-key'
python apikey_example.py /path/to/database.lancedb
Usage:
# Make authenticated request (default key is demo-key-12345)
curl -H "X-API-Key: demo-key-12345" \
-H "Content-Type: application/json" \
-X POST http://localhost:8000/ \
-d '{"jsonrpc":"2.0","method":"message/send","params":{"contextId":"test","message":{"kind":"message","role":"user","messageId":"msg-1","parts":[{"kind":"text","text":"What is Python?"}]}},"id":1}'
"""
import os
from pathlib import Path
from starlette.exceptions import HTTPException
from starlette.responses import JSONResponse
from starlette.status import HTTP_401_UNAUTHORIZED
from haiku.rag.a2a import create_a2a_app
# API Key Configuration - In production, use environment variables or a secure key store
API_KEY_NAME = "X-API-Key"
VALID_API_KEY = os.getenv("API_KEY", "demo-key-12345")
def verify_api_key(api_key: str | None) -> str:
"""Verify API key from request header.
Args:
api_key: API key from X-API-Key header
Returns:
The verified API key
Raises:
HTTPException: If API key is missing or invalid
"""
if not api_key:
raise HTTPException(
status_code=HTTP_401_UNAUTHORIZED,
detail="Missing API key",
headers={"WWW-Authenticate": f'ApiKey realm="{API_KEY_NAME}"'},
)
if api_key != VALID_API_KEY:
raise HTTPException(
status_code=HTTP_401_UNAUTHORIZED,
detail="Invalid API key",
headers={"WWW-Authenticate": f'ApiKey realm="{API_KEY_NAME}"'},
)
return api_key
def create_secure_a2a_app(db_path: Path):
"""Create A2A app with API key authentication.
Args:
db_path: Path to LanceDB database
Returns:
FastA2A application with API key security
"""
# Create app with security declared in AgentCard
app = create_a2a_app(
db_path,
security_schemes={
"apiKeyAuth": {
"type": "apiKey",
"in": "header",
"name": API_KEY_NAME,
"description": "API key authentication",
}
},
security=[{"apiKeyAuth": []}],
)
# Add authentication middleware
@app.middleware("http")
async def authenticate_request(request, call_next):
"""Middleware to verify API key on all requests."""
# Skip authentication for well-known endpoints
if request.url.path in [
"/.well-known/agent-card.json",
"/health",
"/docs",
"/openapi.json",
]:
return await call_next(request)
# Verify API key
api_key = request.headers.get(API_KEY_NAME)
try:
verify_api_key(api_key)
except HTTPException as e:
return JSONResponse(
status_code=e.status_code,
content={"detail": e.detail},
headers=e.headers or {},
)
# Continue with request
return await call_next(request)
return app
if __name__ == "__main__":
import sys
import uvicorn
if len(sys.argv) < 2:
print("Usage: python apikey_example.py <path-to-database.lancedb>")
sys.exit(1)
db_path = Path(sys.argv[1])
app = create_secure_a2a_app(db_path)
uvicorn.run(app, host="127.0.0.1", port=8000)

View file

@ -0,0 +1,222 @@
"""Example: Adding OAuth2 authentication to haiku.rag A2A agent.
This example demonstrates OAuth2 client credentials flow with JWT token verification.
Suitable for enterprise environments with existing OAuth2 infrastructure.
Requirements:
uv pip install python-jose[cryptography]
Setup:
1. Set up an OAuth2 provider (Auth0, Okta, Azure AD, Keycloak, etc.)
2. Create an API and a machine-to-machine application
3. Get the token URL and public key from your provider
4. Set environment variables:
export OAUTH2_TOKEN_URL='https://your-auth.example.com/oauth/token'
export OAUTH2_PUBLIC_KEY='-----BEGIN PUBLIC KEY-----...'
Usage:
python oauth2_example.py /path/to/database.lancedb
# Get access token from your OAuth2 provider:
TOKEN=$(curl -X POST $OAUTH2_TOKEN_URL \
-d "grant_type=client_credentials" \
-d "client_id=your-client-id" \
-d "client_secret=your-client-secret" \
-d "scope=read:documents query:documents" \
| jq -r '.access_token')
# Make authenticated request:
curl -H "Authorization: Bearer $TOKEN" \
-H "Content-Type: application/json" \
-X POST http://localhost:8000/ \
-d '{"jsonrpc":"2.0","method":"message/send","params":{"contextId":"test","message":{"kind":"message","role":"user","messageId":"msg-1","parts":[{"kind":"text","text":"What is Python?"}]}},"id":1}'
"""
import os
from pathlib import Path
from jose import JWTError, jwt
from starlette.exceptions import HTTPException
from starlette.responses import JSONResponse
from starlette.status import (
HTTP_401_UNAUTHORIZED,
HTTP_403_FORBIDDEN,
HTTP_500_INTERNAL_SERVER_ERROR,
)
from haiku.rag.a2a import create_a2a_app
# OAuth2 Configuration
OAUTH2_TOKEN_URL = os.getenv(
"OAUTH2_TOKEN_URL", "https://your-auth.example.com/oauth/token"
)
OAUTH2_AUTH_URL = os.getenv(
"OAUTH2_AUTH_URL", "https://your-auth.example.com/oauth/authorize"
)
OAUTH2_PUBLIC_KEY = os.getenv("OAUTH2_PUBLIC_KEY", "")
OAUTH2_ALGORITHM = os.getenv("OAUTH2_ALGORITHM", "RS256")
# Define required scopes for each skill
SKILL_SCOPES = {
"document-qa": ["read:documents", "query:documents"],
}
def verify_token(token: str) -> dict:
"""Verify JWT token from OAuth2 provider.
Args:
token: JWT token from Authorization header
Returns:
Dictionary with user info and scopes
Raises:
HTTPException: If token is invalid or expired
"""
credentials_exception = HTTPException(
status_code=HTTP_401_UNAUTHORIZED,
detail="Could not validate credentials",
headers={"WWW-Authenticate": "Bearer"},
)
if not OAUTH2_PUBLIC_KEY:
raise HTTPException(
status_code=HTTP_500_INTERNAL_SERVER_ERROR,
detail="OAuth2 public key not configured",
)
try:
payload = jwt.decode(
token,
OAUTH2_PUBLIC_KEY,
algorithms=[OAUTH2_ALGORITHM],
)
username: str | None = payload.get("sub")
scopes: list[str] = (
payload.get("scope", "").split()
if isinstance(payload.get("scope"), str)
else payload.get("scope", [])
)
if username is None:
raise credentials_exception
return {"username": username, "scopes": scopes}
except JWTError as e:
raise HTTPException(
status_code=HTTP_401_UNAUTHORIZED,
detail=f"Invalid token: {str(e)}",
headers={"WWW-Authenticate": "Bearer"},
) from e
def check_skill_permissions(skill_id: str, credentials: dict) -> None:
"""Verify that user has required scopes for a skill.
Args:
skill_id: The skill being accessed
credentials: User credentials with scopes
Raises:
HTTPException: If user lacks required permissions
"""
required_scopes = SKILL_SCOPES.get(skill_id, [])
user_scopes = credentials.get("scopes", [])
missing_scopes = [scope for scope in required_scopes if scope not in user_scopes]
if missing_scopes:
raise HTTPException(
status_code=HTTP_403_FORBIDDEN,
detail=f"Missing required scopes: {', '.join(missing_scopes)} for skill: {skill_id}",
)
def create_secure_a2a_app(db_path: Path):
"""Create A2A app with OAuth2 authentication.
Args:
db_path: Path to LanceDB database
Returns:
FastA2A application with OAuth2 security
"""
# Create app with security declared in AgentCard
app = create_a2a_app(
db_path,
security_schemes={
"oauth2": {
"type": "oauth2",
"flows": {
"clientCredentials": {
"tokenUrl": OAUTH2_TOKEN_URL,
"scopes": {
"read:documents": "Read document content",
"query:documents": "Search and query documents",
},
}
},
"description": "OAuth2 client credentials flow",
}
},
security=[{"oauth2": ["read:documents", "query:documents"]}],
)
# Add authentication middleware
@app.middleware("http")
async def authenticate_request(request, call_next):
"""Middleware to verify OAuth2 token on all requests."""
# Skip authentication for well-known endpoints
if request.url.path in [
"/.well-known/agent-card.json",
"/health",
"/docs",
"/openapi.json",
]:
return await call_next(request)
# Get token from Authorization header
auth_header = request.headers.get("Authorization", "")
if not auth_header.startswith("Bearer "):
return JSONResponse(
status_code=HTTP_401_UNAUTHORIZED,
content={"detail": "Missing or invalid Authorization header"},
headers={"WWW-Authenticate": "Bearer"},
)
token = auth_header[7:] # Remove "Bearer " prefix
# Verify token
try:
credentials = verify_token(token)
# Attach credentials to request state for use in handlers
request.state.credentials = credentials
except HTTPException as e:
return JSONResponse(
status_code=e.status_code,
content={"detail": e.detail},
headers=e.headers or {},
)
# Continue with request
return await call_next(request)
return app
if __name__ == "__main__":
import sys
import uvicorn
if len(sys.argv) < 2:
print("Usage: python oauth2_example.py <path-to-database.lancedb>")
sys.exit(1)
db_path = Path(sys.argv[1])
app = create_secure_a2a_app(db_path)
uvicorn.run(app, host="127.0.0.1", port=8000)

View file

@ -0,0 +1,193 @@
"""Example: Using GitHub Personal Access Tokens for authentication.
This is a simplified OAuth2 example that uses GitHub Personal Access Tokens.
It's much easier to set up than full OAuth2 and perfect for testing.
Setup:
1. Go to https://github.com/settings/tokens
2. Click "Generate new token (classic)"
3. Give it a name and select scopes
4. Copy the generated token
Usage:
export GITHUB_TOKENS="your_github_token_here"
python oauth2_github.py /path/to/database.lancedb
# Make authenticated request:
curl -H "Authorization: Bearer ghp_your_token" \
-H "Content-Type: application/json" \
-X POST http://localhost:8000/ \
-d '{"jsonrpc":"2.0","method":"message/send","params":{"contextId":"test","message":{"kind":"message","role":"user","messageId":"msg-1","parts":[{"kind":"text","text":"What is Python?"}]}},"id":1}'
"""
import os
from pathlib import Path
import httpx
from starlette.exceptions import HTTPException
from starlette.responses import JSONResponse
from starlette.status import HTTP_401_UNAUTHORIZED
from haiku.rag.a2a import create_a2a_app
# Configuration
GITHUB_API_URL = "https://api.github.com"
ALLOWED_TOKENS = (
set(os.getenv("GITHUB_TOKENS", "").split(","))
if os.getenv("GITHUB_TOKENS")
else set()
)
async def verify_github_token(token: str) -> dict:
"""Verify GitHub Personal Access Token by calling GitHub API.
Args:
token: GitHub Personal Access Token (starts with ghp_)
Returns:
Dictionary with user info
Raises:
HTTPException: If token is invalid
"""
if not token.startswith("ghp_") and not token.startswith("github_pat_"):
raise HTTPException(
status_code=HTTP_401_UNAUTHORIZED,
detail="Invalid GitHub token format",
headers={"WWW-Authenticate": "Bearer"},
)
# If we have a list of allowed tokens, check against it
if ALLOWED_TOKENS and token not in ALLOWED_TOKENS:
raise HTTPException(
status_code=HTTP_401_UNAUTHORIZED,
detail="Token not in allowed list",
headers={"WWW-Authenticate": "Bearer"},
)
# Verify token with GitHub API
async with httpx.AsyncClient() as client:
try:
response = await client.get(
f"{GITHUB_API_URL}/user",
headers={
"Authorization": f"Bearer {token}",
"Accept": "application/vnd.github+json",
"X-GitHub-Api-Version": "2022-11-28",
},
timeout=5.0,
)
if response.status_code == 401:
raise HTTPException(
status_code=HTTP_401_UNAUTHORIZED,
detail="Invalid or expired GitHub token",
headers={"WWW-Authenticate": "Bearer"},
)
if response.status_code != 200:
raise HTTPException(
status_code=HTTP_401_UNAUTHORIZED,
detail=f"GitHub API error: {response.status_code}",
headers={"WWW-Authenticate": "Bearer"},
)
user_data = response.json()
return {
"username": user_data.get("login"),
"email": user_data.get("email"),
"name": user_data.get("name"),
"github_id": user_data.get("id"),
}
except httpx.TimeoutException:
raise HTTPException(
status_code=HTTP_401_UNAUTHORIZED,
detail="GitHub API timeout",
headers={"WWW-Authenticate": "Bearer"},
)
except httpx.RequestError as e:
raise HTTPException(
status_code=HTTP_401_UNAUTHORIZED,
detail=f"Failed to verify token with GitHub: {str(e)}",
headers={"WWW-Authenticate": "Bearer"},
)
def create_secure_a2a_app(db_path: Path):
"""Create A2A app with GitHub token authentication.
Args:
db_path: Path to LanceDB database
Returns:
FastA2A application with GitHub authentication
"""
# Create app with security declared in AgentCard
app = create_a2a_app(
db_path,
security_schemes={
"githubAuth": {
"type": "http",
"scheme": "bearer",
"description": "GitHub Personal Access Token authentication",
}
},
security=[{"githubAuth": []}],
)
@app.middleware("http")
async def authenticate_request(request, call_next):
"""Middleware to verify GitHub token on all requests."""
# Skip authentication for well-known endpoints
if request.url.path in [
"/.well-known/agent-card.json",
"/health",
"/docs",
"/openapi.json",
]:
return await call_next(request)
# Get token from Authorization header
auth_header = request.headers.get("Authorization", "")
if not auth_header.startswith("Bearer "):
return JSONResponse(
status_code=HTTP_401_UNAUTHORIZED,
content={"detail": "Missing or invalid Authorization header"},
headers={"WWW-Authenticate": 'Bearer realm="GitHub"'},
)
token = auth_header[7:] # Remove "Bearer " prefix
# Verify token
try:
user_data = await verify_github_token(token)
# Attach user data to request state
request.state.user = user_data
except HTTPException as e:
return JSONResponse(
status_code=e.status_code,
content={"detail": e.detail},
headers=e.headers or {},
)
# Continue with request
return await call_next(request)
return app
if __name__ == "__main__":
import sys
import uvicorn
if len(sys.argv) < 2:
print("Usage: python oauth2_github.py <path-to-database.lancedb>")
sys.exit(1)
db_path = Path(sys.argv[1])
app = create_secure_a2a_app(db_path)
uvicorn.run(app, host="127.0.0.1", port=8000)

View file

@ -64,6 +64,7 @@ nav:
- Agents: agents.md
- Python: python.md
- MCP: mcp.md
- A2A: a2a.md
- Benchmarks: benchmarks.md
markdown_extensions:
- admonition

View file

@ -40,6 +40,7 @@ dependencies = [
[project.optional-dependencies]
voyageai = ["voyageai>=0.3.5"]
mxbai = ["mxbai-rerank>=0.1.6"]
a2a = ["fasta2a>=0.1.0"]
[project.scripts]
haiku-rag = "haiku.rag.cli:cli"
@ -49,7 +50,7 @@ requires = ["hatchling"]
build-backend = "hatchling.build"
[tool.hatch.build]
exclude = ["/docs", "/tests", "/.github"]
exclude = ["/docs", "/examples", "/tests", "/.github"]
[tool.hatch.build.targets.wheel]
packages = ["src/haiku"]

View file

@ -0,0 +1,176 @@
import logging
from contextlib import asynccontextmanager
from pathlib import Path
import logfire
from pydantic_ai import Agent, RunContext
from haiku.rag.config import Config
from haiku.rag.graph.common import get_model
from .context import load_message_history, save_message_history
from .models import AgentDependencies, SearchResult
from .prompts import A2A_SYSTEM_PROMPT
from .skills import extract_question_from_task, get_agent_skills
from .storage import LRUMemoryStorage
from .worker import ConversationalWorker
try:
from fasta2a import FastA2A # type: ignore
from fasta2a.broker import InMemoryBroker # type: ignore
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()
logger = logging.getLogger(__name__)
__all__ = [
"create_a2a_app",
"load_message_history",
"save_message_history",
"extract_question_from_task",
"get_agent_skills",
"LRUMemoryStorage",
]
def create_a2a_app(
db_path: Path,
security_schemes: dict | None = None,
security: list[dict[str, list[str]]] | None = None,
):
"""Create an A2A app for the conversational QA agent.
Args:
db_path: Path to the LanceDB database
security_schemes: Optional security scheme definitions for the AgentCard
security: Optional security requirements for the AgentCard
Returns:
A FastA2A ASGI application
"""
base_storage = InMemoryStorage()
storage = LRUMemoryStorage(
storage=base_storage, max_contexts=Config.A2A_MAX_CONTEXTS
)
broker = InMemoryBroker()
# Create the agent with native search tool
model = get_model(Config.QA_PROVIDER, Config.QA_MODEL)
agent = Agent(
model=model,
deps_type=AgentDependencies,
system_prompt=A2A_SYSTEM_PROMPT,
retries=3,
)
@agent.tool
async def search_documents(
ctx: RunContext[AgentDependencies],
query: str,
limit: int = 3,
) -> list[SearchResult]:
"""Search the knowledge base for relevant documents.
Returns chunks of text with their relevance scores and document URIs.
Use get_full_document if you need to see the complete document content.
"""
search_results = await ctx.deps.client.search(query, limit=limit)
expanded_results = await ctx.deps.client.expand_context(search_results)
return [
SearchResult(
content=chunk.content,
score=score,
document_title=chunk.document_title,
document_uri=(chunk.document_uri or ""),
)
for chunk, score in expanded_results
]
@agent.tool
async def get_full_document(
ctx: RunContext[AgentDependencies],
document_uri: str,
) -> str:
"""Retrieve the complete content of a document by its URI.
Use this when you need more context than what's in a search result chunk.
The document_uri comes from search_documents results.
"""
document = await ctx.deps.client.get_document_by_uri(document_uri)
if document is None:
return f"Document not found: {document_uri}"
return document.content
worker = ConversationalWorker(
storage=storage,
broker=broker,
db_path=db_path,
agent=agent, # type: ignore
)
# Create FastA2A app with custom worker lifecycle
@asynccontextmanager
async def lifespan(app):
logger.info(f"Started A2A server (max contexts: {Config.A2A_MAX_CONTEXTS})")
async with app.task_manager:
async with worker.run():
yield
app = FastA2A(
storage=storage,
broker=broker,
name="haiku-rag",
description="Conversational question answering agent powered by haiku.rag RAG system",
skills=get_agent_skills(),
lifespan=lifespan,
)
# Add security configuration if provided
if security_schemes or security:
# Monkey-patch the agent card endpoint to include security
async def _agent_card_endpoint_with_security(request):
from fasta2a.schema import AgentCapabilities, AgentCard, agent_card_ta
from starlette.responses import Response
if app._agent_card_json_schema is None:
agent_card = AgentCard(
name=app.name,
description=app.description
or "An AI agent exposed as an A2A agent.",
url=app.url,
version=app.version,
protocol_version="0.3.0",
skills=app.skills,
default_input_modes=app.default_input_modes,
default_output_modes=app.default_output_modes,
capabilities=AgentCapabilities(
streaming=False,
push_notifications=False,
state_transition_history=False,
),
)
if app.provider is not None:
agent_card["provider"] = app.provider
if security_schemes:
agent_card["security_schemes"] = security_schemes
if security:
agent_card["security"] = security
app._agent_card_json_schema = agent_card_ta.dump_json(
agent_card, by_alias=True
)
return Response(
content=app._agent_card_json_schema, media_type="application/json"
)
app._agent_card_endpoint = _agent_card_endpoint_with_security
return app

271
src/haiku/rag/a2a/client.py Normal file
View file

@ -0,0 +1,271 @@
import asyncio
import uuid
from typing import Any
import httpx
from rich.console import Console
from rich.markdown import Markdown
from rich.prompt import Prompt
class A2AClient:
"""Simple A2A protocol client."""
def __init__(self, base_url: str = "http://localhost:8000"):
"""Initialize A2A client.
Args:
base_url: Base URL of the A2A server
"""
self.base_url = base_url.rstrip("/")
self.client = httpx.AsyncClient(timeout=60.0)
async def close(self):
"""Close the HTTP client."""
await self.client.aclose()
async def get_agent_card(self) -> dict[str, Any]:
"""Fetch the agent card from the A2A server.
Returns:
Agent card dictionary with agent capabilities and metadata
"""
response = await self.client.get(f"{self.base_url}/.well-known/agent-card.json")
response.raise_for_status()
return response.json()
async def send_message(
self,
text: str,
context_id: str | None = None,
skill_id: str | None = None,
) -> dict[str, Any]:
"""Send a message to the A2A agent and wait for completion.
Args:
text: Message text to send
context_id: Optional conversation context ID (creates new if None)
skill_id: Optional skill ID to use (defaults to document-qa)
Returns:
Completed task with response messages and artifacts
"""
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"},
)
response.raise_for_status()
initial_response = response.json()
# Extract task ID from response
result = initial_response.get("result", {})
task_id = result.get("id")
if not task_id:
return initial_response
# Poll for task completion
return await self.wait_for_task(task_id)
async def wait_for_task(
self, task_id: str, max_wait: int = 60, poll_interval: float = 0.5
) -> dict[str, Any]:
"""Poll for task completion.
Args:
task_id: Task ID to poll for
max_wait: Maximum time to wait in seconds
poll_interval: Interval between polls in seconds
Returns:
Completed task result
"""
import time
start_time = time.time()
while time.time() - start_time < max_wait:
payload = {
"jsonrpc": "2.0",
"method": "tasks/get",
"params": {"id": task_id},
"id": 2,
}
response = await self.client.post(
self.base_url,
json=payload,
headers={"Content-Type": "application/json"},
)
response.raise_for_status()
task = response.json()
result = task.get("result", {})
status = result.get("status", {})
state = status.get("state")
if state == "completed":
return task
elif state == "failed":
raise Exception(f"Task failed: {task}")
await asyncio.sleep(poll_interval)
raise TimeoutError(f"Task {task_id} did not complete within {max_wait}s")
def print_agent_card(card: dict[str, Any], console: Console):
"""Pretty print the agent card using Rich."""
console.print()
console.print("[bold]Agent Card[/bold]")
console.rule()
console.print(f" [repr.attrib_name]name[/repr.attrib_name]: {card.get('name')}")
console.print(
f" [repr.attrib_name]description[/repr.attrib_name]: {card.get('description')}"
)
console.print(
f" [repr.attrib_name]version[/repr.attrib_name]: {card.get('version')}"
)
console.print(
f" [repr.attrib_name]protocol version[/repr.attrib_name]: {card.get('protocolVersion')}"
)
skills = card.get("skills", [])
console.print(f"\n[bold cyan]Skills ({len(skills)}):[/bold cyan]")
for skill in skills:
console.print(f"{skill.get('id')}: {skill.get('name')}")
console.print(f" [dim]{skill.get('description')}[/dim]")
examples = skill.get("examples", [])
if examples:
console.print(f" [dim]Examples: {', '.join(examples[:2])}[/dim]")
console.print()
def print_response(response: dict[str, Any], console: Console):
"""Pretty print the A2A response using Rich."""
if "error" in response:
console.print(f"[red]Error: {response['error']}[/red]")
return
result = response.get("result", {})
# Get messages from history and artifacts from completed task
history = result.get("history", [])
artifacts = result.get("artifacts", [])
# Print agent messages from history with markdown rendering
for msg in history:
if msg.get("role") == "agent":
for part in msg.get("parts", []):
if part.get("kind") == "text":
text = part.get("text", "")
# Render as markdown
console.print()
console.print("[bold green]Answer:[/bold green]")
console.print(Markdown(text))
# Print artifacts summary with details
if artifacts:
summary_lines = []
for artifact in artifacts:
name = artifact.get("name", "")
parts = artifact.get("parts", [])
if name == "search_results" and parts:
data = parts[0].get("data", {})
query = data.get("query", "")
results = data.get("results", [])
summary_lines.append(f"🔍 search: '{query}' ({len(results)} results)")
elif name == "document" and parts:
part = parts[0]
if part.get("kind") == "text":
text = part.get("text", "")
length = len(text)
summary_lines.append(f"📄 document ({length} chars)")
elif name == "qa_result" and parts:
data = parts[0].get("data", {})
skill = data.get("skill", "unknown")
summary_lines.append(f"💬 {skill}")
if summary_lines:
console.print(f"[dim]{''.join(summary_lines)}[/dim]")
console.print()
async def run_interactive_client(url: str = "http://localhost:8000"):
"""Run the interactive A2A client.
Args:
url: Base URL of the A2A server
"""
console = Console()
client = A2AClient(url)
console.print("[bold]haiku.rag A2A interactive client[/bold]")
console.print()
# Fetch and display agent card
console.print("[dim]Fetching agent card...[/dim]")
try:
card = await client.get_agent_card()
print_agent_card(card, console)
except Exception as e:
console.print(f"[red]Error fetching agent card: {e}[/red]")
await client.close()
return
# Create a conversation context
context_id = str(uuid.uuid4())
console.print(f"[dim]context id: {context_id}[/dim]")
console.print("[dim]Type your questions (or 'quit' to exit)[/dim]\n")
try:
while True:
try:
question = Prompt.ask("[bold blue]Question[/bold blue]").strip()
if not question:
continue
if question.lower() in ("quit", "exit", "q"):
console.print("\n[dim]Goodbye![/dim]")
break
response = await client.send_message(question, context_id=context_id)
print_response(response, console)
except KeyboardInterrupt:
console.print("\n\n[dim]Exiting...[/dim]")
break
except Exception as e:
console.print(f"\n[red]Error: {e}[/red]\n")
finally:
await client.close()

View file

@ -0,0 +1,68 @@
import uuid
from pydantic import TypeAdapter
from pydantic_ai.messages import ModelMessage
from pydantic_core import to_jsonable_python
try:
from fasta2a.schema import DataPart, Message # type: ignore
except ImportError as e:
raise ImportError(
"A2A support requires the 'a2a' extra. "
"Install with: uv pip install 'haiku.rag[a2a]'"
) from e
ModelMessagesTypeAdapter = TypeAdapter(list[ModelMessage])
def load_message_history(context: list[Message]) -> list[ModelMessage]:
"""Load pydantic-ai message history from A2A context.
The context stores serialized pydantic-ai message history directly,
which we deserialize and return.
Args:
context: A2A context messages
Returns:
List of pydantic-ai ModelMessage objects
"""
if not context:
return []
# Context should contain a single "state" message with full history
for msg in context:
parts = msg.get("parts", [])
for part in parts:
if part.get("kind") == "data":
metadata = part.get("metadata", {})
if metadata.get("type") == "conversation_state":
stored_history = part.get("data", {}).get("message_history", [])
if stored_history:
return ModelMessagesTypeAdapter.validate_python(stored_history)
return []
def save_message_history(message_history: list[ModelMessage]) -> Message:
"""Save pydantic-ai message history to A2A context format.
Args:
message_history: Full pydantic-ai message history
Returns:
A2A Message containing the serialized state (stored as agent role)
"""
serialized = to_jsonable_python(message_history)
return Message(
role="agent",
parts=[
DataPart(
kind="data",
data={"message_history": serialized},
metadata={"type": "conversation_state"},
)
],
kind="message",
message_id=str(uuid.uuid4()),
)

View file

@ -0,0 +1,21 @@
from pydantic import BaseModel, Field
from haiku.rag.client import HaikuRAG
class SearchResult(BaseModel):
"""Search result with both title and URI for A2A agent."""
content: str = Field(description="The document text content")
score: float = Field(description="Relevance score (higher is more relevant)")
document_title: str | None = Field(
description="Human-readable document title", default=None
)
document_uri: str = Field(description="Document URI/path for get_full_document")
class AgentDependencies(BaseModel):
"""Dependencies for the A2A conversational agent."""
model_config = {"arbitrary_types_allowed": True}
client: HaikuRAG

View file

@ -0,0 +1,59 @@
A2A_SYSTEM_PROMPT = """You are Haiku.rag, an AI assistant that helps users find information from a document knowledge base.
IMPORTANT: You are NOT any person mentioned in the documents. You retrieve and present information about them.
Tools available:
- search_documents: Query for relevant text chunks (returns SearchResult objects with content, score, document_title, document_uri)
- get_full_document: Get complete document content by document_uri
Your behavior depends on the operation:
## For direct search requests:
When the user is explicitly searching (e.g., "search for X", "find documents about Y"):
- Use search_documents tool ONLY
- Format results as a numbered list using markdown formatting
- For each result show:
* First line: *Score in italic* | **source in bold** (title if available, otherwise URI)
* Second line: The FULL chunk content (do not summarize or truncate)
- Present results in order of relevance
- Be concise - just present the search results, do not synthesize or add commentary
Example format:
Found 3 relevant results:
1. *Score: 0.95* | **Python Documentation** (/guides/python.md)
Python is a high-level, general-purpose programming language. Its design philosophy emphasizes code readability with the use of significant indentation.
2. *Score: 0.87* | **/guides/python-basics.md**
Python supports multiple programming paradigms, including structured, object-oriented and functional programming.
## For question-answering:
When the user asks a question (e.g., "What is Python?", "How does X work?"):
- For complex questions, use search_documents MULTIPLE TIMES with DIFFERENT queries to gather comprehensive information
- Example: For "What are the benefits and drawbacks of Python?", search separately for:
* "Python benefits advantages"
* "Python drawbacks disadvantages limitations"
- Synthesize information from all searches into a comprehensive answer
- Include "Sources:" section at the end listing sources used
Sources Format:
List each source with its title/URI and the relevant chunk content (NOT the score).
Format: "- **[title or URI]**: [chunk content]"
Example:
[Your synthesized answer here]
Sources:
- **Python Documentation** (/guides/python.md): Python is a high-level, general-purpose programming language. Its design philosophy emphasizes code readability.
- **/guides/python-basics.md**: Python supports multiple programming paradigms, including structured, object-oriented and functional programming.
Critical rules:
- ONLY answer based on information found via search_documents
- For comprehensive questions, perform MULTIPLE searches with different query angles
- NEVER fabricate or assume information
- If not found, say: "I cannot find information about this in the knowledge base."
- For follow-ups, understand context (pronouns like "he", "it") but always search for facts
- In Sources, include the actual chunk content from your search results, not summaries
Note: When using get_full_document, always use document_uri (not document_title).
"""

View file

@ -0,0 +1,75 @@
try:
from fasta2a.schema import Message, Skill # type: ignore
except ImportError as e:
raise ImportError(
"A2A support requires the 'a2a' extra. "
"Install with: uv pip install 'haiku.rag[a2a]'"
) from e
def get_agent_skills() -> list[Skill]:
"""Define the skills exposed by the haiku.rag A2A agent.
Returns:
List of skills describing the agent's capabilities
"""
return [
Skill(
id="document-qa",
name="Document Question Answering",
description="Answer questions based on a knowledge base of documents using semantic search and retrieval",
tags=["question-answering", "search", "knowledge-base", "rag"],
input_modes=["application/json"],
output_modes=["application/json"],
examples=[
"What does the documentation say about authentication?",
"Find information about Python best practices",
"Show me the full API documentation",
],
),
Skill(
id="document-search",
name="Document Search",
description="Search for relevant document chunks in the knowledge base using hybrid (semantic and BM25) search",
tags=["search", "retrieval", "semantic-search"],
input_modes=["application/json"],
output_modes=["application/json"],
examples=[
"Search for Python best practices",
"Find documents about authentication",
"Look for API documentation",
],
),
Skill(
id="document-retrieve",
name="Document Retrieval",
description="Retrieve the complete content of a specific document by its URI",
tags=["retrieval", "fetch", "document"],
input_modes=["application/json"],
output_modes=["application/json"],
examples=[
"Get the full content of document X",
"Retrieve document by URI",
"Show me the complete document",
],
),
]
def extract_question_from_task(task_history: list[Message]) -> str | None:
"""Extract the user's question from task history.
Args:
task_history: Task history messages
Returns:
The question text if found, None otherwise
"""
for msg in task_history:
if msg.get("role") == "user":
for part in msg.get("parts", []):
if part.get("kind") == "text":
text = part.get("text", "").strip()
if text:
return text
return None

View file

@ -0,0 +1,71 @@
import logging
from collections import OrderedDict
try:
from fasta2a.schema import Artifact, Message, TaskState # type: ignore
from fasta2a.storage import InMemoryStorage, Storage # type: ignore
except ImportError as e:
raise ImportError(
"A2A support requires the 'a2a' extra. "
"Install with: uv pip install 'haiku.rag[a2a]'"
) from e
logger = logging.getLogger(__name__)
class LRUMemoryStorage(Storage[list[Message]]): # type: ignore
"""Storage wrapper with LRU eviction for contexts.
Enforces a maximum context limit using LRU (Least Recently Used) eviction.
"""
def __init__(self, storage: InMemoryStorage, max_contexts: int):
self.storage = storage
self.max_contexts = max_contexts
# Track context access order (LRU cache)
self.context_order: OrderedDict[str, None] = OrderedDict()
async def load_context(self, context_id: str) -> list[Message] | None:
"""Load context and update access order."""
result = await self.storage.load_context(context_id)
if result is not None:
# Move to end (most recently used)
self.context_order.pop(context_id, None)
self.context_order[context_id] = None
return result
async def update_context(self, context_id: str, context: list[Message]) -> None:
"""Update context and enforce LRU limit."""
await self.storage.update_context(context_id, context)
# Move to end (most recently used)
self.context_order.pop(context_id, None)
self.context_order[context_id] = None
# Enforce max contexts limit (LRU eviction)
while len(self.context_order) > self.max_contexts:
# Remove oldest (first item in OrderedDict)
oldest_context_id = next(iter(self.context_order))
self.context_order.pop(oldest_context_id)
logger.debug(
f"Evicted context {oldest_context_id} (LRU, limit={self.max_contexts})"
)
async def load_task(self, task_id: str, history_length: int | None = None):
"""Delegate to underlying storage."""
return await self.storage.load_task(task_id, history_length)
async def update_task(
self,
task_id: str,
state: TaskState,
new_artifacts: list[Artifact] | None = None,
new_messages: list[Message] | None = None,
):
"""Delegate to underlying storage."""
return await self.storage.update_task(
task_id, state, new_artifacts, new_messages
)
async def submit_task(self, context_id: str, message: Message):
"""Delegate to underlying storage."""
return await self.storage.submit_task(context_id, message)

320
src/haiku/rag/a2a/worker.py Normal file
View file

@ -0,0 +1,320 @@
import json
import logging
import uuid
from pathlib import Path
from pydantic_ai import Agent
from haiku.rag.a2a.context import load_message_history, save_message_history
from haiku.rag.a2a.models import AgentDependencies
from haiku.rag.a2a.skills import extract_question_from_task
from haiku.rag.client import HaikuRAG
try:
from fasta2a import Worker # type: ignore
from fasta2a.schema import ( # type: ignore
Artifact,
Message,
TaskIdParams,
TaskSendParams,
TextPart,
)
except ImportError as e:
raise ImportError(
"A2A support requires the 'a2a' extra. "
"Install with: uv pip install 'haiku.rag[a2a]'"
) from e
logger = logging.getLogger(__name__)
class ConversationalWorker(Worker[list[Message]]):
"""Worker that handles conversational QA tasks."""
def __init__(
self,
storage,
broker,
db_path: Path,
agent: "Agent[AgentDependencies, str]",
):
super().__init__(storage=storage, broker=broker)
self.db_path = db_path
self.agent = agent
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")
task_history = task.get("history", [])
question = extract_question_from_task(task_history)
if not question:
await self.storage.update_task(task["id"], state="failed")
return
try:
async with HaikuRAG(self.db_path) as client:
context = await self.storage.load_context(task["context_id"]) or []
message_history = load_message_history(context)
deps = AgentDependencies(client=client)
result = await self.agent.run(
question, deps=deps, message_history=message_history
)
# Detect which skill was used
skill_type = self._detect_skill(result)
# Build messages based on skill type
response_messages = self._build_response_messages(result, skill_type)
# Update context with complete conversation state
updated_history = message_history + result.new_messages()
state_message = save_message_history(updated_history)
await self.storage.update_context(task["context_id"], [state_message])
artifacts = self.build_artifacts(result, skill_type, question)
await self.storage.update_task(
task["id"],
state="completed",
new_messages=response_messages,
new_artifacts=artifacts,
)
except Exception as e:
logger.error(
"Task execution failed: task_id=%s, question=%s, error=%s",
task["id"],
question,
str(e),
exc_info=True,
)
await self.storage.update_task(task["id"], state="failed")
raise
async def cancel_task(self, params: TaskIdParams) -> None:
"""Cancel a task - not implemented for this worker."""
pass
def build_message_history(self, history: list[Message]) -> list[Message]:
"""Required by Worker interface but unused - history stored in context."""
return history
def _detect_skill(self, result) -> str:
"""Detect which skill was used based on tool calls and response pattern.
Returns:
"search", "retrieve", or "qa"
"""
from pydantic_ai.messages import ModelResponse, ToolCallPart
tool_calls = []
for msg in result.new_messages():
if isinstance(msg, ModelResponse):
for part in msg.parts:
if isinstance(part, ToolCallPart):
tool_calls.append(part.tool_name)
# Check if output looks like formatted search results
output_str = str(result.output).strip()
# Check for either format: "Found N relevant results" or "**Search results for"
is_search_format = (
output_str.startswith("Found ") and "relevant results" in output_str[:100]
) or output_str.startswith("**Search results for")
skill_type = "qa"
# If output is in search format and only search tools were used, it's a search
if is_search_format and all(tc == "search_documents" for tc in tool_calls):
skill_type = "search"
elif "get_full_document" in tool_calls and len(tool_calls) == 1:
skill_type = "retrieve"
return skill_type
def _build_response_messages(self, result, skill_type: str) -> list[Message]:
"""Build response messages based on skill type.
All skills return a single text message with LLM's response.
Structured data is provided via artifacts for search/retrieve.
"""
if skill_type == "search":
# Return LLM's formatted response
return [
Message(
role="agent",
parts=[TextPart(kind="text", text=str(result.output))],
kind="message",
message_id=str(uuid.uuid4()),
)
]
elif skill_type == "retrieve":
# Extract document content
from pydantic_ai.messages import ModelRequest, ToolReturnPart
document_content = ""
for msg in result.new_messages():
if isinstance(msg, ModelRequest):
for part in msg.parts:
if (
isinstance(part, ToolReturnPart)
and part.tool_name == "get_full_document"
):
document_content = part.content
break
return [
Message(
role="agent",
parts=[TextPart(kind="text", text=document_content)],
kind="message",
message_id=str(uuid.uuid4()),
)
]
else:
# Conversational Q&A - use agent's answer
return [
Message(
role="agent",
parts=[TextPart(kind="text", text=str(result.output))],
kind="message",
message_id=str(uuid.uuid4()),
)
]
def build_artifacts(
self, result, skill_type: str | None = None, question: str | None = None
) -> list[Artifact]:
"""Build artifacts from agent result based on tool calls.
Creates artifacts for:
- Each tool call (search_documents, get_full_document)
- Q&A operations: additional artifact with question and answer (only if tools were used)
"""
if skill_type is None:
skill_type = self._detect_skill(result)
artifacts = []
# Always create artifacts for all tool calls
tool_artifacts = self._build_all_tool_artifacts(result)
artifacts.extend(tool_artifacts)
# For Q&A, always add a Q&A artifact with question and answer
# This includes follow-up questions, clarifications, and conversational responses
if skill_type == "qa" and question:
from fasta2a.schema import DataPart
artifacts.append(
Artifact(
artifact_id=str(uuid.uuid4()),
name="qa_result",
parts=[
DataPart(
kind="data",
data={
"question": question,
"answer": str(result.output),
"skill": "document-qa",
},
metadata={"skill": "document-qa"},
)
],
)
)
return artifacts
def _build_all_tool_artifacts(self, result) -> list[Artifact]:
"""Build artifacts for all tool calls."""
from pydantic_ai.messages import (
ModelRequest,
ModelResponse,
ToolCallPart,
ToolReturnPart,
)
artifacts = []
# Track tool calls and their returns by call_id
tool_returns = {}
for msg in result.new_messages():
if isinstance(msg, ModelRequest):
for part in msg.parts:
if isinstance(part, ToolReturnPart):
result_count = (
len(part.content) if isinstance(part.content, list) else 1
)
logger.info(
"Tool return: tool_call_id=%s, tool_name=%s, result_count=%s",
part.tool_call_id,
part.tool_name,
result_count,
)
tool_returns[part.tool_call_id] = (part.tool_name, part.content)
# Create artifacts for each tool call
for msg in result.new_messages():
if isinstance(msg, ModelResponse):
for part in msg.parts:
if isinstance(part, ToolCallPart):
tool_name, content = tool_returns.get(
part.tool_call_id, (None, None)
)
if tool_name == "search_documents" and content:
from fasta2a.schema import DataPart
# Extract query from tool call arguments
query = ""
if isinstance(part.args, dict):
query = part.args.get("query", "")
elif isinstance(part.args, str):
# Args is a JSON string - parse it
try:
args_dict = json.loads(part.args)
query = args_dict.get("query", "")
except (json.JSONDecodeError, AttributeError):
query = ""
elif hasattr(part.args, "get") and callable(
getattr(part.args, "get", None)
):
# ArgsDict or dict-like object
query = part.args.get("query", "") # type: ignore
elif hasattr(part.args, "query"):
# Object with query attribute
query = str(part.args.query) # type: ignore
artifacts.append(
Artifact(
artifact_id=str(uuid.uuid4()),
name="search_results",
parts=[
DataPart(
kind="data",
data={"results": content, "query": query},
metadata={"query": query},
)
],
)
)
elif tool_name == "get_full_document" and content:
artifacts.append(
Artifact(
artifact_id=str(uuid.uuid4()),
name="document",
parts=[TextPart(kind="text", text=content)],
)
)
return artifacts

View file

@ -1,5 +1,6 @@
import asyncio
import json
import logging
from importlib.metadata import version as pkg_version
from pathlib import Path
@ -22,6 +23,8 @@ from haiku.rag.research.stream import stream_research_graph
from haiku.rag.store.models.chunk import Chunk
from haiku.rag.store.models.document import Document
logger = logging.getLogger(__name__)
class HaikuRAGApp:
def __init__(self, db_path: Path):
@ -448,23 +451,81 @@ class HaikuRAGApp:
self.console.print(content)
self.console.rule()
async def serve(self, transport: str | None = None):
"""Start the MCP server."""
async def serve(
self,
enable_monitor: bool = True,
enable_mcp: bool = True,
mcp_transport: str | None = None,
mcp_port: int = 8001,
enable_a2a: bool = False,
a2a_host: str = "127.0.0.1",
a2a_port: int = 8000,
):
"""Start the server with selected services."""
async with HaikuRAG(self.db_path) as client:
monitor = FileWatcher(paths=Config.MONITOR_DIRECTORIES, client=client)
monitor_task = asyncio.create_task(monitor.observe())
server = create_mcp_server(self.db_path)
tasks = []
# Start file monitor if enabled
if enable_monitor:
monitor = FileWatcher(paths=Config.MONITOR_DIRECTORIES, client=client)
monitor_task = asyncio.create_task(monitor.observe())
tasks.append(monitor_task)
# Start MCP server if enabled
if enable_mcp:
server = create_mcp_server(self.db_path)
async def run_mcp():
if mcp_transport == "stdio":
await server.run_stdio_async()
else:
logger.info(f"Starting MCP server on port {mcp_port}")
await server.run_http_async(
transport="streamable-http", port=mcp_port
)
mcp_task = asyncio.create_task(run_mcp())
tasks.append(mcp_task)
# Start A2A server if enabled
if enable_a2a:
try:
from haiku.rag.a2a import create_a2a_app
except ImportError as e:
logger.error(f"Failed to import A2A: {e}")
return
import uvicorn
logger.info(f"Starting A2A server on {a2a_host}:{a2a_port}")
async def run_a2a():
app = create_a2a_app(db_path=self.db_path)
config = uvicorn.Config(
app,
host=a2a_host,
port=a2a_port,
log_level="warning",
access_log=False,
)
server = uvicorn.Server(config)
await server.serve()
a2a_task = asyncio.create_task(run_a2a())
tasks.append(a2a_task)
if not tasks:
logger.warning("No services enabled")
return
try:
if transport == "stdio":
await server.run_stdio_async()
else:
await server.run_http_async(transport="streamable-http")
# Wait for any task to complete (or KeyboardInterrupt)
await asyncio.gather(*tasks)
except KeyboardInterrupt:
pass
finally:
monitor_task.cancel()
try:
await monitor_task
except asyncio.CancelledError:
pass
# Cancel all tasks
for task in tasks:
task.cancel()
# Wait for cancellation
await asyncio.gather(*tasks, return_exceptions=True)

View file

@ -366,7 +366,8 @@ def download_models_cmd():
@cli.command(
"serve", help="Start the haiku.rag MCP server (by default in streamable HTTP mode)"
"serve",
help="Start haiku.rag server. Use --monitor, --mcp, and/or --a2a to enable services.",
)
def serve(
db: Path = typer.Option(
@ -374,22 +375,71 @@ def serve(
"--db",
help="Path to the LanceDB database file",
),
monitor: bool = typer.Option(
False,
"--monitor",
help="Enable file monitoring",
),
mcp: bool = typer.Option(
False,
"--mcp",
help="Enable MCP server",
),
stdio: bool = typer.Option(
False,
"--stdio",
help="Run MCP server on stdio Transport",
help="Run MCP server on stdio Transport (requires --mcp)",
),
mcp_port: int = typer.Option(
8001,
"--mcp-port",
help="Port to bind MCP server to (ignored with --stdio)",
),
a2a: bool = typer.Option(
False,
"--a2a",
help="Enable A2A (Agent-to-Agent) server",
),
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:
"""Start the MCP server."""
"""Start the server with selected services."""
# Require at least one service flag
if not (monitor or mcp or a2a):
typer.echo(
"Error: At least one service flag (--monitor, --mcp, or --a2a) must be specified"
)
raise typer.Exit(1)
if stdio and not mcp:
typer.echo("Error: --stdio requires --mcp")
raise typer.Exit(1)
from haiku.rag.app import HaikuRAGApp
app = HaikuRAGApp(db_path=db)
transport = None
if stdio:
transport = "stdio"
transport = "stdio" if stdio else None
asyncio.run(app.serve(transport=transport))
asyncio.run(
app.serve(
enable_monitor=monitor,
enable_mcp=mcp,
mcp_transport=transport,
mcp_port=mcp_port,
enable_a2a=a2a,
a2a_host=a2a_host,
a2a_port=a2a_port,
)
)
@cli.command("migrate", help="Migrate an SQLite database to LanceDB")
@ -410,5 +460,27 @@ def migrate(
raise typer.Exit(1)
@cli.command(
"a2aclient", help="Run interactive client to chat with haiku.rag's A2A server"
)
def a2aclient(
url: str = typer.Option(
"http://localhost:8000",
"--url",
help="Base URL of the A2A server",
),
):
try:
from haiku.rag.a2a.client import run_interactive_client
except ImportError:
typer.echo(
"Error: A2A support requires the 'a2a' extra. "
"Install with: uv pip install 'haiku.rag[a2a]'"
)
raise typer.Exit(1)
asyncio.run(run_interactive_client(url=url))
if __name__ == "__main__":
cli()

View file

@ -62,6 +62,10 @@ class AppConfig(BaseModel):
# to allow concurrent connections to safely use recent versions.
VACUUM_RETENTION_SECONDS: int = 60
# Maximum number of A2A contexts to keep in memory. When exceeded, least
# recently used contexts will be evicted. Default is 1000.
A2A_MAX_CONTEXTS: int = 1000
@field_validator("MONITOR_DIRECTORIES", mode="before")
@classmethod
def parse_monitor_directories(cls, v):

View file

@ -54,9 +54,6 @@ class QuestionAnswerAgent:
limit: int = 3,
) -> list[SearchResult]:
"""Search the knowledge base for relevant documents."""
# Remove quotes from queries as this requires positional indexing in lancedb
query = query.replace('"', "")
search_results = await ctx.deps.client.search(query, limit=limit)
expanded_results = await ctx.deps.client.expand_context(search_results)

642
tests/test_a2a.py Normal file
View file

@ -0,0 +1,642 @@
import uuid
import pytest
from haiku.rag.a2a import (
extract_question_from_task,
get_agent_skills,
load_message_history,
save_message_history,
)
from haiku.rag.a2a.storage import LRUMemoryStorage
from haiku.rag.client import HaikuRAG
pytest.importorskip("fasta2a")
from fasta2a.schema import Message, TextPart # noqa: E402
from pydantic_ai.messages import ( # noqa: E402
ModelMessage,
ModelRequest,
ModelResponse,
ToolCallPart,
ToolReturnPart,
)
from pydantic_ai.messages import (
TextPart as AITextPart,
)
@pytest.mark.asyncio
async def test_save_and_load_message_history():
"""Test round-trip of saving and loading message history."""
# Create sample message history with proper part_kind for ModelRequest
from pydantic_ai.messages import UserPromptPart
original_history: list[ModelMessage] = [
ModelRequest(parts=[UserPromptPart(content="What is Python?")]),
ModelResponse(parts=[AITextPart(content="Python is a programming language")]),
]
# Save to A2A format
saved_message = save_message_history(original_history)
# Verify structure
assert saved_message["role"] == "agent"
assert saved_message["kind"] == "message"
assert len(saved_message["parts"]) == 1
assert saved_message["parts"][0]["kind"] == "data"
metadata = saved_message["parts"][0].get("metadata")
assert metadata is not None
assert metadata.get("type") == "conversation_state"
# Load it back
loaded_history = load_message_history([saved_message])
# Verify it matches
assert len(loaded_history) == len(original_history)
# First message is a request with UserPromptPart
assert isinstance(loaded_history[0], ModelRequest)
first_part = loaded_history[0].parts[0]
assert hasattr(first_part, "content")
assert first_part.content == "What is Python?" # type: ignore
# Second message is a response with TextPart
assert isinstance(loaded_history[1], ModelResponse)
second_part = loaded_history[1].parts[0]
assert hasattr(second_part, "content")
assert second_part.content == "Python is a programming language" # type: ignore
@pytest.mark.asyncio
async def test_save_and_load_message_history_with_tool_calls():
"""Test saving and loading message history that includes tool calls."""
from pydantic_ai.messages import UserPromptPart
original_history: list[ModelMessage] = [
ModelRequest(parts=[UserPromptPart(content="Search for Python")]),
ModelResponse(
parts=[
ToolCallPart(
tool_name="search_documents",
args={"query": "Python", "limit": 3},
tool_call_id="call_1",
)
]
),
ModelRequest(
parts=[
ToolReturnPart(
tool_name="search_documents",
content="Python is a high-level programming language",
tool_call_id="call_1",
)
]
),
ModelResponse(
parts=[AITextPart(content="Based on the search, Python is a language")]
),
]
# Save and load
saved_message = save_message_history(original_history)
loaded_history = load_message_history([saved_message])
# Verify tool calls are preserved
assert len(loaded_history) == 4
assert isinstance(loaded_history[1].parts[0], ToolCallPart)
assert loaded_history[1].parts[0].tool_name == "search_documents"
assert isinstance(loaded_history[2].parts[0], ToolReturnPart)
assert loaded_history[2].parts[0].tool_name == "search_documents"
@pytest.mark.asyncio
async def test_extract_question_from_task():
"""Test extracting user question from task history."""
task_history: list[Message] = [
Message(
role="user",
parts=[TextPart(kind="text", text="What is Python?")],
kind="message",
message_id=str(uuid.uuid4()),
)
]
question = extract_question_from_task(task_history)
assert question == "What is Python?"
@pytest.mark.asyncio
async def test_extract_question_from_task_no_text():
"""Test extracting question when no text part exists."""
task_history: list[Message] = [
Message(
role="user",
parts=[],
kind="message",
message_id=str(uuid.uuid4()),
)
]
question = extract_question_from_task(task_history)
assert question is None
@pytest.mark.asyncio
async def test_lru_memory_storage_lru_eviction():
"""Test that LRUMemoryStorage evicts least recently used contexts."""
from fasta2a.storage import InMemoryStorage
base_storage = InMemoryStorage()
storage = LRUMemoryStorage(storage=base_storage, max_contexts=3)
# Add 3 contexts (at limit)
await storage.update_context("ctx1", [])
await storage.update_context("ctx2", [])
await storage.update_context("ctx3", [])
# All 3 should be tracked
assert len(storage.context_order) == 3
assert "ctx1" in storage.context_order
assert "ctx2" in storage.context_order
assert "ctx3" in storage.context_order
# Add 4th context - should evict ctx1 (oldest)
await storage.update_context("ctx4", [])
assert len(storage.context_order) == 3
assert "ctx1" not in storage.context_order
assert "ctx2" in storage.context_order
assert "ctx3" in storage.context_order
assert "ctx4" in storage.context_order
# Access ctx2 (moves it to end)
await storage.load_context("ctx2")
# Add 5th context - should evict ctx3 (now oldest since ctx2 was accessed)
await storage.update_context("ctx5", [])
assert len(storage.context_order) == 3
assert "ctx3" not in storage.context_order
assert "ctx2" in storage.context_order # Still present (was accessed)
assert "ctx4" in storage.context_order
assert "ctx5" in storage.context_order
@pytest.mark.asyncio
async def test_lru_memory_storage_access_order():
"""Test that accessing contexts updates their order."""
from fasta2a.storage import InMemoryStorage
base_storage = InMemoryStorage()
storage = LRUMemoryStorage(storage=base_storage, max_contexts=2)
# Add 2 contexts
await storage.update_context("ctx1", [])
await storage.update_context("ctx2", [])
# Order should be: ctx1, ctx2
assert list(storage.context_order.keys()) == ["ctx1", "ctx2"]
# Load ctx1 (moves to end)
await storage.load_context("ctx1")
# Order should be: ctx2, ctx1
assert list(storage.context_order.keys()) == ["ctx2", "ctx1"]
# Add ctx3 - should evict ctx2 (oldest)
await storage.update_context("ctx3", [])
assert "ctx2" not in storage.context_order
assert "ctx1" in storage.context_order
assert "ctx3" in storage.context_order
@pytest.mark.asyncio
async def test_a2a_app_creation(temp_db_path):
"""Test that A2A app can be created successfully."""
from haiku.rag.a2a import create_a2a_app
# Create a test database
async with HaikuRAG(temp_db_path) as client:
await client.create_document(
content="Python is a high-level programming language known for its simplicity.",
uri="python_doc",
)
# Create A2A app
app = create_a2a_app(temp_db_path)
# Verify app properties
assert app.name == "haiku-rag"
assert app.description is not None
assert "conversational" in app.description.lower()
@pytest.mark.asyncio
async def test_a2a_app_has_skills(temp_db_path):
"""Test that A2A app exposes skills describing its capabilities."""
from haiku.rag.a2a import create_a2a_app
# Create a test database
async with HaikuRAG(temp_db_path) as client:
await client.create_document(content="Test document", uri="test_doc")
# Create A2A app
app = create_a2a_app(temp_db_path)
# Verify app has skills
assert app.skills is not None
assert len(app.skills) > 0
# Check that at least one skill exists
skill = app.skills[0]
assert "id" in skill
assert "name" in skill
assert "description" in skill
assert "tags" in skill
assert "input_modes" in skill
assert "output_modes" in skill
# Verify the skill describes document search/QA capabilities
skill_text = f"{skill['name']} {skill['description']}".lower()
assert any(
keyword in skill_text
for keyword in ["search", "question", "answer", "document", "knowledge"]
)
def test_get_agent_skills():
"""Test that agent skills include all three skills."""
skills = get_agent_skills()
assert len(skills) == 3
skill_ids = [skill["id"] for skill in skills]
assert "document-qa" in skill_ids
assert "document-search" in skill_ids
assert "document-retrieve" in skill_ids
# Check document-qa skill
doc_qa = next(s for s in skills if s["id"] == "document-qa")
assert "Document Question Answering" in doc_qa["name"]
assert "semantic search" in doc_qa["description"]
assert "question-answering" in doc_qa["tags"]
# Check document-search skill
doc_search = next(s for s in skills if s["id"] == "document-search")
assert "Document Search" in doc_search["name"]
assert "search" in doc_search["tags"]
# Check document-retrieve skill
doc_retrieve = next(s for s in skills if s["id"] == "document-retrieve")
assert "Document Retrieval" in doc_retrieve["name"]
assert "retrieval" in doc_retrieve["tags"]
@pytest.mark.asyncio
async def test_build_artifacts_for_search():
"""Test that search operations produce structured search artifacts."""
from pydantic_ai.messages import (
ModelRequest,
ModelResponse,
ToolCallPart,
ToolReturnPart,
)
from haiku.rag.a2a.worker import ConversationalWorker
class MockResult:
output = "Found 1 relevant results:\n\n1. *Score: 0.9* | **test**\nresult"
def new_messages(self):
return [
ModelResponse(
parts=[
ToolCallPart(
tool_name="search_documents",
args={"query": "test", "limit": 3},
tool_call_id="call_1",
)
]
),
ModelRequest(
parts=[
ToolReturnPart(
tool_name="search_documents",
content=[{"content": "result", "score": 0.9}],
tool_call_id="call_1",
)
]
),
]
from pathlib import Path
from fasta2a.broker import InMemoryBroker
from fasta2a.storage import InMemoryStorage
worker = ConversationalWorker(
storage=InMemoryStorage(),
broker=InMemoryBroker(),
db_path=Path("/tmp/test.db"),
agent=None, # type: ignore
)
artifacts = worker.build_artifacts(MockResult(), "search", "test query")
assert len(artifacts) == 1
assert artifacts[0].get("name") == "search_results"
assert len(artifacts[0]["parts"]) == 1
assert artifacts[0]["parts"][0]["kind"] == "data"
assert "results" in artifacts[0]["parts"][0]["data"]
assert "query" in artifacts[0]["parts"][0]["data"]
@pytest.mark.asyncio
async def test_build_artifacts_for_retrieve():
"""Test that retrieve operations produce document artifacts."""
from pydantic_ai.messages import (
ModelRequest,
ModelResponse,
ToolCallPart,
ToolReturnPart,
)
from haiku.rag.a2a.worker import ConversationalWorker
class MockResult:
output = "Document content"
def new_messages(self):
return [
ModelResponse(
parts=[
ToolCallPart(
tool_name="get_full_document",
args={"document_uri": "test.txt"},
tool_call_id="call_1",
)
]
),
ModelRequest(
parts=[
ToolReturnPart(
tool_name="get_full_document",
content="Full document content here",
tool_call_id="call_1",
)
]
),
]
from pathlib import Path
from fasta2a.broker import InMemoryBroker
from fasta2a.storage import InMemoryStorage
worker = ConversationalWorker(
storage=InMemoryStorage(),
broker=InMemoryBroker(),
db_path=Path("/tmp/test.db"),
agent=None, # type: ignore
)
artifacts = worker.build_artifacts(MockResult(), "retrieve", "test query")
assert len(artifacts) == 1
assert artifacts[0].get("name") == "document"
assert artifacts[0]["parts"][0]["kind"] == "text"
assert artifacts[0]["parts"][0]["text"] == "Full document content here"
@pytest.mark.asyncio
async def test_build_artifacts_for_multiple_searches():
"""Test that multiple searches each get their own artifact with correct results."""
from pydantic_ai.messages import (
ModelRequest,
ModelResponse,
ToolCallPart,
ToolReturnPart,
)
from pydantic_ai.messages import TextPart as AITextPart
from haiku.rag.a2a.worker import ConversationalWorker
class MockResult:
output = "Answer based on multiple searches"
def new_messages(self):
return [
# First search
ModelResponse(
parts=[
ToolCallPart(
tool_name="search_documents",
args={"query": "first query", "limit": 2},
tool_call_id="call_1",
)
]
),
ModelRequest(
parts=[
ToolReturnPart(
tool_name="search_documents",
content=[
{"content": "result 1", "score": 0.9},
{"content": "result 2", "score": 0.8},
],
tool_call_id="call_1",
)
]
),
# Second search
ModelResponse(
parts=[
ToolCallPart(
tool_name="search_documents",
args={"query": "second query", "limit": 2},
tool_call_id="call_2",
)
]
),
ModelRequest(
parts=[
ToolReturnPart(
tool_name="search_documents",
content=[
{"content": "result 3", "score": 0.7},
{"content": "result 4", "score": 0.6},
],
tool_call_id="call_2",
)
]
),
ModelResponse(
parts=[AITextPart(content="Answer based on multiple searches")]
),
]
from pathlib import Path
from fasta2a.broker import InMemoryBroker
from fasta2a.storage import InMemoryStorage
worker = ConversationalWorker(
storage=InMemoryStorage(),
broker=InMemoryBroker(),
db_path=Path("/tmp/test.db"),
agent=None, # type: ignore
)
artifacts = worker.build_artifacts(MockResult(), "qa", "What is the answer?")
# Should have 2 search artifacts + 1 qa_result artifact
assert len(artifacts) == 3
# First search artifact
assert artifacts[0].get("name") == "search_results"
part_0 = artifacts[0]["parts"][0]
assert part_0.get("data", {}).get("query") == "first query"
results_1 = part_0.get("data", {}).get("results", [])
assert len(results_1) == 2
assert results_1[0]["content"] == "result 1"
assert results_1[1]["content"] == "result 2"
# Second search artifact
assert artifacts[1].get("name") == "search_results"
part_1 = artifacts[1]["parts"][0]
assert part_1.get("data", {}).get("query") == "second query"
results_2 = part_1.get("data", {}).get("results", [])
assert len(results_2) == 2
assert results_2[0]["content"] == "result 3"
assert results_2[1]["content"] == "result 4"
# Q&A artifact
assert artifacts[2].get("name") == "qa_result"
@pytest.mark.asyncio
async def test_qa_artifact_for_conversational_messages():
"""Test that conversational Q&A messages always create qa_result artifacts."""
from pydantic_ai.messages import ModelResponse
from pydantic_ai.messages import TextPart as AITextPart
from haiku.rag.a2a.worker import ConversationalWorker
class MockResult:
output = "Hello! How can I help you?"
def new_messages(self):
# No tool calls, just a conversational response
return [
ModelResponse(parts=[AITextPart(content="Hello! How can I help you?")]),
]
from pathlib import Path
from fasta2a.broker import InMemoryBroker
from fasta2a.storage import InMemoryStorage
worker = ConversationalWorker(
storage=InMemoryStorage(),
broker=InMemoryBroker(),
db_path=Path("/tmp/test.db"),
agent=None, # type: ignore
)
artifacts = worker.build_artifacts(MockResult(), "qa", "Hello")
# Should have qa_result artifact (even without tools, for A2A traceability)
assert len(artifacts) == 1
assert artifacts[0].get("name") == "qa_result"
part = artifacts[0]["parts"][0]
assert part.get("data", {}).get("question") == "Hello"
assert part.get("data", {}).get("answer") == "Hello! How can I help you?"
@pytest.mark.asyncio
async def test_build_artifacts_for_qa():
"""Test that Q&A operations produce artifacts for each tool call."""
from pydantic_ai.messages import (
ModelRequest,
ModelResponse,
ToolCallPart,
ToolReturnPart,
)
from pydantic_ai.messages import TextPart as AITextPart
from haiku.rag.a2a.worker import ConversationalWorker
class MockResult:
output = "This is the answer"
def new_messages(self):
# Multiple tool calls indicates Q&A workflow
return [
ModelResponse(
parts=[
ToolCallPart(
tool_name="search_documents",
args={"query": "test", "limit": 3},
tool_call_id="call_1",
)
]
),
ModelRequest(
parts=[
ToolReturnPart(
tool_name="search_documents",
content=[{"content": "result", "score": 0.9}],
tool_call_id="call_1",
)
]
),
ModelResponse(
parts=[
ToolCallPart(
tool_name="get_full_document",
args={"document_uri": "test.txt"},
tool_call_id="call_2",
)
]
),
ModelRequest(
parts=[
ToolReturnPart(
tool_name="get_full_document",
content="Full content",
tool_call_id="call_2",
)
]
),
ModelResponse(parts=[AITextPart(content="This is the answer")]),
]
from pathlib import Path
from fasta2a.broker import InMemoryBroker
from fasta2a.storage import InMemoryStorage
worker = ConversationalWorker(
storage=InMemoryStorage(),
broker=InMemoryBroker(),
db_path=Path("/tmp/test.db"),
agent=None, # type: ignore
)
artifacts = worker.build_artifacts(MockResult(), "qa", "What is Python?")
# Q&A should produce artifacts for each tool call (search + retrieve) + final Q&A artifact
assert len(artifacts) == 3
# First artifact is from search_documents
assert artifacts[0].get("name") == "search_results"
assert artifacts[0]["parts"][0]["kind"] == "data"
assert "results" in artifacts[0]["parts"][0]["data"]
assert artifacts[0]["parts"][0]["data"]["query"] == "test"
# Second artifact is from get_full_document
assert artifacts[1].get("name") == "document"
assert artifacts[1]["parts"][0]["kind"] == "text"
# Third artifact is the Q&A result
assert artifacts[2].get("name") == "qa_result"
assert artifacts[2]["parts"][0]["kind"] == "data"
assert artifacts[2]["parts"][0]["data"]["question"] == "What is Python?"
assert artifacts[2]["parts"][0]["data"]["answer"] == "This is the answer"
assert artifacts[2]["parts"][0]["data"]["skill"] == "document-qa"

View file

@ -181,13 +181,124 @@ async def test_search_no_results(app: HaikuRAGApp, monkeypatch):
@pytest.mark.asyncio
@pytest.mark.parametrize("transport", ["stdio", "http", None])
async def test_serve(app: HaikuRAGApp, monkeypatch, transport):
"""Test the serve method with different transports."""
@pytest.mark.parametrize("transport", ["stdio", None])
async def test_serve_mcp_only(app: HaikuRAGApp, monkeypatch, transport):
"""Test the serve method with MCP server only."""
mock_server = AsyncMock()
mock_watcher = MagicMock()
mock_task = asyncio.create_task(asyncio.sleep(0))
mock_task.cancel = MagicMock()
created_tasks = []
original_create_task = asyncio.create_task
def track_task(coro):
task = original_create_task(coro)
created_tasks.append(task)
task.cancel()
return task
monkeypatch.setattr(
"haiku.rag.app.create_mcp_server", MagicMock(return_value=mock_server)
)
monkeypatch.setattr("haiku.rag.app.asyncio.create_task", track_task)
monkeypatch.setattr(
"haiku.rag.app.asyncio.gather", AsyncMock(side_effect=asyncio.CancelledError)
)
mock_client = AsyncMock()
mock_client.__aenter__.return_value = mock_client
with patch("haiku.rag.app.HaikuRAG", return_value=mock_client):
try:
await app.serve(
enable_monitor=False,
enable_mcp=True,
mcp_transport=transport,
enable_a2a=False,
)
except asyncio.CancelledError:
pass
assert len(created_tasks) == 1
@pytest.mark.asyncio
async def test_serve_monitor_only(app: HaikuRAGApp, monkeypatch):
"""Test the serve method with monitor only."""
mock_watcher = AsyncMock()
created_tasks = []
original_create_task = asyncio.create_task
def track_task(coro):
task = original_create_task(coro)
created_tasks.append(task)
task.cancel()
return task
monkeypatch.setattr(
"haiku.rag.app.FileWatcher", MagicMock(return_value=mock_watcher)
)
monkeypatch.setattr("haiku.rag.app.asyncio.create_task", track_task)
monkeypatch.setattr(
"haiku.rag.app.asyncio.gather", AsyncMock(side_effect=asyncio.CancelledError)
)
mock_client = AsyncMock()
mock_client.__aenter__.return_value = mock_client
with patch("haiku.rag.app.HaikuRAG", return_value=mock_client):
try:
await app.serve(enable_monitor=True, enable_mcp=False, enable_a2a=False)
except asyncio.CancelledError:
pass
assert len(created_tasks) == 1
@pytest.mark.asyncio
async def test_serve_a2a_only(app: HaikuRAGApp, monkeypatch):
"""Test the serve method with A2A server only."""
created_tasks = []
original_create_task = asyncio.create_task
def track_task(coro):
task = original_create_task(coro)
created_tasks.append(task)
task.cancel()
return task
monkeypatch.setattr("haiku.rag.app.asyncio.create_task", track_task)
monkeypatch.setattr(
"haiku.rag.app.asyncio.gather", AsyncMock(side_effect=asyncio.CancelledError)
)
mock_client = AsyncMock()
mock_client.__aenter__.return_value = mock_client
mock_a2a_app = MagicMock()
with patch("haiku.rag.app.HaikuRAG", return_value=mock_client):
with patch("haiku.rag.a2a.create_a2a_app", return_value=mock_a2a_app):
try:
await app.serve(enable_monitor=False, enable_mcp=False, enable_a2a=True)
except asyncio.CancelledError:
pass
assert len(created_tasks) == 1
@pytest.mark.asyncio
async def test_serve_all_services(app: HaikuRAGApp, monkeypatch):
"""Test the serve method with all services enabled."""
created_tasks = []
original_create_task = asyncio.create_task
def track_task(coro):
task = original_create_task(coro)
created_tasks.append(task)
task.cancel()
return task
mock_server = AsyncMock()
mock_watcher = AsyncMock()
mock_a2a_app = MagicMock()
monkeypatch.setattr(
"haiku.rag.app.create_mcp_server", MagicMock(return_value=mock_server)
@ -195,23 +306,22 @@ async def test_serve(app: HaikuRAGApp, monkeypatch, transport):
monkeypatch.setattr(
"haiku.rag.app.FileWatcher", MagicMock(return_value=mock_watcher)
)
monkeypatch.setattr("asyncio.create_task", MagicMock(return_value=mock_task))
monkeypatch.setattr("haiku.rag.app.asyncio.create_task", track_task)
monkeypatch.setattr(
"haiku.rag.app.asyncio.gather", AsyncMock(side_effect=asyncio.CancelledError)
)
mock_client = AsyncMock()
mock_client.__aenter__.return_value = mock_client
with patch("haiku.rag.app.HaikuRAG", return_value=mock_client):
if transport:
await app.serve(transport=transport)
else:
await app.serve()
with patch("haiku.rag.a2a.create_a2a_app", return_value=mock_a2a_app):
try:
await app.serve(enable_monitor=True, enable_mcp=True, enable_a2a=True)
except asyncio.CancelledError:
pass
if transport == "stdio":
mock_server.run_stdio_async.assert_called_once()
else:
mock_server.run_http_async.assert_called_once_with(transport="streamable-http")
mock_task.cancel.assert_called_once()
assert len(created_tasks) == 3
@pytest.mark.asyncio

View file

@ -173,28 +173,123 @@ def test_search():
mock_app_instance.search.assert_called_once_with(query="query", limit=5)
def test_serve():
def test_serve_no_flags():
"""Test serve command fails without flags."""
result = runner.invoke(cli, ["serve"])
assert result.exit_code == 1
assert "At least one service flag" in result.output
def test_serve_mcp_only():
"""Test serve command with MCP only."""
with patch("haiku.rag.app.HaikuRAGApp") as mock_app:
mock_app_instance = MagicMock()
mock_app_instance.serve = AsyncMock()
mock_app.return_value = mock_app_instance
result = runner.invoke(cli, ["serve"])
result = runner.invoke(cli, ["serve", "--mcp"])
assert result.exit_code == 0
mock_app_instance.serve.assert_called_once_with(transport=None)
mock_app_instance.serve.assert_called_once()
_, kwargs = mock_app_instance.serve.call_args
assert kwargs["enable_monitor"] is False
assert kwargs["enable_mcp"] is True
assert kwargs["enable_a2a"] is False
assert kwargs["mcp_transport"] is None
assert kwargs["mcp_port"] == 8001
def test_serve_stdio():
def test_serve_mcp_stdio():
"""Test serve command with MCP stdio transport."""
with patch("haiku.rag.app.HaikuRAGApp") as mock_app:
mock_app_instance = MagicMock()
mock_app_instance.serve = AsyncMock()
mock_app.return_value = mock_app_instance
result = runner.invoke(cli, ["serve", "--stdio"])
result = runner.invoke(cli, ["serve", "--mcp", "--stdio"])
assert result.exit_code == 0
mock_app_instance.serve.assert_called_once_with(transport="stdio")
mock_app_instance.serve.assert_called_once()
_, kwargs = mock_app_instance.serve.call_args
assert kwargs["mcp_transport"] == "stdio"
def test_serve_monitor_only():
"""Test serve command with monitor only."""
with patch("haiku.rag.app.HaikuRAGApp") as mock_app:
mock_app_instance = MagicMock()
mock_app_instance.serve = AsyncMock()
mock_app.return_value = mock_app_instance
result = runner.invoke(cli, ["serve", "--monitor"])
assert result.exit_code == 0
mock_app_instance.serve.assert_called_once()
_, kwargs = mock_app_instance.serve.call_args
assert kwargs["enable_monitor"] is True
assert kwargs["enable_mcp"] is False
assert kwargs["enable_a2a"] is False
def test_serve_a2a_only():
"""Test serve command with A2A only."""
with patch("haiku.rag.app.HaikuRAGApp") as mock_app:
mock_app_instance = MagicMock()
mock_app_instance.serve = AsyncMock()
mock_app.return_value = mock_app_instance
result = runner.invoke(cli, ["serve", "--a2a"])
assert result.exit_code == 0
mock_app_instance.serve.assert_called_once()
_, kwargs = mock_app_instance.serve.call_args
assert kwargs["enable_monitor"] is False
assert kwargs["enable_mcp"] is False
assert kwargs["enable_a2a"] is True
assert kwargs["a2a_host"] == "127.0.0.1"
assert kwargs["a2a_port"] == 8000
def test_serve_all_services():
"""Test serve command with all services."""
with patch("haiku.rag.app.HaikuRAGApp") as mock_app:
mock_app_instance = MagicMock()
mock_app_instance.serve = AsyncMock()
mock_app.return_value = mock_app_instance
result = runner.invoke(cli, ["serve", "--monitor", "--mcp", "--a2a"])
assert result.exit_code == 0
mock_app_instance.serve.assert_called_once()
_, kwargs = mock_app_instance.serve.call_args
assert kwargs["enable_monitor"] is True
assert kwargs["enable_mcp"] is True
assert kwargs["enable_a2a"] is True
def test_serve_custom_ports():
"""Test serve command with custom ports."""
with patch("haiku.rag.app.HaikuRAGApp") as mock_app:
mock_app_instance = MagicMock()
mock_app_instance.serve = AsyncMock()
mock_app.return_value = mock_app_instance
result = runner.invoke(
cli, ["serve", "--mcp", "--mcp-port", "9000", "--a2a", "--a2a-port", "9001"]
)
assert result.exit_code == 0
mock_app_instance.serve.assert_called_once()
_, kwargs = mock_app_instance.serve.call_args
assert kwargs["mcp_port"] == 9000
assert kwargs["a2a_port"] == 9001
def test_serve_stdio_without_mcp():
"""Test serve command fails when --stdio is used without --mcp."""
result = runner.invoke(cli, ["serve", "--stdio", "--monitor"])
assert result.exit_code == 1
assert "--stdio requires --mcp" in result.output
def test_ask():

20
uv.lock
View file

@ -862,6 +862,20 @@ wheels = [
{ url = "https://files.pythonhosted.org/packages/f5/11/02ebebb09ff2104b690457cb7bc6ed700c9e0ce88cf581486bb0a5d3c88b/faker-37.8.0-py3-none-any.whl", hash = "sha256:b08233118824423b5fc239f7dd51f145e7018082b4164f8da6a9994e1f1ae793", size = 1953940, upload-time = "2025-09-15T20:24:11.482Z" },
]
[[package]]
name = "fasta2a"
version = "0.6.0"
source = { registry = "https://pypi.org/simple" }
dependencies = [
{ name = "opentelemetry-api" },
{ name = "pydantic" },
{ name = "starlette" },
]
sdist = { url = "https://files.pythonhosted.org/packages/1b/d1/7a3ab5d4519141978eb47d3f24dff06bc4fa0b39f31e155c1934de95d8e6/fasta2a-0.6.0.tar.gz", hash = "sha256:8078fad9b9dabf7ee4abb3fcb1ca9e5b43bb55c0262be2425bc48cc69f77e963", size = 1436353, upload-time = "2025-10-07T15:08:09.864Z" }
wheels = [
{ url = "https://files.pythonhosted.org/packages/48/14/64899f718727770099f53e8698529fc83ec2a3a4d311270dfb9f6e2bec06/fasta2a-0.6.0-py3-none-any.whl", hash = "sha256:23d49307f6a372e07b9ec9a21187a0864429145e8ded4a41262bd33e2ecaee4c", size = 25403, upload-time = "2025-10-07T15:08:08.196Z" },
]
[[package]]
name = "fastavro"
version = "1.12.0"
@ -1129,6 +1143,9 @@ dependencies = [
]
[package.optional-dependencies]
a2a = [
{ name = "fasta2a" },
]
mxbai = [
{ name = "mxbai-rerank" },
]
@ -1154,6 +1171,7 @@ dev = [
[package.metadata]
requires-dist = [
{ name = "docling", specifier = ">=2.52.0" },
{ name = "fasta2a", marker = "extra == 'a2a'", specifier = ">=0.1.0" },
{ name = "fastmcp", specifier = ">=2.12.3" },
{ name = "httpx", specifier = ">=0.28.1" },
{ name = "lancedb", specifier = ">=0.25.0" },
@ -1168,7 +1186,7 @@ requires-dist = [
{ name = "voyageai", marker = "extra == 'voyageai'", specifier = ">=0.3.5" },
{ name = "watchfiles", specifier = ">=1.1.0" },
]
provides-extras = ["voyageai", "mxbai"]
provides-extras = ["voyageai", "mxbai", "a2a"]
[package.metadata.requires-dev]
dev = [