fix: address PR review feedback

- Remove deprecated custom_update_memory_prompt from MemoryConfig entirely
- Add MAX_BATCH=100 chunking guard to embed_batch in openai.py and azure_openai.py
- Move all inline imports (mem0.utils.*) to top-level in main.py
- Remove custom_update_memory_prompt from test fixture

Co-Authored-By: Claude Opus 4.6 (1M context) <noreply@anthropic.com>
This commit is contained in:
Soumil Rathi
2026-04-13 09:01:07 -07:00
parent c0343aec1b
commit 3402c67805
5 changed files with 34 additions and 44 deletions
-4
View File
@@ -60,10 +60,6 @@ class MemoryConfig(BaseModel):
description="Custom instructions for fact extraction",
default=None,
)
custom_update_memory_prompt: Optional[str] = Field(
description="Custom prompt for the update memory (deprecated: use custom_instructions)",
default=None,
)
class AzureConfig(BaseModel):
+14 -6
View File
@@ -55,10 +55,18 @@ class AzureOpenAIEmbedding(EmbeddingBase):
return self.client.embeddings.create(input=[text], model=self.config.model).data[0].embedding
def embed_batch(self, texts, memory_action="add"):
"""Embed multiple texts in a single Azure OpenAI API call."""
"""Embed multiple texts in a single Azure OpenAI API call.
Automatically chunks into batches of 100 to stay within API limits.
"""
MAX_BATCH = 100
texts = [text.replace("\n", " ") for text in texts]
response = self.client.embeddings.create(
input=texts,
model=self.config.model,
)
return [item.embedding for item in sorted(response.data, key=lambda x: x.index)]
all_embeddings = []
for i in range(0, len(texts), MAX_BATCH):
chunk = texts[i : i + MAX_BATCH]
response = self.client.embeddings.create(
input=chunk,
model=self.config.model,
)
all_embeddings.extend(item.embedding for item in sorted(response.data, key=lambda x: x.index))
return all_embeddings
+18 -10
View File
@@ -55,14 +55,22 @@ class OpenAIEmbedding(EmbeddingBase):
return self.client.embeddings.create(**kwargs).data[0].embedding
def embed_batch(self, texts, memory_action="add"):
"""Embed multiple texts in a single OpenAI API call."""
"""Embed multiple texts in a single OpenAI API call.
Automatically chunks into batches of 100 to stay within API limits.
"""
MAX_BATCH = 100
texts = [text.replace("\n", " ") for text in texts]
kwargs = {
"input": texts,
"model": self.config.model,
"encoding_format": "float",
}
if self._pass_dimensions_to_api:
kwargs["dimensions"] = self.config.embedding_dims
response = self.client.embeddings.create(**kwargs)
return [item.embedding for item in sorted(response.data, key=lambda x: x.index)]
all_embeddings = []
for i in range(0, len(texts), MAX_BATCH):
chunk = texts[i : i + MAX_BATCH]
kwargs = {
"input": chunk,
"model": self.config.model,
"encoding_format": "float",
}
if self._pass_dimensions_to_api:
kwargs["dimensions"] = self.config.embedding_dims
response = self.client.embeddings.create(**kwargs)
all_embeddings.extend(item.embedding for item in sorted(response.data, key=lambda x: x.index))
return all_embeddings
+2 -23
View File
@@ -24,7 +24,8 @@ from mem0.configs.prompts import (
get_update_memory_messages,
)
from mem0.utils.lemmatization import lemmatize_for_bm25
from mem0.utils.entity_extraction import extract_entities_batch
from mem0.utils.entity_extraction import extract_entities, extract_entities_batch
from mem0.utils.scoring import ENTITY_BOOST_WEIGHT, get_bm25_params, normalize_bm25, score_and_rank
from mem0.exceptions import ValidationError as Mem0ValidationError
from mem0.memory.base import MemoryBase
from mem0.memory.setup import mem0_dir, setup_config
@@ -191,7 +192,6 @@ class Memory(MemoryBase):
def __init__(self, config: MemoryConfig = MemoryConfig()):
self.config = config
self.custom_update_memory_prompt = self.config.custom_update_memory_prompt
self.embedding_model = EmbedderFactory.create(
self.config.embedder.provider,
self.config.embedder.config,
@@ -217,15 +217,6 @@ class Memory(MemoryBase):
# Entity store is initialized lazily on first use
self._entity_store = None
if self.config.custom_update_memory_prompt:
import warnings
warnings.warn(
"custom_update_memory_prompt is deprecated and has no effect in the v3 pipeline. "
"Use custom_instructions instead.",
DeprecationWarning,
stacklevel=2,
)
if self.config.graph_store.config:
provider = self.config.graph_store.provider
self.graph = GraphStoreFactory.create(provider, self.config)
@@ -1145,10 +1136,6 @@ class Memory(MemoryBase):
return False
def _search_vector_store(self, query, filters, limit, threshold=0.1):
from mem0.utils.lemmatization import lemmatize_for_bm25
from mem0.utils.entity_extraction import extract_entities
from mem0.utils.scoring import get_bm25_params, normalize_bm25, score_and_rank, ENTITY_BOOST_WEIGHT
# Guard against None threshold (backward compat)
if threshold is None:
threshold = 0.1
@@ -1261,8 +1248,6 @@ class Memory(MemoryBase):
Returns:
Dict mapping memory_id (str) -> max entity boost [0, 0.5].
"""
from mem0.utils.scoring import ENTITY_BOOST_WEIGHT
# Deduplicate entities (max 8)
seen = set()
deduped = []
@@ -2494,10 +2479,6 @@ class AsyncMemory(MemoryBase):
return False
async def _search_vector_store(self, query, filters, limit, threshold=0.1):
from mem0.utils.lemmatization import lemmatize_for_bm25
from mem0.utils.entity_extraction import extract_entities
from mem0.utils.scoring import get_bm25_params, normalize_bm25, score_and_rank, ENTITY_BOOST_WEIGHT
if threshold is None:
threshold = 0.1
@@ -2596,8 +2577,6 @@ class AsyncMemory(MemoryBase):
async def _compute_entity_boosts_async(self, query_entities, filters):
"""Async version of entity boost computation."""
from mem0.utils.scoring import ENTITY_BOOST_WEIGHT
seen = set()
deduped = []
for entity_type, entity_text in query_entities[:8]:
-1
View File
@@ -61,7 +61,6 @@ def memory_custom_instance():
config = MemoryConfig(
version="v1.1",
custom_instructions="custom prompt extracting memory in json format",
custom_update_memory_prompt="custom prompt determining memory update",
)
config.graph_store.config = {"some_config": "value"}
return Memory(config)