diff --git a/README.md b/README.md index c97c15c7..782bee13 100644 --- a/README.md +++ b/README.md @@ -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/ diff --git a/docs/a2a.md b/docs/a2a.md new file mode 100644 index 00000000..c6e6ab2f --- /dev/null +++ b/docs/a2a.md @@ -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 diff --git a/docs/cli.md b/docs/cli.md index afe303a1..3663ea32 100644 --- a/docs/cli.md +++ b/docs/cli.md @@ -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: diff --git a/docs/configuration.md b/docs/configuration.md index 631899ae..c77b140e 100644 --- a/docs/configuration.md +++ b/docs/configuration.md @@ -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" ``` diff --git a/docs/index.md b/docs/index.md index 648eb076..f3bc543d 100644 --- a/docs/index.md +++ b/docs/index.md @@ -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 human‑readable 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 diff --git a/docs/server.md b/docs/server.md index 15f862b1..4a88a478 100644 --- a/docs/server.md +++ b/docs/server.md @@ -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 diff --git a/examples/README.md b/examples/README.md new file mode 100644 index 00000000..2e00c386 --- /dev/null +++ b/examples/README.md @@ -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. diff --git a/examples/a2a-security/apikey_example.py b/examples/a2a-security/apikey_example.py new file mode 100644 index 00000000..c37981ba --- /dev/null +++ b/examples/a2a-security/apikey_example.py @@ -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 ") + 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) diff --git a/examples/a2a-security/oauth2_example.py b/examples/a2a-security/oauth2_example.py new file mode 100644 index 00000000..6caa4627 --- /dev/null +++ b/examples/a2a-security/oauth2_example.py @@ -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 ") + 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) diff --git a/examples/a2a-security/oauth2_github.py b/examples/a2a-security/oauth2_github.py new file mode 100644 index 00000000..f2b10a41 --- /dev/null +++ b/examples/a2a-security/oauth2_github.py @@ -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 ") + 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) diff --git a/mkdocs.yml b/mkdocs.yml index 2edd9719..34b54627 100644 --- a/mkdocs.yml +++ b/mkdocs.yml @@ -64,6 +64,7 @@ nav: - Agents: agents.md - Python: python.md - MCP: mcp.md + - A2A: a2a.md - Benchmarks: benchmarks.md markdown_extensions: - admonition diff --git a/pyproject.toml b/pyproject.toml index 203b0c7d..15caf15c 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -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"] diff --git a/src/haiku/rag/a2a/__init__.py b/src/haiku/rag/a2a/__init__.py new file mode 100644 index 00000000..f4481fbc --- /dev/null +++ b/src/haiku/rag/a2a/__init__.py @@ -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 diff --git a/src/haiku/rag/a2a/client.py b/src/haiku/rag/a2a/client.py new file mode 100644 index 00000000..a522d966 --- /dev/null +++ b/src/haiku/rag/a2a/client.py @@ -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() diff --git a/src/haiku/rag/a2a/context.py b/src/haiku/rag/a2a/context.py new file mode 100644 index 00000000..a91478d2 --- /dev/null +++ b/src/haiku/rag/a2a/context.py @@ -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()), + ) diff --git a/src/haiku/rag/a2a/models.py b/src/haiku/rag/a2a/models.py new file mode 100644 index 00000000..52e9ebfa --- /dev/null +++ b/src/haiku/rag/a2a/models.py @@ -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 diff --git a/src/haiku/rag/a2a/prompts.py b/src/haiku/rag/a2a/prompts.py new file mode 100644 index 00000000..22b07a9f --- /dev/null +++ b/src/haiku/rag/a2a/prompts.py @@ -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). +""" diff --git a/src/haiku/rag/a2a/skills.py b/src/haiku/rag/a2a/skills.py new file mode 100644 index 00000000..4dbaedde --- /dev/null +++ b/src/haiku/rag/a2a/skills.py @@ -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 diff --git a/src/haiku/rag/a2a/storage.py b/src/haiku/rag/a2a/storage.py new file mode 100644 index 00000000..1a8e57e4 --- /dev/null +++ b/src/haiku/rag/a2a/storage.py @@ -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) diff --git a/src/haiku/rag/a2a/worker.py b/src/haiku/rag/a2a/worker.py new file mode 100644 index 00000000..dbda76ef --- /dev/null +++ b/src/haiku/rag/a2a/worker.py @@ -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 diff --git a/src/haiku/rag/app.py b/src/haiku/rag/app.py index ebf18515..8e7c0310 100644 --- a/src/haiku/rag/app.py +++ b/src/haiku/rag/app.py @@ -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) diff --git a/src/haiku/rag/cli.py b/src/haiku/rag/cli.py index ab46e3ac..8ad370e5 100644 --- a/src/haiku/rag/cli.py +++ b/src/haiku/rag/cli.py @@ -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() diff --git a/src/haiku/rag/config.py b/src/haiku/rag/config.py index 40fcc8cd..b108ac32 100644 --- a/src/haiku/rag/config.py +++ b/src/haiku/rag/config.py @@ -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): diff --git a/src/haiku/rag/qa/agent.py b/src/haiku/rag/qa/agent.py index e8157205..f1abe7ab 100644 --- a/src/haiku/rag/qa/agent.py +++ b/src/haiku/rag/qa/agent.py @@ -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) diff --git a/tests/test_a2a.py b/tests/test_a2a.py new file mode 100644 index 00000000..c0b34817 --- /dev/null +++ b/tests/test_a2a.py @@ -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" diff --git a/tests/test_app.py b/tests/test_app.py index 6cd549c9..451e7026 100644 --- a/tests/test_app.py +++ b/tests/test_app.py @@ -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 diff --git a/tests/test_cli.py b/tests/test_cli.py index 72aa361e..a428b5ad 100644 --- a/tests/test_cli.py +++ b/tests/test_cli.py @@ -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(): diff --git a/uv.lock b/uv.lock index 933861b0..d1e71cb0 100644 --- a/uv.lock +++ b/uv.lock @@ -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 = [