fix: declare 'SearchType' for 'search_type' arg

fix: allow 'search_type' only for text searches
This commit is contained in:
Tres Seaver 2026-05-18 06:33:04 -04:00
parent b2ffc50168
commit 7953de46b9
No known key found for this signature in database
GPG key ID: F2A968348913D1D8
5 changed files with 26 additions and 8 deletions

View file

@ -19,6 +19,7 @@ from rich.progress import (
from rich.syntax import Syntax
from haiku.rag.client import HaikuRAG, RebuildMode
from haiku.rag.client.models import SearchType
from haiku.rag.config import AppConfig, Config
from haiku.rag.mcp import create_mcp_server
from haiku.rag.monitor import FileWatcher, S3Watcher
@ -361,7 +362,7 @@ class HaikuRAGApp: # pragma: no cover
query: str | None = None,
limit: int | None = None,
filter: str | None = None,
search_type: str = "hybrid",
search_type: SearchType | None = None,
image: Path | None = None,
):
if query is None and image is None:
@ -373,6 +374,10 @@ class HaikuRAGApp: # pragma: no cover
self.console.print("[red]Pass either a query or --image, not both.[/red]")
return
if query is None and search_type is not None:
self.console.print("[red]Pass --search-type only for text queries[/red]")
return
search_input: str | bytes
if image is not None:
search_input = image.read_bytes()
@ -387,7 +392,10 @@ class HaikuRAGApp: # pragma: no cover
before=self.before,
) as self.client:
results = await self.client.search(
search_input, limit=limit, filter=filter, search_type=search_type,
search_input,
limit=limit,
filter=filter,
search_type=search_type,
)
if not results:
self.console.print("[yellow]No results found.[/yellow]")

View file

@ -15,6 +15,7 @@ from dotenv import find_dotenv, load_dotenv
load_dotenv(find_dotenv(usecwd=True))
from haiku.rag.app import HaikuRAGApp # noqa: E402
from haiku.rag.client.models import SearchType # noqa: E402
from haiku.rag.config import ( # noqa: E402
AppConfig,
find_config_file,
@ -314,11 +315,11 @@ def search( # pragma: no cover
"-f",
help="SQL WHERE clause to filter documents (e.g., \"uri LIKE '%arxiv%'\")",
),
search_type: str | None = typer.Option(
None,
search_type: SearchType = typer.Option(
"hybrid",
"--search-type",
"-s",
help="Type of search: one of 'hybrid' (default) / 'fts' / 'vector' ",
help="Type of search to perform (text searches only)",
),
image: Path | None = typer.Option(
None,
@ -334,7 +335,11 @@ def search( # pragma: no cover
app = create_app(db)
asyncio.run(
app.search(
query=query, limit=limit, filter=filter, search_type=search_type, image=image,
query=query,
limit=limit,
filter=filter,
search_type=search_type,
image=image,
)
)

View file

@ -13,6 +13,7 @@ from urllib.parse import urlparse
import httpx
from haiku.rag.client.models import SearchType
from haiku.rag.config import AppConfig, Config
from haiku.rag.converters import get_converter
from haiku.rag.reranking import get_reranker
@ -349,7 +350,7 @@ class HaikuRAG:
self,
query: "str | bytes | PILImage.Image",
limit: int | None = None,
search_type: str = "hybrid",
search_type: SearchType | None = None,
filter: str | None = None,
include_images: bool = True,
) -> list[SearchResult]:

View file

@ -0,0 +1,3 @@
from typing import Literal
SearchType = Literal["vector", "fts", "hybrid"]

View file

@ -1,6 +1,7 @@
import base64
from typing import TYPE_CHECKING
from haiku.rag.client.models import SearchType
from haiku.rag.reranking import get_reranker
from haiku.rag.store.models.chunk import Chunk, SearchResult
@ -14,7 +15,7 @@ async def search(
client: "HaikuRAG",
query: "str | bytes | PILImage.Image",
limit: int | None = None,
search_type: str = "hybrid",
search_type: SearchType = "hybrid",
filter: str | None = None,
include_images: bool = True,
) -> list[SearchResult]: