fix(python): preserve search query builder types

This commit is contained in:
Gatefixer
2026-08-05 18:10:46 +00:00
parent c7ea91f3ea
commit 5df01b96f4
3 changed files with 176 additions and 4 deletions
+1
View File
@@ -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"
+140 -4
View File
@@ -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,
@@ -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,
)