diff --git a/mock_tests/test_collection.py b/mock_tests/test_collection.py index 4910ae2fd..1cbc5e1b5 100644 --- a/mock_tests/test_collection.py +++ b/mock_tests/test_collection.py @@ -1,6 +1,6 @@ import datetime import json -from typing import Any, Dict, List, Literal +from typing import Any, Dict, List, Literal, Mapping, Optional import grpc import pytest @@ -42,6 +42,7 @@ UnexpectedStatusCodeError, WeaviateStartUpError, ) +from weaviate.proto.v1 import properties_pb2, search_get_pb2, weaviate_pb2_grpc ACCESS_TOKEN = "HELLO!IamAnAccessToken" REFRESH_TOKEN = "UseMeToRefreshYourAccessToken" @@ -639,3 +640,54 @@ def test_grpc_client_version_header( assert "x-weaviate-client" in service.captured_metadata expected = f"weaviate-client-python/{client_version}-sync" assert service.captured_metadata["x-weaviate-client"] == expected + + +class MockSearchCaptureWeaviateService(weaviate_pb2_grpc.WeaviateServicer): + captured_request: Optional[search_get_pb2.SearchRequest] = None + + def Search( + self, request: search_get_pb2.SearchRequest, context: grpc.ServicerContext + ) -> search_get_pb2.SearchReply: + self.captured_request = request + tags: Mapping[str, properties_pb2.Value] = { + "tags": properties_pb2.Value( + list_value=properties_pb2.ListValue( + text_values=properties_pb2.TextValues(values=["tag1", "tag2"]) + ) + ) + } + return search_get_pb2.SearchReply( + results=[ + search_get_pb2.SearchResult( + properties=search_get_pb2.PropertiesResult( + non_ref_props=properties_pb2.Properties(fields=tags) + ) + ) + ] + ) + + +@pytest.mark.asyncio +async def test_async_collection_made_before_connect( + weaviate_mock: HTTPServer, start_grpc_server: grpc.Server +) -> None: + # The async client only learns the server version inside `connect()`, so a collection object + # made beforehand used to send every request as if the server were older than 1.25, see + # issue #1831. The flags on the captured request are what this pins down; the legacy list + # encoding that originally made the reply unparseable no longer exists in the generated + # protos, so the parsed property below is only a round-trip sanity check. + service = MockSearchCaptureWeaviateService() + weaviate_pb2_grpc.add_WeaviateServicer_to_server(service, start_grpc_server) + + client = weaviate.use_async_with_local(port=MOCK_PORT, host=MOCK_IP, grpc_port=MOCK_PORT_GRPC) + collection = client.collections.use("TestCollection") + await client.connect() + try: + objects = (await collection.query.fetch_objects()).objects + finally: + await client.close() + + assert service.captured_request is not None + assert service.captured_request.uses_125_api is True + assert service.captured_request.uses_127_api is True + assert objects[0].properties["tags"] == ["tag1", "tag2"] diff --git a/test/collection/test_bm25_operator.py b/test/collection/test_bm25_operator.py index 9fa188b22..511ee044d 100644 --- a/test/collection/test_bm25_operator.py +++ b/test/collection/test_bm25_operator.py @@ -16,8 +16,6 @@ def _builder(version: str = "1.39.0") -> _QueryGRPC: tenant=None, consistency_level=None, validate_arguments=True, - uses_125_api=True, - uses_127_api=True, ) diff --git a/test/collection/test_hybrid_diversity.py b/test/collection/test_hybrid_diversity.py index 510b4df23..5ebdad950 100644 --- a/test/collection/test_hybrid_diversity.py +++ b/test/collection/test_hybrid_diversity.py @@ -21,8 +21,6 @@ def _builder(version: _ServerVersion = _DEFAULT_VERSION) -> _QueryGRPC: tenant=None, consistency_level=None, validate_arguments=True, - uses_125_api=True, - uses_127_api=True, ) diff --git a/test/collection/test_queries.py b/test/collection/test_queries.py index 513764a17..8fcfcc43a 100644 --- a/test/collection/test_queries.py +++ b/test/collection/test_queries.py @@ -1,10 +1,14 @@ -from typing import Awaitable +from typing import Awaitable, Optional import pytest +from weaviate.collections.classes.internal import _QueryOptions +from weaviate.collections.queries.base_executor import _BaseExecutor from weaviate.collections.query import _QueryCollectionAsync from weaviate.connect import ConnectionV4 from weaviate.exceptions import WeaviateInvalidInputError +from weaviate.proto.v1 import generative_pb2, search_get_pb2 +from weaviate.util import _ServerVersion # TODO: re-enable tests once string syntax is re-enabled in the API @@ -130,3 +134,71 @@ async def test_bad_query_inputs(connection: ConnectionV4) -> None: # near image await _test_query(lambda: query.near_image(42)) + + +@pytest.mark.parametrize( + "version,uses_125_api,uses_127_api", + [ + ("1.24.0", False, False), + ("1.26.0", True, False), + ("1.27.0", True, True), + ("1.32.5", True, True), + ], +) +def test_query_uses_version_learned_after_construction( + connection: ConnectionV4, version: str, uses_125_api: bool, uses_127_api: bool +) -> None: + # The async client only learns the server version inside `connect()`, so a collection object + # made before `await client.connect()` is built against version 0.0.0, see issue #1831. + query = _QueryCollectionAsync(connection, "dummy", None, None, None, None, True) + assert connection._weaviate_version == _ServerVersion(0, 0, 0) + + connection._weaviate_version = _ServerVersion.from_string(version) + + request = query._query.get() + assert request.uses_125_api is uses_125_api + assert request.uses_127_api is uses_127_api + + +def test_query_rereads_version_on_every_access(connection: ConnectionV4) -> None: + # The version is resolved per call rather than cached on first use, so a connection that + # starts talking to a different server is picked up as well. + query = _QueryCollectionAsync(connection, "dummy", None, None, None, None, True) + + connection._weaviate_version = _ServerVersion(1, 26, 0) + assert query._query.get().uses_127_api is False + + connection._weaviate_version = _ServerVersion(1, 32, 5) + assert query._query.get().uses_127_api is True + + +@pytest.mark.parametrize("version,generated", [("1.26.0", None), ("1.32.5", "generated")]) +def test_query_generative_uses_version_learned_after_construction( + connection: ConnectionV4, version: str, generated: Optional[str] +) -> None: + executor = _BaseExecutor(connection, "dummy", None, None, None, None, True) + + connection._weaviate_version = _ServerVersion.from_string(version) + + # The generated text is only present in the generative field, which is the field that servers + # from 1.27 onwards fill, so a stale version reads the empty deprecated metadata field instead. + response = search_get_pb2.SearchReply( + results=[ + search_get_pb2.SearchResult( + generative=generative_pb2.GenerativeResult( + values=[generative_pb2.GenerativeReply(result="generated")] + ) + ) + ] + ) + result = executor._result_to_generative_query_return( + response, + _QueryOptions( + include_metadata=False, + include_properties=False, + include_references=False, + include_vector=False, + is_group_by=False, + ), + ) + assert result.objects[0].generated == generated diff --git a/test/collection/test_target_vectors.py b/test/collection/test_target_vectors.py index 07b4ee74e..f1d0a26e4 100644 --- a/test/collection/test_target_vectors.py +++ b/test/collection/test_target_vectors.py @@ -11,15 +11,12 @@ def _query(version: str) -> _QueryGRPC: - weaviate_version = _ServerVersion.from_string(version) return _QueryGRPC( - weaviate_version=weaviate_version, + weaviate_version=_ServerVersion.from_string(version), name="Documents", tenant=None, consistency_level=None, validate_arguments=True, - uses_125_api=weaviate_version.is_at_least(1, 25, 0), - uses_127_api=weaviate_version.is_at_least(1, 27, 0), ) diff --git a/weaviate/collections/grpc/query.py b/weaviate/collections/grpc/query.py index 281da6e2d..5beef49b8 100644 --- a/weaviate/collections/grpc/query.py +++ b/weaviate/collections/grpc/query.py @@ -84,15 +84,11 @@ def __init__( tenant: Optional[str], consistency_level: Optional[ConsistencyLevel], validate_arguments: bool, - uses_125_api: bool, - uses_127_api: bool, ): super().__init__(weaviate_version, consistency_level, validate_arguments) self._name: str = name self._tenant = tenant self._validate_arguments = validate_arguments - self.__uses_125_api = uses_125_api - self.__uses_127_api = uses_127_api def __parse_near_options( self, @@ -501,8 +497,8 @@ def __create_request( return search_get_pb2.SearchRequest( uses_123_api=True, - uses_125_api=self.__uses_125_api, - uses_127_api=self.__uses_127_api, + uses_125_api=self._weaviate_version.is_at_least(1, 25, 0), + uses_127_api=self._weaviate_version.is_at_least(1, 27, 0), collection=self._name, limit=limit, offset=offset, diff --git a/weaviate/collections/queries/base_executor.py b/weaviate/collections/queries/base_executor.py index 1851d0dc1..72a32b8f0 100644 --- a/weaviate/collections/queries/base_executor.py +++ b/weaviate/collections/queries/base_executor.py @@ -86,16 +86,18 @@ def __init__( self._references = references self._validate_arguments = validate_arguments - self.__uses_125_api = connection._weaviate_version.is_at_least(1, 25, 0) - self.__uses_127_api = connection._weaviate_version.is_at_least(1, 27, 0) - self._query = _QueryGRPC( - connection._weaviate_version, + @property + def _query(self) -> _QueryGRPC: + # The server version is only known once the connection is open, which can happen after + # this object is created, e.g. `client.collections.use(...)` before `await client.connect()` + # with the async client. Build the request factory per call so that it always reflects the + # version of the server that the connection is talking to. + return _QueryGRPC( + self._connection._weaviate_version, self._name, self.__tenant, self.__consistency_level, validate_arguments=self._validate_arguments, - uses_125_api=self.__uses_125_api, - uses_127_api=self.__uses_127_api, ) def __retrieve_timestamp( @@ -391,7 +393,7 @@ def __result_to_generative_object( vector=(self.__extract_vector_for_object(meta) if options.include_vector else {}), generated=( self.__extract_generated_from_generative(gen) - if self.__uses_127_api + if self._connection._weaviate_version.is_at_least(1, 27, 0) else self.__extract_generated_from_metadata(meta) ), generative=self.__extract_generative_single_from_generative(gen),