fix(notices): derive scale counts for Redis and search backends (#5687)

This commit is contained in:
Rod Boev
2026-06-25 06:34:44 -04:00
committed by GitHub
parent 9269a0ad6e
commit 0e02effaf7
2 changed files with 59 additions and 4 deletions
+26 -3
View File
@@ -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
+33 -1
View File
@@ -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