fix(notices): derive scale counts for Redis and search backends (#5687)
This commit is contained in:
+26
-3
@@ -1418,7 +1418,30 @@ def _get_provider_memory_count(memory_instance) -> Optional[int]:
|
||||
try:
|
||||
col_info = getattr(vector_store, "col_info", None)
|
||||
if callable(col_info):
|
||||
return _extract_count(col_info())
|
||||
collection_name = getattr(vector_store, "collection_name", None)
|
||||
if collection_name is None:
|
||||
schema = getattr(vector_store, "schema", None)
|
||||
if isinstance(schema, dict):
|
||||
index = schema.get("index")
|
||||
if isinstance(index, dict):
|
||||
collection_name = index.get("name")
|
||||
if collection_name is not None:
|
||||
try:
|
||||
info = col_info(collection_name)
|
||||
except TypeError:
|
||||
info = col_info()
|
||||
else:
|
||||
info = col_info()
|
||||
value = _extract_count(info)
|
||||
if value is not None:
|
||||
return value
|
||||
|
||||
client = getattr(vector_store, "client", None)
|
||||
client_count = (
|
||||
getattr(client, "count", None) if client is not None and collection_name is not None else None
|
||||
)
|
||||
if callable(client_count):
|
||||
return _extract_count(client_count(index=collection_name))
|
||||
except Exception:
|
||||
return None
|
||||
|
||||
@@ -1430,7 +1453,7 @@ def _extract_count(info: Any) -> Optional[int]:
|
||||
return None
|
||||
|
||||
if isinstance(info, dict):
|
||||
for key in ("count", "points_count", "vectors_count", "indexed_vectors_count"):
|
||||
for key in ("count", "points_count", "vectors_count", "indexed_vectors_count", "num_docs"):
|
||||
value = _coerce_nonnegative_int(info.get(key), None)
|
||||
if value is not None:
|
||||
return value
|
||||
@@ -1445,7 +1468,7 @@ def _extract_count(info: Any) -> Optional[int]:
|
||||
except Exception:
|
||||
return None
|
||||
|
||||
for attr in ("count", "points_count", "vectors_count", "indexed_vectors_count"):
|
||||
for attr in ("count", "points_count", "vectors_count", "indexed_vectors_count", "num_docs"):
|
||||
value = _coerce_nonnegative_int(getattr(info, attr, None), None)
|
||||
if value is not None:
|
||||
return value
|
||||
|
||||
@@ -1403,6 +1403,30 @@ def test_scale_threshold_provider_count_helpers_are_safe():
|
||||
def col_info(self):
|
||||
return {"count": 2300}
|
||||
|
||||
class RedisInfoStore:
|
||||
def __init__(self):
|
||||
self.schema = {"index": {"name": "test_collection"}}
|
||||
|
||||
def count(self):
|
||||
raise RuntimeError("count unavailable")
|
||||
|
||||
def col_info(self, name):
|
||||
assert name == self.schema["index"]["name"]
|
||||
return {"index_name": "test_collection", "num_docs": 2300}
|
||||
|
||||
class SearchMetadataStore:
|
||||
def __init__(self):
|
||||
self.collection_name = "test_collection"
|
||||
self.client = MagicMock()
|
||||
self.client.count.return_value = {"count": 2400}
|
||||
|
||||
def count(self):
|
||||
raise RuntimeError("count unavailable")
|
||||
|
||||
def col_info(self, name):
|
||||
assert name == self.collection_name
|
||||
return {"test_collection": {"settings": {"index": {}}}}
|
||||
|
||||
memory = MagicMock()
|
||||
memory.vector_store = CountStore()
|
||||
assert notices._get_provider_memory_count(memory) == 2100
|
||||
@@ -1410,7 +1434,15 @@ def test_scale_threshold_provider_count_helpers_are_safe():
|
||||
memory.vector_store = FallbackStore()
|
||||
assert notices._get_provider_memory_count(memory) == 2300
|
||||
|
||||
assert notices._extract_count({"points_count": 2400}) == 2400
|
||||
memory.vector_store = RedisInfoStore()
|
||||
assert notices._get_provider_memory_count(memory) == 2300
|
||||
|
||||
search_store = SearchMetadataStore()
|
||||
memory.vector_store = search_store
|
||||
assert notices._get_provider_memory_count(memory) == 2400
|
||||
search_store.client.count.assert_called_once_with(index="test_collection")
|
||||
|
||||
assert notices._extract_count({"points_count": 2500}) == 2500
|
||||
assert notices._extract_count(Info()) == 2200
|
||||
assert notices._extract_count({"count": -1}) is None
|
||||
|
||||
|
||||
Reference in New Issue
Block a user