From 0e02effaf7d6dc87af0e5385f7645d4771511e39 Mon Sep 17 00:00:00 2001 From: Rod Boev Date: Thu, 25 Jun 2026 06:34:44 -0400 Subject: [PATCH] fix(notices): derive scale counts for Redis and search backends (#5687) --- mem0/memory/notices.py | 29 ++++++++++++++++++++++++++--- tests/memory/test_notices.py | 34 +++++++++++++++++++++++++++++++++- 2 files changed, 59 insertions(+), 4 deletions(-) diff --git a/mem0/memory/notices.py b/mem0/memory/notices.py index f2f769d65..d4aeec86b 100644 --- a/mem0/memory/notices.py +++ b/mem0/memory/notices.py @@ -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 diff --git a/tests/memory/test_notices.py b/tests/memory/test_notices.py index f7ac6d1c4..b75b208d3 100644 --- a/tests/memory/test_notices.py +++ b/tests/memory/test_notices.py @@ -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