mirror of
https://github.com/qdrant/qdrant.git
synced 2026-07-29 06:01:05 -05:00
* Return HTTP 405 on cluster endpoints when running in standalone mode * Add test * Skip some tests if not running in distributed mode
223 lines
7.4 KiB
Python
223 lines
7.4 KiB
Python
import pytest
|
|
import json
|
|
import time
|
|
from typing import Any, Dict, List
|
|
import jsonschema
|
|
import requests
|
|
import warnings
|
|
from schemathesis.exceptions import CheckFailed
|
|
from schemathesis.models import APIOperation
|
|
from schemathesis.specs.openapi.references import ConvertingResolver
|
|
from schemathesis.specs.openapi.schemas import OpenApi30
|
|
from functools import lru_cache
|
|
|
|
from .settings import QDRANT_HOST, SCHEMA, QDRANT_HOST_HEADERS
|
|
|
|
|
|
def get_api_string(host, api, path_params):
|
|
"""
|
|
>>> get_api_string('http://localhost:6333', '/collections/{name}', {'name': 'hello', 'a': 'b'})
|
|
'http://localhost:6333/collections/hello'
|
|
"""
|
|
return f"{host}{api}".format(**path_params)
|
|
|
|
|
|
def validate_schema(data, operation_schema: OpenApi30, raw_definitions):
|
|
"""
|
|
:param data: concrete values to validate
|
|
:param operation_schema: operation schema
|
|
:param raw_definitions: definitions to check data with
|
|
:return:
|
|
"""
|
|
resolver = ConvertingResolver(
|
|
operation_schema.location or "",
|
|
operation_schema.raw_schema,
|
|
nullable_name=operation_schema.nullable_name,
|
|
is_response_schema=False
|
|
)
|
|
jsonschema.validate(data, raw_definitions, cls=jsonschema.Draft7Validator, resolver=resolver)
|
|
|
|
|
|
def request_with_validation(
|
|
api: str,
|
|
method: str,
|
|
path_params: dict = None,
|
|
query_params: dict = None,
|
|
body: dict = None
|
|
) -> requests.Response:
|
|
operation: APIOperation = SCHEMA[api][method]
|
|
|
|
assert isinstance(operation.schema, OpenApi30)
|
|
|
|
if body:
|
|
validate_schema(
|
|
data=body,
|
|
operation_schema=operation.schema,
|
|
raw_definitions=operation.definition.raw['requestBody']['content']['application/json']['schema']
|
|
)
|
|
|
|
if path_params is None:
|
|
path_params = {}
|
|
if query_params is None:
|
|
query_params = {}
|
|
action = getattr(requests, method.lower(), None)
|
|
|
|
for param in operation.path_parameters.items:
|
|
if param.is_required:
|
|
assert param.name in path_params
|
|
|
|
for param in operation.query.items:
|
|
if param.is_required:
|
|
assert param.name in query_params
|
|
|
|
for param in path_params.keys():
|
|
assert param in set(p.name for p in operation.path_parameters.items)
|
|
|
|
for param in query_params.keys():
|
|
assert param in set(p.name for p in operation.query.items)
|
|
|
|
if not action:
|
|
raise RuntimeError(f"Method {method} does not exists")
|
|
|
|
if api.endswith("/delete") and method == "POST" and "wait" not in query_params:
|
|
warnings.warn(f"Delete call for {api} missing wait=true param, adding it")
|
|
query_params["wait"] = "true"
|
|
|
|
start_time = time.time()
|
|
response = action(
|
|
url=get_api_string(QDRANT_HOST, api, path_params),
|
|
params=query_params,
|
|
json=body,
|
|
headers=qdrant_host_headers()
|
|
)
|
|
duration = time.time() - start_time
|
|
try:
|
|
operation.validate_response(response)
|
|
except CheckFailed as ex:
|
|
status = response.status_code
|
|
headers_str = "\n".join(f"{k}: {v}" for k, v in response.headers.items())
|
|
body_text = response.text.strip()
|
|
msg = (
|
|
f"Failed validation {ex} for response:\n"
|
|
f"Status: {status}\n"
|
|
f"Headers:\n{headers_str}\n"
|
|
f"Body (decoded text):\n{body_text if body_text else '[Empty Body]'}\n"
|
|
f"Duration: {duration:.2f} seconds\n"
|
|
f"Request was:\n"
|
|
f"Method: {method}\n"
|
|
f"URL: {get_api_string(QDRANT_HOST, api, path_params)}\n"
|
|
f"Query params: {query_params}\n"
|
|
f"Headers: {qdrant_host_headers()}\n"
|
|
f"Body: {body}\n"
|
|
)
|
|
warnings.warn(msg)
|
|
raise
|
|
|
|
return response
|
|
|
|
# from client implementation:
|
|
# https://github.com/qdrant/qdrant-client/blob/d18cb1702f4cf8155766c7b32d1e4a68af11cd6a/qdrant_client/hybrid/fusion.py#L6C1-L31C25
|
|
def reciprocal_rank_fusion(
|
|
responses: List[List[Any]], limit: int = 10, weights: List[float] = None
|
|
) -> List[Any]:
|
|
"""
|
|
Compute RRF scores for multiple results from different sources.
|
|
|
|
Args:
|
|
responses: List of response lists from different sources
|
|
limit: Maximum number of results to return
|
|
weights: Optional weights for each source. Higher weight = more influence.
|
|
If None, all sources are weighted equally (weight = 1.0).
|
|
"""
|
|
ranking_constant = 2 # the constant mitigates the impact of high rankings by outlier systems
|
|
|
|
def compute_score(pos: int, weight: float = 1.0) -> float:
|
|
if weight <= 0:
|
|
return 0.0
|
|
return 1 / ((pos + 1.0) / weight + ranking_constant - 1.0)
|
|
|
|
scores: Dict[Any, float] = {} # id -> score
|
|
point_pile = {}
|
|
for source_idx, response in enumerate(responses):
|
|
weight = weights[source_idx] if weights else 1.0
|
|
for i, scored_point in enumerate(response):
|
|
if scored_point["id"] in scores:
|
|
scores[scored_point["id"]] += compute_score(i, weight)
|
|
else:
|
|
point_pile[scored_point["id"]] = scored_point
|
|
scores[scored_point["id"]] = compute_score(i, weight)
|
|
|
|
sorted_scores = sorted(scores.items(), key=lambda item: item[1], reverse=True)
|
|
sorted_points = []
|
|
for point_id, score in sorted_scores[:limit]:
|
|
point = point_pile[point_id]
|
|
point["score"] = score
|
|
sorted_points.append(point)
|
|
return sorted_points
|
|
|
|
def distribution_based_score_fusion(responses: List[List[Any]], limit: int = 10) -> List[Any]:
|
|
def normalize(response: List[Any]) -> List[Any]:
|
|
total = sum([point["score"] for point in response])
|
|
mean = total / len(response)
|
|
variance = sum([(point["score"] - mean) ** 2 for point in response]) / (len(response) - 1)
|
|
std_dev = variance ** 0.5
|
|
|
|
min = mean - 3 * std_dev
|
|
max = mean + 3 * std_dev
|
|
|
|
for point in response:
|
|
point["score"] = (point["score"] - min) / (max - min)
|
|
|
|
return response
|
|
|
|
points_map = {}
|
|
for response in responses:
|
|
normalized = normalize(response)
|
|
for point in normalized:
|
|
entry = points_map.get(point["id"])
|
|
if entry is None:
|
|
points_map[point["id"]] = point
|
|
else:
|
|
entry["score"] += point["score"]
|
|
|
|
sorted_points = sorted(points_map.values(), key=lambda item: item['score'], reverse=True)
|
|
|
|
return sorted_points[:limit]
|
|
|
|
|
|
@lru_cache
|
|
def qdrant_host_headers():
|
|
headers = json.loads(QDRANT_HOST_HEADERS)
|
|
return headers
|
|
|
|
|
|
def check_feature_enabled(feature) -> bool:
|
|
response = request_with_validation(
|
|
api='/telemetry',
|
|
method="GET",
|
|
query_params={'details_level': 10},
|
|
)
|
|
assert response.ok
|
|
features = response.json()['result']['app']['features']
|
|
result = features[feature]
|
|
return result
|
|
|
|
def skip_if_no_feature(feature):
|
|
feature_is_enabled = check_feature_enabled(feature)
|
|
if not feature_is_enabled:
|
|
pytest.skip(f"Skipping because the feature {feature} is disabled at runtime.")
|
|
|
|
|
|
def is_distributed_mode() -> bool:
|
|
response = request_with_validation(
|
|
api='/cluster',
|
|
method="GET",
|
|
)
|
|
assert response.ok
|
|
# Standalone nodes report `disabled`; any other status means consensus is enabled.
|
|
return response.json()['result']['status'] != 'disabled'
|
|
|
|
|
|
def skip_if_distributed_mode():
|
|
if is_distributed_mode():
|
|
pytest.skip("Skipping because Qdrant is running in distributed mode.") |