mirror of
https://github.com/lancedb/lancedb.git
synced 2026-08-18 12:08:35 +00:00
fix(python): preserve search query builder types
This commit is contained in:
@@ -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"
|
||||
|
||||
@@ -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,
|
||||
)
|
||||
Reference in New Issue
Block a user