Files
qdrant-client/tests/embed_tests/test_schema_parser.py
George ff30799c6d fix: fix embed paths (#1278)
* fix: fix embed paths

* tests: add local inference test for complex prefetch
2026-07-23 22:22:29 +07:00

59 lines
1.7 KiB
Python

import pytest
from qdrant_client import models
from qdrant_client.embed.schema_parser import ModelSchemaParser
from qdrant_client.embed.utils import FieldPath
def check_path_recursive(plain_path_parts: list[str], paths: list[FieldPath]) -> bool:
if not plain_path_parts:
return True
for path in paths:
if path.current == plain_path_parts[0]:
return check_path_recursive(plain_path_parts[1:], path.tail)
return False
@pytest.mark.parametrize(
"model",
[
models.Batch,
models.ContextPair,
models.ContextQuery,
models.DiscoverInput,
models.DiscoverQuery,
models.NearestQuery,
models.PointStruct,
models.PointVectors,
models.PointsBatch,
models.PointsList,
models.RecommendInput,
models.RecommendQuery,
models.UpdateVectors,
models.UpdateVectorsOperation,
models.UpsertOperation,
models.Prefetch,
models.QueryGroupsRequest,
models.QueryRequest,
models.QueryRequestBatch,
models.UpdateOperations,
],
)
def test_parser(model):
parser = ModelSchemaParser()
parser.parse_model(model)
for model_name, plain_paths in parser._cache.items():
count = 0
paths = parser.path_cache[model_name]
for plain_path in plain_paths:
count += check_path_recursive(plain_path.split("."), paths)
assert count == len(plain_paths)
# convert_paths must keep every plain path, including ones that are
# prefixes of longer paths (e.g. "query" and "query.nearest")
flattened = {path for field_path in paths for path in field_path.as_str_list()}
assert set(plain_paths) == flattened