diff --git a/docs/cli.md b/docs/cli.md index afe303a1..57da7f54 100644 --- a/docs/cli.md +++ b/docs/cli.md @@ -125,15 +125,29 @@ 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. + ## Settings View current configuration settings: diff --git a/docs/server.md b/docs/server.md index c1e7d7be..4a88a478 100644 --- a/docs/server.md +++ b/docs/server.md @@ -4,17 +4,20 @@ The server provides automatic file monitoring, MCP functionality, and A2A agent ## Starting the Server -### MCP Server (Default) +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 +### A2A Server Only ```bash haiku-rag serve --a2a @@ -26,6 +29,20 @@ Options: 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 Set `MONITOR_DIRECTORIES` environment variable to enable automatic file monitoring: 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 d89a1a57..06ce8631 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 server (MCP by default, or A2A with --a2a)" + "serve", + help="Start haiku.rag server. Use --monitor, --mcp, and/or --a2a to enable services.", ) def serve( db: Path = typer.Option( @@ -374,15 +375,30 @@ 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="Run A2A (Agent-to-Agent) server instead of MCP", + help="Enable A2A (Agent-to-Agent) server", ), a2a_host: str = typer.Option( "127.0.0.1", @@ -395,29 +411,35 @@ def serve( help="Port to bind A2A server to", ), ) -> None: - """Start the MCP or A2A server.""" - if a2a: - try: - from haiku.rag.a2a import create_a2a_app - except ImportError as e: - typer.echo(f"Error: {e}") - raise typer.Exit(1) + """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) - import uvicorn + if stdio and not mcp: + typer.echo("Error: --stdio requires --mcp") + raise typer.Exit(1) - typer.echo(f"Starting A2A server on {a2a_host}:{a2a_port}") - app = create_a2a_app(db_path=db) - uvicorn.run(app, host=a2a_host, port=a2a_port) - else: - from haiku.rag.app import HaikuRAGApp + from haiku.rag.app import HaikuRAGApp - app = HaikuRAGApp(db_path=db) + 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") 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():