Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
54 changes: 53 additions & 1 deletion mock_tests/test_collection.py
Original file line number Diff line number Diff line change
@@ -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
Expand Down Expand Up @@ -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"
Expand Down Expand Up @@ -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"]
2 changes: 0 additions & 2 deletions test/collection/test_bm25_operator.py
Original file line number Diff line number Diff line change
Expand Up @@ -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,
)


Expand Down
2 changes: 0 additions & 2 deletions test/collection/test_hybrid_diversity.py
Original file line number Diff line number Diff line change
Expand Up @@ -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,
)


Expand Down
74 changes: 73 additions & 1 deletion test/collection/test_queries.py
Original file line number Diff line number Diff line change
@@ -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

Expand Down Expand Up @@ -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
5 changes: 1 addition & 4 deletions test/collection/test_target_vectors.py
Original file line number Diff line number Diff line change
Expand Up @@ -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),
)


Expand Down
8 changes: 2 additions & 6 deletions weaviate/collections/grpc/query.py
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand Down Expand Up @@ -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,
Expand Down
16 changes: 9 additions & 7 deletions weaviate/collections/queries/base_executor.py
Original file line number Diff line number Diff line change
Expand Up @@ -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(
Expand Down Expand Up @@ -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),
Expand Down
Loading