From 5df01b96f412d65ebce084dc77454785d7eaeac9 Mon Sep 17 00:00:00 2001 From: Gatefixer <313497061+lancedb-gatefixer[bot]@users.noreply.github.com> Date: Wed, 5 Aug 2026 18:10:46 +0000 Subject: [PATCH] fix(python): preserve search query builder types --- python/pyproject.toml | 1 + python/python/lancedb/table.py | 144 ++++++++++++++++++++- python/python/typing_tests/table_search.py | 35 +++++ 3 files changed, 176 insertions(+), 4 deletions(-) create mode 100644 python/python/typing_tests/table_search.py diff --git a/python/pyproject.toml b/python/pyproject.toml index cb175bd6d..453bcc809 100644 --- a/python/pyproject.toml +++ b/python/pyproject.toml @@ -140,6 +140,7 @@ include = [ "python/lancedb/remote/errors.py", "python/lancedb/embeddings/__init__.py", "python/lancedb/_lancedb.pyi", + "python/typing_tests/table_search.py", ] exclude = ["python/tests/"] pythonVersion = "3.13" diff --git a/python/python/lancedb/table.py b/python/python/lancedb/table.py index 31a70c298..adceb83a6 100644 --- a/python/python/lancedb/table.py +++ b/python/python/lancedb/table.py @@ -1350,6 +1350,95 @@ class Table(ABC): return LanceMergeInsertBuilder(self, on) + @overload + def search( + self, + query: None = None, + vector_column_name: Optional[str] = None, + query_type: Literal["auto", "vector", "fts"] = "auto", + ordering_field_name: Optional[str] = None, + fts_columns: Optional[Union[str, List[str]]] = None, + ) -> LanceEmptyQueryBuilder: ... + + @overload + def search( + self, + query: str, + vector_column_name: Optional[str] = None, + query_type: Literal["auto"] = "auto", + ordering_field_name: Optional[str] = None, + fts_columns: Optional[Union[str, List[str]]] = None, + ) -> Union[LanceFtsQueryBuilder, LanceVectorQueryBuilder]: ... + + @overload + def search( + self, + query: FullTextQuery, + vector_column_name: Optional[str] = None, + query_type: Literal["auto"] = "auto", + ordering_field_name: Optional[str] = None, + fts_columns: Optional[Union[str, List[str]]] = None, + ) -> LanceFtsQueryBuilder: ... + + @overload + def search( + self, + query: Union[VEC, "PIL.Image.Image", Tuple], + vector_column_name: Optional[str] = None, + query_type: Literal["auto"] = "auto", + ordering_field_name: Optional[str] = None, + fts_columns: Optional[Union[str, List[str]]] = None, + ) -> LanceVectorQueryBuilder: ... + + @overload + def search( + self, + query: Optional[Union[VEC, str, "PIL.Image.Image", Tuple]] = None, + vector_column_name: Optional[str] = None, + query_type: Literal["vector"] = "vector", + ordering_field_name: Optional[str] = None, + fts_columns: Optional[Union[str, List[str]]] = None, + ) -> LanceVectorQueryBuilder: ... + + @overload + def search( + self, + query: Optional[Union[str, FullTextQuery]] = None, + vector_column_name: Optional[str] = None, + query_type: Literal["fts"] = "fts", + ordering_field_name: Optional[str] = None, + fts_columns: Optional[Union[str, List[str]]] = None, + ) -> LanceFtsQueryBuilder: ... + + @overload + def search( + self, + query: Optional[ + Union[VEC, str, "PIL.Image.Image", Tuple, FullTextQuery] + ] = None, + vector_column_name: Optional[str] = None, + query_type: Literal["hybrid"] = "hybrid", + ordering_field_name: Optional[str] = None, + fts_columns: Optional[Union[str, List[str]]] = None, + ) -> LanceHybridQueryBuilder: ... + + @overload + def search( + self, + query: Optional[ + Union[VEC, str, "PIL.Image.Image", Tuple, FullTextQuery] + ] = None, + vector_column_name: Optional[str] = None, + query_type: QueryType = "auto", + ordering_field_name: Optional[str] = None, + fts_columns: Optional[Union[str, List[str]]] = None, + ) -> Union[ + LanceEmptyQueryBuilder, + LanceFtsQueryBuilder, + LanceHybridQueryBuilder, + LanceVectorQueryBuilder, + ]: ... + @abstractmethod def search( self, @@ -3379,7 +3468,47 @@ class LanceTable(Table): ) @overload - def search( # type: ignore + def search( + self, + query: None = None, + vector_column_name: Optional[str] = None, + query_type: Literal["auto", "vector", "fts"] = "auto", + ordering_field_name: Optional[str] = None, + fts_columns: Optional[Union[str, List[str]]] = None, + ) -> LanceEmptyQueryBuilder: ... + + @overload + def search( + self, + query: str, + vector_column_name: Optional[str] = None, + query_type: Literal["auto"] = "auto", + ordering_field_name: Optional[str] = None, + fts_columns: Optional[Union[str, List[str]]] = None, + ) -> Union[LanceFtsQueryBuilder, LanceVectorQueryBuilder]: ... + + @overload + def search( + self, + query: FullTextQuery, + vector_column_name: Optional[str] = None, + query_type: Literal["auto"] = "auto", + ordering_field_name: Optional[str] = None, + fts_columns: Optional[Union[str, List[str]]] = None, + ) -> LanceFtsQueryBuilder: ... + + @overload + def search( + self, + query: Union[VEC, "PIL.Image.Image", Tuple], + vector_column_name: Optional[str] = None, + query_type: Literal["auto"] = "auto", + ordering_field_name: Optional[str] = None, + fts_columns: Optional[Union[str, List[str]]] = None, + ) -> LanceVectorQueryBuilder: ... + + @overload + def search( self, query: Optional[Union[VEC, str, "PIL.Image.Image", Tuple]] = None, vector_column_name: Optional[str] = None, @@ -3391,7 +3520,7 @@ class LanceTable(Table): @overload def search( self, - query: Optional[Union[VEC, str, "PIL.Image.Image", Tuple]] = None, + query: Optional[Union[str, FullTextQuery]] = None, vector_column_name: Optional[str] = None, query_type: Literal["fts"] = "fts", ordering_field_name: Optional[str] = None, @@ -3413,12 +3542,19 @@ class LanceTable(Table): @overload def search( self, - query: None = None, + query: Optional[ + Union[VEC, str, "PIL.Image.Image", Tuple, FullTextQuery] + ] = None, vector_column_name: Optional[str] = None, query_type: QueryType = "auto", ordering_field_name: Optional[str] = None, fts_columns: Optional[Union[str, List[str]]] = None, - ) -> LanceEmptyQueryBuilder: ... + ) -> Union[ + LanceEmptyQueryBuilder, + LanceFtsQueryBuilder, + LanceHybridQueryBuilder, + LanceVectorQueryBuilder, + ]: ... def search( self, diff --git a/python/python/typing_tests/table_search.py b/python/python/typing_tests/table_search.py new file mode 100644 index 000000000..a21206c13 --- /dev/null +++ b/python/python/typing_tests/table_search.py @@ -0,0 +1,35 @@ +# SPDX-License-Identifier: Apache-2.0 +# SPDX-FileCopyrightText: Copyright The LanceDB Authors + +from typing import assert_type + +from lancedb.db import DBConnection +from lancedb.query import ( + LanceEmptyQueryBuilder, + LanceFtsQueryBuilder, + LanceHybridQueryBuilder, + LanceVectorQueryBuilder, +) +from lancedb.table import LanceTable + + +def check_table_search_types(connection: DBConnection, lance_table: LanceTable) -> None: + table = connection.open_table("table") + + assert_type(table.search(), LanceEmptyQueryBuilder) + assert_type(table.search([1.0, 2.0]), LanceVectorQueryBuilder) + assert_type( + table.search("query"), + LanceFtsQueryBuilder | LanceVectorQueryBuilder, + ) + assert_type(table.search("query", query_type="vector"), LanceVectorQueryBuilder) + assert_type(table.search("query", query_type="fts"), LanceFtsQueryBuilder) + assert_type(table.search("query", query_type="hybrid"), LanceHybridQueryBuilder) + assert_type(table.search("query", None, "vector"), LanceVectorQueryBuilder) + assert_type(table.search("query", None, "fts"), LanceFtsQueryBuilder) + assert_type(table.search("query", None, "hybrid"), LanceHybridQueryBuilder) + + assert_type( + lance_table.search("query"), + LanceFtsQueryBuilder | LanceVectorQueryBuilder, + )