diff --git a/README.md b/README.md index 681625c..2c18642 100644 --- a/README.md +++ b/README.md @@ -140,6 +140,9 @@ dremio query run "SELECT * FROM myspace.orders LIMIT 5" --output pretty # Search the catalog for anything matching "revenue" dremio search "revenue" +# Search for anything matching either term +dremio search "revenue" "sales" + # Search only jobs and limit the first page size dremio search "revenue" --filter 'category in ["JOB"]' --max-results 20 diff --git a/src/drs/cli.py b/src/drs/cli.py index 6e2f171..ece2758 100644 --- a/src/drs/cli.py +++ b/src/drs/cli.py @@ -171,7 +171,10 @@ def get_client() -> DremioClient: @app.command("search") def search_command( - term: str = typer.Argument(help="Search term (matches table names, view names, source names)"), + terms: list[str] = typer.Argument( + help="One or more search terms (matches table names, view names, source names). " + "Multiple terms are searched as alternatives." + ), filter_: str | None = typer.Option(None, "--filter", help="CEL filter expression to refine search results"), max_results: int | None = typer.Option(None, "--max-results", min=1, help="Maximum results to return per page"), next_page_token: str | None = typer.Option( @@ -194,7 +197,7 @@ async def _execute(): } if next_page_token is not None: search_kwargs["next_page_token"] = next_page_token - return await client.search(term, **search_kwargs) + return await client.search(terms, **search_kwargs) except httpx.HTTPStatusError as exc: raise handle_api_error(exc) from exc finally: diff --git a/src/drs/client.py b/src/drs/client.py index 5cfc19a..623d665 100644 --- a/src/drs/client.py +++ b/src/drs/client.py @@ -271,12 +271,12 @@ async def get_catalog_by_path(self, path_parts: list[str]) -> dict: async def search( self, - query: str, + queries: list[str], filter_: str | None = None, max_results: int | None = None, next_page_token: str | None = None, ) -> dict: - body: dict[str, Any] = {"query": query} + body: dict[str, Any] = {"queries": queries} if filter_: body["filter"] = filter_ if max_results is not None: diff --git a/src/drs/introspect.py b/src/drs/introspect.py index c88d174..b2ee9e3 100644 --- a/src/drs/introspect.py +++ b/src/drs/introspect.py @@ -211,7 +211,13 @@ "mechanism": "REST", "endpoints": ["POST /v0/projects/{pid}/search"], "parameters": [ - {"name": "term", "type": "string", "required": True, "positional": True, "description": "Search term"}, + { + "name": "terms", + "type": "string", + "required": True, + "positional": True, + "description": "One or more search terms, searched as alternatives", + }, { "name": "filter", "type": "string", diff --git a/tests/test_cli.py b/tests/test_cli.py index 85ca6d4..add1c9c 100644 --- a/tests/test_cli.py +++ b/tests/test_cli.py @@ -231,7 +231,7 @@ def test_search_command_passes_filter_and_max_results(monkeypatch) -> None: result = runner.invoke(app, ["search", "revenue", "--filter", 'category in ["JOB"]', "--max-results", "20"]) assert result.exit_code == 0 - search_mock.assert_awaited_once_with("revenue", filter_='category in ["JOB"]', max_results=20) + search_mock.assert_awaited_once_with(["revenue"], filter_='category in ["JOB"]', max_results=20) close_mock.assert_awaited_once() @@ -247,5 +247,27 @@ def test_search_command_passes_next_page_token(monkeypatch) -> None: result = runner.invoke(app, ["search", "revenue", "--next-page-token", "token-123"]) assert result.exit_code == 0 - search_mock.assert_awaited_once_with("revenue", filter_=None, max_results=None, next_page_token="token-123") + search_mock.assert_awaited_once_with(["revenue"], filter_=None, max_results=None, next_page_token="token-123") close_mock.assert_awaited_once() + + +def test_search_command_passes_multiple_terms(monkeypatch) -> None: + search_mock = AsyncMock(return_value={"results": []}) + close_mock = AsyncMock() + client = MagicMock() + client.search = search_mock + client.close = close_mock + + monkeypatch.setattr("drs.cli.get_client", lambda: client) + + result = runner.invoke(app, ["search", "query a", "query b"]) + + assert result.exit_code == 0 + search_mock.assert_awaited_once_with(["query a", "query b"], filter_=None, max_results=None) + close_mock.assert_awaited_once() + + +def test_search_command_requires_a_term() -> None: + result = runner.invoke(app, ["search"]) + + assert result.exit_code != 0 diff --git a/tests/test_client.py b/tests/test_client.py index 4a38c6f..3f221de 100644 --- a/tests/test_client.py +++ b/tests/test_client.py @@ -116,10 +116,10 @@ async def _capture(request: httpx.Request) -> httpx.Response: client._client = httpx.AsyncClient(transport=httpx.MockTransport(_capture)) - await client.search("revenue", filter_='category in ["JOB"]', max_results=20) + await client.search(["revenue"], filter_='category in ["JOB"]', max_results=20) assert captured["body"] == { - "query": "revenue", + "queries": ["revenue"], "filter": 'category in ["JOB"]', "maxResults": 20, } @@ -136,13 +136,40 @@ async def _capture(request: httpx.Request) -> httpx.Response: client._client = httpx.AsyncClient(transport=httpx.MockTransport(_capture)) - await client.search("revenue", next_page_token="token-123") + await client.search(["revenue"], next_page_token="token-123") assert captured["body"] == { - "query": "revenue", + "queries": ["revenue"], "pageToken": "token-123", } + @pytest.mark.parametrize( + ("queries", "expected_body"), + [ + (["revenue"], {"queries": ["revenue"]}), + (["revenue", "sales"], {"queries": ["revenue", "sales"]}), + ([" ", "sales"], {"queries": [" ", "sales"]}), + ([""], {"queries": [""]}), + ], + ) + @pytest.mark.asyncio + async def test_search_sends_queries_list_and_legacy_query( + self, client: DremioClient, queries: list[str], expected_body: dict + ) -> None: + captured: dict = {} + + async def _capture(request: httpx.Request) -> httpx.Response: + import json + + captured["body"] = json.loads(request.content) + return httpx.Response(200, json={"results": []}) + + client._client = httpx.AsyncClient(transport=httpx.MockTransport(_capture)) + + await client.search(queries) + + assert captured["body"] == expected_body + class TestSQLBreadcrumb: @pytest.mark.asyncio